From d75451a3cae235bd995072d6827849d343e806be Mon Sep 17 00:00:00 2001 From: RCD <90105158+Finesssee@users.noreply.github.com> Date: Tue, 29 Sep 2026 15:34:04 +0700 Subject: [PATCH] Port upstream 0.68.0: Abacus concurrent credits and billing, session candidates, strict parsing --- rust/src/providers/abacus/mod.rs | 475 ++++++++++++++++------- rust/src/providers/abacus/tests.rs | 590 +++++++++++++++++++++++++++++ 2 files changed, 922 insertions(+), 143 deletions(-) create mode 100644 rust/src/providers/abacus/tests.rs diff --git a/rust/src/providers/abacus/mod.rs b/rust/src/providers/abacus/mod.rs index b59dd8eb32..43fcb572c7 100644 --- a/rust/src/providers/abacus/mod.rs +++ b/rust/src/providers/abacus/mod.rs @@ -3,53 +3,128 @@ //! Fetches compute-point usage and billing info via apps.abacus.ai web APIs. //! Uses browser cookies for authentication. +use std::future::Future; +use std::time::Duration; + use async_trait::async_trait; use chrono::{DateTime, Utc}; use reqwest::Client; -use serde::Deserialize; +use serde::de::DeserializeOwned; +use serde::{Deserialize, Deserializer}; +use serde_json::Value; use crate::core::{ FetchContext, Provider, ProviderError, ProviderFetchResult, ProviderId, ProviderMetadata, - RateWindow, SourceMode, UsageSnapshot, + ProviderStateKind, RateWindow, SourceMode, UsageSnapshot, }; -const COMPUTE_URL: &str = "https://apps.abacus.ai/api/_getOrganizationComputePoints"; -const BILLING_URL: &str = "https://apps.abacus.ai/api/_getBillingInfo"; +#[cfg(test)] +mod tests; + +const ORIGIN: &str = "https://apps.abacus.ai"; +const COMPUTE_PATH: &str = "/api/_getOrganizationComputePoints"; +const BILLING_PATH: &str = "/api/_getBillingInfo"; +/// Parent domain of `apps.abacus.ai`; one browser query covers both hosts. +const COOKIE_DOMAIN: &str = "abacus.ai"; const CREDITS_LABEL: &str = "Credits"; const FALLBACK_MONTHLY_WINDOW_MINUTES: u32 = 30 * 24 * 60; - +const MAX_COOKIE_CANDIDATES: u32 = 5; +const MAX_BODY_BYTES: usize = 1024 * 1024; +const BILLING_BUDGET_CAP: Duration = Duration::from_secs(5); +const REFRESH_DEADLINE_CAP: Duration = Duration::from_secs(90); +const MISSING_SESSION_MESSAGE: &str = "No Abacus AI session found. Please log in to apps.abacus.ai in your browser or paste a Cookie header in manual mode."; + +/// Exact cookie names that carry Abacus session state. CSRF tokens are +/// excluded on purpose: anonymous jars contain them. +const KNOWN_SESSION_COOKIE_NAMES: [&str; 5] = [ + "sessionid", + "session_id", + "session_token", + "auth_token", + "access_token", +]; +/// Substrings that mark a session cookie when no exact name matches. +const SESSION_COOKIE_SUBSTRINGS: [&str; 4] = ["session", "auth", "sid", "jwt"]; +/// Prefixes that mark a non-session cookie even if a substring matches. +const EXCLUDED_COOKIE_PREFIXES: [&str; 5] = ["csrf", "_ga", "_gid", "tracking", "analytics"]; +/// `success:false` messages containing one of these mean the session is bad. +const AUTH_ERROR_KEYWORDS: [&str; 7] = [ + "expired", + "session", + "login", + "authenticate", + "unauthorized", + "unauthenticated", + "forbidden", +]; + +/// `{"success": true, "result": {...}}`. Every field tolerates a wrong type so +/// a malformed envelope becomes a classified error rather than a serde error. #[derive(Debug, Deserialize)] -struct ApiEnvelope { - #[serde(default)] +struct ApiEnvelope { + #[serde(default, deserialize_with = "lenient_true")] success: bool, - result: Option, + #[serde(default)] + result: Option, + #[serde(default, deserialize_with = "lenient_string")] + error: Option, } #[derive(Debug, Deserialize)] #[serde(rename_all = "camelCase")] +struct RawComputePoints { + #[serde(default, deserialize_with = "lenient_finite_number")] + total_compute_points: Option, + #[serde(default, deserialize_with = "lenient_finite_number")] + compute_points_left: Option, +} + +#[derive(Debug, PartialEq)] struct ComputePoints { - #[serde(default)] - total_compute_points: f64, - #[serde(default)] - compute_points_left: f64, + total: f64, + left: f64, } #[derive(Debug, Deserialize)] #[serde(rename_all = "camelCase")] struct BillingInfo { - #[serde(default)] + #[serde(default, deserialize_with = "lenient_string")] next_billing_date: Option, - #[serde(default)] + #[serde(default, deserialize_with = "lenient_string")] current_tier: Option, } +fn lenient_true<'de, D: Deserializer<'de>>(deserializer: D) -> Result { + Ok(Value::deserialize(deserializer)? == Value::Bool(true)) +} + +fn lenient_string<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { + Ok(match Value::deserialize(deserializer)? { + Value::String(text) => Some(text), + _ => None, + }) +} + +fn lenient_finite_number<'de, D: Deserializer<'de>>( + deserializer: D, +) -> Result, D::Error> { + Ok(Value::deserialize(deserializer)? + .as_f64() + .filter(|number| number.is_finite())) +} + pub struct AbacusProvider { metadata: ProviderMetadata, client: Client, + origin: String, } impl AbacusProvider { pub fn new() -> Self { + Self::with_origin(ORIGIN) + } + + fn with_origin(origin: &str) -> Self { Self { metadata: ProviderMetadata { id: ProviderId::Abacus, @@ -64,10 +139,12 @@ impl AbacusProvider { status_page_url: None, tertiary_label_key: None, }, + // Every request sets its own timeout; the client has none so a + // configured web timeout above 30 s is honored. client: crate::core::credentialed_http_client_builder() - .timeout(std::time::Duration::from_secs(30)) .build() .unwrap_or_else(|_| Client::new()), + origin: origin.to_string(), } } @@ -75,9 +152,8 @@ impl AbacusProvider { compute: ComputePoints, billing: Option, ) -> Result { - let total = compute.total_compute_points.max(0.0); - let left = compute.compute_points_left.max(0.0); - let used = (total - left).max(0.0); + let ComputePoints { total, left } = compute; + let used = total - left; let percent = if total > 0.0 { ((used / total) * 100.0).clamp(0.0, 100.0) } else { @@ -87,8 +163,7 @@ impl AbacusProvider { let resets_at = billing .as_ref() .and_then(|b| b.next_billing_date.as_deref()) - .and_then(|ts| DateTime::parse_from_rfc3339(ts).ok()) - .map(|dt| dt.with_timezone(&Utc)); + .and_then(parse_billing_date); let primary = RateWindow::with_details( percent, RateWindow::monthly_window_minutes(resets_at).or(Some(FALLBACK_MONTHLY_WINDOW_MINUTES)), @@ -106,66 +181,251 @@ impl AbacusProvider { Ok(snapshot) } - async fn fetch_compute(&self, cookie_header: &str) -> Result { - let resp = self + async fn fetch_compute( + &self, + cookie_header: &str, + timeout: Duration, + ) -> Result { + let request = self .client - .get(COMPUTE_URL) + .get(format!("{}{COMPUTE_PATH}", self.origin)) .header("Cookie", cookie_header) .header("Accept", "application/json") - .send() - .await?; - - let status = resp.status(); - if status.as_u16() == 401 || status.as_u16() == 403 { - return Err(ProviderError::AuthRequired); - } - if !status.is_success() { - return Err(ProviderError::Other(format!( - "Abacus compute API returned {}", - status - ))); - } - - let body = resp.text().await?; - let env: ApiEnvelope = serde_json::from_str(&body) - .map_err(|e| ProviderError::Parse(format!("Failed to parse compute points: {}", e)))?; - - if !env.success { - return Err(ProviderError::AuthRequired); + .header("Content-Type", "application/json") + .timeout(timeout); + let raw: RawComputePoints = send_envelope(request).await?; + match (raw.total_compute_points, raw.compute_points_left) { + (Some(total), Some(left)) => Ok(ComputePoints { total, left }), + _ => Err(parse_failure( + "Missing credit fields in compute points response", + )), } - env.result - .ok_or_else(|| ProviderError::Parse("Missing compute points result".to_string())) } - async fn fetch_billing(&self, cookie_header: &str) -> Option { - let resp = self + /// Billing only enriches the credits result, so every failure is `None`. + async fn fetch_billing(&self, cookie_header: &str, budget: Duration) -> Option { + let request = self .client - .post(BILLING_URL) + .post(format!("{}{BILLING_PATH}", self.origin)) .header("Cookie", cookie_header) - .header("Content-Type", "application/json") .header("Accept", "application/json") + .header("Content-Type", "application/json") .body("{}") - .send() - .await - .ok()?; - - if !resp.status().is_success() { - return None; + .timeout(budget); + match send_envelope(request).await { + Ok(billing) => Some(billing), + Err(error) => { + tracing::debug!(%error, "Abacus billing info unavailable; using fallback window"); + None + } } - - let body = resp.text().await.ok()?; - let env: ApiEnvelope = serde_json::from_str(&body).ok()?; - env.result } async fn fetch_with_cookies( &self, cookie_header: &str, + request_timeout: Duration, ) -> Result { - let compute = self.fetch_compute(cookie_header).await?; - let billing = self.fetch_billing(cookie_header).await; + let budget = request_timeout.min(BILLING_BUDGET_CAP); + let (compute, billing) = credits_with_billing( + self.fetch_compute(cookie_header, request_timeout), + self.fetch_billing(cookie_header, budget), + budget, + ) + .await?; Self::build_snapshot(compute, billing) } + + /// Try each imported session in order. Any failure moves on to the next + /// one; the last failure is reported when none succeeds. + async fn fetch_with_candidates( + &self, + candidates: &[(String, String)], + request_timeout: Duration, + ) -> Result { + let mut last_error = None; + for (label, cookie_header) in candidates { + match self + .fetch_with_cookies(cookie_header, request_timeout) + .await + { + Ok(snapshot) => return Ok(snapshot), + Err(error) => { + tracing::debug!(browser = %label, %error, "Abacus session candidate failed"); + last_error = Some(error); + } + } + } + Err(last_error.unwrap_or_else(|| ProviderError::Other(MISSING_SESSION_MESSAGE.to_string()))) + } + + async fn fetch_web(&self, ctx: &FetchContext) -> Result { + let request_timeout = request_timeout(ctx.web_timeout); + let refresh = async { + // A pasted header is exclusive: never fall back to browser cookies. + if let Some(cookie_header) = ctx.manual_cookie_header.as_deref() { + return self + .fetch_with_cookies(cookie_header, request_timeout) + .await; + } + let candidates = session_candidates( + crate::providers::browser_cookie_headers_for_domain(COOKIE_DOMAIN), + )?; + self.fetch_with_candidates(&candidates, request_timeout) + .await + }; + tokio::time::timeout(refresh_timeout(request_timeout), refresh) + .await + .map_err(|_| ProviderError::Timeout)? + } +} + +/// Web timeout clamped to the 1..=90 s request range. +fn request_timeout(web_timeout: u64) -> Duration { + Duration::from_secs(web_timeout.clamp(1, 90)) +} + +/// Total refresh deadline: room for every candidate plus one billing budget, +/// never more than 90 s. +fn refresh_timeout(request_timeout: Duration) -> Duration { + (request_timeout * MAX_COOKIE_CANDIDATES + request_timeout.min(BILLING_BUDGET_CAP)) + .min(REFRESH_DEADLINE_CAP) +} + +/// Run the required credits request and the optional billing request +/// concurrently. Billing gets `budget` from the start; when it errors or runs +/// out it is dropped, and a credits failure cancels it immediately. +async fn credits_with_billing( + credits: impl Future>, + billing: impl Future>, + budget: Duration, +) -> Result<(T, Option), ProviderError> { + let billing = async { tokio::time::timeout(budget, billing).await.ok().flatten() }; + tokio::pin!(credits, billing); + let mut billing_result = None; + let credits = loop { + tokio::select! { + result = &mut credits => break result?, + result = &mut billing, if billing_result.is_none() => billing_result = Some(result), + } + }; + let billing = match billing_result { + Some(result) => result, + None => billing.await, + }; + Ok((credits, billing)) +} + +/// Send `request` and decode the `{success, result}` envelope into `T`. +async fn send_envelope( + request: reqwest::RequestBuilder, +) -> Result { + let response = request.send().await?; + let status = response.status().as_u16(); + if status == 401 || status == 403 { + return Err(ProviderError::AuthRequired); + } + if status != 200 { + return Err(ProviderError::Other(format!( + "Abacus AI API error: HTTP {status}" + ))); + } + let body = crate::providers::read_bounded_response(response, MAX_BODY_BYTES) + .await + .map_err(|error| match error { + crate::providers::BoundedBodyError::TooLarge => parse_failure("response is too large"), + crate::providers::BoundedBodyError::Read(error) => ProviderError::Network(error), + })?; + decode_envelope(&body) +} + +fn decode_envelope(body: &[u8]) -> Result { + let root: Value = serde_json::from_slice(body).map_err(|_| parse_failure("invalid JSON"))?; + if !root.is_object() { + return Err(parse_failure("invalid response")); + } + let envelope: ApiEnvelope = + serde_json::from_value(root).map_err(|_| parse_failure("invalid response"))?; + match envelope.result { + Some(result) if envelope.success && result.is_object() => { + serde_json::from_value(result).map_err(|_| parse_failure("invalid result")) + } + _ => { + let message = envelope + .error + .map_or_else(|| "unknown error".to_string(), |text| text.to_lowercase()); + if AUTH_ERROR_KEYWORDS + .iter() + .any(|keyword| message.contains(keyword)) + { + Err(ProviderError::AuthRequired) + } else { + Err(parse_failure(&message)) + } + } + } +} + +fn parse_failure(message: &str) -> ProviderError { + ProviderError::Parse(format!("Could not parse Abacus AI usage: {message}")) +} + +/// Billing dates must look like ISO 8601 date-times before they are trusted. +fn parse_billing_date(value: &str) -> Option> { + let bytes = value.as_bytes(); + let looks_iso = bytes.len() > 10 + && bytes[..10] + .iter() + .enumerate() + .all(|(index, byte)| match index { + 4 | 7 => *byte == b'-', + _ => byte.is_ascii_digit(), + }) + && bytes[10] == b'T'; + if !looks_iso { + return None; + } + DateTime::parse_from_rfc3339(value) + .ok() + .map(|date| date.with_timezone(&Utc)) +} + +/// Browser cookie sets worth trying, Chrome first, at most five. Sets without +/// a session cookie (anonymous or marketing-only jars) are skipped. +fn session_candidates( + headers: Result, ProviderError>, +) -> Result, ProviderError> { + let mut candidates = match headers { + Ok(headers) => headers, + Err(ProviderError::NoCookies) => Vec::new(), + Err(error) => return Err(error), + }; + candidates.retain(|(_, header)| has_session_cookie(header)); + candidates.sort_by_key(|(label, _)| label != "Google Chrome"); + candidates.truncate(MAX_COOKIE_CANDIDATES as usize); + Ok(candidates) +} + +fn has_session_cookie(cookie_header: &str) -> bool { + cookie_header.split(';').any(|pair| { + let name = pair + .split_once('=') + .map_or(pair, |(name, _)| name) + .trim() + .to_ascii_lowercase(); + if KNOWN_SESSION_COOKIE_NAMES.contains(&name.as_str()) { + return true; + } + if EXCLUDED_COOKIE_PREFIXES + .iter() + .any(|prefix| name.starts_with(prefix)) + { + return false; + } + SESSION_COOKIE_SUBSTRINGS + .iter() + .any(|needle| name.contains(needle)) + }) } fn format_credit_detail(used: f64, total: f64) -> String { @@ -235,28 +495,23 @@ impl Provider for AbacusProvider { match ctx.source_mode { SourceMode::Auto | SourceMode::Web => { - if let Some(ref cookie_header) = ctx.manual_cookie_header { - let usage = self.fetch_with_cookies(cookie_header).await?; - return Ok(ProviderFetchResult::new(usage, "web")); - } - - match crate::providers::browser_cookie_header(&["apps.abacus.ai"]) { - Ok(cookie_header) => match self.fetch_with_cookies(&cookie_header).await { - Ok(usage) => return Ok(ProviderFetchResult::new(usage, "web")), - Err(ProviderError::AuthRequired) => {} - Err(e) => return Err(e), - }, - Err(ProviderError::NoCookies) => {} - Err(e) => return Err(e), - } - - Err(ProviderError::AuthRequired) + let usage = self.fetch_web(ctx).await?; + Ok(ProviderFetchResult::new(usage, "web")) } SourceMode::Cli => Err(ProviderError::UnsupportedSource(SourceMode::Cli)), SourceMode::OAuth => Err(ProviderError::UnsupportedSource(SourceMode::OAuth)), } } + fn error_state_kind(&self, error: &ProviderError) -> ProviderStateKind { + match error { + ProviderError::Other(message) if message == MISSING_SESSION_MESSAGE => { + ProviderStateKind::NeedsAuthentication + } + _ => error.state_kind(), + } + } + fn available_sources(&self) -> Vec { vec![SourceMode::Auto, SourceMode::Web] } @@ -269,69 +524,3 @@ impl Provider for AbacusProvider { false } } - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn parses_compute_points_and_tier() { - let compute = ComputePoints { - total_compute_points: 1000.0, - compute_points_left: 750.0, - }; - let billing = BillingInfo { - next_billing_date: Some("2025-03-01T00:00:00Z".into()), - current_tier: Some("Pro".into()), - }; - let snap = AbacusProvider::build_snapshot(compute, Some(billing)).unwrap(); - assert!((snap.primary.used_percent - 25.0).abs() < 0.001); - assert_eq!( - snap.primary.reset_description.as_deref(), - Some("250 / 1,000 credits") - ); - assert_eq!( - snap.primary.resets_at, - Some( - DateTime::parse_from_rfc3339("2025-03-01T00:00:00Z") - .unwrap() - .with_timezone(&Utc) - ) - ); - assert_eq!(snap.primary.window_minutes, Some(28 * 24 * 60)); - assert_eq!(snap.primary_label.as_deref(), Some(CREDITS_LABEL)); - assert_eq!(snap.login_method.as_deref(), Some("Pro")); - assert!(snap.account_email.is_none()); - assert!(snap.account_organization.is_none()); - } - - #[test] - fn handles_missing_billing() { - let compute = ComputePoints { - total_compute_points: 500.0, - compute_points_left: 500.0, - }; - let snap = AbacusProvider::build_snapshot(compute, None).unwrap(); - assert!((snap.primary.used_percent - 0.0).abs() < f64::EPSILON); - assert_eq!( - snap.primary.reset_description.as_deref(), - Some("0 / 500 credits") - ); - assert_eq!( - snap.primary.window_minutes, - Some(FALLBACK_MONTHLY_WINDOW_MINUTES) - ); - assert!(snap.primary.resets_at.is_none()); - assert_eq!(snap.primary_label.as_deref(), Some(CREDITS_LABEL)); - assert!(snap.login_method.is_none()); - } - - #[test] - fn formats_credit_details_with_grouping_and_fraction() { - assert_eq!( - format_credit_detail(12_345.0, 50_000.0), - "12,345 / 50,000 credits" - ); - assert_eq!(format_credit_detail(42.5, 100.0), "42.5 / 100 credits"); - } -} diff --git a/rust/src/providers/abacus/tests.rs b/rust/src/providers/abacus/tests.rs new file mode 100644 index 0000000000..6bbf77c33b --- /dev/null +++ b/rust/src/providers/abacus/tests.rs @@ -0,0 +1,590 @@ +use std::time::Duration; + +use mockito::Matcher; +use tokio::time::Instant; + +use super::*; + +// Wire fixtures copied from upstream `TestsPlugin/AbacusPluginTests.swift` (v0.68.0). +const POINTS: &str = + r#"{"success":true,"result":{"totalComputePoints":1000,"computePointsLeft":750}}"#; +const BILLING: &str = + r#"{"success":true,"result":{"currentTier":"Pro","nextBillingDate":"2024-03-31T12:30:00Z"}}"#; +const JSON: &str = "application/json"; + +fn compute(total: f64, left: f64) -> ComputePoints { + ComputePoints { total, left } +} + +fn candidate(label: &str, header: &str) -> (String, String) { + (label.to_string(), header.to_string()) +} + +fn web_context(manual_cookie_header: Option<&str>) -> FetchContext { + FetchContext { + source_mode: SourceMode::Web, + web_timeout: 2, + manual_cookie_header: manual_cookie_header.map(str::to_string), + ..FetchContext::default() + } +} + +#[test] +fn parses_compute_points_and_tier() { + let billing = BillingInfo { + next_billing_date: Some("2025-03-01T00:00:00Z".into()), + current_tier: Some("Pro".into()), + }; + let snap = AbacusProvider::build_snapshot(compute(1000.0, 750.0), Some(billing)).unwrap(); + assert!((snap.primary.used_percent - 25.0).abs() < 0.001); + assert_eq!( + snap.primary.reset_description.as_deref(), + Some("250 / 1,000 credits") + ); + assert_eq!( + snap.primary.resets_at, + Some( + DateTime::parse_from_rfc3339("2025-03-01T00:00:00Z") + .unwrap() + .with_timezone(&Utc) + ) + ); + assert_eq!(snap.primary.window_minutes, Some(28 * 24 * 60)); + assert_eq!(snap.primary_label.as_deref(), Some(CREDITS_LABEL)); + assert_eq!(snap.login_method.as_deref(), Some("Pro")); + assert!(snap.account_email.is_none()); + assert!(snap.account_organization.is_none()); +} + +#[test] +fn handles_missing_billing() { + let snap = AbacusProvider::build_snapshot(compute(500.0, 500.0), None).unwrap(); + assert!((snap.primary.used_percent - 0.0).abs() < f64::EPSILON); + assert_eq!( + snap.primary.reset_description.as_deref(), + Some("0 / 500 credits") + ); + assert_eq!( + snap.primary.window_minutes, + Some(FALLBACK_MONTHLY_WINDOW_MINUTES) + ); + assert!(snap.primary.resets_at.is_none()); + assert_eq!(snap.primary_label.as_deref(), Some(CREDITS_LABEL)); + assert!(snap.login_method.is_none()); +} + +#[test] +fn upstream_credit_fixtures_match_percent_and_detail() { + // (total, left, percent, detail) from the upstream fixture matrix. + for (total, left, percent, detail) in [ + (1000.0, 750.0, 25.0, "250 / 1,000 credits"), + (500.0, 500.0, 0.0, "0 / 500 credits"), + (1000.0, -500.0, 100.0, "1,500 / 1,000 credits"), + (0.0, 0.0, 0.0, "0 / 0 credits"), + (100.0, 57.5, 42.5, "42.5 / 100 credits"), + // 1000.5 rounds half-even, like NumberFormatter and the plugin helper. + (2000.0, 999.5, 50.025, "1,000 / 2,000 credits"), + ] { + let snap = AbacusProvider::build_snapshot(compute(total, left), None).unwrap(); + assert!( + (snap.primary.used_percent - percent).abs() < 1e-9, + "{total}/{left}" + ); + assert_eq!(snap.primary.reset_description.as_deref(), Some(detail)); + } +} + +#[test] +fn formats_credit_details_with_grouping_and_fraction() { + assert_eq!( + format_credit_detail(12_345.0, 50_000.0), + "12,345 / 50,000 credits" + ); + assert_eq!(format_credit_detail(42.5, 100.0), "42.5 / 100 credits"); +} + +#[test] +fn billing_date_must_look_like_an_iso_date_time() { + assert!(parse_billing_date("2024-03-31T12:30:00Z").is_some()); + assert!(parse_billing_date("2024-03-31T12:30:00+02:00").is_some()); + for rejected in [ + "not-a-date", + "2024-03-31", + "2024/03/31T12:30:00Z", + "", + "2024-03-31 12:30:00Z", + ] { + assert!(parse_billing_date(rejected).is_none(), "{rejected}"); + } +} + +#[test] +fn envelope_success_decodes_credit_fields() { + let points: RawComputePoints = decode_envelope(POINTS.as_bytes()).unwrap(); + assert_eq!(points.total_compute_points, Some(1000.0)); + assert_eq!(points.compute_points_left, Some(750.0)); + let billing: BillingInfo = decode_envelope(BILLING.as_bytes()).unwrap(); + assert_eq!(billing.current_tier.as_deref(), Some("Pro")); + assert_eq!( + billing.next_billing_date.as_deref(), + Some("2024-03-31T12:30:00Z") + ); +} + +#[test] +fn envelope_failure_messages_split_auth_from_parse() { + for message in [ + "session expired", + "Please LOGIN first", + "could not authenticate", + "Unauthorized", + "user is unauthenticated", + "Forbidden", + ] { + let body = format!(r#"{{"success":false,"error":"{message}"}}"#); + assert!( + matches!( + decode_envelope::(body.as_bytes()), + Err(ProviderError::AuthRequired) + ), + "{message}" + ); + } + for body in [ + r#"{"success":false,"error":"internal failure"}"#, + r#"{"success":false}"#, + r#"{"success":true}"#, + r#"{"success":true,"result":[]}"#, + r#"{"success":"true","result":{}}"#, + r#"{"success":true,"result":null,"error":5}"#, + ] { + assert!( + matches!( + decode_envelope::(body.as_bytes()), + Err(ProviderError::Parse(_)) + ), + "{body}" + ); + } +} + +#[test] +fn malformed_bodies_are_parse_failures() { + for body in ["error", "", "[]", "\"text\"", "null", "7"] { + assert!( + matches!( + decode_envelope::(body.as_bytes()), + Err(ProviderError::Parse(_)) + ), + "{body}" + ); + } +} + +#[test] +fn non_numeric_credit_fields_are_missing() { + for body in [ + r#"{"success":true,"result":{}}"#, + r#"{"success":true,"result":{"totalComputePoints":"1000","computePointsLeft":750}}"#, + r#"{"success":true,"result":{"totalComputePoints":1000,"computePointsLeft":null}}"#, + ] { + let points: RawComputePoints = decode_envelope(body.as_bytes()).unwrap(); + assert!( + points.total_compute_points.is_none() || points.compute_points_left.is_none(), + "{body}" + ); + } +} + +#[test] +fn session_cookie_names_follow_upstream_filter() { + for accepted in [ + "sessionid=1", + "SESSION_ID=1", + "_ga=1; session_token=2", + "access_token=1", + "foo=1; abacus_session=2", + "userAuth=1", + "connect.sid=1", + "jwt_value=1", + ] { + assert!(has_session_cookie(accepted), "{accepted}"); + } + for rejected in [ + "csrftoken=1", + "_ga=1; _gid=2", + "tracking_session=1", + "analytics_sid=1", + "theme=dark; locale=en", + "csrf_session=1", + "", + ] { + assert!(!has_session_cookie(rejected), "{rejected}"); + } +} + +#[test] +fn candidates_are_chrome_first_filtered_and_bounded() { + let mut headers = vec![ + candidate("Firefox", "sessionid=firefox"), + candidate("Microsoft Edge", "csrftoken=anon"), + candidate("Google Chrome", "sessionid=chrome"), + ]; + for index in 0..5 { + headers.push(candidate("Brave", &format!("sessionid=extra{index}"))); + } + let candidates = session_candidates(Ok(headers)).unwrap(); + assert_eq!(candidates.len(), MAX_COOKIE_CANDIDATES as usize); + assert_eq!( + candidates[0], + candidate("Google Chrome", "sessionid=chrome") + ); + assert_eq!(candidates[1], candidate("Firefox", "sessionid=firefox")); + assert!( + candidates + .iter() + .all(|(label, _)| label != "Microsoft Edge") + ); +} + +#[test] +fn candidate_lookup_errors_are_classified() { + assert!( + session_candidates(Err(ProviderError::NoCookies)) + .unwrap() + .is_empty() + ); + assert!(matches!( + session_candidates(Err(ProviderError::Other("locked".into()))), + Err(ProviderError::Other(_)) + )); +} + +#[test] +fn timeouts_match_upstream_budget() { + for (web, request, refresh) in [ + (0, 1, 6), + (1, 1, 6), + (2, 2, 12), + (15, 15, 80), + (60, 60, 90), + (500, 90, 90), + ] { + assert_eq!(request_timeout(web), Duration::from_secs(request), "{web}"); + assert_eq!( + refresh_timeout(request_timeout(web)), + Duration::from_secs(refresh), + "{web}" + ); + } +} + +#[tokio::test(start_paused = true)] +async fn slow_billing_is_bounded_and_keeps_credits() { + let started = Instant::now(); + let (points, billing) = credits_with_billing( + async { Ok::<_, ProviderError>(1) }, + async { + tokio::time::sleep(Duration::from_secs(30)).await; + Some("late") + }, + Duration::from_secs(5), + ) + .await + .unwrap(); + assert_eq!((points, billing), (1, None)); + assert_eq!(started.elapsed(), Duration::from_secs(5)); +} + +#[tokio::test(start_paused = true)] +async fn billing_runs_concurrently_with_slow_credits() { + let started = Instant::now(); + let (points, billing) = credits_with_billing( + async { + tokio::time::sleep(Duration::from_secs(8)).await; + Ok::<_, ProviderError>(1) + }, + async { + tokio::time::sleep(Duration::from_secs(3)).await; + Some("billing") + }, + Duration::from_secs(5), + ) + .await + .unwrap(); + assert_eq!((points, billing), (1, Some("billing"))); + assert_eq!(started.elapsed(), Duration::from_secs(8)); +} + +#[tokio::test(start_paused = true)] +async fn credits_failure_cancels_billing_immediately() { + let started = Instant::now(); + let result = credits_with_billing( + async { Err::(ProviderError::AuthRequired) }, + async { + tokio::time::sleep(Duration::from_secs(30)).await; + Some(()) + }, + Duration::from_secs(5), + ) + .await; + assert!(matches!(result, Err(ProviderError::AuthRequired))); + assert_eq!(started.elapsed(), Duration::ZERO); +} + +#[tokio::test] +async fn sends_upstream_wire_requests_and_builds_snapshot() { + let mut server = mockito::Server::new_async().await; + let credits = server + .mock("GET", "/api/_getOrganizationComputePoints") + .match_header("cookie", "sessionid=fixture") + .match_header("accept", JSON) + .match_header("content-type", JSON) + .with_body(POINTS) + .expect(1) + .create_async() + .await; + let billing = server + .mock("POST", "/api/_getBillingInfo") + .match_header("cookie", "sessionid=fixture") + .match_header("accept", JSON) + .match_header("content-type", JSON) + .match_body("{}") + .with_body(BILLING) + .expect(1) + .create_async() + .await; + + let provider = AbacusProvider::with_origin(&server.url()); + let result = provider + .fetch_usage(&web_context(Some("sessionid=fixture"))) + .await + .unwrap(); + + credits.assert_async().await; + billing.assert_async().await; + assert_eq!(result.source_label, "web"); + let usage = result.usage; + assert!((usage.primary.used_percent - 25.0).abs() < 1e-9); + assert_eq!(usage.login_method.as_deref(), Some("Pro")); + assert_eq!( + usage.primary.resets_at, + Some( + DateTime::parse_from_rfc3339("2024-03-31T12:30:00Z") + .unwrap() + .with_timezone(&Utc) + ) + ); +} + +#[tokio::test] +async fn billing_failures_keep_credits_and_fallback_window() { + let cases: [(&str, u16, &str); 5] = [ + ("status", 500, "unavailable"), + ( + "auth envelope", + 200, + r#"{"success":false,"error":"session expired"}"#, + ), + ("unauthorized", 401, ""), + ("html", 200, "error"), + ( + "bad date", + 200, + r#"{"success":true,"result":{"currentTier":"Pro","nextBillingDate":"not-a-date"}}"#, + ), + ]; + for (name, status, body) in cases { + let mut server = mockito::Server::new_async().await; + server + .mock("GET", "/api/_getOrganizationComputePoints") + .with_body(POINTS) + .create_async() + .await; + server + .mock("POST", "/api/_getBillingInfo") + .with_status(status.into()) + .with_body(body) + .create_async() + .await; + + let provider = AbacusProvider::with_origin(&server.url()); + let usage = provider + .fetch_with_cookies("sessionid=fixture", Duration::from_secs(2)) + .await + .unwrap(); + assert!((usage.primary.used_percent - 25.0).abs() < 1e-9, "{name}"); + assert_eq!( + usage.primary.reset_description.as_deref(), + Some("250 / 1,000 credits"), + "{name}" + ); + assert!(usage.primary.resets_at.is_none(), "{name}"); + assert_eq!( + usage.primary.window_minutes, + Some(FALLBACK_MONTHLY_WINDOW_MINUTES), + "{name}" + ); + let expected_tier = (name == "bad date").then_some("Pro"); + assert_eq!(usage.login_method.as_deref(), expected_tier, "{name}"); + } +} + +type ErrorCheck = fn(&ProviderError) -> bool; + +#[tokio::test] +async fn required_failures_are_classified() { + let cases: [(u16, &str, ErrorCheck); 5] = [ + (401, "", |e| matches!(e, ProviderError::AuthRequired)), + (403, "", |e| matches!(e, ProviderError::AuthRequired)), + ( + 500, + "boom", + |e| matches!(e, ProviderError::Other(m) if m == "Abacus AI API error: HTTP 500"), + ), + ( + 200, + r#"{"success":true,"result":{}}"#, + |e| matches!(e, ProviderError::Parse(m) if m.contains("Missing credit fields")), + ), + (200, "[]", |e| matches!(e, ProviderError::Parse(_))), + ]; + for (status, body, check) in cases { + let mut server = mockito::Server::new_async().await; + server + .mock("GET", "/api/_getOrganizationComputePoints") + .with_status(status.into()) + .with_body(body) + .create_async() + .await; + server + .mock("POST", "/api/_getBillingInfo") + .with_body(BILLING) + .create_async() + .await; + let provider = AbacusProvider::with_origin(&server.url()); + let error = provider + .fetch_with_cookies("sessionid=fixture", Duration::from_secs(2)) + .await + .unwrap_err(); + assert!(check(&error), "{status} {body}: {error:?}"); + } +} + +/// Mount a credits endpoint that answers `status` for one cookie header. +async fn credits_for( + server: &mut mockito::ServerGuard, + cookie: &str, + status: usize, + body: &str, +) -> mockito::Mock { + server + .mock("GET", "/api/_getOrganizationComputePoints") + .match_header("cookie", cookie) + .with_status(status) + .with_body(body) + .expect(1) + .create_async() + .await +} + +#[tokio::test] +async fn failed_candidates_advance_to_the_next_session() { + for (stale_status, stale_body) in [ + (401, ""), + (200, "[]"), + (200, r#"{"success":true,"result":{}}"#), + (500, "boom"), + ] { + let mut server = mockito::Server::new_async().await; + let stale = credits_for(&mut server, "session=stale", stale_status, stale_body).await; + let fresh = credits_for(&mut server, "session=fresh", 200, POINTS).await; + server + .mock("POST", "/api/_getBillingInfo") + .with_body(BILLING) + .create_async() + .await; + + let provider = AbacusProvider::with_origin(&server.url()); + let usage = provider + .fetch_with_candidates( + &[ + candidate("Google Chrome", "session=stale"), + candidate("Firefox", "session=fresh"), + ], + Duration::from_secs(2), + ) + .await + .unwrap(); + assert!((usage.primary.used_percent - 25.0).abs() < 1e-9); + stale.assert_async().await; + fresh.assert_async().await; + } +} + +#[tokio::test] +async fn exhausted_candidates_report_the_last_error() { + let mut server = mockito::Server::new_async().await; + let first = credits_for(&mut server, "session=one", 500, "boom").await; + let second = credits_for(&mut server, "session=two", 401, "").await; + server + .mock("POST", "/api/_getBillingInfo") + .with_body(BILLING) + .create_async() + .await; + + let provider = AbacusProvider::with_origin(&server.url()); + let error = provider + .fetch_with_candidates( + &[candidate("A", "session=one"), candidate("B", "session=two")], + Duration::from_secs(2), + ) + .await + .unwrap_err(); + assert!(matches!(error, ProviderError::AuthRequired)); + first.assert_async().await; + second.assert_async().await; +} + +#[tokio::test] +async fn no_candidates_report_the_missing_session_message() { + let provider = AbacusProvider::with_origin("http://127.0.0.1:9"); + let error = provider + .fetch_with_candidates(&[], Duration::from_secs(2)) + .await + .unwrap_err(); + assert!(matches!(&error, ProviderError::Other(m) if m == MISSING_SESSION_MESSAGE)); + assert_eq!( + provider.error_state_kind(&error), + ProviderStateKind::NeedsAuthentication + ); + assert_eq!( + provider.error_state_kind(&ProviderError::Other("other".into())), + ProviderStateKind::Unknown + ); +} + +#[tokio::test] +async fn manual_cookie_is_exclusive_and_errors_propagate() { + let mut server = mockito::Server::new_async().await; + let credits = server + .mock("GET", "/api/_getOrganizationComputePoints") + .match_header("cookie", Matcher::Exact("session=stale".into())) + .with_status(401) + .expect(1) + .create_async() + .await; + server + .mock("POST", "/api/_getBillingInfo") + .with_body(BILLING) + .create_async() + .await; + + let provider = AbacusProvider::with_origin(&server.url()); + let error = provider + .fetch_usage(&web_context(Some("session=stale"))) + .await + .unwrap_err(); + assert!(matches!(error, ProviderError::AuthRequired)); + credits.assert_async().await; +}