diff --git a/rust/src/providers/alibabatokenplan/cli.rs b/rust/src/providers/alibabatokenplan/cli.rs index 15966ff499..400bd1397a 100644 --- a/rust/src/providers/alibabatokenplan/cli.rs +++ b/rust/src/providers/alibabatokenplan/cli.rs @@ -192,66 +192,33 @@ mod tests { #[test] fn regional_cli_arguments_match_bailian_contract() { - assert_eq!( - cli_arguments(AlibabaTokenPlanRegion::Cn), - vec![ - "usage", - "token-plan", - "--console-region", - "cn-beijing", - "--console-site", - "domestic", - "--output", - "json" - ] - ); - assert_eq!( - cli_arguments(AlibabaTokenPlanRegion::CnPersonal), - vec![ - "console", - "call", - "--api", - "zeldaHttp.apikeyMgr./tokenplan/personal/api/v2/usage", - "--data", - "{}", - "--console-region", - "cn-beijing", - "--console-site", - "domestic", - "--output", - "json" - ] - ); - assert_eq!( - cli_arguments(AlibabaTokenPlanRegion::IntlPersonal), - vec![ - "console", - "call", - "--api", - "zeldaHttp.apikeyMgr./tokenplan/personal/api/v2/usage", - "--data", - "{}", - "--console-region", - "ap-southeast-1", - "--console-site", - "international", - "--output", - "json" - ] - ); - assert_eq!( - cli_arguments(AlibabaTokenPlanRegion::Intl), - vec![ - "usage", - "token-plan", + use AlibabaTokenPlanRegion::*; + const TEAM: &[&str] = &["usage", "token-plan"]; + const PERSONAL: &[&str] = &[ + "console", + "call", + "--api", + "zeldaHttp.apikeyMgr./tokenplan/personal/api/v2/usage", + "--data", + "{}", + ]; + for (region, head, console_region, console_site) in [ + (Cn, TEAM, "cn-beijing", "domestic"), + (CnPersonal, PERSONAL, "cn-beijing", "domestic"), + (IntlPersonal, PERSONAL, "ap-southeast-1", "international"), + (Intl, TEAM, "ap-southeast-1", "international"), + ] { + let mut expected = head.to_vec(); + expected.extend([ "--console-region", - "ap-southeast-1", + console_region, "--console-site", - "international", + console_site, "--output", - "json" - ] - ); + "json", + ]); + assert_eq!(cli_arguments(region), expected, "{region:?}"); + } } #[test] diff --git a/rust/src/providers/alibabatokenplan/fields.rs b/rust/src/providers/alibabatokenplan/fields.rs new file mode 100644 index 0000000000..dc3d793cba --- /dev/null +++ b/rust/src/providers/alibabatokenplan/fields.rs @@ -0,0 +1,341 @@ +//! Lenient lookups over the expanded Alibaba console JSON. + +use chrono::{DateTime, NaiveDate, NaiveDateTime, TimeZone, Utc}; +use serde_json::Value; + +pub(super) fn expand_json_strings(value: Value) -> Value { + match value { + Value::Array(values) => Value::Array(values.into_iter().map(expand_json_strings).collect()), + Value::Object(map) => Value::Object( + map.into_iter() + .map(|(key, value)| (key, expand_json_strings(value))) + .collect(), + ), + Value::String(text) => serde_json::from_str::(&text) + .ok() + .filter(|nested| nested.is_object() || nested.is_array()) + .map(expand_json_strings) + .unwrap_or(Value::String(text)), + other => other, + } +} + +pub(super) fn percentage_points(ratio: Option) -> Option { + let ratio = ratio.filter(|v| v.is_finite())?; + Some((ratio.clamp(0.0, 1.0) * 100.0).clamp(0.0, 100.0)) +} + +pub(super) fn number_field(value: &Value, key: &str) -> Option { + value.as_object().and_then(|map| parse_f64(map.get(key))) +} + +pub(super) fn date_field(value: &Value, key: &str) -> Option> { + value.as_object().and_then(|map| parse_date(map.get(key))) +} + +/// Depth-first search: `probe` runs on each node before its children, and +/// the first `Some` wins. Probes return `None` for arrays and scalars. +pub(super) fn deep_find<'a, T>( + value: &'a Value, + probe: &impl Fn(&'a Value) -> Option, +) -> Option { + if let Some(found) = probe(value) { + return Some(found); + } + match value { + Value::Object(map) => map.values().find_map(|nested| deep_find(nested, probe)), + Value::Array(values) => values.iter().find_map(|nested| deep_find(nested, probe)), + _ => None, + } +} + +pub(super) fn find_object_containing_any_of(value: &Value, keys: &[&str]) -> Option { + deep_find(value, &|node| { + let map = node.as_object()?; + keys.iter() + .any(|key| map.contains_key(*key)) + .then(|| node.clone()) + }) +} + +const PLAN_NAME_KEYS: &[&str] = &[ + "planName", + "plan_name", + "packageName", + "package_name", + "commodityName", + "commodity_name", + "instanceName", + "instance_name", + "displayName", + "display_name", + "name", + "title", + "planType", + "plan_type", + "ProductName", + "productName", +]; +pub(super) const USED_QUOTA_KEYS: &[&str] = &[ + "usedQuota", + "used_quota", + "usedCredits", + "usedCredit", + "consumedCredits", + "usage", + "used", + "usedAmount", + "consumeAmount", + "usedValue", + "UsedValue", + "consumedValue", + "ConsumedValue", +]; +pub(super) const TOTAL_QUOTA_KEYS: &[&str] = &[ + "totalQuota", + "total_quota", + "totalCredits", + "totalCredit", + "quota", + "creditLimit", + "creditsTotal", + "monthlyTotalQuota", + "amount", + "totalValue", + "TotalValue", + "totalCount", + "TotalCount", + "subscriptionTotalNumber", + "SubscriptionTotalNumber", +]; +pub(super) const REMAINING_QUOTA_KEYS: &[&str] = &[ + "remainingQuota", + "remainQuota", + "remainingCredits", + "remainingCredit", + "availableCredits", + "balance", + "remaining", + "availableAmount", + "remainAmount", + "totalSurplusValue", + "TotalSurplusValue", + "surplusValue", + "SurplusValue", +]; +const RESET_DATE_KEYS: &[&str] = &[ + "nextRefreshTime", + "resetTime", + "periodEndTime", + "billingCycleEnd", + "billCycleEndTime", + "expireTime", + "expirationTime", + "endTime", + "validEndTime", + "instanceEndTime", + "nearestExpireDate", + "NearestExpireDate", +]; + +pub(super) fn find_token_plan_instance(value: &Value) -> Option { + find_first_object( + value, + &[ + "tokenPlanInstanceInfo", + "token_plan_instance_info", + "instanceInfo", + "instance_info", + ], + ) + .or_else(|| { + find_first_array( + value, + &[ + "tokenPlanInstanceInfos", + "token_plan_instance_infos", + "instanceInfos", + "instances", + "Data", + "data", + "successResponse", + ], + ) + .and_then(|values| { + values + .into_iter() + .filter(Value::is_object) + .max_by_key(active_signal_score) + }) + }) +} + +pub(super) fn find_plan_name(value: &Value) -> Option { + first_string(value, PLAN_NAME_KEYS).or_else(|| find_first_string(value, PLAN_NAME_KEYS)) +} + +pub(super) fn find_quota_info(value: &Value) -> Option { + find_first_object( + value, + &[ + "quotaInfo", + "quota_info", + "tokenPlanQuotaInfo", + "token_plan_quota_info", + ], + ) + .or_else(|| { + find_object_containing_any_of( + value, + &[USED_QUOTA_KEYS, TOTAL_QUOTA_KEYS, REMAINING_QUOTA_KEYS].concat(), + ) + }) +} + +pub(super) fn find_reset_date(value: &Value) -> Option> { + first_date(value, RESET_DATE_KEYS).or_else(|| find_first_date(value, RESET_DATE_KEYS)) +} + +fn find_first_object(value: &Value, keys: &[&str]) -> Option { + deep_find(value, &|node| { + let map = node.as_object()?; + keys.iter() + .find_map(|key| map.get(*key).filter(|v| v.is_object()).cloned()) + }) +} + +fn find_first_array(value: &Value, keys: &[&str]) -> Option> { + deep_find(value, &|node| { + let map = node.as_object()?; + keys.iter() + .find_map(|key| map.get(*key).and_then(Value::as_array).cloned()) + }) +} + +fn first_string(value: &Value, keys: &[&str]) -> Option { + let map = value.as_object()?; + keys.iter().find_map(|key| parse_string(map.get(*key))) +} + +pub(super) fn find_first_string(value: &Value, keys: &[&str]) -> Option { + deep_find(value, &|node| first_string(node, keys)) +} + +pub(super) fn first_f64(value: &Value, keys: &[&str]) -> Option { + let map = value.as_object()?; + keys.iter().find_map(|key| parse_f64(map.get(*key))) +} + +pub(super) fn find_first_i64(value: &Value, keys: &[&str]) -> Option { + deep_find(value, &|node| { + let map = node.as_object()?; + keys.iter().find_map(|key| parse_i64(map.get(*key))) + }) +} + +fn first_date(value: &Value, keys: &[&str]) -> Option> { + let map = value.as_object()?; + keys.iter().find_map(|key| parse_date(map.get(*key))) +} + +fn find_first_date(value: &Value, keys: &[&str]) -> Option> { + deep_find(value, &|node| first_date(node, keys)) +} + +fn parse_string(value: Option<&Value>) -> Option { + value? + .as_str() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) +} + +fn parse_f64(value: Option<&Value>) -> Option { + match value? { + Value::Number(number) => number.as_f64(), + Value::String(text) => text.trim().replace(',', "").parse().ok(), + _ => None, + } +} + +fn parse_i64(value: Option<&Value>) -> Option { + match value? { + Value::Number(number) => number.as_i64().or_else(|| { + // Quota/timestamp JSON floats are whole numbers; the fractional + // part is rounding noise from the upstream API. + let v = number.as_f64()?; + #[expect(clippy::cast_possible_truncation, reason = "quota/timestamp JSON floats are whole numbers; fractional part is rounding noise")] + let whole = v as i64; + Some(whole) + }), + Value::String(text) => text.trim().replace(',', "").parse().ok(), + _ => None, + } +} + +pub(super) fn parse_bool(value: Option<&Value>) -> Option { + match value? { + Value::Bool(flag) => Some(*flag), + Value::Number(number) => number.as_i64().map(|v| v != 0), + Value::String(text) => match text.trim().to_lowercase().as_str() { + "true" | "1" | "yes" | "active" | "valid" | "normal" => Some(true), + "false" | "0" | "no" | "inactive" | "invalid" | "expired" => Some(false), + _ => None, + }, + _ => None, + } +} + +fn parse_date(value: Option<&Value>) -> Option> { + if let Some(raw) = parse_i64(value) { + if raw > 1_000_000_000_000 { + return Utc.timestamp_opt(raw / 1000, 0).single(); + } + if raw > 1_000_000_000 { + return Utc.timestamp_opt(raw, 0).single(); + } + } + let text = parse_string(value)?; + if let Ok(date) = DateTime::parse_from_rfc3339(&text) { + return Some(date.with_timezone(&Utc)); + } + if let Ok(date) = NaiveDate::parse_from_str(&text, "%Y-%m-%d") + && let Some(date_time) = date.and_hms_opt(0, 0, 0) + { + return Some(date_time.and_utc()); + } + for format in ["%Y-%m-%d %H:%M", "%Y-%m-%d %H:%M:%S"] { + if let Ok(date) = NaiveDateTime::parse_from_str(&text, format) { + return Some(date.and_utc()); + } + } + None +} + +fn active_signal_score(value: &Value) -> i32 { + let status = first_string(value, &["status", "instanceStatus", "state"]) + .unwrap_or_default() + .to_uppercase(); + if ["VALID", "ACTIVE", "NORMAL"].contains(&status.as_str()) { + return 3; + } + if [ + "EXPIRED", + "INVALID", + "INACTIVE", + "DISABLED", + "TERMINATED", + "STOPPED", + ] + .contains(&status.as_str()) + { + return -1; + } + parse_bool( + value + .as_object() + .and_then(|map| map.get("isActive").or_else(|| map.get("active"))), + ) + .map(|active| if active { 3 } else { -1 }) + .unwrap_or(0) +} diff --git a/rust/src/providers/alibabatokenplan/mod.rs b/rust/src/providers/alibabatokenplan/mod.rs index fb21b27ec6..6f529bf523 100644 --- a/rust/src/providers/alibabatokenplan/mod.rs +++ b/rust/src/providers/alibabatokenplan/mod.rs @@ -8,6 +8,7 @@ //! Personal/Solo path: OneConsole personal token-plan APIs (+ best-effort sec_token). mod cli; +mod fields; #[cfg(test)] mod monthly_tests; mod personal; @@ -16,7 +17,7 @@ mod region; pub use region::AlibabaTokenPlanRegion; use async_trait::async_trait; -use chrono::{DateTime, NaiveDate, NaiveDateTime, TimeZone, Utc}; +use chrono::{DateTime, Utc}; use regex_lite::Regex; use serde_json::Value; @@ -28,6 +29,13 @@ use crate::providers::{browser_cookie_header, strip_cookie_prefix}; use region::AlibabaTokenPlanRegion as Region; +use fields::{ + REMAINING_QUOTA_KEYS, TOTAL_QUOTA_KEYS, USED_QUOTA_KEYS, date_field, deep_find, + expand_json_strings, find_first_i64, find_first_string, find_object_containing_any_of, + find_plan_name, find_quota_info, find_reset_date, find_token_plan_instance, first_f64, + number_field, parse_bool, percentage_points, +}; + pub(super) const USER_AGENT: &str = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/143.0.0.0 Safari/537.36"; pub(super) const LANGUAGE: &str = "en-US"; pub(super) const PERSONAL_CONSOLE_PRODUCT: &str = "sfm_bailian"; @@ -77,6 +85,19 @@ impl AlibabaTokenPlanProvider { Self::snapshot_to_usage(snapshot) } + async fn fetch_via( + &self, + source: &'static str, + ctx: &FetchContext, + ) -> Result { + let usage = if source == "web" { + self.fetch_via_web(ctx).await? + } else { + self.fetch_via_cli(ctx).await? + }; + Ok(ProviderFetchResult::new(usage, source)) + } + async fn fetch_via_web(&self, ctx: &FetchContext) -> Result { let region = Self::resolve_region(ctx); let cookie_header = Self::resolve_cookie_header(ctx, region)?; @@ -117,45 +138,16 @@ impl AlibabaTokenPlanProvider { ("params", Self::team_request_params(region)), ("region", region.current_region_id().to_string()), ]; - if let Some(token) = sec_token - .as_deref() - .filter(|token| !token.trim().is_empty()) - { - form.push(("sec_token", token.to_string())); - } - let mut request = client - .post(Self::team_quota_url(region)) - .header("Cookie", cookie_header) - .header("Accept", "*/*") - .header("Content-Type", "application/x-www-form-urlencoded") - .header("Origin", region.gateway_base_url()) - .header("Referer", region.dashboard_url()) - .header("User-Agent", USER_AGENT) - .header("X-Requested-With", "XMLHttpRequest") - .form(&form); - - if let Some(csrf) = cookie_value("login_aliyunid_csrf", cookie_header) - .or_else(|| cookie_value("csrf", cookie_header)) - { - request = request - .header("x-xsrf-token", csrf.clone()) - .header("x-csrf-token", csrf); - } - - let response = request.send().await?; - let status = response.status(); - let body = response.bytes().await?; - if !status.is_success() { - if status == reqwest::StatusCode::UNAUTHORIZED - || status == reqwest::StatusCode::FORBIDDEN - { - return Err(ProviderError::AuthRequired); - } - return Err(ProviderError::Other(format!( - "Alibaba Token Plan API error: HTTP {status}" - ))); - } - + push_sec_token(&mut form, sec_token.as_deref()); + let body = send_console_form( + client.post(Self::team_quota_url(region)), + cookie_header, + region, + "*/*", + &form, + "Alibaba Token Plan", + ) + .await?; Self::parse_usage_snapshot(&body) } @@ -227,20 +219,7 @@ impl AlibabaTokenPlanProvider { } fn parse_usage_snapshot(data: &[u8]) -> Result { - if data.is_empty() { - return Err(ProviderError::Parse( - "Empty Alibaba Token Plan response".into(), - )); - } - let value: Value = serde_json::from_slice(data).map_err(|_| { - if is_likely_login_html(data) { - ProviderError::AuthRequired - } else { - ProviderError::Parse("Invalid Alibaba Token Plan JSON response".into()) - } - })?; - let expanded = expand_json_strings(value); - throw_if_error_payload(&expanded)?; + let expanded = decode_console_payload(data, "Alibaba Token Plan")?; let instance = find_token_plan_instance(&expanded); let plan_name = instance @@ -376,30 +355,16 @@ impl Provider for AlibabaTokenPlanProvider { } async fn fetch_usage(&self, ctx: &FetchContext) -> Result { - match ctx.source_mode { - SourceMode::Auto if ctx.auto_prefer_web => match self.fetch_via_web(ctx).await { - Ok(usage) => Ok(ProviderFetchResult::new(usage, "web")), - Err(_) => { - let usage = self.fetch_via_cli(ctx).await?; - Ok(ProviderFetchResult::new(usage, "cli")) - } - }, - SourceMode::Auto => match self.fetch_via_cli(ctx).await { - Ok(usage) => Ok(ProviderFetchResult::new(usage, "cli")), - Err(_) => { - let usage = self.fetch_via_web(ctx).await?; - Ok(ProviderFetchResult::new(usage, "web")) - } - }, - SourceMode::Cli => { - let usage = self.fetch_via_cli(ctx).await?; - Ok(ProviderFetchResult::new(usage, "cli")) - } - SourceMode::Web => { - let usage = self.fetch_via_web(ctx).await?; - Ok(ProviderFetchResult::new(usage, "web")) - } - SourceMode::OAuth => Err(ProviderError::UnsupportedSource(ctx.source_mode)), + let (first, fallback) = match ctx.source_mode { + SourceMode::Auto if ctx.auto_prefer_web => ("web", Some("cli")), + SourceMode::Auto => ("cli", Some("web")), + SourceMode::Cli => ("cli", None), + SourceMode::Web => ("web", None), + SourceMode::OAuth => return Err(ProviderError::UnsupportedSource(ctx.source_mode)), + }; + match (self.fetch_via(first, ctx).await, fallback) { + (Err(_), Some(fallback)) => self.fetch_via(fallback, ctx).await, + (result, _) => result, } } @@ -417,6 +382,72 @@ impl Provider for AlibabaTokenPlanProvider { } } +pub(super) fn push_sec_token(form: &mut Vec<(&'static str, String)>, sec_token: Option<&str>) { + if let Some(token) = sec_token.filter(|token| !token.trim().is_empty()) { + form.push(("sec_token", token.to_string())); + } +} + +/// POST a console gateway form. The CSRF headers go after the form so the +/// header order matches the browser capture. +pub(super) async fn send_console_form( + request: reqwest::RequestBuilder, + cookie_header: &str, + region: Region, + accept: &str, + form: &[(&'static str, String)], + scope: &str, +) -> Result, ProviderError> { + let mut request = request + .header("Cookie", cookie_header) + .header("Accept", accept) + .header("Content-Type", "application/x-www-form-urlencoded") + .header("Origin", region.gateway_base_url()) + .header("Referer", region.dashboard_url()) + .header("User-Agent", USER_AGENT) + .header("X-Requested-With", "XMLHttpRequest") + .form(form); + + if let Some(csrf) = cookie_value("login_aliyunid_csrf", cookie_header) + .or_else(|| cookie_value("csrf", cookie_header)) + { + request = request + .header("x-xsrf-token", csrf.clone()) + .header("x-csrf-token", csrf); + } + + let response = request.send().await?; + let status = response.status(); + let body = response.bytes().await?; + if !status.is_success() { + if status == reqwest::StatusCode::UNAUTHORIZED || status == reqwest::StatusCode::FORBIDDEN { + return Err(ProviderError::AuthRequired); + } + return Err(ProviderError::Other(format!( + "{scope} API error: HTTP {status}" + ))); + } + Ok(body.to_vec()) +} + +/// Decode a console response body: reject empty bodies and login HTML, +/// expand JSON-in-string fields, then surface gateway error payloads. +pub(super) fn decode_console_payload(data: &[u8], scope: &str) -> Result { + if data.is_empty() { + return Err(ProviderError::Parse(format!("Empty {scope} response"))); + } + let value: Value = serde_json::from_slice(data).map_err(|_| { + if is_likely_login_html(data) { + ProviderError::AuthRequired + } else { + ProviderError::Parse(format!("Invalid {scope} JSON response")) + } + })?; + let expanded = expand_json_strings(value); + throw_if_error_payload(&expanded)?; + Ok(expanded) +} + pub(super) fn throw_if_error_payload(value: &Value) -> Result<(), ProviderError> { if let Some(status) = find_first_i64(value, &["statusCode", "status_code", "code"]) && status != 0 @@ -470,21 +501,15 @@ pub(super) fn throw_if_error_payload(value: &Value) -> Result<(), ProviderError> } fn find_failing_success_frame(value: &Value) -> Option<&Value> { - match value { - Value::Object(map) => { - let failed_here = ["success", "Success"] - .iter() - .filter_map(|key| map.get(*key)) - .filter_map(|value| parse_bool(Some(value))) - .any(|success| !success); - if failed_here { - return Some(value); - } - map.values().find_map(find_failing_success_frame) - } - Value::Array(values) => values.iter().find_map(find_failing_success_frame), - _ => None, - } + deep_find(value, &|node| { + let map = node.as_object()?; + ["success", "Success"] + .iter() + .filter_map(|key| map.get(*key)) + .filter_map(|value| parse_bool(Some(value))) + .any(|success| !success) + .then_some(node) + }) } fn normalize_cookie_header(raw: &str) -> Option { @@ -534,394 +559,6 @@ pub(super) fn is_likely_login_html(data: &[u8]) -> bool { && (text.contains("login") || text.contains("sign in") || text.contains("signin")) } -pub(super) fn expand_json_strings(value: Value) -> Value { - match value { - Value::Array(values) => Value::Array(values.into_iter().map(expand_json_strings).collect()), - Value::Object(map) => Value::Object( - map.into_iter() - .map(|(key, value)| (key, expand_json_strings(value))) - .collect(), - ), - Value::String(text) => serde_json::from_str::(&text) - .ok() - .filter(|nested| nested.is_object() || nested.is_array()) - .map(expand_json_strings) - .unwrap_or(Value::String(text)), - other => other, - } -} - -pub(super) fn percentage_points(ratio: Option) -> Option { - let ratio = ratio.filter(|v| v.is_finite())?; - Some((ratio.clamp(0.0, 1.0) * 100.0).clamp(0.0, 100.0)) -} - -pub(super) fn number_field(value: &Value, key: &str) -> Option { - value.as_object().and_then(|map| parse_f64(map.get(key))) -} - -pub(super) fn date_field(value: &Value, key: &str) -> Option> { - value.as_object().and_then(|map| parse_date(map.get(key))) -} - -pub(super) fn find_object_containing_any_of(value: &Value, keys: &[&str]) -> Option { - match value { - Value::Object(map) => { - if keys.iter().any(|key| map.contains_key(*key)) { - return Some(Value::Object(map.clone())); - } - map.values() - .find_map(|nested| find_object_containing_any_of(nested, keys)) - } - Value::Array(values) => values - .iter() - .find_map(|nested| find_object_containing_any_of(nested, keys)), - _ => None, - } -} - -const PLAN_NAME_KEYS: &[&str] = &[ - "planName", - "plan_name", - "packageName", - "package_name", - "commodityName", - "commodity_name", - "instanceName", - "instance_name", - "displayName", - "display_name", - "name", - "title", - "planType", - "plan_type", - "ProductName", - "productName", -]; -const USED_QUOTA_KEYS: &[&str] = &[ - "usedQuota", - "used_quota", - "usedCredits", - "usedCredit", - "consumedCredits", - "usage", - "used", - "usedAmount", - "consumeAmount", - "usedValue", - "UsedValue", - "consumedValue", - "ConsumedValue", -]; -const TOTAL_QUOTA_KEYS: &[&str] = &[ - "totalQuota", - "total_quota", - "totalCredits", - "totalCredit", - "quota", - "creditLimit", - "creditsTotal", - "monthlyTotalQuota", - "amount", - "totalValue", - "TotalValue", - "totalCount", - "TotalCount", - "subscriptionTotalNumber", - "SubscriptionTotalNumber", -]; -const REMAINING_QUOTA_KEYS: &[&str] = &[ - "remainingQuota", - "remainQuota", - "remainingCredits", - "remainingCredit", - "availableCredits", - "balance", - "remaining", - "availableAmount", - "remainAmount", - "totalSurplusValue", - "TotalSurplusValue", - "surplusValue", - "SurplusValue", -]; -const RESET_DATE_KEYS: &[&str] = &[ - "nextRefreshTime", - "resetTime", - "periodEndTime", - "billingCycleEnd", - "billCycleEndTime", - "expireTime", - "expirationTime", - "endTime", - "validEndTime", - "instanceEndTime", - "nearestExpireDate", - "NearestExpireDate", -]; - -fn find_token_plan_instance(value: &Value) -> Option { - find_first_object( - value, - &[ - "tokenPlanInstanceInfo", - "token_plan_instance_info", - "instanceInfo", - "instance_info", - ], - ) - .or_else(|| { - find_first_array( - value, - &[ - "tokenPlanInstanceInfos", - "token_plan_instance_infos", - "instanceInfos", - "instances", - "Data", - "data", - "successResponse", - ], - ) - .and_then(|values| { - values - .into_iter() - .filter(Value::is_object) - .max_by_key(active_signal_score) - }) - }) -} - -fn find_plan_name(value: &Value) -> Option { - first_string(value, PLAN_NAME_KEYS).or_else(|| find_first_string(value, PLAN_NAME_KEYS)) -} - -fn find_quota_info(value: &Value) -> Option { - find_first_object( - value, - &[ - "quotaInfo", - "quota_info", - "tokenPlanQuotaInfo", - "token_plan_quota_info", - ], - ) - .or_else(|| { - find_first_object_with_any_key( - value, - &[USED_QUOTA_KEYS, TOTAL_QUOTA_KEYS, REMAINING_QUOTA_KEYS].concat(), - ) - }) -} - -fn find_reset_date(value: &Value) -> Option> { - first_date(value, RESET_DATE_KEYS).or_else(|| find_first_date(value, RESET_DATE_KEYS)) -} - -fn find_first_object(value: &Value, keys: &[&str]) -> Option { - match value { - Value::Object(map) => { - for key in keys { - if let Some(nested) = map.get(*key).filter(|v| v.is_object()) { - return Some(nested.clone()); - } - } - map.values() - .find_map(|nested| find_first_object(nested, keys)) - } - Value::Array(values) => values - .iter() - .find_map(|nested| find_first_object(nested, keys)), - _ => None, - } -} - -fn find_first_object_with_any_key(value: &Value, keys: &[&str]) -> Option { - match value { - Value::Object(map) => { - if keys.iter().any(|key| map.contains_key(*key)) { - return Some(value.clone()); - } - map.values() - .find_map(|nested| find_first_object_with_any_key(nested, keys)) - } - Value::Array(values) => values - .iter() - .find_map(|nested| find_first_object_with_any_key(nested, keys)), - _ => None, - } -} - -fn find_first_array(value: &Value, keys: &[&str]) -> Option> { - match value { - Value::Object(map) => { - for key in keys { - if let Some(values) = map.get(*key).and_then(Value::as_array) { - return Some(values.clone()); - } - } - map.values() - .find_map(|nested| find_first_array(nested, keys)) - } - Value::Array(values) => values - .iter() - .find_map(|nested| find_first_array(nested, keys)), - _ => None, - } -} - -fn first_string(value: &Value, keys: &[&str]) -> Option { - let map = value.as_object()?; - keys.iter().find_map(|key| parse_string(map.get(*key))) -} - -pub(super) fn find_first_string(value: &Value, keys: &[&str]) -> Option { - match value { - Value::Object(map) => first_string(value, keys).or_else(|| { - map.values() - .find_map(|nested| find_first_string(nested, keys)) - }), - Value::Array(values) => values - .iter() - .find_map(|nested| find_first_string(nested, keys)), - _ => None, - } -} - -fn first_f64(value: &Value, keys: &[&str]) -> Option { - let map = value.as_object()?; - keys.iter().find_map(|key| parse_f64(map.get(*key))) -} - -fn find_first_i64(value: &Value, keys: &[&str]) -> Option { - match value { - Value::Object(map) => keys - .iter() - .find_map(|key| parse_i64(map.get(*key))) - .or_else(|| map.values().find_map(|nested| find_first_i64(nested, keys))), - Value::Array(values) => values - .iter() - .find_map(|nested| find_first_i64(nested, keys)), - _ => None, - } -} - -fn first_date(value: &Value, keys: &[&str]) -> Option> { - let map = value.as_object()?; - keys.iter().find_map(|key| parse_date(map.get(*key))) -} - -fn find_first_date(value: &Value, keys: &[&str]) -> Option> { - match value { - Value::Object(map) => first_date(value, keys).or_else(|| { - map.values() - .find_map(|nested| find_first_date(nested, keys)) - }), - Value::Array(values) => values - .iter() - .find_map(|nested| find_first_date(nested, keys)), - _ => None, - } -} - -fn parse_string(value: Option<&Value>) -> Option { - value? - .as_str() - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) -} - -fn parse_f64(value: Option<&Value>) -> Option { - match value? { - Value::Number(number) => number.as_f64(), - Value::String(text) => text.trim().replace(',', "").parse().ok(), - _ => None, - } -} - -fn parse_i64(value: Option<&Value>) -> Option { - match value? { - Value::Number(number) => number.as_i64().or_else(|| { - // Quota/timestamp JSON floats are whole numbers; the fractional - // part is rounding noise from the upstream API. - let v = number.as_f64()?; - #[expect(clippy::cast_possible_truncation, reason = "quota/timestamp JSON floats are whole numbers; fractional part is rounding noise")] - let whole = v as i64; - Some(whole) - }), - Value::String(text) => text.trim().replace(',', "").parse().ok(), - _ => None, - } -} - -fn parse_bool(value: Option<&Value>) -> Option { - match value? { - Value::Bool(flag) => Some(*flag), - Value::Number(number) => number.as_i64().map(|v| v != 0), - Value::String(text) => match text.trim().to_lowercase().as_str() { - "true" | "1" | "yes" | "active" | "valid" | "normal" => Some(true), - "false" | "0" | "no" | "inactive" | "invalid" | "expired" => Some(false), - _ => None, - }, - _ => None, - } -} - -fn parse_date(value: Option<&Value>) -> Option> { - if let Some(raw) = parse_i64(value) { - if raw > 1_000_000_000_000 { - return Utc.timestamp_opt(raw / 1000, 0).single(); - } - if raw > 1_000_000_000 { - return Utc.timestamp_opt(raw, 0).single(); - } - } - let text = parse_string(value)?; - if let Ok(date) = DateTime::parse_from_rfc3339(&text) { - return Some(date.with_timezone(&Utc)); - } - if let Ok(date) = NaiveDate::parse_from_str(&text, "%Y-%m-%d") - && let Some(date_time) = date.and_hms_opt(0, 0, 0) - { - return Some(date_time.and_utc()); - } - for format in ["%Y-%m-%d %H:%M", "%Y-%m-%d %H:%M:%S"] { - if let Ok(date) = NaiveDateTime::parse_from_str(&text, format) { - return Some(date.and_utc()); - } - } - None -} - -fn active_signal_score(value: &Value) -> i32 { - let status = first_string(value, &["status", "instanceStatus", "state"]) - .unwrap_or_default() - .to_uppercase(); - if ["VALID", "ACTIVE", "NORMAL"].contains(&status.as_str()) { - return 3; - } - if [ - "EXPIRED", - "INVALID", - "INACTIVE", - "DISABLED", - "TERMINATED", - "STOPPED", - ] - .contains(&status.as_str()) - { - return -1; - } - parse_bool( - value - .as_object() - .and_then(|map| map.get("isActive").or_else(|| map.get("active"))), - ) - .map(|active| if active { 3 } else { -1 }) - .unwrap_or(0) -} - fn used_percent(used: Option, total: Option, remaining: Option) -> Option { let total = total.filter(|total| *total > 0.0)?; let used = used.or_else(|| remaining.map(|remaining| total - remaining))?; @@ -1007,150 +644,4 @@ fn payload_diagnostics(value: &Value) -> String { } #[cfg(test)] -mod tests { - use super::*; - - #[test] - fn missing_bailian_cli_is_local_runtime_offline() { - let provider = AlibabaTokenPlanProvider::new(); - let error = ProviderError::NotInstalled( - "Bailian CLI 'bl' is not installed or not on PATH.".to_string(), - ); - assert_eq!( - provider.error_state_kind(&error), - crate::core::ProviderStateKind::LocalRuntimeOffline - ); - } - #[test] - fn parses_token_plan_instance_payload() { - let payload = serde_json::json!({ - "data": { - "tokenPlanInstanceInfo": { - "commodityName": "Token Plan Pro", - "quotaInfo": { - "usedQuota": "1250", - "totalQuota": "5000" - }, - "nextRefreshTime": 1780763009000_i64 - } - } - }); - let snapshot = - AlibabaTokenPlanProvider::parse_usage_snapshot(payload.to_string().as_bytes()).unwrap(); - assert_eq!(snapshot.plan_name.as_deref(), Some("Token Plan Pro")); - assert_eq!(snapshot.used_quota, Some(1250.0)); - assert_eq!(snapshot.total_quota, Some(5000.0)); - - let usage = AlibabaTokenPlanProvider::snapshot_to_usage(snapshot).unwrap(); - assert_eq!(usage.primary.used_percent, 25.0); - assert_eq!( - usage.primary.reset_description.as_deref(), - Some("1,250 / 5,000 credits used") - ); - assert_eq!(usage.login_method.as_deref(), Some("Token Plan Pro")); - } - - #[test] - fn expands_nested_string_payloads_and_uses_remaining_quota() { - let nested = serde_json::json!({ - "successResponse": serde_json::json!({ - "instances": [ - {"status": "EXPIRED", "quota": 1000, "remaining": 1000}, - {"status": "ACTIVE", "packageName": "Team", "quota": 1000, "remaining": 250} - ] - }).to_string() - }); - let snapshot = - AlibabaTokenPlanProvider::parse_usage_snapshot(nested.to_string().as_bytes()).unwrap(); - assert_eq!(snapshot.plan_name.as_deref(), Some("Team")); - assert_eq!( - used_percent( - snapshot.used_quota, - snapshot.total_quota, - snapshot.remaining_quota - ), - Some(75.0) - ); - } - - #[test] - fn parses_new_subscription_summary_payload() { - let payload = serde_json::json!({ - "success": true, - "Data": { - "ProductName": "Token Plan Team", - "TotalValue": "1000000", - "TotalSurplusValue": "250000", - "NearestExpireDate": "2026-06-30" - } - }); - let snapshot = - AlibabaTokenPlanProvider::parse_usage_snapshot(payload.to_string().as_bytes()).unwrap(); - assert_eq!(snapshot.plan_name.as_deref(), Some("Token Plan Team")); - assert_eq!(snapshot.total_quota, Some(1_000_000.0)); - assert_eq!(snapshot.remaining_quota, Some(250_000.0)); - assert_eq!( - used_percent( - snapshot.used_quota, - snapshot.total_quota, - snapshot.remaining_quota - ), - Some(75.0) - ); - assert!(snapshot.resets_at.is_some()); - } - - #[test] - fn detects_login_payloads() { - let err = AlibabaTokenPlanProvider::parse_usage_snapshot( - br#"{"code":"NeedLogin","message":"please login"}"#, - ) - .unwrap_err(); - assert!(matches!(err, ProviderError::AuthRequired)); - } - - #[test] - fn extracts_sec_token_from_html_or_cookie() { - assert_eq!( - extract_sec_token(r#""#).as_deref(), - Some("abc123") - ); - assert_eq!( - extract_sec_token( - r#""# - ) - .as_deref(), - Some("upper123") - ); - assert_eq!( - cookie_value("sec_token", "foo=bar; sec_token=xyz"), - Some("xyz".to_string()) - ); - } - - #[test] - fn sec_token_shell_referer_uses_same_origin_root() { - assert_eq!( - dashboard_referer(Region::CnPersonal), - "https://bailian.console.aliyun.com/" - ); - assert_eq!( - dashboard_referer(Region::IntlPersonal), - "https://modelstudio.console.alibabacloud.com/" - ); - } - - #[test] - fn default_region_cn_team_urls_match_legacy() { - let region = Region::Cn; - assert_eq!( - AlibabaTokenPlanProvider::team_quota_url(region), - "https://bailian.console.aliyun.com/data/api.json?action=GetSubscriptionSummary&product=BssOpenAPI-V3&_tag=" - ); - assert_eq!( - AlibabaTokenPlanProvider::team_request_params(region), - serde_json::json!({"ProductCode": "sfm_tokenplanteams_dp_cn"}).to_string() - ); - assert_eq!(region.current_region_id(), "cn-beijing"); - } -} +mod tests; diff --git a/rust/src/providers/alibabatokenplan/personal.rs b/rust/src/providers/alibabatokenplan/personal.rs index 13389a2686..abef795cd9 100644 --- a/rust/src/providers/alibabatokenplan/personal.rs +++ b/rust/src/providers/alibabatokenplan/personal.rs @@ -7,9 +7,9 @@ use uuid::Uuid; use super::region::AlibabaTokenPlanRegion; use super::{ LANGUAGE, PERSONAL_CONSOLE_PRODUCT, PERSONAL_QUOTA_CONFIG_API, PERSONAL_SUBSCRIPTION_API, - PERSONAL_USAGE_API, TokenPlanSnapshot, USER_AGENT, cookie_value, date_field, - expand_json_strings, find_object_containing_any_of, is_likely_login_html, number_field, - percentage_points, throw_if_error_payload, + PERSONAL_USAGE_API, TokenPlanSnapshot, cookie_value, date_field, decode_console_payload, + deep_find, expand_json_strings, find_object_containing_any_of, number_field, percentage_points, + push_sec_token, send_console_form, }; use crate::core::{FetchContext, ProviderError}; @@ -66,10 +66,12 @@ pub(super) async fn fetch_personal_usage( Value::String(region.product_code().to_string()), ); let subscription_body = - post_personal_api_optional(&context, PERSONAL_SUBSCRIPTION_API, subscription_params).await; - - let quota_config_body = - post_personal_api_optional(&context, PERSONAL_QUOTA_CONFIG_API, Map::new()).await; + post_personal_api(&context, PERSONAL_SUBSCRIPTION_API, subscription_params) + .await + .ok(); + let quota_config_body = post_personal_api(&context, PERSONAL_QUOTA_CONFIG_API, Map::new()) + .await + .ok(); const MAX_USAGE_ATTEMPTS: usize = 3; const RETRY_DELAY: std::time::Duration = std::time::Duration::from_millis(400); @@ -84,7 +86,7 @@ pub(super) async fn fetch_personal_usage( quota_config_body.as_deref(), ) { Ok(snapshot) => return Ok(snapshot), - Err(error) if personal_usage_success_without_windows(&usage_body) => { + Err(_) if personal_usage_success_without_windows(&usage_body) => { tracing::info!( attempt = attempt + 1, max_attempts = MAX_USAGE_ATTEMPTS, @@ -96,7 +98,6 @@ pub(super) async fn fetch_personal_usage( .into(), )); } - let _ = error; } Err(error) => return Err(error), } @@ -144,49 +145,21 @@ async fn post_personal_api( context.sec_token, ); - let mut request = context + let request = context .client .post(&url) .timeout(std::time::Duration::from_secs( context.fetch_context.web_timeout.max(1), - )) - .header("Cookie", context.cookie_header) - .header("Accept", "application/json, text/plain, */*") - .header("Content-Type", "application/x-www-form-urlencoded") - .header("Origin", context.region.gateway_base_url()) - .header("Referer", context.region.dashboard_url()) - .header("User-Agent", USER_AGENT) - .header("X-Requested-With", "XMLHttpRequest") - .form(&form); - - if let Some(csrf) = cookie_value("login_aliyunid_csrf", context.cookie_header) - .or_else(|| cookie_value("csrf", context.cookie_header)) - { - request = request - .header("x-xsrf-token", csrf.clone()) - .header("x-csrf-token", csrf); - } - - let response = request.send().await?; - let status = response.status(); - let body = response.bytes().await?; - if !status.is_success() { - if status == reqwest::StatusCode::UNAUTHORIZED || status == reqwest::StatusCode::FORBIDDEN { - return Err(ProviderError::AuthRequired); - } - return Err(ProviderError::Other(format!( - "Alibaba Token Plan Personal API error: HTTP {status}" - ))); - } - Ok(body.to_vec()) -} - -async fn post_personal_api_optional( - context: &PersonalApiContext<'_>, - api: &str, - data_parameters: Map, -) -> Option> { - post_personal_api(context, api, data_parameters).await.ok() + )); + send_console_form( + request, + context.cookie_header, + context.region, + "application/json, text/plain, */*", + &form, + "Alibaba Token Plan Personal", + ) + .await } fn build_personal_form( @@ -204,9 +177,7 @@ fn build_personal_form( ("language", LANGUAGE.to_string()), ("params", params_json), ]; - if let Some(token) = sec_token.filter(|token| !token.trim().is_empty()) { - form.push(("sec_token", token.to_string())); - } + push_sec_token(&mut form, sec_token); form } @@ -233,31 +204,26 @@ fn build_personal_params_json( .next() .unwrap_or_default(); - let mut cornerstone = Map::new(); - cornerstone.insert( - "feTraceId".into(), - Value::String(Uuid::new_v4().to_string().to_lowercase()), - ); - cornerstone.insert("feURL".into(), Value::String(dashboard.to_string())); - cornerstone.insert("protocol".into(), Value::String("V2".into())); - cornerstone.insert("console".into(), Value::String("ONE_CONSOLE".into())); - cornerstone.insert("productCode".into(), Value::String("p_efm".into())); - // Let the gateway resolve the Personal/Solo session workspace. A captured - // Teams switchAgent is workspace-bound and rejects other accounts. - cornerstone.insert("switchUserType".into(), json!(3)); - cornerstone.insert("domain".into(), Value::String(domain.to_string())); - cornerstone.insert( - "consoleSite".into(), - Value::String(region.personal_console_site().to_string()), - ); - cornerstone.insert("userNickName".into(), Value::String(String::new())); - cornerstone.insert("userPrincipalName".into(), Value::String(String::new())); - cornerstone.insert("xsp_lang".into(), Value::String(LANGUAGE.into())); + let mut cornerstone = json!({ + "feTraceId": Uuid::new_v4().to_string().to_lowercase(), + "feURL": dashboard, + "protocol": "V2", + "console": "ONE_CONSOLE", + "productCode": "p_efm", + // Let the gateway resolve the Personal/Solo session workspace. A captured + // Teams switchAgent is workspace-bound and rejects other accounts. + "switchUserType": 3, + "domain": domain, + "consoleSite": region.personal_console_site(), + "userNickName": "", + "userPrincipalName": "", + "xsp_lang": LANGUAGE, + }); if let Some(cna) = cookie_value("cna", cookie_header) { - cornerstone.insert("X-Anonymous-Id".into(), Value::String(cna)); + cornerstone["X-Anonymous-Id"] = Value::String(cna); } - data_parameters.insert("cornerstoneParam".into(), Value::Object(cornerstone)); + data_parameters.insert("cornerstoneParam".into(), cornerstone); json!({ "Api": api, @@ -272,22 +238,7 @@ pub(super) fn parse_personal_usage( subscription_data: Option<&[u8]>, quota_config_data: Option<&[u8]>, ) -> Result { - if usage_data.is_empty() { - return Err(ProviderError::Parse( - "Empty Alibaba Token Plan Personal response".into(), - )); - } - - let value: Value = serde_json::from_slice(usage_data).map_err(|_| { - if is_likely_login_html(usage_data) { - ProviderError::AuthRequired - } else { - ProviderError::Parse("Invalid Alibaba Token Plan Personal JSON response".into()) - } - })?; - let expanded = expand_json_strings(value); - throw_if_error_payload(&expanded)?; - + let expanded = decode_console_payload(usage_data, "Alibaba Token Plan Personal")?; personal_usage_snapshot( &expanded, subscription_data, @@ -439,19 +390,7 @@ fn quota_totals_from_bytes(data: &[u8], plan_code: &str) -> Option } fn find_first_value_for_key(value: &Value, key: &str) -> Option { - match value { - Value::Object(map) => { - if let Some(nested) = map.get(key) { - return Some(nested.clone()); - } - map.values() - .find_map(|nested| find_first_value_for_key(nested, key)) - } - Value::Array(values) => values - .iter() - .find_map(|nested| find_first_value_for_key(nested, key)), - _ => None, - } + deep_find(value, &|node| node.as_object()?.get(key).cloned()) } #[cfg(test)] @@ -504,72 +443,110 @@ mod tests { assert_eq!(cornerstone.get("switchUserType"), Some(&json!(3))); } + /// Pins the whole personal `params` payload, including the console domain + /// taken from the region's dashboard URL and the optional `cna` id. #[test] - fn nested_workspace_error_surfaces_real_code_without_auth_eviction() { - let payload = json!({ - "code": "200", - "successResponse": true, - "data": { - "success": false, - "httpStatus": 200, - "errorCode": "BailianGateway.Workspace.NotAuthorised" + fn personal_params_json_wraps_data_with_cornerstone_fields() { + for (region, cookie, domain, site, anonymous_id) in [ + ( + AlibabaTokenPlanRegion::CnPersonal, + "cna=anon-1; other=x", + "bailian.console.aliyun.com", + "BAILIAN_ALIYUN", + Some("anon-1"), + ), + ( + AlibabaTokenPlanRegion::IntlPersonal, + "other=x", + "modelstudio.console.alibabacloud.com", + "MODELSTUDIO_ALBABACLOUD", + None, + ), + ] { + let mut data = Map::new(); + data.insert("commodityCode".into(), json!("test-code")); + let params = build_personal_params_json(PERSONAL_USAGE_API, data, cookie, region); + let value: Value = serde_json::from_str(¶ms).unwrap(); + let trace = value["Data"]["cornerstoneParam"]["feTraceId"] + .as_str() + .unwrap(); + assert_eq!(trace.len(), 36); + assert_eq!(trace, trace.to_lowercase()); + let mut cornerstone = json!({ + "feTraceId": trace, + "feURL": region.dashboard_url(), + "protocol": "V2", + "console": "ONE_CONSOLE", + "productCode": "p_efm", + "switchUserType": 3, + "domain": domain, + "consoleSite": site, + "userNickName": "", + "userPrincipalName": "", + "xsp_lang": "en-US", + }); + if let Some(id) = anonymous_id { + cornerstone["X-Anonymous-Id"] = json!(id); } - }); - - let error = - crate::providers::alibabatokenplan::throw_if_error_payload(&payload).unwrap_err(); - assert!(matches!( - error, - ProviderError::Other(message) - if message.contains("BailianGateway.Workspace.NotAuthorised") - )); - } - - #[test] - fn nested_gateway_error_prefers_error_message() { - let payload = json!({ - "code": "200", - "successResponse": true, - "data": { - "success": false, - "httpStatus": 200, - "errorCode": "BailianGateway.Quota.ServiceUnavailable", - "errorMsg": "quota service unavailable" - } - }); - - let error = - crate::providers::alibabatokenplan::throw_if_error_payload(&payload).unwrap_err(); - assert!(matches!( - error, - ProviderError::Other(message) if message.contains("quota service unavailable") - )); + let expected = json!({ + "Api": PERSONAL_USAGE_API, + "V": "1.0", + "Data": {"commodityCode": "test-code", "cornerstoneParam": cornerstone}, + }); + assert_eq!(params, expected.to_string()); + } } #[test] - fn success_envelope_without_windows_is_transient() { - let payload = json!({ - "code": "SUCCESS", - "successResponse": true, - "errorCode": "", - "data": {"success": true, "httpStatus": 200} - }); - assert!(personal_usage_success_without_windows( - payload.to_string().as_bytes() - )); + fn nested_gateway_errors_surface_without_auth_eviction() { + // (data frame, expected message fragment): the message wins over the code. + for (data, expected) in [ + ( + json!({ + "success": false, + "httpStatus": 200, + "errorCode": "BailianGateway.Workspace.NotAuthorised" + }), + "BailianGateway.Workspace.NotAuthorised", + ), + ( + json!({ + "success": false, + "httpStatus": 200, + "errorCode": "BailianGateway.Quota.ServiceUnavailable", + "errorMsg": "quota service unavailable" + }), + "quota service unavailable", + ), + ] { + let payload = json!({"code": "200", "successResponse": true, "data": data}); + let error = + crate::providers::alibabatokenplan::throw_if_error_payload(&payload).unwrap_err(); + assert!( + matches!(&error, ProviderError::Other(message) if message.contains(expected)), + "{error:?}" + ); + } } #[test] - fn success_envelope_with_windows_is_not_transient() { - let payload = json!({ - "code": "SUCCESS", - "successResponse": true, - "errorCode": "", - "data": {"per5HourPercentage": 0.5} - }); - assert!(!personal_usage_success_without_windows( - payload.to_string().as_bytes() - )); + fn success_envelope_is_transient_only_without_windows() { + for (data, transient) in [ + (json!({"success": true, "httpStatus": 200}), true), + (json!({"per5HourPercentage": 0.5}), false), + ] { + let payload = json!({ + "code": "SUCCESS", + "successResponse": true, + "errorCode": "", + "data": data + }); + assert_eq!( + personal_usage_success_without_windows(payload.to_string().as_bytes()), + transient, + "{payload}" + ); + } } #[test] diff --git a/rust/src/providers/alibabatokenplan/region.rs b/rust/src/providers/alibabatokenplan/region.rs index 2082746b07..e3d2aa9a66 100644 --- a/rust/src/providers/alibabatokenplan/region.rs +++ b/rust/src/providers/alibabatokenplan/region.rs @@ -179,60 +179,54 @@ mod tests { #[test] fn region_gateway_and_commodity_mapping() { - let cn = AlibabaTokenPlanRegion::Cn; - assert_eq!(cn.gateway_base_url(), "https://bailian.console.aliyun.com"); - assert_eq!(cn.quota_base_url(), cn.gateway_base_url()); - assert_eq!(cn.product_code(), "sfm_tokenplanteams_dp_cn"); - assert_eq!(cn.current_region_id(), "cn-beijing"); - assert!(!cn.uses_personal_api()); - - let intl = AlibabaTokenPlanRegion::Intl; - assert_eq!( - intl.gateway_base_url(), - "https://modelstudio.console.alibabacloud.com" - ); - assert_eq!(intl.quota_base_url(), intl.gateway_base_url()); - assert_eq!(intl.product_code(), "sfm_tokenplanteams_dp_intl"); - assert_eq!(intl.current_region_id(), "ap-southeast-1"); - assert!(!intl.uses_personal_api()); - - let cn_personal = AlibabaTokenPlanRegion::CnPersonal; - assert_eq!( - cn_personal.gateway_base_url(), - "https://bailian.console.aliyun.com" - ); - assert_eq!( - cn_personal.quota_base_url(), - "https://bailian-cs.console.aliyun.com" - ); - assert_eq!(cn_personal.product_code(), "sfm_tokenplansolo_public_cn"); - assert_eq!(cn_personal.current_region_id(), "cn-beijing"); - assert!(cn_personal.uses_personal_api()); - assert_eq!(cn_personal.personal_api_action(), "BroadScopeAspnGateway"); - assert_eq!(cn_personal.personal_console_site(), "BAILIAN_ALIYUN"); - - let intl_personal = AlibabaTokenPlanRegion::IntlPersonal; - assert_eq!( - intl_personal.gateway_base_url(), - "https://modelstudio.console.alibabacloud.com" - ); - assert_eq!( - intl_personal.quota_base_url(), - "https://bailian-singapore-cs.alibabacloud.com" - ); - assert_eq!( - intl_personal.product_code(), - "sfm_tokenplansolo_public_intl" - ); - assert_eq!(intl_personal.current_region_id(), "ap-southeast-1"); - assert!(intl_personal.uses_personal_api()); - assert_eq!( - intl_personal.personal_api_action(), - "IntlBroadScopeAspnGateway" - ); - assert_eq!( - intl_personal.personal_console_site(), - "MODELSTUDIO_ALBABACLOUD" - ); + use AlibabaTokenPlanRegion::*; + const CN_GATEWAY: &str = "https://bailian.console.aliyun.com"; + const INTL_GATEWAY: &str = "https://modelstudio.console.alibabacloud.com"; + // (region, gateway, quota base, product code, region id, personal (action, site)) + let rows = [ + ( + Cn, + CN_GATEWAY, + CN_GATEWAY, + "sfm_tokenplanteams_dp_cn", + "cn-beijing", + None, + ), + ( + Intl, + INTL_GATEWAY, + INTL_GATEWAY, + "sfm_tokenplanteams_dp_intl", + "ap-southeast-1", + None, + ), + ( + CnPersonal, + CN_GATEWAY, + "https://bailian-cs.console.aliyun.com", + "sfm_tokenplansolo_public_cn", + "cn-beijing", + Some(("BroadScopeAspnGateway", "BAILIAN_ALIYUN")), + ), + ( + IntlPersonal, + INTL_GATEWAY, + "https://bailian-singapore-cs.alibabacloud.com", + "sfm_tokenplansolo_public_intl", + "ap-southeast-1", + Some(("IntlBroadScopeAspnGateway", "MODELSTUDIO_ALBABACLOUD")), + ), + ]; + for (region, gateway, quota, product, region_id, personal) in rows { + assert_eq!(region.gateway_base_url(), gateway, "{region:?}"); + assert_eq!(region.quota_base_url(), quota, "{region:?}"); + assert_eq!(region.product_code(), product, "{region:?}"); + assert_eq!(region.current_region_id(), region_id, "{region:?}"); + assert_eq!(region.uses_personal_api(), personal.is_some(), "{region:?}"); + if let Some((action, site)) = personal { + assert_eq!(region.personal_api_action(), action, "{region:?}"); + assert_eq!(region.personal_console_site(), site, "{region:?}"); + } + } } } diff --git a/rust/src/providers/alibabatokenplan/tests.rs b/rust/src/providers/alibabatokenplan/tests.rs new file mode 100644 index 0000000000..a05002186f --- /dev/null +++ b/rust/src/providers/alibabatokenplan/tests.rs @@ -0,0 +1,145 @@ +use super::*; + +#[test] +fn missing_bailian_cli_is_local_runtime_offline() { + let provider = AlibabaTokenPlanProvider::new(); + let error = ProviderError::NotInstalled( + "Bailian CLI 'bl' is not installed or not on PATH.".to_string(), + ); + assert_eq!( + provider.error_state_kind(&error), + crate::core::ProviderStateKind::LocalRuntimeOffline + ); +} +#[test] +fn parses_token_plan_instance_payload() { + let payload = serde_json::json!({ + "data": { + "tokenPlanInstanceInfo": { + "commodityName": "Token Plan Pro", + "quotaInfo": { + "usedQuota": "1250", + "totalQuota": "5000" + }, + "nextRefreshTime": 1780763009000_i64 + } + } + }); + let snapshot = + AlibabaTokenPlanProvider::parse_usage_snapshot(payload.to_string().as_bytes()).unwrap(); + assert_eq!(snapshot.plan_name.as_deref(), Some("Token Plan Pro")); + assert_eq!(snapshot.used_quota, Some(1250.0)); + assert_eq!(snapshot.total_quota, Some(5000.0)); + + let usage = AlibabaTokenPlanProvider::snapshot_to_usage(snapshot).unwrap(); + assert_eq!(usage.primary.used_percent, 25.0); + assert_eq!( + usage.primary.reset_description.as_deref(), + Some("1,250 / 5,000 credits used") + ); + assert_eq!(usage.login_method.as_deref(), Some("Token Plan Pro")); +} + +#[test] +fn expands_nested_string_payloads_and_uses_remaining_quota() { + let nested = serde_json::json!({ + "successResponse": serde_json::json!({ + "instances": [ + {"status": "EXPIRED", "quota": 1000, "remaining": 1000}, + {"status": "ACTIVE", "packageName": "Team", "quota": 1000, "remaining": 250} + ] + }).to_string() + }); + let snapshot = + AlibabaTokenPlanProvider::parse_usage_snapshot(nested.to_string().as_bytes()).unwrap(); + assert_eq!(snapshot.plan_name.as_deref(), Some("Team")); + assert_eq!( + used_percent( + snapshot.used_quota, + snapshot.total_quota, + snapshot.remaining_quota + ), + Some(75.0) + ); +} + +#[test] +fn parses_new_subscription_summary_payload() { + let payload = serde_json::json!({ + "success": true, + "Data": { + "ProductName": "Token Plan Team", + "TotalValue": "1000000", + "TotalSurplusValue": "250000", + "NearestExpireDate": "2026-06-30" + } + }); + let snapshot = + AlibabaTokenPlanProvider::parse_usage_snapshot(payload.to_string().as_bytes()).unwrap(); + assert_eq!(snapshot.plan_name.as_deref(), Some("Token Plan Team")); + assert_eq!(snapshot.total_quota, Some(1_000_000.0)); + assert_eq!(snapshot.remaining_quota, Some(250_000.0)); + assert_eq!( + used_percent( + snapshot.used_quota, + snapshot.total_quota, + snapshot.remaining_quota + ), + Some(75.0) + ); + assert!(snapshot.resets_at.is_some()); +} + +#[test] +fn detects_login_payloads() { + let err = AlibabaTokenPlanProvider::parse_usage_snapshot( + br#"{"code":"NeedLogin","message":"please login"}"#, + ) + .unwrap_err(); + assert!(matches!(err, ProviderError::AuthRequired)); +} + +#[test] +fn extracts_sec_token_from_html_or_cookie() { + assert_eq!( + extract_sec_token(r#""#).as_deref(), + Some("abc123") + ); + assert_eq!( + extract_sec_token( + r#""# + ) + .as_deref(), + Some("upper123") + ); + assert_eq!( + cookie_value("sec_token", "foo=bar; sec_token=xyz"), + Some("xyz".to_string()) + ); +} + +#[test] +fn sec_token_shell_referer_uses_same_origin_root() { + assert_eq!( + dashboard_referer(Region::CnPersonal), + "https://bailian.console.aliyun.com/" + ); + assert_eq!( + dashboard_referer(Region::IntlPersonal), + "https://modelstudio.console.alibabacloud.com/" + ); +} + +#[test] +fn default_region_cn_team_urls_match_legacy() { + let region = Region::Cn; + assert_eq!( + AlibabaTokenPlanProvider::team_quota_url(region), + "https://bailian.console.aliyun.com/data/api.json?action=GetSubscriptionSummary&product=BssOpenAPI-V3&_tag=" + ); + assert_eq!( + AlibabaTokenPlanProvider::team_request_params(region), + serde_json::json!({"ProductCode": "sfm_tokenplanteams_dp_cn"}).to_string() + ); + assert_eq!(region.current_region_id(), "cn-beijing"); +} diff --git a/rust/src/providers/bedrock/mod.rs b/rust/src/providers/bedrock/mod.rs index 17c02d870d..21af95bf21 100644 --- a/rust/src/providers/bedrock/mod.rs +++ b/rust/src/providers/bedrock/mod.rs @@ -135,28 +135,23 @@ impl BedrockProvider { if json_profile_name(&json).is_some() { return None; } - let access_key_id = json - .get("access_key_id") - .or_else(|| json.get("accessKeyId")) - .or_else(|| json.get("AWS_ACCESS_KEY_ID")) - .and_then(|v| v.as_str()) - .map(str::trim) - .filter(|v| !v.is_empty())?; - let secret_access_key = json - .get("secret_access_key") - .or_else(|| json.get("secretAccessKey")) - .or_else(|| json.get("AWS_SECRET_ACCESS_KEY")) - .and_then(|v| v.as_str()) - .map(str::trim) - .filter(|v| !v.is_empty())?; - let session_token = json - .get("session_token") - .or_else(|| json.get("sessionToken")) - .or_else(|| json.get("AWS_SESSION_TOKEN")) - .and_then(|v| v.as_str()) - .map(str::trim) - .filter(|v| !v.is_empty()) - .map(str::to_string); + let access_key_id = json_str( + &json, + &["access_key_id", "accessKeyId", "AWS_ACCESS_KEY_ID"], + )?; + let secret_access_key = json_str( + &json, + &[ + "secret_access_key", + "secretAccessKey", + "AWS_SECRET_ACCESS_KEY", + ], + )?; + let session_token = json_str( + &json, + &["session_token", "sessionToken", "AWS_SESSION_TOKEN"], + ) + .map(str::to_string); return Some(AwsCredentials { access_key_id: access_key_id.to_string(), @@ -202,18 +197,14 @@ impl BedrockProvider { } fn credentials_from_env() -> Result { - let access_key_id = cleaned_env("AWS_ACCESS_KEY_ID").ok_or_else(|| { + let missing = || { ProviderError::NotInstalled( "AWS credentials not configured. Set AWS_ACCESS_KEY_ID and AWS_SECRET_ACCESS_KEY." .to_string(), ) - })?; - let secret_access_key = cleaned_env("AWS_SECRET_ACCESS_KEY").ok_or_else(|| { - ProviderError::NotInstalled( - "AWS credentials not configured. Set AWS_ACCESS_KEY_ID and AWS_SECRET_ACCESS_KEY." - .to_string(), - ) - })?; + }; + let access_key_id = cleaned_env("AWS_ACCESS_KEY_ID").ok_or_else(missing)?; + let secret_access_key = cleaned_env("AWS_SECRET_ACCESS_KEY").ok_or_else(missing)?; Ok(AwsCredentials { access_key_id, @@ -240,7 +231,7 @@ impl BedrockProvider { } fn credentials_from_profile(profile: &str) -> Result { - let aws = aws_cli_path()?; + let aws = aws_cli_path(); let mut command = std::process::Command::new(&aws); command .args([ @@ -360,41 +351,16 @@ impl BedrockProvider { region: &str, ) -> Result { let endpoint = format!("https://monitoring.{region}.amazonaws.com"); - let body_bytes = cloudwatch_request_body()?; - let body_hash = sha256_hex(&body_bytes); - let now = Utc::now(); - let amz_date = now.format("%Y%m%dT%H%M%SZ").to_string(); - let date_stamp = now.format("%Y%m%d").to_string(); - let authorization = sign_authorization_for( - credentials, - AwsSigningRequest { - date_stamp: &date_stamp, - amz_date: &amz_date, - body_hash: &body_hash, - url: &endpoint, - body: &body_bytes, - target: CLOUDWATCH_TARGET, + let response = self + .signed_post( + credentials, + &endpoint, + cloudwatch_request_body()?, + CLOUDWATCH_TARGET, region, - service: CLOUDWATCH_SERVICE, - }, - )?; - let host = reqwest::Url::parse(&endpoint) - .ok() - .and_then(|u| u.host_str().map(str::to_string)) - .unwrap_or_else(|| format!("monitoring.{region}.amazonaws.com")); - let mut request = self - .client - .post(endpoint) - .header("Content-Type", "application/x-amz-json-1.1") - .header("Host", host) - .header("X-Amz-Target", CLOUDWATCH_TARGET) - .header("X-Amz-Date", amz_date) - .header("x-amz-content-sha256", body_hash) - .header("Authorization", authorization); - if let Some(token) = &credentials.session_token { - request = request.header("X-Amz-Security-Token", token); - } - let response = request.body(body_bytes).send().await?; + CLOUDWATCH_SERVICE, + ) + .await?; let status = response.status(); let text = response.text().await?; if !status.is_success() { @@ -418,49 +384,65 @@ impl BedrockProvider { granularity: &str, next_page_token: Option<&str>, ) -> Result { - let body_bytes = cost_request_body(start_date, end_date, granularity, next_page_token)?; - let body_hash = sha256_hex(&body_bytes); - let now = Utc::now(); - let amz_date = now.format("%Y%m%dT%H%M%SZ").to_string(); - let date_stamp = now.format("%Y%m%d").to_string(); - let authorization = sign_authorization( - credentials, - &date_stamp, - &amz_date, - &body_hash, - COST_EXPLORER_URL, - &body_bytes, - )?; - + let body = cost_request_body(start_date, end_date, granularity, next_page_token)?; let response = self - .signed_cost_request(credentials, amz_date, body_hash, authorization) - .body(body_bytes) - .send() + .signed_post( + credentials, + COST_EXPLORER_URL, + body, + COST_EXPLORER_TARGET, + SIGNING_REGION, + SERVICE, + ) .await?; parse_cost_response(response).await } - fn signed_cost_request( + /// SigV4-signs `body` for `target` and POSTs it to `endpoint`. + async fn signed_post( &self, credentials: &AwsCredentials, - amz_date: String, - body_hash: String, - authorization: String, - ) -> reqwest::RequestBuilder { - let request = self + endpoint: &str, + body: Vec, + target: &str, + region: &str, + service: &str, + ) -> Result { + let body_hash = sha256_hex(&body); + let now = Utc::now(); + let amz_date = now.format("%Y%m%dT%H%M%SZ").to_string(); + let date_stamp = now.format("%Y%m%d").to_string(); + let authorization = sign_authorization_for( + credentials, + AwsSigningRequest { + date_stamp: &date_stamp, + amz_date: &amz_date, + body_hash: &body_hash, + url: endpoint, + body: &body, + target, + region, + service, + }, + )?; + // The signer already rejected an unparseable endpoint. + let host = reqwest::Url::parse(endpoint) + .ok() + .and_then(|u| u.host_str().map(str::to_string)) + .unwrap_or_default(); + let mut request = self .client - .post(COST_EXPLORER_URL) + .post(endpoint) .header("Content-Type", "application/x-amz-json-1.1") - .header("Host", "ce.us-east-1.amazonaws.com") - .header("X-Amz-Target", COST_EXPLORER_TARGET) + .header("Host", host) + .header("X-Amz-Target", target) .header("X-Amz-Date", amz_date) .header("x-amz-content-sha256", body_hash) .header("Authorization", authorization); - - match &credentials.session_token { - Some(token) => request.header("X-Amz-Security-Token", token), - None => request, + if let Some(token) = &credentials.session_token { + request = request.header("X-Amz-Security-Token", token); } + Ok(request.body(body).send().await?) } async fn fetch_via_api( @@ -623,20 +605,23 @@ fn cleaned_env(key: &str) -> Option { .filter(|value| !value.is_empty()) } -fn json_profile_name(json: &Value) -> Option { - json.get("profile") - .or_else(|| json.get("aws_profile")) - .or_else(|| json.get("AWS_PROFILE")) - .and_then(|v| v.as_str()) +/// Trimmed non-empty string under the first of `keys` present in `json`. +fn json_str<'a>(json: &'a Value, keys: &[&str]) -> Option<&'a str> { + keys.iter() + .find_map(|key| json.get(*key)) + .and_then(Value::as_str) .map(str::trim) .filter(|v| !v.is_empty()) - .map(str::to_string) } -fn aws_cli_path() -> Result { - Ok(cleaned_env("CODEXBAR_AWS_CLI_PATH") +fn json_profile_name(json: &Value) -> Option { + json_str(json, &["profile", "aws_profile", "AWS_PROFILE"]).map(str::to_string) +} + +fn aws_cli_path() -> String { + cleaned_env("CODEXBAR_AWS_CLI_PATH") .or_else(|| cleaned_env("AWS_CLI_PATH")) - .unwrap_or_else(|| "aws".to_string())) + .unwrap_or_else(|| "aws".to_string()) } fn map_aws_profile_error(profile: &str, stderr: &str) -> ProviderError { @@ -660,28 +645,14 @@ fn parse_aws_profile_credentials(stdout: &[u8]) -> Result f64 { .sum() } -fn sign_authorization( - credentials: &AwsCredentials, - date_stamp: &str, - amz_date: &str, - body_hash: &str, - url: &str, - body: &[u8], -) -> Result { - sign_authorization_for( - credentials, - AwsSigningRequest { - date_stamp, - amz_date, - body_hash, - url, - body, - target: COST_EXPLORER_TARGET, - region: SIGNING_REGION, - service: SERVICE, - }, - ) -} - fn sign_authorization_for( credentials: &AwsCredentials, request: AwsSigningRequest<'_>, @@ -764,25 +712,21 @@ fn sign_authorization_for( let parsed = reqwest::Url::parse(request.url) .map_err(|e| ProviderError::Other(format!("Invalid AWS endpoint URL: {e}")))?; let host = parsed.host_str().unwrap_or("ce.us-east-1.amazonaws.com"); - let (canonical_headers, signed_headers) = if let Some(session_token) = - &credentials.session_token - { - ( - format!( - "content-type:application/x-amz-json-1.1\nhost:{host}\nx-amz-content-sha256:{}\nx-amz-date:{}\nx-amz-security-token:{session_token}\nx-amz-target:{}\n", - request.body_hash, request.amz_date, request.target - ), + // The security token header sorts between x-amz-date and x-amz-target. + let (token_header, signed_headers) = match &credentials.session_token { + Some(token) => ( + format!("x-amz-security-token:{token}\n"), "content-type;host;x-amz-content-sha256;x-amz-date;x-amz-security-token;x-amz-target", - ) - } else { - ( - format!( - "content-type:application/x-amz-json-1.1\nhost:{host}\nx-amz-content-sha256:{}\nx-amz-date:{}\nx-amz-target:{}\n", - request.body_hash, request.amz_date, request.target - ), + ), + None => ( + String::new(), "content-type;host;x-amz-content-sha256;x-amz-date;x-amz-target", - ) + ), }; + let canonical_headers = format!( + "content-type:application/x-amz-json-1.1\nhost:{host}\nx-amz-content-sha256:{}\nx-amz-date:{}\n{token_header}x-amz-target:{}\n", + request.body_hash, request.amz_date, request.target + ); let canonical_request = [ "POST", "/", @@ -826,113 +770,4 @@ fn sanitized_body(body: &str) -> String { } #[cfg(test)] -mod tests { - use super::*; - - #[test] - fn parses_bedrock_cost_only() { - let page = json!({ - "ResultsByTime": [{ - "Groups": [ - { - "Keys": ["Amazon Bedrock"], - "Metrics": { "UnblendedCost": { "Amount": "12.34" } } - }, - { - "Keys": ["Amazon S3"], - "Metrics": { "UnblendedCost": { "Amount": "99.00" } } - } - ] - }] - }); - assert_eq!(parse_bedrock_cost(&page), 12.34); - } - - #[test] - fn parses_cloudwatch_claude_activity() { - let activity = parse_claude_activity(&json!({ - "MetricDataResults": [ - {"Id": "input", "Values": [10, 15]}, - {"Id": "output", "Values": [7]}, - {"Id": "requests", "Values": [2, 3]} - ] - })); - assert_eq!(activity.input_tokens, 25.0); - assert_eq!(activity.output_tokens, 7.0); - assert_eq!(activity.request_count, 5.0); - } - - #[test] - fn parses_context_credentials_from_json() { - let credentials = BedrockProvider::credentials_from_context(Some( - r#"{ - "access_key_id": "AKIAEXAMPLE", - "secret_access_key": "secret", - "session_token": "session" - }"#, - )) - .expect("credentials"); - - assert_eq!(credentials.access_key_id, "AKIAEXAMPLE"); - assert_eq!(credentials.secret_access_key, "secret"); - assert_eq!(credentials.session_token.as_deref(), Some("session")); - } - - #[test] - fn parses_context_credentials_from_colon_delimited_value() { - let credentials = - BedrockProvider::credentials_from_context(Some("AKIAEXAMPLE:secret:session")) - .expect("credentials"); - - assert_eq!(credentials.access_key_id, "AKIAEXAMPLE"); - assert_eq!(credentials.secret_access_key, "secret"); - assert_eq!(credentials.session_token.as_deref(), Some("session")); - } - - #[test] - fn parses_profile_from_context_prefix() { - assert_eq!( - BedrockProvider::profile_from_context(Some("profile:production")).as_deref(), - Some("production") - ); - assert!(BedrockProvider::credentials_from_context(Some("profile:production")).is_none()); - } - - #[test] - fn parses_profile_from_context_json() { - assert_eq!( - BedrockProvider::profile_from_context(Some(r#"{"aws_profile":"sso-dev"}"#)).as_deref(), - Some("sso-dev") - ); - assert!( - BedrockProvider::credentials_from_context(Some(r#"{"aws_profile":"sso-dev"}"#)) - .is_none() - ); - } - - #[test] - fn parses_aws_cli_export_credentials_output() { - let credentials = parse_aws_profile_credentials( - br#"{ - "Version": 1, - "AccessKeyId": "ASIAEXAMPLE", - "SecretAccessKey": "secret", - "SessionToken": "session" - }"#, - ) - .expect("aws profile credentials"); - - assert_eq!(credentials.access_key_id, "ASIAEXAMPLE"); - assert_eq!(credentials.secret_access_key, "secret"); - assert_eq!(credentials.session_token.as_deref(), Some("session")); - } - - #[test] - fn hmac_sha256_matches_rfc_4231_case_1() { - let digest = hmac_sha256(&[0x0b; 20], b"Hi There"); - assert_eq!( - hex(&digest), - "b0344c61d8db38535ca8afceaf0bf12b881dc200c9833da726e9376c2e32cff7" - ); - } -} +mod tests; diff --git a/rust/src/providers/bedrock/tests.rs b/rust/src/providers/bedrock/tests.rs new file mode 100644 index 0000000000..31ea5c7784 --- /dev/null +++ b/rust/src/providers/bedrock/tests.rs @@ -0,0 +1,151 @@ +use super::*; + +#[test] +fn parses_bedrock_cost_only() { + let page = json!({ + "ResultsByTime": [{ + "Groups": [ + { + "Keys": ["Amazon Bedrock"], + "Metrics": { "UnblendedCost": { "Amount": "12.34" } } + }, + { + "Keys": ["Amazon S3"], + "Metrics": { "UnblendedCost": { "Amount": "99.00" } } + } + ] + }] + }); + assert_eq!(parse_bedrock_cost(&page), 12.34); +} + +#[test] +fn parses_cloudwatch_claude_activity() { + let activity = parse_claude_activity(&json!({ + "MetricDataResults": [ + {"Id": "input", "Values": [10, 15]}, + {"Id": "output", "Values": [7]}, + {"Id": "requests", "Values": [2, 3]} + ] + })); + assert_eq!(activity.input_tokens, 25.0); + assert_eq!(activity.output_tokens, 7.0); + assert_eq!(activity.request_count, 5.0); +} + +#[test] +fn parses_context_credentials_from_json() { + let credentials = BedrockProvider::credentials_from_context(Some( + r#"{ + "access_key_id": "AKIAEXAMPLE", + "secret_access_key": "secret", + "session_token": "session" + }"#, + )) + .expect("credentials"); + + assert_eq!(credentials.access_key_id, "AKIAEXAMPLE"); + assert_eq!(credentials.secret_access_key, "secret"); + assert_eq!(credentials.session_token.as_deref(), Some("session")); +} + +#[test] +fn parses_context_credentials_from_colon_delimited_value() { + let credentials = BedrockProvider::credentials_from_context(Some("AKIAEXAMPLE:secret:session")) + .expect("credentials"); + + assert_eq!(credentials.access_key_id, "AKIAEXAMPLE"); + assert_eq!(credentials.secret_access_key, "secret"); + assert_eq!(credentials.session_token.as_deref(), Some("session")); +} + +#[test] +fn parses_profile_from_context_prefix() { + assert_eq!( + BedrockProvider::profile_from_context(Some("profile:production")).as_deref(), + Some("production") + ); + assert!(BedrockProvider::credentials_from_context(Some("profile:production")).is_none()); +} + +#[test] +fn parses_profile_from_context_json() { + assert_eq!( + BedrockProvider::profile_from_context(Some(r#"{"aws_profile":"sso-dev"}"#)).as_deref(), + Some("sso-dev") + ); + assert!( + BedrockProvider::credentials_from_context(Some(r#"{"aws_profile":"sso-dev"}"#)).is_none() + ); +} + +#[test] +fn parses_aws_cli_export_credentials_output() { + let credentials = parse_aws_profile_credentials( + br#"{ + "Version": 1, + "AccessKeyId": "ASIAEXAMPLE", + "SecretAccessKey": "secret", + "SessionToken": "session" + }"#, + ) + .expect("aws profile credentials"); + + assert_eq!(credentials.access_key_id, "ASIAEXAMPLE"); + assert_eq!(credentials.secret_access_key, "secret"); + assert_eq!(credentials.session_token.as_deref(), Some("session")); +} + +/// Golden SigV4 Authorization values pinned from the current signer, one +/// per canonical-header shape (with and without a session token). +#[test] +fn sigv4_authorization_matches_golden_values() { + let body = br#"{"Granularity":"MONTHLY"}"#; + let body_hash = sha256_hex(body); + let mut credentials = AwsCredentials { + access_key_id: "AKIDEXAMPLE".to_string(), + secret_access_key: "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY".to_string(), + session_token: None, + }; + let request = AwsSigningRequest { + date_stamp: "20260115", + amz_date: "20260115T123456Z", + body_hash: &body_hash, + url: COST_EXPLORER_URL, + body, + target: COST_EXPLORER_TARGET, + region: SIGNING_REGION, + service: SERVICE, + }; + let cost_explorer = sign_authorization_for(&credentials, request).unwrap(); + assert_eq!( + cost_explorer, + "AWS4-HMAC-SHA256 Credential=AKIDEXAMPLE/20260115/us-east-1/ce/aws4_request, SignedHeaders=content-type;host;x-amz-content-sha256;x-amz-date;x-amz-target, Signature=bd41e425427d67f7b3e3979d3f1f617285addc2af80ebd24990b1be60c2cbaa2" + ); + + credentials.session_token = Some("session-token-example".to_string()); + let request = AwsSigningRequest { + date_stamp: "20260115", + amz_date: "20260115T123456Z", + body_hash: &body_hash, + url: "https://monitoring.eu-west-1.amazonaws.com", + body, + target: CLOUDWATCH_TARGET, + region: "eu-west-1", + service: CLOUDWATCH_SERVICE, + }; + let cloudwatch = sign_authorization_for(&credentials, request).unwrap(); + assert_eq!( + cloudwatch, + "AWS4-HMAC-SHA256 Credential=AKIDEXAMPLE/20260115/eu-west-1/monitoring/aws4_request, SignedHeaders=content-type;host;x-amz-content-sha256;x-amz-date;x-amz-security-token;x-amz-target, Signature=ed9dd8a45f767b557cc7f86b1870456e2fe9c93d908b35d12e1024165cd263ca" + ); +} + +#[test] +fn hmac_sha256_matches_rfc_4231_case_1() { + let digest = hmac_sha256(&[0x0b; 20], b"Hi There"); + assert_eq!( + hex(&digest), + "b0344c61d8db38535ca8afceaf0bf12b881dc200c9833da726e9376c2e32cff7" + ); +} diff --git a/rust/src/providers/codebuddy/mod.rs b/rust/src/providers/codebuddy/mod.rs index ef1d040694..f1bfabda95 100644 --- a/rust/src/providers/codebuddy/mod.rs +++ b/rust/src/providers/codebuddy/mod.rs @@ -377,34 +377,23 @@ fn validated_package_codes(items: &[Value]) -> Vec { .collect() } +fn env_path(key: &str) -> Option { + std::env::var(key) + .ok() + .filter(|value| !value.is_empty()) + .map(PathBuf::from) +} + fn codebuddy_home() -> Option { - if let Ok(home) = std::env::var("CODEBUDDY_HOME") { - let p = PathBuf::from(home); - if !p.as_os_str().is_empty() { - return Some(p); - } - } - dirs::home_dir().map(|h| h.join(".codebuddy")) + env_path("CODEBUDDY_HOME").or_else(|| dirs::home_dir().map(|h| h.join(".codebuddy"))) } fn cookie_file_path() -> Option { - if let Ok(p) = std::env::var("CB_COOKIE_FILE") { - let path = PathBuf::from(p); - if !path.as_os_str().is_empty() { - return Some(path); - } - } - codebuddy_home().map(|h| h.join("cb_cookie.txt")) + env_path("CB_COOKIE_FILE").or_else(|| codebuddy_home().map(|h| h.join("cb_cookie.txt"))) } fn credits_cache_path() -> Option { - if let Ok(p) = std::env::var("CB_CREDITS_FILE") { - let path = PathBuf::from(p); - if !path.as_os_str().is_empty() { - return Some(path); - } - } - codebuddy_home().map(|h| h.join("cb_credits.json")) + env_path("CB_CREDITS_FILE").or_else(|| codebuddy_home().map(|h| h.join("cb_credits.json"))) } fn read_cookie_file() -> Option { @@ -788,474 +777,4 @@ impl Provider for CodeBuddyProvider { } #[cfg(test)] -mod tests { - use super::*; - - fn api_payload() -> Value { - serde_json::json!({ - "code": 0, - "msg": "ok", - "data": { - "Response": { - "Data": { - "Accounts": [ - { - "CapacitySizePrecise": "2000", - "CapacityUsedPrecise": "100", - "CapacityRemainPrecise": "1900", - "ExpireTime": "2026-09-01T00:00:00Z" - }, - { - "CapacitySize": 1100, - "CapacityUsed": 11, - "CapacityRemain": 1089 - } - ] - } - } - } - }) - } - - #[test] - fn parses_get_user_resource_payload_into_typed_totals() { - let totals = totals_from_api_payload(&api_payload()).unwrap(); - assert_eq!(totals.total, 3100.0); - assert_eq!(totals.used, 111.0); - assert_eq!(totals.remaining, 2989.0); - assert!(totals.reset.is_some()); - - let snapshot = snapshot_from_totals(&totals); - assert!((snapshot.primary.used_percent - (111.0 / 3100.0 * 100.0)).abs() < 0.01); - assert_eq!( - snapshot.primary.reset_description.as_deref(), - Some("2,989 / 3,100 left") - ); - assert!(snapshot.primary.resets_at.is_some()); - } - - #[test] - fn auth_flavoured_payload_maps_to_auth_required() { - for msg in ["未登录", "登录已过期", "auth token expired"] { - let payload = serde_json::json!({ "code": 14001, "msg": msg }); - let err = totals_from_api_payload(&payload).unwrap_err(); - assert!( - matches!(err, ProviderError::AuthRequired), - "expected AuthRequired for msg={msg:?}, got {err}" - ); - } - } - - #[test] - fn empty_accounts_errors_with_hint() { - let payload = serde_json::json!({ - "code": 0, - "data": { "Response": { "Data": { "Accounts": [] } } } - }); - let err = totals_from_api_payload(&payload).unwrap_err(); - assert!(format!("{err}").contains("PackageCodes") || format!("{err}").contains("package")); - } - - #[test] - fn non_zero_code_errors() { - let payload = serde_json::json!({ "code": 14001, "msg": "quote exceeded" }); - let err = totals_from_api_payload(&payload).unwrap_err(); - assert!(matches!(err, ProviderError::Other(_)), "got {err}"); - } - - #[test] - fn payloads_with_non_finite_values_are_rejected() { - // serde_json cannot carry NaN/inf, but an absurd +/-1e308 pair can - // still overflow the sum to inf — that must be rejected, not cached. - let payload = serde_json::json!({ - "code": 0, - "data": { "Response": { "Data": { "Accounts": [ - { "CapacitySize": 1e308, "CapacityUsed": 1e308, "CapacityRemain": 1e308 }, - { "CapacitySize": 1e308, "CapacityUsed": 0, "CapacityRemain": 0 } - ] } } } - }); - assert!(totals_from_api_payload(&payload).is_err()); - } - - #[test] - fn formats_compact_credit_labels() { - assert_eq!(format_credits_short(1989.0, 3100.0), "1,989 / 3,100 left"); - assert_eq!(format_credits_short(12.5, 100.0), "12.5 / 100 left"); - } - - #[test] - fn normalizes_cookie_with_caret_escapes() { - assert_eq!( - normalize_cookie_header("Cookie: a=1^|2; b=3").as_deref(), - Some("a=1|2; b=3") - ); - assert_eq!(normalize_cookie_header(" "), None); - assert_eq!(normalize_cookie_header("Cookie:"), None); - } - - #[test] - fn cache_round_trip_preserves_typed_totals_exactly() { - let totals = CreditTotals { - total: 3100.5, - used: 111.25, - remaining: 2989.25, - reset: Some( - DateTime::parse_from_rfc3339("2026-09-01T08:30:00Z") - .unwrap() - .into(), - ), - }; - let json = cache_json_from_totals(&totals, Some("0123456789abcdef")); - let parsed = totals_from_cache_json(&json).unwrap(); - assert_eq!(parsed, totals); - // The cache carries the fingerprint verbatim. - assert_eq!( - json.get("accountHash").and_then(|v| v.as_str()), - Some("0123456789abcdef") - ); - } - - #[test] - fn cache_file_round_trip_via_disk() { - let dir = tempfile::tempdir().unwrap(); - let path = dir.path().join("cb_credits.json"); - let totals = CreditTotals { - total: 2000.0, - used: 100.25, - remaining: 1899.75, - reset: None, - }; - write_credits_cache(&path, &totals, Some("feedbeefcafe0001")).unwrap(); - - let value = read_credits_cache(&path).unwrap(); - let parsed = totals_from_cache_json(&value).unwrap(); - assert_eq!(parsed, totals); - - // Raw file exposes typed JSON numbers, not display text. - let raw = std::fs::read_to_string(&path).unwrap(); - assert!(raw.contains("\"total\": 2000.0")); - assert!(raw.contains("feedbeefcafe0001")); - } - - #[test] - fn hostile_cache_inputs_are_rejected() { - // Missing total. - let err = totals_from_cache_json(&serde_json::json!({"used": 1})).unwrap_err(); - assert!(format!("{err}").contains("missing total")); - // Negative total. - assert!(totals_from_cache_json(&serde_json::json!({"total": -5})).is_err()); - // Non-JSON is a Parse error at the read layer. - let dir = tempfile::tempdir().unwrap(); - let path = dir.path().join("bad.json"); - std::fs::write(&path, b"{not json").unwrap(); - assert!(matches!( - read_credits_cache(&path), - Err(ProviderError::Parse(_)) - )); - // Oversized cache is refused before parsing. - let big = dir.path().join("big.json"); - #[allow( - clippy::cast_possible_truncation, - reason = "MAX_CACHE_BYTES is 1 MiB; +1 stays far inside usize on any supported target" - )] - std::fs::write(&big, vec![b' '; (MAX_CACHE_BYTES + 1) as usize]).unwrap(); - assert!(read_credits_cache(&big).is_err()); - } - - #[test] - fn validated_package_codes_filters_and_caps() { - let long_code = "x".repeat(MAX_PACKAGE_CODE_LEN + 1); - let items = vec![ - json!("TCACA_code_001_ok"), - json!(" "), - json!(""), - json!(42), - json!("aB\u{0007}c"), - json!(long_code), - json!(" TCACA_code_002_trimmed "), - ]; - let codes = validated_package_codes(&items); - assert_eq!( - codes, - vec![ - "TCACA_code_001_ok".to_string(), - "TCACA_code_002_trimmed".to_string(), - ] - ); - - // Cap: more than MAX_PACKAGE_CODES valid entries are truncated. - let many: Vec = (0..MAX_PACKAGE_CODES + 10) - .map(|i| json!(format!("code_{i}"))) - .collect(); - assert_eq!(validated_package_codes(&many).len(), MAX_PACKAGE_CODES); - } - - #[test] - fn validate_api_url_rules() { - assert_eq!( - validate_api_url("https://www.codebuddy.cn/x").unwrap(), - "https://www.codebuddy.cn/x" - ); - assert!(validate_api_url("http://127.0.0.1:8080/x").is_ok()); - assert!(validate_api_url("http://localhost:8080/x").is_ok()); - assert!(validate_api_url("http://[::1]:8080/x").is_ok()); - assert!(validate_api_url("http://example.com/x").is_err()); - assert!(validate_api_url("ftp://example.com/x").is_err()); - assert!(validate_api_url("not a url").is_err()); - assert!(validate_api_url(" ").is_err()); - } - - #[test] - fn failure_classification_for_http_statuses() { - assert!(failure_for_status(StatusCode::OK).is_none()); - - for status in [StatusCode::UNAUTHORIZED, StatusCode::FORBIDDEN] { - let fail = failure_for_status(status).unwrap(); - assert!(!fail.is_transient(), "{status} must be permanent"); - assert!(matches!(fail.into_error(), ProviderError::AuthRequired)); - } - for status in [ - StatusCode::TOO_MANY_REQUESTS, - StatusCode::INTERNAL_SERVER_ERROR, - StatusCode::BAD_GATEWAY, - StatusCode::SERVICE_UNAVAILABLE, - ] { - assert!( - failure_for_status(status).unwrap().is_transient(), - "{status} must be transient" - ); - } - // Other failures (400 etc.) are permanent but not auth errors. - let fail = failure_for_status(StatusCode::BAD_REQUEST).unwrap(); - assert!(!fail.is_transient()); - assert!(matches!(fail.into_error(), ProviderError::Other(_))); - } - - #[test] - fn cache_fallback_requires_auto_mode_and_transient_failure() { - let transient = FetchFailure::Transient(ProviderError::Other("flake".into())); - let permanent_auth = FetchFailure::Permanent(ProviderError::AuthRequired); - let transient_ref = &transient; - let auth_ref = &permanent_auth; - - // Auth never masks behind a stale cache. - assert!(!cache_fallback_allowed(SourceMode::Auto, auth_ref)); - // Transient failures may fall back in Auto only. - assert!(cache_fallback_allowed(SourceMode::Auto, transient_ref)); - assert!(!cache_fallback_allowed(SourceMode::Web, transient_ref)); - assert!(!cache_fallback_allowed(SourceMode::Cli, transient_ref)); - } - - #[test] - fn cookie_fingerprint_is_stable_per_account_and_secret_free() { - let fp_a = cookie_fingerprint("session=aaa; uid=1"); - let fp_b = cookie_fingerprint("session=bbb; uid=2"); - assert_eq!(fp_a, cookie_fingerprint("session=aaa; uid=1")); - assert_ne!(fp_a, fp_b); - assert_eq!(fp_a.len(), 16); - assert!(fp_a.chars().all(|c| c.is_ascii_hexdigit())); - } - - #[test] - fn parse_datetime_accepts_rfc3339_naive_and_epochs() { - assert!(parse_datetime("2026-09-01T00:00:00Z").is_some()); - assert!(parse_datetime("2026-09-01 12:30:00").is_some()); - assert_eq!( - parse_datetime("1767225600").unwrap().to_rfc3339(), - DateTime::parse_from_rfc3339("2026-01-01T00:00:00Z") - .unwrap() - .to_rfc3339() - ); - // Millis precision is divided down to seconds. - assert_eq!( - parse_datetime("1767225600000"), - parse_datetime("1767225600") - ); - assert!(parse_datetime("garbage").is_none()); - } - - #[test] - fn number_field_does_not_short_circuit_on_missing_keys() { - let obj = serde_json::json!({"b": "2.5"}); - assert_eq!(number_field(&obj, &["missing", "b"]), Some(2.5)); - assert_eq!(number_field(&obj, &["missing", "other"]), None); - } - - fn provider_at(url: &str) -> CodeBuddyProvider { - let mut provider = CodeBuddyProvider::new(); - provider.api_url = url.to_string(); - provider - } - - #[tokio::test] - async fn transient_failures_are_retried_exactly_once() { - let mut server = mockito::Server::new_async().await; - let mock = server - .mock("POST", "/billing/meter/get-user-resource") - .with_status(500) - .expect(2) // one initial attempt + one retry - .create_async() - .await; - let provider = provider_at(&format!("{}/billing/meter/get-user-resource", server.url())); - let fail = provider.fetch_web("a=1").await.unwrap_err(); - assert!(fail.is_transient()); - mock.assert_async().await; - } - - #[tokio::test] - async fn auth_required_is_permanent_and_never_retried() { - let mut server = mockito::Server::new_async().await; - let mock = server - .mock("POST", "/billing/meter/get-user-resource") - .with_status(401) - .expect(1) // auth failures must not be retried - .create_async() - .await; - let provider = provider_at(&format!("{}/billing/meter/get-user-resource", server.url())); - let fail = provider.fetch_web("a=1").await.unwrap_err(); - assert!(!fail.is_transient()); - assert!(matches!(fail.into_error(), ProviderError::AuthRequired)); - mock.assert_async().await; - } - - #[tokio::test] - async fn waf_html_body_is_treated_as_transient() { - let mut server = mockito::Server::new_async().await; - let mock = server - .mock("POST", "/billing/meter/get-user-resource") - .with_status(200) - .with_header("content-type", "text/html") - .with_body("edgeone block") - .expect(2) - .create_async() - .await; - let provider = provider_at(&format!("{}/billing/meter/get-user-resource", server.url())); - let fail = provider.fetch_web("a=1").await.unwrap_err(); - assert!(fail.is_transient()); - mock.assert_async().await; - } - - /// End-to-end Auto behaviour with a real on-disk cache: - /// success persists typed totals; transient failure falls back to them; - /// auth failure surfaces even with a valid cache; foreign cache rejected. - #[tokio::test] - async fn auto_mode_cache_semantics_end_to_end() { - let dir = tempfile::tempdir().unwrap(); - let cache_path = dir.path().join("cb_credits.json"); - // SAFETY: single test using this env var; restored at the end. - unsafe { - std::env::set_var("CB_CREDITS_FILE", &cache_path); - } - - let payload_body = r#"{"code":0,"msg":"ok","data":{"Response":{"Data":{"Accounts":[{"CapacitySize":2000,"CapacityUsed":100,"CapacityRemain":1900,"ExpireTime":"2026-09-01T00:00:00Z"}]}}}}"#; - let ctx = |mode: SourceMode| FetchContext { - source_mode: mode, - manual_cookie_header: Some("session=abc; uid=42".to_string()), - ..Default::default() - }; - - // Phase A: web success persists typed totals + fingerprint, source=web. - { - let mut server = mockito::Server::new_async().await; - let mock = server - .mock("POST", "/billing/meter/get-user-resource") - .with_status(200) - .with_header("content-type", "application/json") - .with_body(payload_body) - .expect(1) - .create_async() - .await; - let provider = - provider_at(&format!("{}/billing/meter/get-user-resource", server.url())); - let result = provider.fetch_usage(&ctx(SourceMode::Auto)).await.unwrap(); - assert_eq!(result.source_label, "web"); - mock.assert_async().await; - - let file = std::fs::read_to_string(&cache_path).unwrap(); - let json: Value = serde_json::from_str(&file).unwrap(); - assert_eq!(json.get("total").and_then(|v| v.as_f64()), Some(2000.0)); - assert_eq!(json.get("used").and_then(|v| v.as_f64()), Some(100.0)); - assert_eq!(json.get("remaining").and_then(|v| v.as_f64()), Some(1900.0)); - assert_eq!( - json.get("accountHash").and_then(|v| v.as_str()), - Some(cookie_fingerprint("session=abc; uid=42").as_str()) - ); - assert!(json.get("resetsAt").and_then(|v| v.as_str()).is_some()); - } - - // Phase B: auth failure must surface — never masked by the valid cache. - { - let mut server = mockito::Server::new_async().await; - let mock = server - .mock("POST", "/billing/meter/get-user-resource") - .with_status(401) - .expect(1) - .create_async() - .await; - let provider = - provider_at(&format!("{}/billing/meter/get-user-resource", server.url())); - let err = provider - .fetch_usage(&ctx(SourceMode::Auto)) - .await - .unwrap_err(); - assert!(matches!(err, ProviderError::AuthRequired), "got {err}"); - mock.assert_async().await; - } - - // Phase C: transient failure in Auto falls back to the cache (source=cli). - { - let mut server = mockito::Server::new_async().await; - let mock = server - .mock("POST", "/billing/meter/get-user-resource") - .with_status(500) - .expect(2) - .create_async() - .await; - let provider = - provider_at(&format!("{}/billing/meter/get-user-resource", server.url())); - let result = provider.fetch_usage(&ctx(SourceMode::Auto)).await.unwrap(); - assert_eq!(result.source_label, "cli"); - assert_eq!( - result.usage.primary.reset_description.as_deref(), - Some("1,900 / 2,000 left") - ); - assert!( - (result.usage.primary.used_percent - 5.0).abs() < 0.01, - "used_percent={}", - result.usage.primary.used_percent - ); - mock.assert_async().await; - } - - // Phase D: a cache belonging to another account is rejected. - { - let mut json: Value = - serde_json::from_str(&std::fs::read_to_string(&cache_path).unwrap()).unwrap(); - json["accountHash"] = json!("deadbeefdeadbeef"); - std::fs::write(&cache_path, serde_json::to_string_pretty(&json).unwrap()).unwrap(); - - let mut server = mockito::Server::new_async().await; - server - .mock("POST", "/billing/meter/get-user-resource") - .with_status(500) - .create_async() - .await; - let provider = - provider_at(&format!("{}/billing/meter/get-user-resource", server.url())); - assert!(provider.fetch_usage(&ctx(SourceMode::Auto)).await.is_err()); - - // Web mode with a transient failure must not fall back at all. - let err = provider - .fetch_usage(&ctx(SourceMode::Web)) - .await - .unwrap_err(); - assert!(matches!(err, ProviderError::Other(_)), "got {err}"); - } - - // SAFETY: this test set CB_CREDITS_FILE at its start under the same - // single-test ownership; removing it restores the shared environment. - unsafe { - std::env::remove_var("CB_CREDITS_FILE"); - } - } -} +mod tests; diff --git a/rust/src/providers/codebuddy/tests.rs b/rust/src/providers/codebuddy/tests.rs new file mode 100644 index 0000000000..a19cd0d697 --- /dev/null +++ b/rust/src/providers/codebuddy/tests.rs @@ -0,0 +1,436 @@ +use super::*; + +fn api_payload() -> Value { + serde_json::json!({ + "code": 0, + "msg": "ok", + "data": { + "Response": { + "Data": { + "Accounts": [ + { + "CapacitySizePrecise": "2000", + "CapacityUsedPrecise": "100", + "CapacityRemainPrecise": "1900", + "ExpireTime": "2026-09-01T00:00:00Z" + }, + { + "CapacitySize": 1100, + "CapacityUsed": 11, + "CapacityRemain": 1089 + } + ] + } + } + } + }) +} + +#[test] +fn parses_get_user_resource_payload_into_typed_totals() { + let totals = totals_from_api_payload(&api_payload()).unwrap(); + assert_eq!(totals.total, 3100.0); + assert_eq!(totals.used, 111.0); + assert_eq!(totals.remaining, 2989.0); + assert!(totals.reset.is_some()); + + let snapshot = snapshot_from_totals(&totals); + assert!((snapshot.primary.used_percent - (111.0 / 3100.0 * 100.0)).abs() < 0.01); + assert_eq!( + snapshot.primary.reset_description.as_deref(), + Some("2,989 / 3,100 left") + ); + assert!(snapshot.primary.resets_at.is_some()); +} + +#[test] +fn auth_flavoured_payload_maps_to_auth_required() { + for msg in ["未登录", "登录已过期", "auth token expired"] { + let payload = serde_json::json!({ "code": 14001, "msg": msg }); + let err = totals_from_api_payload(&payload).unwrap_err(); + assert!( + matches!(err, ProviderError::AuthRequired), + "expected AuthRequired for msg={msg:?}, got {err}" + ); + } +} + +#[test] +fn empty_accounts_errors_with_hint() { + let payload = serde_json::json!({ + "code": 0, + "data": { "Response": { "Data": { "Accounts": [] } } } + }); + let err = totals_from_api_payload(&payload).unwrap_err(); + assert!(format!("{err}").contains("PackageCodes") || format!("{err}").contains("package")); +} + +#[test] +fn non_zero_code_errors() { + let payload = serde_json::json!({ "code": 14001, "msg": "quote exceeded" }); + let err = totals_from_api_payload(&payload).unwrap_err(); + assert!(matches!(err, ProviderError::Other(_)), "got {err}"); +} + +#[test] +fn payloads_with_non_finite_values_are_rejected() { + // serde_json cannot carry NaN/inf, but an absurd +/-1e308 pair can + // still overflow the sum to inf — that must be rejected, not cached. + let payload = serde_json::json!({ + "code": 0, + "data": { "Response": { "Data": { "Accounts": [ + { "CapacitySize": 1e308, "CapacityUsed": 1e308, "CapacityRemain": 1e308 }, + { "CapacitySize": 1e308, "CapacityUsed": 0, "CapacityRemain": 0 } + ] } } } + }); + assert!(totals_from_api_payload(&payload).is_err()); +} + +#[test] +fn formats_compact_credit_labels() { + assert_eq!(format_credits_short(1989.0, 3100.0), "1,989 / 3,100 left"); + assert_eq!(format_credits_short(12.5, 100.0), "12.5 / 100 left"); +} + +#[test] +fn normalizes_cookie_with_caret_escapes() { + assert_eq!( + normalize_cookie_header("Cookie: a=1^|2; b=3").as_deref(), + Some("a=1|2; b=3") + ); + assert_eq!(normalize_cookie_header(" "), None); + assert_eq!(normalize_cookie_header("Cookie:"), None); +} + +#[test] +fn cache_round_trip_preserves_typed_totals_exactly() { + let totals = CreditTotals { + total: 3100.5, + used: 111.25, + remaining: 2989.25, + reset: Some( + DateTime::parse_from_rfc3339("2026-09-01T08:30:00Z") + .unwrap() + .into(), + ), + }; + let json = cache_json_from_totals(&totals, Some("0123456789abcdef")); + let parsed = totals_from_cache_json(&json).unwrap(); + assert_eq!(parsed, totals); + // The cache carries the fingerprint verbatim. + assert_eq!( + json.get("accountHash").and_then(|v| v.as_str()), + Some("0123456789abcdef") + ); +} + +#[test] +fn cache_file_round_trip_via_disk() { + let dir = tempfile::tempdir().unwrap(); + let path = dir.path().join("cb_credits.json"); + let totals = CreditTotals { + total: 2000.0, + used: 100.25, + remaining: 1899.75, + reset: None, + }; + write_credits_cache(&path, &totals, Some("feedbeefcafe0001")).unwrap(); + + let value = read_credits_cache(&path).unwrap(); + let parsed = totals_from_cache_json(&value).unwrap(); + assert_eq!(parsed, totals); + + // Raw file exposes typed JSON numbers, not display text. + let raw = std::fs::read_to_string(&path).unwrap(); + assert!(raw.contains("\"total\": 2000.0")); + assert!(raw.contains("feedbeefcafe0001")); +} + +#[test] +fn hostile_cache_inputs_are_rejected() { + // Missing total. + let err = totals_from_cache_json(&serde_json::json!({"used": 1})).unwrap_err(); + assert!(format!("{err}").contains("missing total")); + // Negative total. + assert!(totals_from_cache_json(&serde_json::json!({"total": -5})).is_err()); + // Non-JSON is a Parse error at the read layer. + let dir = tempfile::tempdir().unwrap(); + let path = dir.path().join("bad.json"); + std::fs::write(&path, b"{not json").unwrap(); + assert!(matches!( + read_credits_cache(&path), + Err(ProviderError::Parse(_)) + )); + // Oversized cache is refused before parsing. + let big = dir.path().join("big.json"); + #[allow( + clippy::cast_possible_truncation, + reason = "MAX_CACHE_BYTES is 1 MiB; +1 stays far inside usize on any supported target" + )] + std::fs::write(&big, vec![b' '; (MAX_CACHE_BYTES + 1) as usize]).unwrap(); + assert!(read_credits_cache(&big).is_err()); +} + +#[test] +fn validated_package_codes_filters_and_caps() { + let long_code = "x".repeat(MAX_PACKAGE_CODE_LEN + 1); + let items = vec![ + json!("TCACA_code_001_ok"), + json!(" "), + json!(""), + json!(42), + json!("aB\u{0007}c"), + json!(long_code), + json!(" TCACA_code_002_trimmed "), + ]; + let codes = validated_package_codes(&items); + assert_eq!( + codes, + vec![ + "TCACA_code_001_ok".to_string(), + "TCACA_code_002_trimmed".to_string(), + ] + ); + + // Cap: more than MAX_PACKAGE_CODES valid entries are truncated. + let many: Vec = (0..MAX_PACKAGE_CODES + 10) + .map(|i| json!(format!("code_{i}"))) + .collect(); + assert_eq!(validated_package_codes(&many).len(), MAX_PACKAGE_CODES); +} + +#[test] +fn validate_api_url_rules() { + assert_eq!( + validate_api_url("https://www.codebuddy.cn/x").unwrap(), + "https://www.codebuddy.cn/x" + ); + assert!(validate_api_url("http://127.0.0.1:8080/x").is_ok()); + assert!(validate_api_url("http://localhost:8080/x").is_ok()); + assert!(validate_api_url("http://[::1]:8080/x").is_ok()); + assert!(validate_api_url("http://example.com/x").is_err()); + assert!(validate_api_url("ftp://example.com/x").is_err()); + assert!(validate_api_url("not a url").is_err()); + assert!(validate_api_url(" ").is_err()); +} + +#[test] +fn failure_classification_for_http_statuses() { + assert!(failure_for_status(StatusCode::OK).is_none()); + + for status in [StatusCode::UNAUTHORIZED, StatusCode::FORBIDDEN] { + let fail = failure_for_status(status).unwrap(); + assert!(!fail.is_transient(), "{status} must be permanent"); + assert!(matches!(fail.into_error(), ProviderError::AuthRequired)); + } + for status in [ + StatusCode::TOO_MANY_REQUESTS, + StatusCode::INTERNAL_SERVER_ERROR, + StatusCode::BAD_GATEWAY, + StatusCode::SERVICE_UNAVAILABLE, + ] { + assert!( + failure_for_status(status).unwrap().is_transient(), + "{status} must be transient" + ); + } + // Other failures (400 etc.) are permanent but not auth errors. + let fail = failure_for_status(StatusCode::BAD_REQUEST).unwrap(); + assert!(!fail.is_transient()); + assert!(matches!(fail.into_error(), ProviderError::Other(_))); +} + +#[test] +fn cache_fallback_requires_auto_mode_and_transient_failure() { + let transient = FetchFailure::Transient(ProviderError::Other("flake".into())); + let permanent_auth = FetchFailure::Permanent(ProviderError::AuthRequired); + let transient_ref = &transient; + let auth_ref = &permanent_auth; + + // Auth never masks behind a stale cache. + assert!(!cache_fallback_allowed(SourceMode::Auto, auth_ref)); + // Transient failures may fall back in Auto only. + assert!(cache_fallback_allowed(SourceMode::Auto, transient_ref)); + assert!(!cache_fallback_allowed(SourceMode::Web, transient_ref)); + assert!(!cache_fallback_allowed(SourceMode::Cli, transient_ref)); +} + +#[test] +fn cookie_fingerprint_is_stable_per_account_and_secret_free() { + let fp_a = cookie_fingerprint("session=aaa; uid=1"); + let fp_b = cookie_fingerprint("session=bbb; uid=2"); + assert_eq!(fp_a, cookie_fingerprint("session=aaa; uid=1")); + assert_ne!(fp_a, fp_b); + assert_eq!(fp_a.len(), 16); + assert!(fp_a.chars().all(|c| c.is_ascii_hexdigit())); +} + +#[test] +fn parse_datetime_accepts_rfc3339_naive_and_epochs() { + assert!(parse_datetime("2026-09-01T00:00:00Z").is_some()); + assert!(parse_datetime("2026-09-01 12:30:00").is_some()); + assert_eq!( + parse_datetime("1767225600").unwrap().to_rfc3339(), + DateTime::parse_from_rfc3339("2026-01-01T00:00:00Z") + .unwrap() + .to_rfc3339() + ); + // Millis precision is divided down to seconds. + assert_eq!( + parse_datetime("1767225600000"), + parse_datetime("1767225600") + ); + assert!(parse_datetime("garbage").is_none()); +} + +#[test] +fn number_field_does_not_short_circuit_on_missing_keys() { + let obj = serde_json::json!({"b": "2.5"}); + assert_eq!(number_field(&obj, &["missing", "b"]), Some(2.5)); + assert_eq!(number_field(&obj, &["missing", "other"]), None); +} + +const USAGE_PATH: &str = "/billing/meter/get-user-resource"; + +/// Mock the usage endpoint once; `body` is `(content-type, body)`. The +/// server guard is returned so the mock outlives the call. +async fn mock_usage( + status: usize, + body: Option<(&str, &str)>, + expect: Option, +) -> (mockito::ServerGuard, mockito::Mock, CodeBuddyProvider) { + let mut server = mockito::Server::new_async().await; + let mut mock = server.mock("POST", USAGE_PATH).with_status(status); + if let Some((content_type, body)) = body { + mock = mock + .with_header("content-type", content_type) + .with_body(body); + } + if let Some(hits) = expect { + mock = mock.expect(hits); + } + let mock = mock.create_async().await; + let mut provider = CodeBuddyProvider::new(); + provider.api_url = format!("{}{USAGE_PATH}", server.url()); + (server, mock, provider) +} + +#[tokio::test] +async fn transient_failures_are_retried_exactly_once() { + // One initial attempt + one retry. + let (_server, mock, provider) = mock_usage(500, None, Some(2)).await; + let fail = provider.fetch_web("a=1").await.unwrap_err(); + assert!(fail.is_transient()); + mock.assert_async().await; +} + +#[tokio::test] +async fn auth_required_is_permanent_and_never_retried() { + // Auth failures must not be retried. + let (_server, mock, provider) = mock_usage(401, None, Some(1)).await; + let fail = provider.fetch_web("a=1").await.unwrap_err(); + assert!(!fail.is_transient()); + assert!(matches!(fail.into_error(), ProviderError::AuthRequired)); + mock.assert_async().await; +} + +#[tokio::test] +async fn waf_html_body_is_treated_as_transient() { + let html = ("text/html", "edgeone block"); + let (_server, mock, provider) = mock_usage(200, Some(html), Some(2)).await; + let fail = provider.fetch_web("a=1").await.unwrap_err(); + assert!(fail.is_transient()); + mock.assert_async().await; +} + +/// End-to-end Auto behaviour with a real on-disk cache: +/// success persists typed totals; transient failure falls back to them; +/// auth failure surfaces even with a valid cache; foreign cache rejected. +#[tokio::test] +async fn auto_mode_cache_semantics_end_to_end() { + let dir = tempfile::tempdir().unwrap(); + let cache_path = dir.path().join("cb_credits.json"); + // SAFETY: single test using this env var; restored at the end. + unsafe { + std::env::set_var("CB_CREDITS_FILE", &cache_path); + } + + let payload_body = r#"{"code":0,"msg":"ok","data":{"Response":{"Data":{"Accounts":[{"CapacitySize":2000,"CapacityUsed":100,"CapacityRemain":1900,"ExpireTime":"2026-09-01T00:00:00Z"}]}}}}"#; + let ctx = |mode: SourceMode| FetchContext { + source_mode: mode, + manual_cookie_header: Some("session=abc; uid=42".to_string()), + ..Default::default() + }; + + // Phase A: web success persists typed totals + fingerprint, source=web. + { + let json_body = ("application/json", payload_body); + let (_server, mock, provider) = mock_usage(200, Some(json_body), Some(1)).await; + let result = provider.fetch_usage(&ctx(SourceMode::Auto)).await.unwrap(); + assert_eq!(result.source_label, "web"); + mock.assert_async().await; + + let file = std::fs::read_to_string(&cache_path).unwrap(); + let json: Value = serde_json::from_str(&file).unwrap(); + assert_eq!(json.get("total").and_then(|v| v.as_f64()), Some(2000.0)); + assert_eq!(json.get("used").and_then(|v| v.as_f64()), Some(100.0)); + assert_eq!(json.get("remaining").and_then(|v| v.as_f64()), Some(1900.0)); + assert_eq!( + json.get("accountHash").and_then(|v| v.as_str()), + Some(cookie_fingerprint("session=abc; uid=42").as_str()) + ); + assert!(json.get("resetsAt").and_then(|v| v.as_str()).is_some()); + } + + // Phase B: auth failure must surface — never masked by the valid cache. + { + let (_server, mock, provider) = mock_usage(401, None, Some(1)).await; + let err = provider + .fetch_usage(&ctx(SourceMode::Auto)) + .await + .unwrap_err(); + assert!(matches!(err, ProviderError::AuthRequired), "got {err}"); + mock.assert_async().await; + } + + // Phase C: transient failure in Auto falls back to the cache (source=cli). + { + let (_server, mock, provider) = mock_usage(500, None, Some(2)).await; + let result = provider.fetch_usage(&ctx(SourceMode::Auto)).await.unwrap(); + assert_eq!(result.source_label, "cli"); + assert_eq!( + result.usage.primary.reset_description.as_deref(), + Some("1,900 / 2,000 left") + ); + assert!( + (result.usage.primary.used_percent - 5.0).abs() < 0.01, + "used_percent={}", + result.usage.primary.used_percent + ); + mock.assert_async().await; + } + + // Phase D: a cache belonging to another account is rejected. + { + let mut json: Value = + serde_json::from_str(&std::fs::read_to_string(&cache_path).unwrap()).unwrap(); + json["accountHash"] = json!("deadbeefdeadbeef"); + std::fs::write(&cache_path, serde_json::to_string_pretty(&json).unwrap()).unwrap(); + + let (_server, _mock, provider) = mock_usage(500, None, None).await; + assert!(provider.fetch_usage(&ctx(SourceMode::Auto)).await.is_err()); + + // Web mode with a transient failure must not fall back at all. + let err = provider + .fetch_usage(&ctx(SourceMode::Web)) + .await + .unwrap_err(); + assert!(matches!(err, ProviderError::Other(_)), "got {err}"); + } + + // SAFETY: this test set CB_CREDITS_FILE at its start under the same + // single-test ownership; removing it restores the shared environment. + unsafe { + std::env::remove_var("CB_CREDITS_FILE"); + } +} diff --git a/rust/src/providers/cursor/api.rs b/rust/src/providers/cursor/api.rs index 85cc73f379..94f8c1d03c 100755 --- a/rust/src/providers/cursor/api.rs +++ b/rust/src/providers/cursor/api.rs @@ -4,12 +4,10 @@ use super::team_budget::CursorMemberBudget; use crate::core::{CostSnapshot, NamedRateWindow, ProviderError, RateWindow}; -use crate::providers::browser_cookie_header; use chrono::{DateTime, Utc}; use serde::Deserialize; const BASE_URL: &str = "https://cursor.com"; -const COOKIE_DOMAINS: [&str; 2] = ["cursor.com", "cursor.sh"]; #[derive(Debug)] pub struct CursorUsageResult { @@ -56,13 +54,6 @@ impl CursorApi { &self.base_url } - /// Fetch usage information from Cursor API - pub async fn fetch_usage(&self) -> Result { - // Try to get cookies from browser - let cookie_header = self.get_cookie_header()?; - self.fetch_usage_with_cookie_header(&cookie_header).await - } - /// Fetch usage information with an already resolved Cookie header. pub async fn fetch_usage_with_cookie_header( &self, @@ -80,16 +71,11 @@ impl CursorApi { let team_budget = self .resolve_team_budget(&usage_summary, user_info.as_ref(), cookie_header) .await; - let mut result = - self.build_result_with_team_budget(usage_summary, user_info, team_budget)?; + let mut result = self.build_result_with_team_budget(usage_summary, user_info, team_budget); result.grok_bot = sand_result.ok().flatten(); Ok(result) } - fn get_cookie_header(&self) -> Result { - browser_cookie_header(&COOKIE_DOMAINS) - } - async fn fetch_usage_summary( &self, cookie_header: &str, @@ -185,7 +171,7 @@ impl CursorApi { summary: UsageSummary, user_info: Option, team_budget: Option, - ) -> Result { + ) -> CursorUsageResult { let billing_end = summary .billing_cycle_end .as_ref() @@ -194,20 +180,7 @@ impl CursorApi { let (percent_used, secondary, model_specific, cost_snapshot) = if let Some(team_budget) = team_budget { let percent = clamp_percent(team_budget.used_usd / team_budget.limit_usd * 100.0); - let cost = Self::on_demand_cost( - summary - .individual_usage - .as_ref() - .and_then(|individual| individual.on_demand.as_ref()), - billing_end, - ) - .or_else(|| { - summary - .team_usage - .as_ref() - .and_then(|team| Self::on_demand_cost(team.on_demand.as_ref(), billing_end)) - }) - .or_else(|| { + let cost = Self::summary_on_demand_cost(&summary, billing_end).or_else(|| { Some(Self::plan_cost( team_budget.used_usd, team_budget.limit_usd, @@ -242,13 +215,8 @@ impl CursorApi { RateWindow::with_details(clamp_percent(v), None, billing_end, None) }); - let cost = Self::on_demand_cost(individual.on_demand.as_ref(), billing_end) - .or_else(|| { - summary.team_usage.as_ref().and_then(|team| { - Self::on_demand_cost(team.on_demand.as_ref(), billing_end) - }) - }) - .unwrap_or_else(|| { + let cost = + Self::summary_on_demand_cost(&summary, billing_end).unwrap_or_else(|| { // Plan-included spend (cents → USD) when on-demand is off. Self::plan_cost( used_cents / 100.0, @@ -293,7 +261,7 @@ impl CursorApi { let email = user_info.as_ref().and_then(|u| u.email.clone()); - Ok(CursorUsageResult { + CursorUsageResult { primary, secondary, model_specific, @@ -301,7 +269,7 @@ impl CursorApi { email, plan_type, grok_bot: None, - }) + } } fn plan_cost( @@ -320,6 +288,23 @@ impl CursorApi { cost } + /// Individual on-demand spend, else the team's. + fn summary_on_demand_cost( + summary: &UsageSummary, + billing_end: Option>, + ) -> Option { + let individual = summary + .individual_usage + .as_ref() + .and_then(|individual| individual.on_demand.as_ref()); + Self::on_demand_cost(individual, billing_end).or_else(|| { + summary + .team_usage + .as_ref() + .and_then(|team| Self::on_demand_cost(team.on_demand.as_ref(), billing_end)) + }) + } + fn on_demand_cost( on_demand: Option<&OnDemandUsage>, billing_end: Option>, @@ -329,15 +314,7 @@ impl CursorApi { return None; } - let used_cents = usage.used.unwrap_or(0) as f64; - let limit_cents = usage - .limit - .or_else(|| { - usage - .remaining - .map(|remaining| remaining + usage.used.unwrap_or(0)) - }) - .unwrap_or(0) as f64; + let (used_cents, limit_cents) = usage.used_and_limit_cents(); if used_cents <= 0.0 && limit_cents <= 0.0 { return None; @@ -356,15 +333,7 @@ impl CursorApi { } fn usage_percent(usage: &OnDemandUsage) -> Option { - let used = usage.used.unwrap_or(0) as f64; - let limit = usage - .limit - .or_else(|| { - usage - .remaining - .map(|remaining| remaining + usage.used.unwrap_or(0)) - }) - .unwrap_or(0) as f64; + let (used, limit) = usage.used_and_limit_cents(); (limit > 0.0).then_some(clamp_percent(used / limit * 100.0)) } } @@ -458,6 +427,18 @@ struct OnDemandUsage { remaining: Option, } +impl OnDemandUsage { + /// (used, limit) in cents; a missing limit falls back to remaining + used. + fn used_and_limit_cents(&self) -> (f64, f64) { + let used = self.used.unwrap_or(0); + let limit = self + .limit + .or_else(|| self.remaining.map(|remaining| remaining + used)) + .unwrap_or(0); + (used as f64, limit as f64) + } +} + #[derive(Debug, Deserialize)] #[serde(rename_all = "camelCase")] struct TeamUsage { @@ -548,17 +529,9 @@ impl UserInfo { // --- Helper functions --- fn parse_iso_date(s: &str) -> Option> { - // Try with fractional seconds - if let Ok(dt) = DateTime::parse_from_rfc3339(s) { - return Some(dt.with_timezone(&Utc)); - } - - // Try without fractional seconds - if let Ok(dt) = chrono::DateTime::parse_from_str(s, "%Y-%m-%dT%H:%M:%SZ") { - return Some(dt.with_timezone(&Utc)); - } - - None + DateTime::parse_from_rfc3339(s) + .ok() + .map(|dt| dt.with_timezone(&Utc)) } fn capitalize(s: &str) -> String { @@ -570,387 +543,4 @@ fn capitalize(s: &str) -> String { } #[cfg(test)] -mod tests { - use super::*; - - fn api() -> CursorApi { - CursorApi::new() - } - - fn parse_summary(json: &str) -> UsageSummary { - serde_json::from_str(json).expect("fixture should parse") - } - - #[test] - fn sand_usage_maps_to_weekly_extra_window() { - let status = SandUsageStatus { - current_period_start: Some("2026-08-18T00:00:00Z".into()), - next_reset_timestamp_utc: Some("2026-08-25T00:00:00Z".into()), - usage_percent: Some(37.5), - has_non_zero_included_limit: Some(true), - included_limit_zero: None, - sand_trial_expires_at: None, - }; - let row = status - .to_window("2026-08-20T00:00:00Z".parse().unwrap()) - .expect("grok bot window"); - assert_eq!(row.id, "cursor-grok-bot"); - assert_eq!(row.title, "Grok Bot"); - assert!((row.window.used_percent - 37.5).abs() < 0.001); - assert_eq!(row.window.window_minutes, Some(10080)); - } - - #[test] - fn sand_usage_hides_accounts_without_included_allowance() { - let status = SandUsageStatus { - current_period_start: None, - next_reset_timestamp_utc: None, - usage_percent: Some(0.0), - has_non_zero_included_limit: Some(false), - included_limit_zero: None, - sand_trial_expires_at: None, - }; - assert!( - status - .to_window("2026-08-20T00:00:00Z".parse().unwrap()) - .is_none() - ); - } - - #[test] - fn sand_usage_maps_paid_allowance_from_explicit_zero_flag() { - let status = SandUsageStatus { - current_period_start: Some("2026-08-18T00:00:00Z".into()), - next_reset_timestamp_utc: Some("2026-08-25T00:00:00Z".into()), - usage_percent: Some(37.5), - has_non_zero_included_limit: Some(false), - included_limit_zero: Some(false), - sand_trial_expires_at: None, - }; - let row = status - .to_window("2026-08-20T00:00:00Z".parse().unwrap()) - .expect("paid Grok Bot window"); - assert_eq!(row.window.window_minutes, Some(10080)); - assert_eq!( - row.window.resets_at, - Some("2026-08-25T00:00:00Z".parse().unwrap()) - ); - } - - #[test] - fn sand_usage_maps_an_active_trial_without_a_recurring_reset() { - let status = SandUsageStatus { - current_period_start: Some("2026-08-18T00:00:00Z".into()), - next_reset_timestamp_utc: Some("2026-08-25T00:00:00Z".into()), - usage_percent: Some(12.5), - has_non_zero_included_limit: Some(false), - included_limit_zero: Some(true), - sand_trial_expires_at: Some("2026-08-28T00:00:00Z".into()), - }; - let row = status - .to_window("2026-08-20T00:00:00Z".parse().unwrap()) - .expect("active trial window"); - assert_eq!(row.window.resets_at, None); - assert_eq!(row.window.window_minutes, None); - assert!((row.window.used_percent - 12.5).abs() < 0.001); - } - - #[test] - fn sand_usage_hides_expired_trial() { - let status = SandUsageStatus { - current_period_start: None, - next_reset_timestamp_utc: None, - usage_percent: Some(12.5), - has_non_zero_included_limit: Some(false), - included_limit_zero: Some(true), - sand_trial_expires_at: Some("2026-08-19T00:00:00Z".into()), - }; - assert!( - status - .to_window("2026-08-20T00:00:00Z".parse().unwrap()) - .is_none() - ); - } - - #[test] - fn test_cursor_build_result_with_lanes() { - let json = r#"{ - "billingCycleStart": "2026-03-01T00:00:00Z", - "billingCycleEnd": "2026-04-01T00:00:00Z", - "membershipType": "pro", - "individualUsage": { - "plan": { - "used": 1500, - "limit": 5000, - "totalPercentUsed": 30.0, - "autoPercentUsed": 20.0, - "apiPercentUsed": 10.0 - } - } - }"#; - - let summary = parse_summary(json); - let result = api() - .build_result_with_team_budget(summary, None, None) - .unwrap(); - - assert!((result.primary.used_percent - 30.0).abs() < 0.01); - - let sec = result.secondary.expect("secondary should be present"); - assert!((sec.used_percent - 20.0).abs() < 0.01); - assert!(sec.resets_at.is_some()); - - let ms = result - .model_specific - .expect("model_specific should be present"); - assert!((ms.used_percent - 10.0).abs() < 0.01); - assert!(ms.resets_at.is_some()); - - assert!(result.cost.is_some()); - assert_eq!(result.plan_type.as_deref(), Some("Cursor Pro")); - } - - #[test] - fn clamps_plan_usage_percent_at_100_when_over_limit() { - // Upstream #2255: included usage past limit must not paint >100%. - let json = r#"{ - "membershipType": "pro", - "individualUsage": { - "plan": { - "used": 6000, - "limit": 5000, - "totalPercentUsed": 120.0, - "autoPercentUsed": 110.0, - "apiPercentUsed": 105.0 - } - } - }"#; - let summary = parse_summary(json); - let result = api() - .build_result_with_team_budget(summary, None, None) - .unwrap(); - assert!((result.primary.used_percent - 100.0).abs() < 0.01); - assert!((result.secondary.unwrap().used_percent - 100.0).abs() < 0.01); - assert!((result.model_specific.unwrap().used_percent - 100.0).abs() < 0.01); - } - - #[test] - fn test_cursor_build_result_prefers_api_percent_fields() { - let json = r#"{ - "membershipType": "pro", - "autoModelSelectedDisplayMessage": "You've used 13% of your included total usage", - "individualUsage": { - "plan": { - "used": 2000, - "limit": 2000, - "breakdown": { - "included": 2000, - "bonus": 580, - "total": 2580 - }, - "autoPercentUsed": 17.2, - "apiPercentUsed": 0, - "totalPercentUsed": 13.230769230769232 - } - } - }"#; - - let summary = parse_summary(json); - let result = api() - .build_result_with_team_budget(summary, None, None) - .unwrap(); - - assert!((result.primary.used_percent - 13.230769230769232).abs() < 0.01); - assert!((result.secondary.unwrap().used_percent - 17.2).abs() < 0.01); - assert!((result.model_specific.unwrap().used_percent - 0.0).abs() < 0.01); - - let cost = result - .cost - .expect("plan usage should still produce cost snapshot"); - assert!((cost.used - 20.0).abs() < 0.01); - assert_eq!(cost.limit, Some(20.0)); - assert_eq!(result.plan_type.as_deref(), Some("Cursor Pro")); - } - - #[test] - fn test_cursor_build_result_cents_only() { - let json = r#"{ - "billingCycleEnd": "2026-04-01T00:00:00Z", - "membershipType": "pro", - "individualUsage": { - "plan": { - "used": 2500, - "limit": 5000 - } - } - }"#; - - let summary = parse_summary(json); - let result = api() - .build_result_with_team_budget(summary, None, None) - .unwrap(); - - assert!((result.primary.used_percent - 50.0).abs() < 0.01); - assert!(result.secondary.is_none(), "no autoPercentUsed in payload"); - assert!( - result.model_specific.is_none(), - "no apiPercentUsed in payload" - ); - assert!(result.cost.is_some()); - } - - #[test] - fn test_cursor_build_result_missing_plan() { - let json = r#"{ - "membershipType": "hobby", - "individualUsage": {} - }"#; - - let summary = parse_summary(json); - let result = api() - .build_result_with_team_budget(summary, None, None) - .unwrap(); - - assert!((result.primary.used_percent).abs() < 0.01); - assert!(result.secondary.is_none()); - assert!(result.model_specific.is_none()); - assert!(result.cost.is_none()); - } - - #[test] - fn test_cursor_on_demand_as_cost() { - let json = r#"{ - "billingCycleEnd": "2026-04-01T00:00:00Z", - "membershipType": "pro", - "individualUsage": { - "plan": { - "used": 800, - "limit": 5000, - "totalPercentUsed": 16.0 - }, - "onDemand": { - "enabled": true, - "used": 350, - "limit": 1000 - } - } - }"#; - - let summary = parse_summary(json); - let result = api() - .build_result_with_team_budget(summary, None, None) - .unwrap(); - - assert!((result.primary.used_percent - 16.0).abs() < 0.01); - let cost = result.cost.expect("cost should exist from on-demand usage"); - assert!((cost.used - 3.5).abs() < 0.01); - assert_eq!(cost.limit, Some(10.0)); - assert_eq!(cost.period, "On-demand (billing cycle)"); - } - - #[test] - fn plan_cost_period_uses_billing_cycle_start() { - let json = r#"{ - "billingCycleStart": "2026-03-01T00:00:00Z", - "billingCycleEnd": "2026-04-01T00:00:00Z", - "membershipType": "pro", - "individualUsage": { - "plan": { - "used": 2500, - "limit": 5000 - } - } - }"#; - let summary = parse_summary(json); - let result = api() - .build_result_with_team_budget(summary, None, None) - .unwrap(); - let cost = result.cost.expect("plan cost"); - assert!((cost.used - 25.0).abs() < 0.01); - assert_eq!(cost.limit, Some(50.0)); - assert_eq!( - cost.period, - "Cursor and Third Party (since 2026-03-01T00:00:00Z)" - ); - } - - #[test] - fn test_cursor_individual_overall_fallback() { - let summary = - parse_summary(r#"{"individualUsage":{"overall":{"used":2500,"limit":10000}}}"#); - let result = api() - .build_result_with_team_budget(summary, None, None) - .unwrap(); - assert!((result.primary.used_percent - 25.0).abs() < 0.01); - assert_eq!(result.cost.unwrap().limit, Some(100.0)); - } - - #[test] - fn test_cursor_team_pooled_fallback() { - let summary = parse_summary(r#"{"teamUsage":{"pooled":{"used":5000,"limit":10000}}}"#); - let result = api() - .build_result_with_team_budget(summary, None, None) - .unwrap(); - assert!((result.primary.used_percent - 50.0).abs() < 0.01); - assert_eq!(result.cost.unwrap().used, 50.0); - } - - #[test] - fn member_lookup_requires_nonempty_authenticated_email() { - for email in [None, Some(String::new()), Some(" ".to_string())] { - let user = UserInfo { - email, - email_verified: None, - name: None, - sub: None, - created_at: None, - updated_at: None, - picture: None, - }; - assert!(user.verified_email().is_none()); - } - } - - #[test] - fn verified_team_budget_replaces_summary_plan_and_keeps_zero_summary_fallback() { - let summary = parse_summary( - r#"{ - "billingCycleStart":"2026-09-01T00:00:00Z", - "billingCycleEnd":"2026-10-01T00:00:00Z", - "membershipType":"enterprise", - "individualUsage":{"plan":{"used":0,"limit":2000,"totalPercentUsed":0}} - }"#, - ); - let result = api() - .build_result_with_team_budget( - summary, - None, - Some(CursorMemberBudget { - used_usd: 13.12, - limit_usd: 150.0, - }), - ) - .unwrap(); - assert!((result.primary.used_percent - 8.7466666667).abs() < 0.00001); - let cost = result.cost.expect("verified member budget cost"); - assert!((cost.used - 13.12).abs() < 0.00001); - assert_eq!(cost.limit, Some(150.0)); - - let fallback_summary = parse_summary( - r#"{ - "billingCycleStart":"2026-09-01T00:00:00Z", - "billingCycleEnd":"2026-10-01T00:00:00Z", - "membershipType":"enterprise", - "individualUsage":{"plan":{"used":0,"limit":2000,"totalPercentUsed":0}} - }"#, - ); - let fallback = api() - .build_result_with_team_budget(fallback_summary, None, None) - .unwrap(); - assert_eq!(fallback.primary.used_percent, 0.0); - assert_eq!( - fallback.cost.expect("summary fallback cost").limit, - Some(20.0) - ); - } -} +mod tests; diff --git a/rust/src/providers/cursor/api/tests.rs b/rust/src/providers/cursor/api/tests.rs new file mode 100644 index 0000000000..c9fd34417a --- /dev/null +++ b/rust/src/providers/cursor/api/tests.rs @@ -0,0 +1,359 @@ +use super::*; + +fn api() -> CursorApi { + CursorApi::new() +} + +fn parse_summary(json: &str) -> UsageSummary { + serde_json::from_str(json).expect("fixture should parse") +} + +#[test] +fn sand_usage_maps_to_weekly_extra_window() { + let status = SandUsageStatus { + current_period_start: Some("2026-08-18T00:00:00Z".into()), + next_reset_timestamp_utc: Some("2026-08-25T00:00:00Z".into()), + usage_percent: Some(37.5), + has_non_zero_included_limit: Some(true), + included_limit_zero: None, + sand_trial_expires_at: None, + }; + let row = status + .to_window("2026-08-20T00:00:00Z".parse().unwrap()) + .expect("grok bot window"); + assert_eq!(row.id, "cursor-grok-bot"); + assert_eq!(row.title, "Grok Bot"); + assert!((row.window.used_percent - 37.5).abs() < 0.001); + assert_eq!(row.window.window_minutes, Some(10080)); +} + +#[test] +fn sand_usage_hides_accounts_without_included_allowance() { + let status = SandUsageStatus { + current_period_start: None, + next_reset_timestamp_utc: None, + usage_percent: Some(0.0), + has_non_zero_included_limit: Some(false), + included_limit_zero: None, + sand_trial_expires_at: None, + }; + assert!( + status + .to_window("2026-08-20T00:00:00Z".parse().unwrap()) + .is_none() + ); +} + +#[test] +fn sand_usage_maps_paid_allowance_from_explicit_zero_flag() { + let status = SandUsageStatus { + current_period_start: Some("2026-08-18T00:00:00Z".into()), + next_reset_timestamp_utc: Some("2026-08-25T00:00:00Z".into()), + usage_percent: Some(37.5), + has_non_zero_included_limit: Some(false), + included_limit_zero: Some(false), + sand_trial_expires_at: None, + }; + let row = status + .to_window("2026-08-20T00:00:00Z".parse().unwrap()) + .expect("paid Grok Bot window"); + assert_eq!(row.window.window_minutes, Some(10080)); + assert_eq!( + row.window.resets_at, + Some("2026-08-25T00:00:00Z".parse().unwrap()) + ); +} + +#[test] +fn sand_usage_maps_an_active_trial_without_a_recurring_reset() { + let status = SandUsageStatus { + current_period_start: Some("2026-08-18T00:00:00Z".into()), + next_reset_timestamp_utc: Some("2026-08-25T00:00:00Z".into()), + usage_percent: Some(12.5), + has_non_zero_included_limit: Some(false), + included_limit_zero: Some(true), + sand_trial_expires_at: Some("2026-08-28T00:00:00Z".into()), + }; + let row = status + .to_window("2026-08-20T00:00:00Z".parse().unwrap()) + .expect("active trial window"); + assert_eq!(row.window.resets_at, None); + assert_eq!(row.window.window_minutes, None); + assert!((row.window.used_percent - 12.5).abs() < 0.001); +} + +#[test] +fn sand_usage_hides_expired_trial() { + let status = SandUsageStatus { + current_period_start: None, + next_reset_timestamp_utc: None, + usage_percent: Some(12.5), + has_non_zero_included_limit: Some(false), + included_limit_zero: Some(true), + sand_trial_expires_at: Some("2026-08-19T00:00:00Z".into()), + }; + assert!( + status + .to_window("2026-08-20T00:00:00Z".parse().unwrap()) + .is_none() + ); +} + +#[test] +fn test_cursor_build_result_with_lanes() { + let json = r#"{ + "billingCycleStart": "2026-03-01T00:00:00Z", + "billingCycleEnd": "2026-04-01T00:00:00Z", + "membershipType": "pro", + "individualUsage": { + "plan": { + "used": 1500, + "limit": 5000, + "totalPercentUsed": 30.0, + "autoPercentUsed": 20.0, + "apiPercentUsed": 10.0 + } + } + }"#; + + let summary = parse_summary(json); + let result = api().build_result_with_team_budget(summary, None, None); + + assert!((result.primary.used_percent - 30.0).abs() < 0.01); + + let sec = result.secondary.expect("secondary should be present"); + assert!((sec.used_percent - 20.0).abs() < 0.01); + assert!(sec.resets_at.is_some()); + + let ms = result + .model_specific + .expect("model_specific should be present"); + assert!((ms.used_percent - 10.0).abs() < 0.01); + assert!(ms.resets_at.is_some()); + + assert!(result.cost.is_some()); + assert_eq!(result.plan_type.as_deref(), Some("Cursor Pro")); +} + +#[test] +fn clamps_plan_usage_percent_at_100_when_over_limit() { + // Upstream #2255: included usage past limit must not paint >100%. + let json = r#"{ + "membershipType": "pro", + "individualUsage": { + "plan": { + "used": 6000, + "limit": 5000, + "totalPercentUsed": 120.0, + "autoPercentUsed": 110.0, + "apiPercentUsed": 105.0 + } + } + }"#; + let summary = parse_summary(json); + let result = api().build_result_with_team_budget(summary, None, None); + assert!((result.primary.used_percent - 100.0).abs() < 0.01); + assert!((result.secondary.unwrap().used_percent - 100.0).abs() < 0.01); + assert!((result.model_specific.unwrap().used_percent - 100.0).abs() < 0.01); +} + +#[test] +fn test_cursor_build_result_prefers_api_percent_fields() { + let json = r#"{ + "membershipType": "pro", + "autoModelSelectedDisplayMessage": "You've used 13% of your included total usage", + "individualUsage": { + "plan": { + "used": 2000, + "limit": 2000, + "breakdown": { + "included": 2000, + "bonus": 580, + "total": 2580 + }, + "autoPercentUsed": 17.2, + "apiPercentUsed": 0, + "totalPercentUsed": 13.230769230769232 + } + } + }"#; + + let summary = parse_summary(json); + let result = api().build_result_with_team_budget(summary, None, None); + + assert!((result.primary.used_percent - 13.230769230769232).abs() < 0.01); + assert!((result.secondary.unwrap().used_percent - 17.2).abs() < 0.01); + assert!((result.model_specific.unwrap().used_percent - 0.0).abs() < 0.01); + + let cost = result + .cost + .expect("plan usage should still produce cost snapshot"); + assert!((cost.used - 20.0).abs() < 0.01); + assert_eq!(cost.limit, Some(20.0)); + assert_eq!(result.plan_type.as_deref(), Some("Cursor Pro")); +} + +#[test] +fn test_cursor_build_result_cents_only() { + let json = r#"{ + "billingCycleEnd": "2026-04-01T00:00:00Z", + "membershipType": "pro", + "individualUsage": { + "plan": { + "used": 2500, + "limit": 5000 + } + } + }"#; + + let summary = parse_summary(json); + let result = api().build_result_with_team_budget(summary, None, None); + + assert!((result.primary.used_percent - 50.0).abs() < 0.01); + assert!(result.secondary.is_none(), "no autoPercentUsed in payload"); + assert!( + result.model_specific.is_none(), + "no apiPercentUsed in payload" + ); + assert!(result.cost.is_some()); +} + +#[test] +fn test_cursor_build_result_missing_plan() { + let json = r#"{ + "membershipType": "hobby", + "individualUsage": {} + }"#; + + let summary = parse_summary(json); + let result = api().build_result_with_team_budget(summary, None, None); + + assert!((result.primary.used_percent).abs() < 0.01); + assert!(result.secondary.is_none()); + assert!(result.model_specific.is_none()); + assert!(result.cost.is_none()); +} + +#[test] +fn test_cursor_on_demand_as_cost() { + let json = r#"{ + "billingCycleEnd": "2026-04-01T00:00:00Z", + "membershipType": "pro", + "individualUsage": { + "plan": { + "used": 800, + "limit": 5000, + "totalPercentUsed": 16.0 + }, + "onDemand": { + "enabled": true, + "used": 350, + "limit": 1000 + } + } + }"#; + + let summary = parse_summary(json); + let result = api().build_result_with_team_budget(summary, None, None); + + assert!((result.primary.used_percent - 16.0).abs() < 0.01); + let cost = result.cost.expect("cost should exist from on-demand usage"); + assert!((cost.used - 3.5).abs() < 0.01); + assert_eq!(cost.limit, Some(10.0)); + assert_eq!(cost.period, "On-demand (billing cycle)"); +} + +#[test] +fn plan_cost_period_uses_billing_cycle_start() { + let json = r#"{ + "billingCycleStart": "2026-03-01T00:00:00Z", + "billingCycleEnd": "2026-04-01T00:00:00Z", + "membershipType": "pro", + "individualUsage": { + "plan": { + "used": 2500, + "limit": 5000 + } + } + }"#; + let summary = parse_summary(json); + let result = api().build_result_with_team_budget(summary, None, None); + let cost = result.cost.expect("plan cost"); + assert!((cost.used - 25.0).abs() < 0.01); + assert_eq!(cost.limit, Some(50.0)); + assert_eq!( + cost.period, + "Cursor and Third Party (since 2026-03-01T00:00:00Z)" + ); +} + +#[test] +fn test_cursor_individual_overall_fallback() { + let summary = parse_summary(r#"{"individualUsage":{"overall":{"used":2500,"limit":10000}}}"#); + let result = api().build_result_with_team_budget(summary, None, None); + assert!((result.primary.used_percent - 25.0).abs() < 0.01); + assert_eq!(result.cost.unwrap().limit, Some(100.0)); +} + +#[test] +fn test_cursor_team_pooled_fallback() { + let summary = parse_summary(r#"{"teamUsage":{"pooled":{"used":5000,"limit":10000}}}"#); + let result = api().build_result_with_team_budget(summary, None, None); + assert!((result.primary.used_percent - 50.0).abs() < 0.01); + assert_eq!(result.cost.unwrap().used, 50.0); +} + +#[test] +fn member_lookup_requires_nonempty_authenticated_email() { + for email in [None, Some(String::new()), Some(" ".to_string())] { + let user = UserInfo { + email, + email_verified: None, + name: None, + sub: None, + created_at: None, + updated_at: None, + picture: None, + }; + assert!(user.verified_email().is_none()); + } +} + +#[test] +fn verified_team_budget_replaces_summary_plan_and_keeps_zero_summary_fallback() { + let summary = parse_summary( + r#"{ + "billingCycleStart":"2026-09-01T00:00:00Z", + "billingCycleEnd":"2026-10-01T00:00:00Z", + "membershipType":"enterprise", + "individualUsage":{"plan":{"used":0,"limit":2000,"totalPercentUsed":0}} + }"#, + ); + let result = api().build_result_with_team_budget( + summary, + None, + Some(CursorMemberBudget { + used_usd: 13.12, + limit_usd: 150.0, + }), + ); + assert!((result.primary.used_percent - 8.7466666667).abs() < 0.00001); + let cost = result.cost.expect("verified member budget cost"); + assert!((cost.used - 13.12).abs() < 0.00001); + assert_eq!(cost.limit, Some(150.0)); + + let fallback_summary = parse_summary( + r#"{ + "billingCycleStart":"2026-09-01T00:00:00Z", + "billingCycleEnd":"2026-10-01T00:00:00Z", + "membershipType":"enterprise", + "individualUsage":{"plan":{"used":0,"limit":2000,"totalPercentUsed":0}} + }"#, + ); + let fallback = api().build_result_with_team_budget(fallback_summary, None, None); + assert_eq!(fallback.primary.used_percent, 0.0); + assert_eq!( + fallback.cost.expect("summary fallback cost").limit, + Some(20.0) + ); +} diff --git a/rust/src/providers/cursor/app_auth.rs b/rust/src/providers/cursor/app_auth.rs index 5e195cad9b..22278774a6 100644 --- a/rust/src/providers/cursor/app_auth.rs +++ b/rust/src/providers/cursor/app_auth.rs @@ -35,7 +35,9 @@ pub fn load_app_auth_access_token() -> Option { // WAL database can retain WAL mode in its header after the // sidecars disappear, and immutable mode reads the main file // without recreating them. - let wal_missing = !wal_sidecar(&db_path).exists() && !shm_sidecar(&db_path).exists(); + let wal_missing = ["-wal", "-shm"] + .into_iter() + .all(|suffix| !sidecar(&db_path, suffix).exists()); if !wal_missing { tracing::debug!("Cursor app auth read failed: {err}"); return None; @@ -51,15 +53,9 @@ pub fn load_app_auth_access_token() -> Option { } } -fn wal_sidecar(db_path: &std::path::Path) -> std::path::PathBuf { +fn sidecar(db_path: &std::path::Path, suffix: &str) -> std::path::PathBuf { let mut name = db_path.as_os_str().to_os_string(); - name.push("-wal"); - std::path::PathBuf::from(name) -} - -fn shm_sidecar(db_path: &std::path::Path) -> std::path::PathBuf { - let mut name = db_path.as_os_str().to_os_string(); - name.push("-shm"); + name.push(suffix); std::path::PathBuf::from(name) } diff --git a/rust/src/providers/cursor/mod.rs b/rust/src/providers/cursor/mod.rs index 443b15a52d..bf09384808 100755 --- a/rust/src/providers/cursor/mod.rs +++ b/rust/src/providers/cursor/mod.rs @@ -23,6 +23,12 @@ pub use api::CursorApi; use cost_cooldown::{CostCooldown, credential_fingerprint}; use token_cost::TokenCostError; +/// Quota usage plus the best-effort token-cost report. +type CursorFetch = ( + api::CursorUsageResult, + Option, +); + /// Cursor provider for fetching AI usage limits pub struct CursorProvider { api: CursorApi, @@ -42,16 +48,7 @@ impl CursorProvider { } } - async fn fetch_web_usage( - &self, - ctx: &FetchContext, - ) -> Result< - ( - api::CursorUsageResult, - Option, - ), - ProviderError, - > { + async fn fetch_web_usage(&self, ctx: &FetchContext) -> Result { let cookie_header = if let Some(cookie_header) = ctx.manual_cookie_header.as_deref() { cookie_header.to_string() } else { @@ -67,17 +64,17 @@ impl CursorProvider { crate::providers::browser_cookie_header(&["cursor.com", "cursor.sh"])? }; - self.fetch_usage_and_token_report(&cookie_header).await + let usage = self + .api + .fetch_usage_with_cookie_header(&cookie_header) + .await?; + let token_report = self.fetch_token_report_best_effort(&cookie_header).await; + Ok((usage, token_report)) } /// One usage pass with the app's local session; `None` means the app /// session was unavailable or rejected (caller falls back to cookies). - async fn fetch_via_app_session( - &self, - ) -> Option<( - api::CursorUsageResult, - Option, - )> { + async fn fetch_via_app_session(&self) -> Option { let app_cookie = app_auth::preferred_auto_cookie_header()?; let usage = match self.api.fetch_usage_with_cookie_header(&app_cookie).await { Ok(usage) => usage, @@ -93,24 +90,6 @@ impl CursorProvider { Some((usage, token_report)) } - async fn fetch_usage_and_token_report( - &self, - cookie_header: &str, - ) -> Result< - ( - api::CursorUsageResult, - Option, - ), - ProviderError, - > { - let usage = self - .api - .fetch_usage_with_cookie_header(cookie_header) - .await?; - let token_report = self.fetch_token_report_best_effort(cookie_header).await; - Ok((usage, token_report)) - } - /// Best-effort token-cost page; never fail the main usage fetch. /// /// A 403 is a cost-only rejection: the events request is skipped for six diff --git a/rust/src/providers/cursor/token_cost.rs b/rust/src/providers/cursor/token_cost.rs index a5f3965cb2..c802912401 100644 --- a/rust/src/providers/cursor/token_cost.rs +++ b/rust/src/providers/cursor/token_cost.rs @@ -644,33 +644,34 @@ mod tests { )); } + /// One event with no cache tokens. + fn event( + timestamp_ms: i64, + model: &str, + input_tokens: i64, + output_tokens: i64, + cost: EventCost, + charged_cents: Option, + ) -> UsageEvent { + UsageEvent { + timestamp_ms: Some(timestamp_ms), + model: Some(model.into()), + token_usage: Some(EventTokenUsage { + input_tokens, + output_tokens, + cache_write_tokens: 0, + cache_read_tokens: 0, + cost, + }), + charged_cents, + } + } + #[test] fn summarizes_per_model_and_metered() { let events = vec![ - UsageEvent { - timestamp_ms: Some(1), - model: Some("gpt-5".into()), - token_usage: Some(EventTokenUsage { - input_tokens: 10, - output_tokens: 5, - cache_write_tokens: 0, - cache_read_tokens: 0, - cost: EventCost::Valid(25.0), - }), - charged_cents: Some(10.0), - }, - UsageEvent { - timestamp_ms: Some(2), - model: Some("claude-4".into()), - token_usage: Some(EventTokenUsage { - input_tokens: 1, - output_tokens: 1, - cache_write_tokens: 0, - cache_read_tokens: 0, - cost: EventCost::Valid(75.0), - }), - charged_cents: Some(40.0), - }, + event(1, "gpt-5", 10, 5, EventCost::Valid(25.0), Some(10.0)), + event(2, "claude-4", 1, 1, EventCost::Valid(75.0), Some(40.0)), ]; let report = summarize_events(&events); assert!((report.api_rate_usd - 1.0).abs() < 0.001); @@ -684,30 +685,8 @@ mod tests { #[test] fn missing_model_cost_keeps_priced_siblings() { let events = vec![ - UsageEvent { - timestamp_ms: Some(1), - model: Some("gpt-5".into()), - token_usage: Some(EventTokenUsage { - input_tokens: 1, - output_tokens: 1, - cache_write_tokens: 0, - cache_read_tokens: 0, - cost: EventCost::Valid(25.0), - }), - charged_cents: Some(0.0), - }, - UsageEvent { - timestamp_ms: Some(2), - model: Some("gpt-5".into()), - token_usage: Some(EventTokenUsage { - input_tokens: 1, - output_tokens: 1, - cache_write_tokens: 0, - cache_read_tokens: 0, - cost: EventCost::Omitted, - }), - charged_cents: Some(0.0), - }, + event(1, "gpt-5", 1, 1, EventCost::Valid(25.0), Some(0.0)), + event(2, "gpt-5", 1, 1, EventCost::Omitted, Some(0.0)), ]; let report = summarize_events(&events); assert!(report.api_rate_usd > 0.25); @@ -719,30 +698,8 @@ mod tests { #[test] fn invalid_model_cost_latches_and_cannot_revive() { let events = vec![ - UsageEvent { - timestamp_ms: Some(1), - model: Some("gpt-5".into()), - token_usage: Some(EventTokenUsage { - input_tokens: 1, - output_tokens: 1, - cache_write_tokens: 0, - cache_read_tokens: 0, - cost: EventCost::Invalid, - }), - charged_cents: Some(0.0), - }, - UsageEvent { - timestamp_ms: Some(2), - model: Some("gpt-5".into()), - token_usage: Some(EventTokenUsage { - input_tokens: 1, - output_tokens: 1, - cache_write_tokens: 0, - cache_read_tokens: 0, - cost: EventCost::Valid(50.0), - }), - charged_cents: Some(0.0), - }, + event(1, "gpt-5", 1, 1, EventCost::Invalid, Some(0.0)), + event(2, "gpt-5", 1, 1, EventCost::Valid(50.0), Some(0.0)), ]; let report = summarize_events(&events); assert_eq!(report.api_rate_usd, 0.0); @@ -751,18 +708,7 @@ mod tests { #[test] fn incomplete_metered_is_none() { - let events = vec![UsageEvent { - timestamp_ms: Some(1), - model: Some("gpt-5".into()), - token_usage: Some(EventTokenUsage { - input_tokens: 1, - output_tokens: 1, - cache_write_tokens: 0, - cache_read_tokens: 0, - cost: EventCost::Valid(10.0), - }), - charged_cents: None, - }]; + let events = vec![event(1, "gpt-5", 1, 1, EventCost::Valid(10.0), None)]; let report = summarize_events(&events); assert!(report.metered_usd.is_none()); assert!((report.api_rate_usd - 0.1).abs() < 0.001); @@ -803,30 +749,15 @@ mod tests { fn omitted_known_cost_is_estimated_but_unknown_model_stays_unpriced() { let timestamp = 1_700_000_000_000; let events = vec![ - UsageEvent { - timestamp_ms: Some(timestamp), - model: Some("gpt-5".into()), - token_usage: Some(EventTokenUsage { - input_tokens: 200, - output_tokens: 20, - cache_write_tokens: 0, - cache_read_tokens: 0, - cost: EventCost::Omitted, - }), - charged_cents: Some(10.0), - }, - UsageEvent { - timestamp_ms: Some(timestamp + 1), - model: Some("fixture-model".into()), - token_usage: Some(EventTokenUsage { - input_tokens: 7, - output_tokens: 0, - cache_write_tokens: 0, - cache_read_tokens: 0, - cost: EventCost::Omitted, - }), - charged_cents: Some(10.0), - }, + event(timestamp, "gpt-5", 200, 20, EventCost::Omitted, Some(10.0)), + event( + timestamp + 1, + "fixture-model", + 7, + 0, + EventCost::Omitted, + Some(10.0), + ), ]; let report = summarize_events(&events); assert!((report.api_rate_usd - 0.00045).abs() < 1e-9); diff --git a/rust/src/providers/doubao/arkcli.rs b/rust/src/providers/doubao/arkcli.rs new file mode 100644 index 0000000000..013ae08d4a --- /dev/null +++ b/rust/src/providers/doubao/arkcli.rs @@ -0,0 +1,239 @@ +//! `arkcli usage plan` fallback for Doubao Coding Plan usage (upstream 0.45 #2221). + +use chrono::{DateTime, TimeZone, Utc}; +use serde::Deserialize; +use std::path::PathBuf; +use std::process::Command; +use std::time::Duration; + +use super::{CodingPlanQuota, CodingPlanResult, coding_plan_snapshot}; +use crate::core::{ProviderError, UsageSnapshot}; + +#[derive(Debug, Deserialize)] +struct ArkcliUsageResponse { + #[serde(default)] + viewer: Option, + #[serde(default)] + items: Vec, +} + +#[derive(Debug, Deserialize)] +struct ArkcliViewer { + #[serde(default, rename = "auth_method")] + auth_method: Option, +} + +#[derive(Debug, Deserialize)] +struct ArkcliUsageItem { + product: String, + #[serde(default)] + subscribed: Option, + #[serde(default)] + periods: Option>, + #[serde(default, rename = "updated_at")] + updated_at: Option, + #[serde(default)] + error: Option, +} + +#[derive(Debug, Deserialize)] +struct ArkcliPeriod { + label: String, + percent: f64, + #[serde(default, rename = "reset_at")] + reset_at: Option, +} + +pub(super) fn resolve_arkcli_binary() -> Option { + if let Ok(path) = std::env::var("ARKCLI_PATH") { + let p = PathBuf::from(path.trim()); + if p.is_file() { + return Some(p); + } + } + which::which("arkcli").ok() +} + +fn is_arkcli_auth_error(message: &str) -> bool { + let n = message.to_ascii_lowercase(); + [ + "not logged in", + "not authenticated", + "authentication required", + "login required", + "please login", + "please log in", + ] + .iter() + .any(|s| n.contains(s)) +} + +fn run_arkcli_usage_plan() -> Result, ProviderError> { + let bin = resolve_arkcli_binary().ok_or_else(|| { + ProviderError::NotInstalled( + "arkcli was not found. Install arkcli, run 'arkcli auth login', or configure Doubao API credentials." + .into(), + ) + })?; + let mut command = Command::new(&bin); + command + .args(["usage", "plan", "--format", "json"]) + .stdout(std::process::Stdio::piped()) + .stderr(std::process::Stdio::piped()); + // Runs during background refreshes: keep the CLI's console window hidden + // so it does not flash up or take focus. + #[cfg(windows)] + { + use std::os::windows::process::CommandExt; + const CREATE_NO_WINDOW: u32 = 0x0800_0000; + command.creation_flags(CREATE_NO_WINDOW); + } + let mut child = command + .spawn() + .map_err(|e| ProviderError::Other(format!("Failed to launch arkcli: {e}")))?; + + // ponytail: 15s wall-clock via join timeout isn't available on std Command; + // kill after wait timeout via a simple timed poll loop. + let deadline = std::time::Instant::now() + Duration::from_secs(15); + loop { + match child.try_wait() { + Ok(Some(_)) => break, + Ok(None) if std::time::Instant::now() >= deadline => { + // Best-effort teardown of the timed-out child; the outcome is already + // reported as timed out. + let _killed = child.kill(); + let _reaped = child.wait(); + return Err(ProviderError::Other( + "arkcli usage timed out. Check arkcli authentication and try again.".into(), + )); + } + Ok(None) => std::thread::sleep(Duration::from_millis(50)), + Err(e) => { + return Err(ProviderError::Other(format!("arkcli wait failed: {e}"))); + } + } + } + let output = child + .wait_with_output() + .map_err(|e| ProviderError::Other(format!("arkcli wait failed: {e}")))?; + if !output.status.success() { + let stderr = String::from_utf8_lossy(&output.stderr); + let message = stderr.split_whitespace().collect::>().join(" "); + if is_arkcli_auth_error(&message) { + return Err(ProviderError::AuthRequired); + } + let code = output.status.code().unwrap_or(-1); + return Err(ProviderError::Other(format!( + "arkcli usage failed ({code}): {}", + if message.is_empty() { + "unknown error" + } else { + &message + } + ))); + } + if output.stdout.len() > 256 * 1024 { + return Err(ProviderError::Other( + "arkcli returned too much output. Update arkcli and try again.".into(), + )); + } + Ok(output.stdout) +} + +pub(super) fn decode_arkcli_usage(bytes: &[u8]) -> Result { + let response: ArkcliUsageResponse = serde_json::from_slice(bytes) + .map_err(|e| ProviderError::Parse(format!("Failed to parse arkcli usage: {e}")))?; + + if let Some(method) = response + .viewer + .as_ref() + .and_then(|v| v.auth_method.as_deref()) + .map(str::trim) + && method.eq_ignore_ascii_case("none") + { + return Err(ProviderError::AuthRequired); + } + + let mut quotas = Vec::new(); + let mut update_ts: Option = None; + let mut status = response + .viewer + .and_then(|v| v.auth_method) + .filter(|s| !s.trim().is_empty()); + + for item in response.items { + let product = item.product.to_ascii_lowercase(); + let level_prefix = match product.as_str() { + "agent-plan" => "agent_", + "coding-plan" => "", + "agent-plan-team" => "agent_team_", + "coding-plan-team" => "coding_team_", + _ => continue, + }; + if item.subscribed == Some(false) { + continue; + } + let periods = item.periods.unwrap_or_default(); + if !periods.is_empty() + && let Some(updated_at) = item.updated_at.filter(|v| *v > 0.0) + { + // arkcli may emit ms or seconds; 1e11 is the unit threshold. + let seconds = if updated_at >= 1e11 { + updated_at / 1000.0 + } else { + updated_at + }; + if update_ts.map(|t| seconds > t).unwrap_or(true) { + update_ts = Some(seconds); + } + } + for period in periods { + let level = format!("{level_prefix}{}", period.label); + let reset_timestamp = period + .reset_at + .as_deref() + .and_then(|raw| DateTime::parse_from_rfc3339(raw.trim()).ok()) + .map(|d| d.timestamp() as f64); + quotas.push(CodingPlanQuota { + level, + percent: period.percent, + reset_timestamp, + }); + } + if status.is_none() { + status = Some(product); + } + } + + if quotas.is_empty() { + return Err(ProviderError::Parse( + "arkcli returned no active Coding or Agent Plan usage.".into(), + )); + } + + Ok(CodingPlanResult { + status, + update_timestamp: update_ts, + quota_usage: quotas, + }) +} + +pub(super) fn fetch_arkcli_usage() -> Result { + let stdout = run_arkcli_usage_plan()?; + let usage = decode_arkcli_usage(&stdout)?; + Ok(coding_plan_snapshot(usage)) +} + +pub(super) fn datetime_from_epoch(timestamp: f64) -> Option> { + if !timestamp.is_finite() || timestamp <= 0.0 { + return None; + } + // Epoch seconds guarded finite and positive; the sub-second fraction is + // below timestamp resolution. + #[expect( + clippy::cast_possible_truncation, + reason = "epoch seconds; sub-second fraction below timestamp resolution" + )] + let whole_seconds = timestamp as i64; + Utc.timestamp_opt(whole_seconds, 0).single() +} diff --git a/rust/src/providers/doubao/mod.rs b/rust/src/providers/doubao/mod.rs index 84bb2f6ab8..bc1476641b 100644 --- a/rust/src/providers/doubao/mod.rs +++ b/rust/src/providers/doubao/mod.rs @@ -3,19 +3,23 @@ //! Probes Ark chat-completions with a one-token request and reads rate-limit headers. //! Also supports signed Coding Plan API credentials and `arkcli usage plan` (0.45). +mod arkcli; + use async_trait::async_trait; use chrono::{DateTime, TimeZone, Utc}; use reqwest::Client; use serde::Deserialize; use serde_json::json; -use std::path::PathBuf; -use std::process::Command; -use std::time::Duration; use crate::core::{ FetchContext, IconLane, NamedRateWindow, Provider, ProviderError, ProviderFetchResult, ProviderId, RateWindow, SourceMode, UsageSnapshot, hex, hmac_sha256, sha256_hex, }; +use crate::providers::resolve_api_key; + +#[cfg(test)] +use arkcli::decode_arkcli_usage; +use arkcli::{datetime_from_epoch, fetch_arkcli_usage, resolve_arkcli_binary}; const DOUBAO_API_URL: &str = "https://ark.cn-beijing.volces.com/api/coding/v3/chat/completions"; const DOUBAO_CODING_PLAN_URL: &str = @@ -514,265 +518,6 @@ fn coding_plan_window( )) } -// --- arkcli usage plan (upstream 0.45 #2221) --- - -#[derive(Debug, Deserialize)] -struct ArkcliUsageResponse { - #[serde(default)] - viewer: Option, - #[serde(default)] - items: Vec, -} - -#[derive(Debug, Deserialize)] -struct ArkcliViewer { - #[serde(default, rename = "auth_method")] - auth_method: Option, -} - -#[derive(Debug, Deserialize)] -struct ArkcliUsageItem { - product: String, - #[serde(default)] - subscribed: Option, - #[serde(default)] - periods: Option>, - #[serde(default, rename = "updated_at")] - updated_at: Option, - #[serde(default)] - error: Option, -} - -#[derive(Debug, Deserialize)] -struct ArkcliPeriod { - label: String, - percent: f64, - #[serde(default, rename = "reset_at")] - reset_at: Option, -} - -fn resolve_arkcli_binary() -> Option { - if let Ok(path) = std::env::var("ARKCLI_PATH") { - let p = PathBuf::from(path.trim()); - if p.is_file() { - return Some(p); - } - } - which::which("arkcli").ok() -} - -fn is_arkcli_auth_error(message: &str) -> bool { - let n = message.to_ascii_lowercase(); - [ - "not logged in", - "not authenticated", - "authentication required", - "login required", - "please login", - "please log in", - ] - .iter() - .any(|s| n.contains(s)) -} - -fn run_arkcli_usage_plan() -> Result, ProviderError> { - let bin = resolve_arkcli_binary().ok_or_else(|| { - ProviderError::NotInstalled( - "arkcli was not found. Install arkcli, run 'arkcli auth login', or configure Doubao API credentials." - .into(), - ) - })?; - let mut command = Command::new(&bin); - command - .args(["usage", "plan", "--format", "json"]) - .stdout(std::process::Stdio::piped()) - .stderr(std::process::Stdio::piped()); - // Runs during background refreshes: keep the CLI's console window hidden - // so it does not flash up or take focus. - #[cfg(windows)] - { - use std::os::windows::process::CommandExt; - const CREATE_NO_WINDOW: u32 = 0x0800_0000; - command.creation_flags(CREATE_NO_WINDOW); - } - let mut child = command - .spawn() - .map_err(|e| ProviderError::Other(format!("Failed to launch arkcli: {e}")))?; - - // ponytail: 15s wall-clock via join timeout isn't available on std Command; - // kill after wait timeout via a simple timed poll loop. - let deadline = std::time::Instant::now() + Duration::from_secs(15); - loop { - match child.try_wait() { - Ok(Some(_)) => break, - Ok(None) if std::time::Instant::now() >= deadline => { - // Best-effort teardown of the timed-out child; the outcome is already - // reported as timed out. - let _killed = child.kill(); - let _reaped = child.wait(); - return Err(ProviderError::Other( - "arkcli usage timed out. Check arkcli authentication and try again.".into(), - )); - } - Ok(None) => std::thread::sleep(Duration::from_millis(50)), - Err(e) => { - return Err(ProviderError::Other(format!("arkcli wait failed: {e}"))); - } - } - } - let output = child - .wait_with_output() - .map_err(|e| ProviderError::Other(format!("arkcli wait failed: {e}")))?; - if !output.status.success() { - let stderr = String::from_utf8_lossy(&output.stderr); - let message = stderr.split_whitespace().collect::>().join(" "); - if is_arkcli_auth_error(&message) { - return Err(ProviderError::AuthRequired); - } - let code = output.status.code().unwrap_or(-1); - return Err(ProviderError::Other(format!( - "arkcli usage failed ({code}): {}", - if message.is_empty() { - "unknown error" - } else { - &message - } - ))); - } - if output.stdout.len() > 256 * 1024 { - return Err(ProviderError::Other( - "arkcli returned too much output. Update arkcli and try again.".into(), - )); - } - Ok(output.stdout) -} - -fn decode_arkcli_usage(bytes: &[u8]) -> Result { - let response: ArkcliUsageResponse = serde_json::from_slice(bytes) - .map_err(|e| ProviderError::Parse(format!("Failed to parse arkcli usage: {e}")))?; - - if let Some(method) = response - .viewer - .as_ref() - .and_then(|v| v.auth_method.as_deref()) - .map(str::trim) - && method.eq_ignore_ascii_case("none") - { - return Err(ProviderError::AuthRequired); - } - - let supported = [ - "agent-plan", - "coding-plan", - "agent-plan-team", - "coding-plan-team", - ]; - for item in &response.items { - let product = item.product.to_ascii_lowercase(); - if !supported.iter().any(|p| *p == product) { - continue; - } - if item.subscribed == Some(false) { - continue; - } - let periods_empty = item.periods.as_ref().map(|p| p.is_empty()).unwrap_or(true); - if periods_empty { - let message = item - .error - .as_deref() - .map(str::trim) - .filter(|s| !s.is_empty()) - .map(|s| s.to_string()) - .unwrap_or_else(|| format!("{product} has no usage periods")); - // Incomplete but keep scanning other products. - let _ = message; - } - } - - let mut quotas = Vec::new(); - let mut update_ts: Option = None; - let mut status = response - .viewer - .and_then(|v| v.auth_method) - .filter(|s| !s.trim().is_empty()); - - for item in response.items { - let product = item.product.to_ascii_lowercase(); - let level_prefix = match product.as_str() { - "agent-plan" => "agent_", - "coding-plan" => "", - "agent-plan-team" => "agent_team_", - "coding-plan-team" => "coding_team_", - _ => continue, - }; - if item.subscribed == Some(false) { - continue; - } - let periods = item.periods.unwrap_or_default(); - if !periods.is_empty() - && let Some(updated_at) = item.updated_at.filter(|v| *v > 0.0) - { - // arkcli may emit ms or seconds; 1e11 is the unit threshold. - let seconds = if updated_at >= 1e11 { - updated_at / 1000.0 - } else { - updated_at - }; - if update_ts.map(|t| seconds > t).unwrap_or(true) { - update_ts = Some(seconds); - } - } - for period in periods { - let level = format!("{level_prefix}{}", period.label); - let reset_timestamp = period - .reset_at - .as_deref() - .and_then(|raw| DateTime::parse_from_rfc3339(raw.trim()).ok()) - .map(|d| d.timestamp() as f64); - quotas.push(CodingPlanQuota { - level, - percent: period.percent, - reset_timestamp, - }); - } - if status.is_none() { - status = Some(product); - } - } - - if quotas.is_empty() { - return Err(ProviderError::Parse( - "arkcli returned no active Coding or Agent Plan usage.".into(), - )); - } - - Ok(CodingPlanResult { - status, - update_timestamp: update_ts, - quota_usage: quotas, - }) -} - -fn fetch_arkcli_usage() -> Result { - let stdout = run_arkcli_usage_plan()?; - let usage = decode_arkcli_usage(&stdout)?; - Ok(coding_plan_snapshot(usage)) -} - -fn datetime_from_epoch(timestamp: f64) -> Option> { - if !timestamp.is_finite() || timestamp <= 0.0 { - return None; - } - // Epoch seconds guarded finite and positive; the sub-second fraction is - // below timestamp resolution. - #[expect( - clippy::cast_possible_truncation, - reason = "epoch seconds; sub-second fraction below timestamp resolution" - )] - let whole_seconds = timestamp as i64; - Utc.timestamp_opt(whole_seconds, 0).single() -} - struct SignedVolcengineRequest { content_type: &'static str, host: String, @@ -974,267 +719,8 @@ fn selected_ark_api_key(ctx: &FetchContext) -> Result { .ok_or(ProviderError::AuthRequired) } -fn resolve_api_key( - explicit: Option<&str>, - credential_target: &str, - env_names: &[&str], -) -> Result { - if let Some(key) = explicit - && !key.trim().is_empty() - { - return Ok(key.trim().to_string()); - } - if let Ok(entry) = keyring::Entry::new(credential_target, "api_key") - && let Ok(key) = entry.get_password() - && !key.trim().is_empty() - { - return Ok(key); - } - for env in env_names { - if let Ok(key) = std::env::var(env) - && !key.trim().is_empty() - { - return Ok(key); - } - } - Err(ProviderError::NotInstalled(format!( - "API key not found. Set {} in Preferences or environment.", - env_names.join(" / ") - ))) -} - #[cfg(test)] mod icon_lane_tests; #[cfg(test)] -mod tests { - use super::*; - - #[test] - fn selected_ark_account_requires_its_projected_key() { - let isolated = FetchContext { - token_account_isolated: true, - token_account_kind: Some(crate::core::TokenAccountKind::ApiKey), - ..FetchContext::default() - }; - assert!(matches!( - selected_ark_api_key(&isolated), - Err(ProviderError::AuthRequired) - )); - - let selected = FetchContext { - api_key: Some(" selected-key ".into()), - ..isolated - }; - assert_eq!(selected_ark_api_key(&selected).unwrap(), "selected-key"); - } - use reqwest::header::{HeaderMap, HeaderValue}; - - #[test] - fn doubao_snapshot_uses_rate_limit_headers() { - let mut headers = HeaderMap::new(); - headers.insert( - "x-ratelimit-remaining-requests", - HeaderValue::from_static("25"), - ); - headers.insert( - "x-ratelimit-limit-requests", - HeaderValue::from_static("100"), - ); - let snapshot = - probe_result_from_response(reqwest::StatusCode::OK, &headers, &json!({})).snapshot; - assert_eq!(snapshot.primary.used_percent, 75.0); - } - - #[test] - fn doubao_repeated_successful_zero_remaining_falls_back_to_active() { - let mut headers = HeaderMap::new(); - headers.insert( - "x-ratelimit-remaining-requests", - HeaderValue::from_static("0"), - ); - headers.insert( - "x-ratelimit-limit-requests", - HeaderValue::from_static("1000"), - ); - - let result = probe_result_from_response(reqwest::StatusCode::OK, &headers, &json!({})); - assert!(result.has_ambiguous_zero_remaining()); - - let snapshot = snapshot_from_parts( - result.remaining, - result.limit, - result.resets_at, - result.total_tokens, - false, - ); - assert_eq!(snapshot.primary.used_percent, 0.0); - assert_eq!( - snapshot.primary.reset_description.as_deref(), - Some("Active - check dashboard for details") - ); - } - - #[test] - fn doubao_rate_limit_with_limit_header_reports_exhausted() { - let mut headers = HeaderMap::new(); - headers.insert( - "x-ratelimit-limit-requests", - HeaderValue::from_static("1000"), - ); - let snapshot = probe_result_from_response( - reqwest::StatusCode::TOO_MANY_REQUESTS, - &headers, - &json!({}), - ) - .snapshot; - - assert_eq!(snapshot.primary.used_percent, 100.0); - assert_eq!( - snapshot.primary.reset_description.as_deref(), - Some("1000/1000 requests") - ); - } - - #[test] - fn doubao_bare_rate_limit_uses_active_fallback() { - let snapshot = probe_result_from_response( - reqwest::StatusCode::TOO_MANY_REQUESTS, - &HeaderMap::new(), - &json!({}), - ) - .snapshot; - - assert_eq!(snapshot.primary.used_percent, 0.0); - assert_eq!( - snapshot.primary.reset_description.as_deref(), - Some("Active - check dashboard for details") - ); - } - - #[test] - fn doubao_parses_coding_plan_usage() { - let body = br#"{ - "Result": { - "Status": "active", - "UpdateTimestamp": 1783036800, - "QuotaUsage": [ - {"Level": "session", "Percent": 12.5, "ResetTimestamp": 1783040400}, - {"Level": "weekly", "Percent": 50.0, "ResetTimestamp": 1783641600}, - {"Level": "monthly", "Percent": 75.0, "ResetTimestamp": 1785628800} - ] - } - }"#; - let snapshot = coding_plan_snapshot(decode_coding_plan_usage(body).unwrap()); - assert_eq!(snapshot.primary.used_percent, 12.5); - assert_eq!(snapshot.secondary.unwrap().used_percent, 50.0); - let monthly = snapshot.tertiary.expect("monthly"); - assert_eq!(monthly.used_percent, 75.0); - // ResetTimestamp 1785628800 = 2026-08-02 → prior month is 31 days. - assert_eq!(monthly.window_minutes, Some(31 * 24 * 60)); - assert_eq!(snapshot.login_method.as_deref(), Some("active")); - } - - #[test] - fn doubao_parses_coding_plan_credentials() { - let creds = - DoubaoCodingPlanCredentials::parse("ak-test|sk-test|cn-shanghai").expect("creds"); - assert_eq!(creds.access_key_id, "ak-test"); - assert_eq!(creds.secret_access_key, "sk-test"); - assert_eq!(creds.region, "cn-shanghai"); - - let json = r#"{"accessKeyId":"ak-json","secretAccessKey":"sk-json","region":"cn-beijing"}"#; - let creds = DoubaoCodingPlanCredentials::parse(json).expect("json creds"); - assert_eq!(creds.access_key_id, "ak-json"); - assert_eq!(creds.secret_access_key, "sk-json"); - } - - #[test] - fn doubao_signer_sets_required_volcengine_headers() { - let creds = DoubaoCodingPlanCredentials { - access_key_id: "AKID".into(), - secret_access_key: "SECRET".into(), - region: "cn-beijing".into(), - }; - let signed = sign_volcengine_request( - &creds, - b"", - Utc.with_ymd_and_hms(2026, 7, 3, 0, 0, 0).unwrap(), - ) - .unwrap(); - assert_eq!(signed.host, "open.volcengineapi.com"); - assert_eq!(signed.timestamp, "20260703T000000Z"); - assert_eq!(signed.payload_hash, sha256_hex(b"")); - assert!( - signed - .authorization - .starts_with("HMAC-SHA256 Credential=AKID/20260703/cn-beijing/ark/request") - ); - } - - #[test] - fn arkcli_usage_parses_coding_and_agent_plans() { - let raw = br#"{ - "viewer": { "auth_method": "oauth" }, - "items": [ - { - "product": "coding-plan", - "subscribed": true, - "updated_at": 1720000000, - "periods": [ - { "label": "session", "percent": 12.5, "reset_at": "2026-07-21T10:00:00Z" }, - { "label": "weekly", "percent": 40.0, "reset_at": "2026-07-28T00:00:00Z" } - ] - }, - { - "product": "agent-plan", - "subscribed": true, - "periods": [ - { "label": "session", "percent": 5.0 }, - { "label": "weekly", "percent": 15.0 } - ] - } - ] - }"#; - let usage = decode_arkcli_usage(raw).expect("arkcli json"); - let snap = coding_plan_snapshot(usage); - assert!((snap.primary.used_percent - 12.5).abs() < 0.01); - assert!((snap.secondary.as_ref().unwrap().used_percent - 40.0).abs() < 0.01); - assert!( - snap.extra_rate_windows - .iter() - .any(|w| w.id == "doubao-agent-session") - ); - assert!( - snap.extra_rate_windows - .iter() - .any(|w| (w.window.used_percent - 5.0).abs() < 0.01) - ); - } - - #[test] - fn arkcli_auth_none_is_auth_required() { - let raw = br#"{ "viewer": { "auth_method": "none" }, "items": [] }"#; - let err = decode_arkcli_usage(raw).unwrap_err(); - assert!(matches!(err, ProviderError::AuthRequired)); - } - - #[test] - fn arkcli_presence_maps_to_local_runtime_offline_but_api_key_stays_default() { - assert_eq!( - DoubaoProvider::new().error_state_kind(&ProviderError::NotInstalled( - "arkcli was not found. Install arkcli, run 'arkcli auth login', or configure \ - Doubao API credentials." - .into(), - )), - crate::core::ProviderStateKind::LocalRuntimeOffline - ); - // The shared API-key producer keeps the default mapping. - assert_eq!( - DoubaoProvider::new().error_state_kind(&ProviderError::NotInstalled( - "API key not found. Set ARK_API_KEY in Preferences or environment.".into(), - )), - crate::core::ProviderStateKind::NeedsAuthentication - ); - } -} +mod tests; diff --git a/rust/src/providers/doubao/tests.rs b/rust/src/providers/doubao/tests.rs new file mode 100644 index 0000000000..584af48c11 --- /dev/null +++ b/rust/src/providers/doubao/tests.rs @@ -0,0 +1,211 @@ +use super::*; + +#[test] +fn selected_ark_account_requires_its_projected_key() { + let isolated = FetchContext { + token_account_isolated: true, + token_account_kind: Some(crate::core::TokenAccountKind::ApiKey), + ..FetchContext::default() + }; + assert!(matches!( + selected_ark_api_key(&isolated), + Err(ProviderError::AuthRequired) + )); + + let selected = FetchContext { + api_key: Some(" selected-key ".into()), + ..isolated + }; + assert_eq!(selected_ark_api_key(&selected).unwrap(), "selected-key"); +} +use reqwest::header::{HeaderMap, HeaderValue}; + +/// Probe result for `status` with `(remaining, limit)` rate-limit headers. +fn probe_with( + status: reqwest::StatusCode, + remaining: Option<&'static str>, + limit: Option<&'static str>, +) -> DoubaoProbeResult { + let mut headers = HeaderMap::new(); + for (name, value) in [ + ("x-ratelimit-remaining-requests", remaining), + ("x-ratelimit-limit-requests", limit), + ] { + if let Some(value) = value { + headers.insert(name, HeaderValue::from_static(value)); + } + } + probe_result_from_response(status, &headers, &json!({})) +} + +#[test] +fn doubao_snapshot_uses_rate_limit_headers() { + let snapshot = probe_with(reqwest::StatusCode::OK, Some("25"), Some("100")).snapshot; + assert_eq!(snapshot.primary.used_percent, 75.0); +} + +#[test] +fn doubao_repeated_successful_zero_remaining_falls_back_to_active() { + let result = probe_with(reqwest::StatusCode::OK, Some("0"), Some("1000")); + assert!(result.has_ambiguous_zero_remaining()); + + let snapshot = snapshot_from_parts( + result.remaining, + result.limit, + result.resets_at, + result.total_tokens, + false, + ); + assert_eq!(snapshot.primary.used_percent, 0.0); + assert_eq!( + snapshot.primary.reset_description.as_deref(), + Some("Active - check dashboard for details") + ); +} + +#[test] +fn doubao_rate_limit_with_limit_header_reports_exhausted() { + let snapshot = probe_with(reqwest::StatusCode::TOO_MANY_REQUESTS, None, Some("1000")).snapshot; + + assert_eq!(snapshot.primary.used_percent, 100.0); + assert_eq!( + snapshot.primary.reset_description.as_deref(), + Some("1000/1000 requests") + ); +} + +#[test] +fn doubao_bare_rate_limit_uses_active_fallback() { + let snapshot = probe_with(reqwest::StatusCode::TOO_MANY_REQUESTS, None, None).snapshot; + + assert_eq!(snapshot.primary.used_percent, 0.0); + assert_eq!( + snapshot.primary.reset_description.as_deref(), + Some("Active - check dashboard for details") + ); +} + +#[test] +fn doubao_parses_coding_plan_usage() { + let body = br#"{ + "Result": { + "Status": "active", + "UpdateTimestamp": 1783036800, + "QuotaUsage": [ + {"Level": "session", "Percent": 12.5, "ResetTimestamp": 1783040400}, + {"Level": "weekly", "Percent": 50.0, "ResetTimestamp": 1783641600}, + {"Level": "monthly", "Percent": 75.0, "ResetTimestamp": 1785628800} + ] + } + }"#; + let snapshot = coding_plan_snapshot(decode_coding_plan_usage(body).unwrap()); + assert_eq!(snapshot.primary.used_percent, 12.5); + assert_eq!(snapshot.secondary.unwrap().used_percent, 50.0); + let monthly = snapshot.tertiary.expect("monthly"); + assert_eq!(monthly.used_percent, 75.0); + // ResetTimestamp 1785628800 = 2026-08-02 → prior month is 31 days. + assert_eq!(monthly.window_minutes, Some(31 * 24 * 60)); + assert_eq!(snapshot.login_method.as_deref(), Some("active")); +} + +#[test] +fn doubao_parses_coding_plan_credentials() { + let creds = DoubaoCodingPlanCredentials::parse("ak-test|sk-test|cn-shanghai").expect("creds"); + assert_eq!(creds.access_key_id, "ak-test"); + assert_eq!(creds.secret_access_key, "sk-test"); + assert_eq!(creds.region, "cn-shanghai"); + + let json = r#"{"accessKeyId":"ak-json","secretAccessKey":"sk-json","region":"cn-beijing"}"#; + let creds = DoubaoCodingPlanCredentials::parse(json).expect("json creds"); + assert_eq!(creds.access_key_id, "ak-json"); + assert_eq!(creds.secret_access_key, "sk-json"); +} + +#[test] +fn doubao_signer_sets_required_volcengine_headers() { + let creds = DoubaoCodingPlanCredentials { + access_key_id: "AKID".into(), + secret_access_key: "SECRET".into(), + region: "cn-beijing".into(), + }; + let signed = sign_volcengine_request( + &creds, + b"", + Utc.with_ymd_and_hms(2026, 7, 3, 0, 0, 0).unwrap(), + ) + .unwrap(); + assert_eq!(signed.host, "open.volcengineapi.com"); + assert_eq!(signed.timestamp, "20260703T000000Z"); + assert_eq!(signed.payload_hash, sha256_hex(b"")); + assert!( + signed + .authorization + .starts_with("HMAC-SHA256 Credential=AKID/20260703/cn-beijing/ark/request") + ); +} + +#[test] +fn arkcli_usage_parses_coding_and_agent_plans() { + let raw = br#"{ + "viewer": { "auth_method": "oauth" }, + "items": [ + { + "product": "coding-plan", + "subscribed": true, + "updated_at": 1720000000, + "periods": [ + { "label": "session", "percent": 12.5, "reset_at": "2026-07-21T10:00:00Z" }, + { "label": "weekly", "percent": 40.0, "reset_at": "2026-07-28T00:00:00Z" } + ] + }, + { + "product": "agent-plan", + "subscribed": true, + "periods": [ + { "label": "session", "percent": 5.0 }, + { "label": "weekly", "percent": 15.0 } + ] + } + ] + }"#; + let usage = decode_arkcli_usage(raw).expect("arkcli json"); + let snap = coding_plan_snapshot(usage); + assert!((snap.primary.used_percent - 12.5).abs() < 0.01); + assert!((snap.secondary.as_ref().unwrap().used_percent - 40.0).abs() < 0.01); + assert!( + snap.extra_rate_windows + .iter() + .any(|w| w.id == "doubao-agent-session") + ); + assert!( + snap.extra_rate_windows + .iter() + .any(|w| (w.window.used_percent - 5.0).abs() < 0.01) + ); +} + +#[test] +fn arkcli_auth_none_is_auth_required() { + let raw = br#"{ "viewer": { "auth_method": "none" }, "items": [] }"#; + let err = decode_arkcli_usage(raw).unwrap_err(); + assert!(matches!(err, ProviderError::AuthRequired)); +} + +#[test] +fn arkcli_presence_maps_to_local_runtime_offline_but_api_key_stays_default() { + assert_eq!( + DoubaoProvider::new().error_state_kind(&ProviderError::NotInstalled( + "arkcli was not found. Install arkcli, run 'arkcli auth login', or configure \ + Doubao API credentials." + .into(), + )), + crate::core::ProviderStateKind::LocalRuntimeOffline + ); + // The shared API-key producer keeps the default mapping. + assert_eq!( + DoubaoProvider::new().error_state_kind(&ProviderError::NotInstalled( + "API key not found. Set ARK_API_KEY in Preferences or environment.".into(), + )), + crate::core::ProviderStateKind::NeedsAuthentication + ); +} diff --git a/rust/src/providers/factory/mod.rs b/rust/src/providers/factory/mod.rs index fdb027f7b9..768a04ec7e 100755 --- a/rust/src/providers/factory/mod.rs +++ b/rust/src/providers/factory/mod.rs @@ -403,16 +403,16 @@ impl FactoryProvider { .header("x-factory-client", FACTORY_CLIENT_HEADER) } - /// Fetch auth info with optional cookie and/or bearer token. - async fn fetch_auth_info( - &self, + /// GET `url` with optional cookie and/or bearer token; `label` names the + /// API in the non-success error. + async fn get_json( client: &reqwest::Client, - base: &str, + url: &str, cookies: Option<&str>, bearer: Option<&str>, - ) -> Result { - let url = format!("{base}/api/app/auth/me"); - let mut req = Self::apply_factory_headers(client.get(&url)); + label: &str, + ) -> Result { + let mut req = Self::apply_factory_headers(client.get(url)); if let Some(c) = cookies.filter(|s| !s.is_empty()) { req = req.header("Cookie", c); } @@ -422,15 +422,12 @@ impl FactoryProvider { let resp = req.send().await?; let status = resp.status(); - if status == reqwest::StatusCode::UNAUTHORIZED { - return Err(ProviderError::AuthRequired); - } - if status == reqwest::StatusCode::FORBIDDEN { + if status == reqwest::StatusCode::UNAUTHORIZED || status == reqwest::StatusCode::FORBIDDEN { return Err(ProviderError::AuthRequired); } if !status.is_success() { return Err(ProviderError::Other(format!( - "Factory auth API returned status {status}" + "Factory {label} API returned status {status}" ))); } @@ -439,6 +436,17 @@ impl FactoryProvider { .map_err(|e| ProviderError::Parse(e.to_string())) } + async fn fetch_auth_info( + &self, + client: &reqwest::Client, + base: &str, + cookies: Option<&str>, + bearer: Option<&str>, + ) -> Result { + let url = format!("{base}/api/app/auth/me"); + Self::get_json(client, &url, cookies, bearer, "auth").await + } + /// Fetch legacy subscription usage. async fn fetch_usage_api( &self, @@ -448,28 +456,7 @@ impl FactoryProvider { bearer: Option<&str>, ) -> Result { let url = format!("{base}/api/organization/subscription/usage?useCache=true"); - let mut req = Self::apply_factory_headers(client.get(&url)); - if let Some(c) = cookies.filter(|s| !s.is_empty()) { - req = req.header("Cookie", c); - } - if let Some(token) = bearer.filter(|s| !s.is_empty()) { - req = req.header("Authorization", format!("Bearer {token}")); - } - - let resp = req.send().await?; - let status = resp.status(); - if status == reqwest::StatusCode::UNAUTHORIZED || status == reqwest::StatusCode::FORBIDDEN { - return Err(ProviderError::AuthRequired); - } - if !status.is_success() { - return Err(ProviderError::Other(format!( - "Factory usage API returned status {status}" - ))); - } - - resp.json() - .await - .map_err(|e| ProviderError::Parse(e.to_string())) + Self::get_json(client, &url, cookies, bearer, "usage").await } /// Optional billing-limits probe (token-rate-limits accounts). @@ -519,7 +506,7 @@ impl FactoryProvider { for base in [FACTORY_API_BASE, FACTORY_APP_BASE] { match self - .fetch_auth_and_usage_bearer(&client, base, api_key) + .fetch_auth_and_usage(&client, base, None, Some(api_key)) .await { Ok(snapshot) => return Ok(snapshot), @@ -538,19 +525,19 @@ impl FactoryProvider { .unwrap_or(ProviderError::AuthRequired)) } - async fn fetch_auth_and_usage_bearer( + /// Best-effort auth info plus required legacy usage from one host. + async fn fetch_auth_and_usage( &self, client: &reqwest::Client, base: &str, - api_key: &str, + cookies: Option<&str>, + bearer: Option<&str>, ) -> Result { let auth_info = self - .fetch_auth_info(client, base, None, Some(api_key)) + .fetch_auth_info(client, base, cookies, bearer) .await .ok(); - let usage_data = self - .fetch_usage_api(client, base, None, Some(api_key)) - .await?; + let usage_data = self.fetch_usage_api(client, base, cookies, bearer).await?; Ok(Self::apply_auth_info( Self::usage_snapshot_from_response(&usage_data), auth_info, @@ -561,18 +548,8 @@ impl FactoryProvider { async fn fetch_via_web(&self, ctx: &FetchContext) -> Result { let cookies = self.get_cookies(ctx)?; let client = Self::build_client()?; - let auth_info = self - .fetch_auth_info(&client, FACTORY_APP_BASE, Some(&cookies), None) + self.fetch_auth_and_usage(&client, FACTORY_APP_BASE, Some(&cookies), None) .await - .ok(); - let usage_data = self - .fetch_usage_api(&client, FACTORY_APP_BASE, Some(&cookies), None) - .await?; - - Ok(Self::apply_auth_info( - Self::usage_snapshot_from_response(&usage_data), - auth_info, - )) } fn usage_snapshot_from_response(usage_data: &FactoryUsageResponse) -> UsageSnapshot { @@ -761,240 +738,4 @@ impl Provider for FactoryProvider { } #[cfg(test)] -mod tests { - use super::*; - use std::fs; - - #[test] - fn cleans_quotes_and_whitespace() { - assert_eq!( - clean_factory_secret(Some(" \"fk-quoted\" ")).as_deref(), - Some("fk-quoted") - ); - assert_eq!( - clean_factory_secret(Some("'fk-single'")).as_deref(), - Some("fk-single") - ); - assert_eq!(clean_factory_secret(Some(" ")), None); - assert_eq!(clean_factory_secret(None), None); - } - - #[test] - fn parses_factory_dotenv_variants() { - assert_eq!( - parse_factory_dotenv_key("FACTORY_API_KEY=fk-plain").as_deref(), - Some("fk-plain") - ); - assert_eq!( - parse_factory_dotenv_key("export FACTORY_API_KEY='fk-single'").as_deref(), - Some("fk-single") - ); - assert_eq!( - parse_factory_dotenv_key("# comment\nFACTORY_API_KEY=\"fk-double\"").as_deref(), - Some("fk-double") - ); - assert_eq!(parse_factory_dotenv_key("OTHER=1\n"), None); - assert_eq!( - parse_factory_dotenv_key( - "export FACTORY_API_KEY='fk-quoted'\nMALFORMED\nFACTORY_API_KEY=\n" - ) - .as_deref(), - Some("fk-quoted") - ); - // Empty assignment must not short-circuit a later real key. - assert_eq!( - parse_factory_dotenv_key("FACTORY_API_KEY=\nFACTORY_API_KEY=fk-real\n").as_deref(), - Some("fk-real") - ); - assert_eq!( - parse_factory_dotenv_key("FACTORY_API_KEY=\"\"\nFACTORY_API_KEY='fk-later'\n") - .as_deref(), - Some("fk-later") - ); - } - - #[test] - fn key_resolution_honors_explicit_env_dotenv_precedence() { - let home = tempfile::tempdir().unwrap(); - fs::create_dir(home.path().join(".factory")).unwrap(); - fs::write( - home.path().join(".factory").join(".env"), - "FACTORY_API_KEY=fk-dotenv\n", - ) - .unwrap(); - - let env_with_key = HashMap::from([ - ( - FACTORY_API_KEY_ENV.to_string(), - " \"fk-env\" ".to_string(), - ), - ("USERPROFILE".to_string(), home.path().display().to_string()), - ]); - let env_dotenv_only = - HashMap::from([("USERPROFILE".to_string(), home.path().display().to_string())]); - let env_home_fallback = - HashMap::from([("HOME".to_string(), home.path().display().to_string())]); - - assert_eq!( - resolve_factory_api_key_from(Some(" 'fk-saved' "), &env_with_key, None).as_deref(), - Some("fk-saved") - ); - assert_eq!( - resolve_factory_api_key_from(None, &env_with_key, Some(" fk-store ")).as_deref(), - Some("fk-store") - ); - assert_eq!( - resolve_factory_api_key_from(None, &env_with_key, None).as_deref(), - Some("fk-env") - ); - assert_eq!( - resolve_factory_api_key_from(None, &env_dotenv_only, None).as_deref(), - Some("fk-dotenv") - ); - assert_eq!( - resolve_factory_api_key_from(None, &env_home_fallback, None).as_deref(), - Some("fk-dotenv") - ); - assert_eq!( - resolve_factory_api_key_from(None, &HashMap::new(), None), - None - ); - } - - #[test] - fn parses_billing_limits_fixture_json() { - let body = r#"{ - "usesTokenRateLimitsBilling": true, - "limits": { - "standard": { - "fiveHour": { "usedPercent": 12, "secondsRemaining": 3600 }, - "weekly": { "usedPercent": 34, "secondsRemaining": 86400 }, - "monthly": { "usedPercent": 56, "secondsRemaining": 604800 } - } - }, - "extraUsageBalanceCents": 0, - "extraUsageAllowed": false, - "tokenRateLimitsRolloutEligible": true - }"#; - let parsed: FactoryBillingLimitsResponse = serde_json::from_str(body).unwrap(); - assert!(parsed.uses_token_rate_limits_billing); - let limits = parsed.limits.unwrap(); - let snap = snapshot_from_billing_limits(&limits, None); - assert!((snap.primary.used_percent - 12.0).abs() < f64::EPSILON); - assert!((snap.secondary.unwrap().used_percent - 34.0).abs() < f64::EPSILON); - assert!((snap.tertiary.unwrap().used_percent - 56.0).abs() < f64::EPSILON); - } - - #[test] - fn parses_legacy_usage_fixture_json() { - let body = r#"{ - "standard": { "used": 25.0, "allowance": 100.0 }, - "premium": { "used": 10.0, "allowance": 50.0 } - }"#; - let parsed: FactoryUsageResponse = serde_json::from_str(body).unwrap(); - let snap = FactoryProvider::usage_snapshot_from_response(&parsed); - assert!((snap.primary.used_percent - 25.0).abs() < f64::EPSILON); - assert!((snap.secondary.unwrap().used_percent - 20.0).abs() < f64::EPSILON); - } - - #[test] - fn parses_nested_usage_fixture_json() { - let body = r#"{ - "usage": { - "standard": { "userTokens": 1200, "totalAllowance": 4000, "usedRatio": 0.3 }, - "premium": { "userTokens": 100, "totalAllowance": 1000, "usedRatio": 0.1 } - } - }"#; - let parsed: FactoryUsageResponse = serde_json::from_str(body).unwrap(); - let snap = FactoryProvider::usage_snapshot_from_response(&parsed); - assert!((snap.primary.used_percent - 30.0).abs() < f64::EPSILON); - assert!((snap.secondary.unwrap().used_percent - 10.0).abs() < f64::EPSILON); - } - - #[test] - fn empty_nested_usage_falls_through_to_top_level() { - let body = r#"{ - "usage": {}, - "standard": { "used": 40.0, "allowance": 100.0 }, - "premium": { "used": 5.0, "allowance": 50.0 } - }"#; - let parsed: FactoryUsageResponse = serde_json::from_str(body).unwrap(); - let snap = FactoryProvider::usage_snapshot_from_response(&parsed); - assert!((snap.primary.used_percent - 40.0).abs() < f64::EPSILON); - assert!((snap.secondary.unwrap().used_percent - 10.0).abs() < f64::EPSILON); - } - - #[test] - fn available_sources_do_not_advertise_cli() { - let sources = FactoryProvider::new().available_sources(); - assert!(!sources.contains(&SourceMode::Cli)); - assert!(sources.contains(&SourceMode::Auto)); - assert!(sources.contains(&SourceMode::OAuth)); - assert!(sources.contains(&SourceMode::Web)); - } - - #[test] - fn parses_auth_fixture_json() { - let body = r#"{ - "organization": { - "id": "org_1", - "name": "Acme", - "subscription": { - "factoryTier": "team", - "orbSubscription": { - "plan": { "name": "Team", "id": "plan_1" }, - "status": "active" - } - } - }, - "userProfile": { "id": "u1", "email": "user@example.com" } - }"#; - let auth: FactoryAuthResponse = serde_json::from_str(body).unwrap(); - let snap = - FactoryProvider::apply_auth_info(UsageSnapshot::new(RateWindow::new(0.0)), Some(auth)); - assert_eq!(snap.account_email.as_deref(), Some("user@example.com")); - assert_eq!(snap.account_organization.as_deref(), Some("Acme")); - assert!( - snap.login_method - .as_deref() - .is_some_and(|m| m.contains("team") || m.contains("Team")) - ); - } - - #[test] - fn auto_api_errors_are_recoverable() { - assert!(factory_api_error_is_recoverable( - &ProviderError::AuthRequired - )); - assert!(factory_api_error_is_recoverable(&ProviderError::Timeout)); - assert!(factory_api_error_is_recoverable(&ProviderError::Parse( - "bad json".into() - ))); - assert!(factory_api_error_is_recoverable(&ProviderError::Other( - "HTTP 500".into() - ))); - assert!(factory_api_error_is_recoverable( - &ProviderError::NotInstalled("missing".into()) - )); - assert!(!factory_api_error_is_recoverable( - &ProviderError::UnsupportedSource(SourceMode::Web) - )); - } - - #[test] - fn available_sources_include_auto_api_and_web() { - let sources = FactoryProvider::new().available_sources(); - assert!(sources.contains(&SourceMode::Auto)); - assert!(sources.contains(&SourceMode::OAuth)); // explicit API - assert!(sources.contains(&SourceMode::Web)); - } - - #[test] - fn secret_redactor_covers_factory_keys() { - let redacted = crate::core::SecretRedactor::redact("Factory key fk-test-key-abcdef"); - assert!( - !redacted.contains("fk-test-key"), - "factory key must not appear: {redacted}" - ); - } -} +mod tests; diff --git a/rust/src/providers/factory/tests.rs b/rust/src/providers/factory/tests.rs new file mode 100644 index 0000000000..f5948f1968 --- /dev/null +++ b/rust/src/providers/factory/tests.rs @@ -0,0 +1,226 @@ +use super::*; +use std::fs; + +#[test] +fn cleans_quotes_and_whitespace() { + assert_eq!( + clean_factory_secret(Some(" \"fk-quoted\" ")).as_deref(), + Some("fk-quoted") + ); + assert_eq!( + clean_factory_secret(Some("'fk-single'")).as_deref(), + Some("fk-single") + ); + assert_eq!(clean_factory_secret(Some(" ")), None); + assert_eq!(clean_factory_secret(None), None); +} + +#[test] +fn parses_factory_dotenv_variants() { + assert_eq!( + parse_factory_dotenv_key("FACTORY_API_KEY=fk-plain").as_deref(), + Some("fk-plain") + ); + assert_eq!( + parse_factory_dotenv_key("export FACTORY_API_KEY='fk-single'").as_deref(), + Some("fk-single") + ); + assert_eq!( + parse_factory_dotenv_key("# comment\nFACTORY_API_KEY=\"fk-double\"").as_deref(), + Some("fk-double") + ); + assert_eq!(parse_factory_dotenv_key("OTHER=1\n"), None); + assert_eq!( + parse_factory_dotenv_key( + "export FACTORY_API_KEY='fk-quoted'\nMALFORMED\nFACTORY_API_KEY=\n" + ) + .as_deref(), + Some("fk-quoted") + ); + // Empty assignment must not short-circuit a later real key. + assert_eq!( + parse_factory_dotenv_key("FACTORY_API_KEY=\nFACTORY_API_KEY=fk-real\n").as_deref(), + Some("fk-real") + ); + assert_eq!( + parse_factory_dotenv_key("FACTORY_API_KEY=\"\"\nFACTORY_API_KEY='fk-later'\n").as_deref(), + Some("fk-later") + ); +} + +#[test] +fn key_resolution_honors_explicit_env_dotenv_precedence() { + let home = tempfile::tempdir().unwrap(); + fs::create_dir(home.path().join(".factory")).unwrap(); + fs::write( + home.path().join(".factory").join(".env"), + "FACTORY_API_KEY=fk-dotenv\n", + ) + .unwrap(); + + let env_with_key = HashMap::from([ + ( + FACTORY_API_KEY_ENV.to_string(), + " \"fk-env\" ".to_string(), + ), + ("USERPROFILE".to_string(), home.path().display().to_string()), + ]); + let env_dotenv_only = + HashMap::from([("USERPROFILE".to_string(), home.path().display().to_string())]); + let env_home_fallback = + HashMap::from([("HOME".to_string(), home.path().display().to_string())]); + + assert_eq!( + resolve_factory_api_key_from(Some(" 'fk-saved' "), &env_with_key, None).as_deref(), + Some("fk-saved") + ); + assert_eq!( + resolve_factory_api_key_from(None, &env_with_key, Some(" fk-store ")).as_deref(), + Some("fk-store") + ); + assert_eq!( + resolve_factory_api_key_from(None, &env_with_key, None).as_deref(), + Some("fk-env") + ); + assert_eq!( + resolve_factory_api_key_from(None, &env_dotenv_only, None).as_deref(), + Some("fk-dotenv") + ); + assert_eq!( + resolve_factory_api_key_from(None, &env_home_fallback, None).as_deref(), + Some("fk-dotenv") + ); + assert_eq!( + resolve_factory_api_key_from(None, &HashMap::new(), None), + None + ); +} + +#[test] +fn parses_billing_limits_fixture_json() { + let body = r#"{ + "usesTokenRateLimitsBilling": true, + "limits": { + "standard": { + "fiveHour": { "usedPercent": 12, "secondsRemaining": 3600 }, + "weekly": { "usedPercent": 34, "secondsRemaining": 86400 }, + "monthly": { "usedPercent": 56, "secondsRemaining": 604800 } + } + }, + "extraUsageBalanceCents": 0, + "extraUsageAllowed": false, + "tokenRateLimitsRolloutEligible": true + }"#; + let parsed: FactoryBillingLimitsResponse = serde_json::from_str(body).unwrap(); + assert!(parsed.uses_token_rate_limits_billing); + let limits = parsed.limits.unwrap(); + let snap = snapshot_from_billing_limits(&limits, None); + assert!((snap.primary.used_percent - 12.0).abs() < f64::EPSILON); + assert!((snap.secondary.unwrap().used_percent - 34.0).abs() < f64::EPSILON); + assert!((snap.tertiary.unwrap().used_percent - 56.0).abs() < f64::EPSILON); +} + +#[test] +fn parses_legacy_usage_fixture_json() { + let body = r#"{ + "standard": { "used": 25.0, "allowance": 100.0 }, + "premium": { "used": 10.0, "allowance": 50.0 } + }"#; + let parsed: FactoryUsageResponse = serde_json::from_str(body).unwrap(); + let snap = FactoryProvider::usage_snapshot_from_response(&parsed); + assert!((snap.primary.used_percent - 25.0).abs() < f64::EPSILON); + assert!((snap.secondary.unwrap().used_percent - 20.0).abs() < f64::EPSILON); +} + +#[test] +fn parses_nested_usage_fixture_json() { + let body = r#"{ + "usage": { + "standard": { "userTokens": 1200, "totalAllowance": 4000, "usedRatio": 0.3 }, + "premium": { "userTokens": 100, "totalAllowance": 1000, "usedRatio": 0.1 } + } + }"#; + let parsed: FactoryUsageResponse = serde_json::from_str(body).unwrap(); + let snap = FactoryProvider::usage_snapshot_from_response(&parsed); + assert!((snap.primary.used_percent - 30.0).abs() < f64::EPSILON); + assert!((snap.secondary.unwrap().used_percent - 10.0).abs() < f64::EPSILON); +} + +#[test] +fn empty_nested_usage_falls_through_to_top_level() { + let body = r#"{ + "usage": {}, + "standard": { "used": 40.0, "allowance": 100.0 }, + "premium": { "used": 5.0, "allowance": 50.0 } + }"#; + let parsed: FactoryUsageResponse = serde_json::from_str(body).unwrap(); + let snap = FactoryProvider::usage_snapshot_from_response(&parsed); + assert!((snap.primary.used_percent - 40.0).abs() < f64::EPSILON); + assert!((snap.secondary.unwrap().used_percent - 10.0).abs() < f64::EPSILON); +} + +#[test] +fn available_sources_do_not_advertise_cli() { + let sources = FactoryProvider::new().available_sources(); + assert!(!sources.contains(&SourceMode::Cli)); + assert!(sources.contains(&SourceMode::Auto)); + assert!(sources.contains(&SourceMode::OAuth)); + assert!(sources.contains(&SourceMode::Web)); +} + +#[test] +fn parses_auth_fixture_json() { + let body = r#"{ + "organization": { + "id": "org_1", + "name": "Acme", + "subscription": { + "factoryTier": "team", + "orbSubscription": { + "plan": { "name": "Team", "id": "plan_1" }, + "status": "active" + } + } + }, + "userProfile": { "id": "u1", "email": "user@example.com" } + }"#; + let auth: FactoryAuthResponse = serde_json::from_str(body).unwrap(); + let snap = + FactoryProvider::apply_auth_info(UsageSnapshot::new(RateWindow::new(0.0)), Some(auth)); + assert_eq!(snap.account_email.as_deref(), Some("user@example.com")); + assert_eq!(snap.account_organization.as_deref(), Some("Acme")); + assert!( + snap.login_method + .as_deref() + .is_some_and(|m| m.contains("team") || m.contains("Team")) + ); +} + +#[test] +fn auto_api_errors_are_recoverable() { + assert!(factory_api_error_is_recoverable( + &ProviderError::AuthRequired + )); + assert!(factory_api_error_is_recoverable(&ProviderError::Timeout)); + assert!(factory_api_error_is_recoverable(&ProviderError::Parse( + "bad json".into() + ))); + assert!(factory_api_error_is_recoverable(&ProviderError::Other( + "HTTP 500".into() + ))); + assert!(factory_api_error_is_recoverable( + &ProviderError::NotInstalled("missing".into()) + )); + assert!(!factory_api_error_is_recoverable( + &ProviderError::UnsupportedSource(SourceMode::Web) + )); +} + +#[test] +fn secret_redactor_covers_factory_keys() { + let redacted = crate::core::SecretRedactor::redact("Factory key fk-test-key-abcdef"); + assert!( + !redacted.contains("fk-test-key"), + "factory key must not appear: {redacted}" + ); +} diff --git a/rust/src/providers/gemini/api.rs b/rust/src/providers/gemini/api.rs index 25a43abb6d..0d88e5e61b 100755 --- a/rust/src/providers/gemini/api.rs +++ b/rust/src/providers/gemini/api.rs @@ -2,7 +2,7 @@ //! //! Uses Google Cloud Code Private API with OAuth tokens from ~/.gemini/oauth_creds.json -use crate::core::{FetchContext, ProviderError, RateWindow}; +use crate::core::{FetchContext, ProviderError, RateWindow, UsageSnapshot}; use chrono::{DateTime, Utc}; use serde::{Deserialize, Serialize}; use std::path::{Path, PathBuf}; @@ -26,20 +26,8 @@ impl GeminiApi { } /// Fetch quota information from the Gemini API - /// Returns (primary RateWindow, optional model-specific RateWindow, optional email, optional plan) /// Note: Gemini quota API requires OAuth tokens, not API keys - pub async fn fetch_quota( - &self, - _ctx: &FetchContext, - ) -> Result< - ( - RateWindow, - Option, - Option, - Option, - ), - ProviderError, - > { + pub async fn fetch_quota(&self, _ctx: &FetchContext) -> Result { // Gemini quota endpoint requires OAuth credentials (not API keys) // Always load OAuth credentials from ~/.gemini/oauth_creds.json let mut creds = self.load_credentials()?; @@ -99,7 +87,14 @@ impl GeminiApi { self.parse_quota_response(quota_response, Some(&creds))?; let plan = resolve_account_plan(&code_assist, hosted_domain.as_deref()); - Ok((primary, model_specific, email, plan)) + let mut usage = UsageSnapshot::new(primary); + if let Some(ms) = model_specific { + usage = usage.with_model_specific(ms); + } + if let Some(e) = email { + usage = usage.with_email(e); + } + Ok(usage.with_login_method(plan.unwrap_or_else(|| "Gemini CLI".to_string()))) } async fn load_code_assist_status(&self, access_token: &str) -> CodeAssistStatus { @@ -380,26 +375,12 @@ impl GeminiApi { None } - #[cfg(windows)] fn fnm_oauth_credentials() -> Option { #[cfg(windows)] - if let Some(local_appdata) = dirs::data_local_dir() { - let fnm_versions = local_appdata.join("fnm").join("node-versions"); - return Self::fnm_oauth_credentials_from(&fnm_versions); - } - - None - } - - #[cfg(not(windows))] - fn fnm_oauth_credentials() -> Option { + let fnm_root = dirs::data_local_dir()?; #[cfg(not(windows))] - if let Some(data_dir) = dirs::data_dir() { - let fnm_versions = data_dir.join("fnm").join("node-versions"); - return Self::fnm_oauth_credentials_from(&fnm_versions); - } - - None + let fnm_root = dirs::data_dir()?; + Self::fnm_oauth_credentials_from(&fnm_root.join("fnm").join("node-versions")) } fn fnm_oauth_credentials_from(fnm_versions: &Path) -> Option { @@ -489,56 +470,38 @@ impl GeminiApi { } } - // Find Flash and Pro quotas - let flash_quota = model_quotas - .iter() - .filter(|(k, _)| k.to_lowercase().contains("flash")) - .min_by(|a, b| { - a.1.0 - .partial_cmp(&b.1.0) - .unwrap_or(std::cmp::Ordering::Equal) - }); - - let pro_quota = model_quotas - .iter() - .filter(|(k, _)| k.to_lowercase().contains("pro")) - .min_by(|a, b| { - a.1.0 - .partial_cmp(&b.1.0) - .unwrap_or(std::cmp::Ordering::Equal) - }); - - // Build primary RateWindow from the most constrained quota - let (primary_fraction, primary_reset) = if let Some((_, (frac, reset))) = pro_quota { - (*frac, reset.clone()) - } else if let Some((_, (frac, reset))) = flash_quota { - (*frac, reset.clone()) - } else if let Some((_, (frac, reset))) = model_quotas.iter().next() { - (*frac, reset.clone()) - } else { - (1.0, None) + // Most constrained Flash / Pro quota. + let lowest = |family: &str| { + model_quotas + .iter() + .filter(|(k, _)| k.to_lowercase().contains(family)) + .min_by(|a, b| { + a.1.0 + .partial_cmp(&b.1.0) + .unwrap_or(std::cmp::Ordering::Equal) + }) }; - - let primary_percent_used = (1.0 - primary_fraction) * 100.0; - let primary_reset_at = primary_reset.as_ref().and_then(|s| parse_iso_date(s)); - - let primary = RateWindow::with_details( - primary_percent_used, - Some(1440), // 24 hours - primary_reset_at, - None, - ); - - // Model-specific window for Flash if Pro is primary - let model_specific = if pro_quota.is_some() { - flash_quota.map(|(_, (frac, reset))| { - let percent_used = (1.0 - frac) * 100.0; - let reset_at = reset.as_ref().and_then(|s| parse_iso_date(s)); - RateWindow::with_details(percent_used, Some(1440), reset_at, None) - }) - } else { - None + let flash_quota = lowest("flash"); + let pro_quota = lowest("pro"); + + let window = |(_, (fraction, reset)): (&String, &(f64, Option))| { + RateWindow::with_details( + (1.0 - fraction) * 100.0, + Some(1440), // 24 hours + reset.as_ref().and_then(|s| parse_iso_date(s)), + None, + ) }; + // Primary is Pro, else Flash, else any model; Flash gets its own + // window only when Pro is primary. + let primary = pro_quota + .or(flash_quota) + .or_else(|| model_quotas.iter().next()) + .map_or_else( + || RateWindow::with_details(0.0, Some(1440), None, None), + window, + ); + let model_specific = pro_quota.and(flash_quota).map(window); // Extract email from ID token let email = creds @@ -693,17 +656,9 @@ fn resolve_account_plan(status: &CodeAssistStatus, hosted_domain: Option<&str>) // --- Helper functions --- fn parse_iso_date(s: &str) -> Option> { - // Try with fractional seconds first - if let Ok(dt) = DateTime::parse_from_rfc3339(s) { - return Some(dt.with_timezone(&Utc)); - } - - // Try without fractional seconds - if let Ok(dt) = chrono::DateTime::parse_from_str(s, "%Y-%m-%dT%H:%M:%SZ") { - return Some(dt.with_timezone(&Utc)); - } - - None + DateTime::parse_from_rfc3339(s) + .ok() + .map(|dt| dt.with_timezone(&Utc)) } fn extract_email_from_jwt(token: &str) -> Option { @@ -742,209 +697,4 @@ fn jwt_payload(token: &str) -> Option { } #[cfg(test)] -mod tests { - use super::*; - - #[test] - fn bundled_cli_layout_yields_oauth_client_credentials() { - // npm global layout on Windows: %APPDATA%\npm\gemini.cmd next to - // node_modules\@google\gemini-cli\bundle\chunk-*.js (no gemini-cli-core/dist). - let dir = tempfile::tempdir().unwrap(); - let bin_dir = dir.path(); - let bundle = bin_dir - .join("node_modules") - .join("@google") - .join("gemini-cli") - .join("bundle"); - std::fs::create_dir_all(&bundle).unwrap(); - std::fs::write(bundle.join("chunk-AAA.js"), "var x = 1;").unwrap(); - std::fs::write( - bundle.join("chunk-BBB.js"), - r#"var OAUTH_CLIENT_ID = "id-123.apps.googleusercontent.com"; var OAUTH_CLIENT_SECRET = 'secret-xyz';"#, - ) - .unwrap(); - - assert!( - GeminiApi::oauth_credentials_from_candidates(GeminiApi::binary_oauth_candidates( - bin_dir - )) - .is_none(), - "legacy dist layout must not match" - ); - let creds = GeminiApi::bundled_cli_oauth_credentials(bin_dir) - .expect("bundle chunks should be scanned"); - assert_eq!(creds.client_id, "id-123.apps.googleusercontent.com"); - assert_eq!(creds.client_secret, "secret-xyz"); - } - - const BUNDLE_CHUNK_WITH_CONSTANTS: &str = r#"var OAUTH_CLIENT_ID = "id-456.apps.googleusercontent.com"; var OAUTH_CLIENT_SECRET = "secret-abc";"#; - - fn write_gemini_bundle(node_modules: &Path) -> PathBuf { - let bundle = node_modules - .join("@google") - .join("gemini-cli") - .join("bundle"); - std::fs::create_dir_all(&bundle).unwrap(); - std::fs::write(bundle.join("gemini.js"), "import './chunk-A.js';").unwrap(); - std::fs::write(bundle.join("chunk-A.js"), BUNDLE_CHUNK_WITH_CONSTANTS).unwrap(); - bundle - } - - #[test] - fn symlinked_binary_inside_bundle_yields_oauth_client_credentials() { - // Unix npm/Homebrew: bin/gemini canonicalizes to .../gemini-cli/bundle/gemini.js. - let dir = tempfile::tempdir().unwrap(); - let bundle = write_gemini_bundle(&dir.path().join("lib").join("node_modules")); - - let creds = GeminiApi::bundled_cli_oauth_credentials(&bundle) - .expect("the bundle that holds the binary should be scanned"); - assert_eq!(creds.client_id, "id-456.apps.googleusercontent.com"); - assert_eq!(creds.client_secret, "secret-abc"); - } - - #[test] - fn unrelated_bundle_directory_is_not_scanned() { - let dir = tempfile::tempdir().unwrap(); - let other = dir.path().join("other-tool").join("bundle"); - std::fs::create_dir_all(&other).unwrap(); - std::fs::write(other.join("chunk.js"), BUNDLE_CHUNK_WITH_CONSTANTS).unwrap(); - - assert!(GeminiApi::bundled_cli_oauth_credentials(&other).is_none()); - } - - #[test] - fn npm_global_node_modules_bundle_yields_oauth_client_credentials() { - // %APPDATA%\npm\node_modules fallback when `gemini` is not on PATH. - let dir = tempfile::tempdir().unwrap(); - let node_modules = dir.path().join("npm").join("node_modules"); - write_gemini_bundle(&node_modules); - - let creds = GeminiApi::node_modules_oauth_credentials(&node_modules) - .expect("bundle under the npm global node_modules should be scanned"); - assert_eq!(creds.client_secret, "secret-abc"); - } - - #[test] - fn fnm_windows_and_unix_layouts_yield_bundle_credentials() { - let dir = tempfile::tempdir().unwrap(); - let versions = dir.path().join("node-versions"); - // Windows fnm keeps global packages directly under installation\node_modules. - write_gemini_bundle( - &versions - .join("v22.0.0") - .join("installation") - .join("node_modules"), - ); - let creds = GeminiApi::fnm_oauth_credentials_from(&versions) - .expect("Windows fnm layout should be scanned"); - assert_eq!(creds.client_id, "id-456.apps.googleusercontent.com"); - - let unix_dir = tempfile::tempdir().unwrap(); - let unix_versions = unix_dir.path().join("node-versions"); - write_gemini_bundle( - &unix_versions - .join("v22.0.0") - .join("installation") - .join("lib") - .join("node_modules"), - ); - assert!(GeminiApi::fnm_oauth_credentials_from(&unix_versions).is_some()); - } - - #[test] - fn paid_tier_name_overrides_generic_tier_fallbacks() { - let status = parse_code_assist_status( - r#"{ - "currentTier": { "id": "free-tier" }, - "paidTier": { "name": "Gemini Code Assist in Google One AI Pro" } - }"#, - ); - - assert_eq!( - resolve_account_plan(&status, Some("example.com")), - Some("Gemini Code Assist in Google One AI Pro".to_string()) - ); - - let standard = parse_code_assist_status( - r#"{ - "currentTier": { "id": "standard-tier" }, - "paidTier": { "name": "Plus" } - }"#, - ); - - assert_eq!( - resolve_account_plan(&standard, None), - Some("Plus".to_string()) - ); - } - - #[test] - fn consumer_shutdown_signal_excludes_paid_and_workspace_accounts() { - let shutdown = parse_code_assist_status( - r#"{ - "ineligibleTiers": [ - {"tier":{"id":"free-tier"},"reason":"UNSUPPORTED_CLIENT"} - ] - }"#, - ); - assert!(is_consumer_client_unsupported(&shutdown, None)); - assert!(!is_consumer_client_unsupported( - &shutdown, - Some("example.com") - )); - - let paid = parse_code_assist_status( - r#"{ - "paidTier":{"name":"Gemini Code Assist Standard"}, - "ineligibleTiers":[ - {"tier":{"id":"free-tier"},"reason":"UNSUPPORTED_CLIENT"} - ] - }"#, - ); - assert!(!is_consumer_client_unsupported(&paid, None)); - - let standard = parse_code_assist_status( - r#"{ - "currentTier":{"id":"standard-tier"}, - "ineligibleTiers":[ - {"tier":{"id":"free-tier"},"reason":"UNSUPPORTED_CLIENT"} - ] - }"#, - ); - assert!(!is_consumer_client_unsupported(&standard, None)); - } - - #[test] - fn generic_tier_fallbacks_remain_when_paid_tier_is_absent() { - let free_tier = parse_code_assist_status(r#"{"currentTier":{"id":"free-tier"}}"#); - let paid = parse_code_assist_status(r#"{"currentTier":{"id":"standard-tier"}}"#); - - assert_eq!( - resolve_account_plan(&free_tier, Some("example.com")), - Some("Workspace".to_string()) - ); - assert_eq!( - resolve_account_plan(&free_tier, None), - Some("Free".to_string()) - ); - assert_eq!(resolve_account_plan(&paid, None), Some("Paid".to_string())); - } - - #[test] - fn invalid_code_assist_response_does_not_create_a_generic_plan() { - let status = parse_code_assist_status("not json"); - - assert_eq!(resolve_account_plan(&status, Some("example.com")), None); - } - - #[test] - fn malformed_paid_tier_preserves_current_tier_fallback() { - let status = - parse_code_assist_status(r#"{"currentTier":{"id":"free-tier"},"paidTier":[]}"#); - - assert_eq!( - resolve_account_plan(&status, Some("example.com")), - Some("Workspace".to_string()) - ); - } -} +mod tests; diff --git a/rust/src/providers/gemini/api/tests.rs b/rust/src/providers/gemini/api/tests.rs new file mode 100644 index 0000000000..49bd5a4f5f --- /dev/null +++ b/rust/src/providers/gemini/api/tests.rs @@ -0,0 +1,343 @@ +use super::*; + +fn bucket(model: Option<&str>, fraction: Option, reset: Option<&str>) -> QuotaBucket { + QuotaBucket { + remaining_fraction: fraction, + reset_time: reset.map(str::to_string), + model_id: model.map(str::to_string), + token_type: None, + } +} + +fn parse_buckets( + buckets: Vec, + creds: Option<&OAuthCredentials>, +) -> Result<(RateWindow, Option, Option), ProviderError> { + GeminiApi::new().parse_quota_response( + QuotaResponse { + buckets: Some(buckets), + }, + creds, + ) +} + +fn at(rfc3339: &str) -> Option> { + Some( + DateTime::parse_from_rfc3339(rfc3339) + .unwrap() + .with_timezone(&Utc), + ) +} + +#[test] +fn quota_pro_is_primary_and_flash_is_model_specific() { + use base64::Engine; + let payload = + base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(br#"{"email":"user@example.com"}"#); + let creds = OAuthCredentials { + access_token: None, + id_token: Some(format!("header.{payload}.sig")), + refresh_token: None, + expiry_date: None, + }; + let (primary, model_specific, email) = parse_buckets( + vec![ + bucket( + Some("gemini-2.5-pro"), + Some(0.6), + Some("2026-01-14T00:00:00Z"), + ), + bucket( + Some("gemini-2.5-pro"), + Some(0.4), + Some("2026-01-15T00:00:00Z"), + ), + bucket( + Some("gemini-2.5-flash"), + Some(0.9), + Some("2026-01-16T00:00:00.5Z"), + ), + bucket(Some("gemini-2.0-flash"), Some(0.95), None), + bucket(None, Some(0.0), None), + bucket(Some("gemini-2.5-pro"), None, None), + ], + Some(&creds), + ) + .unwrap(); + assert_eq!(primary.used_percent, (1.0 - 0.4) * 100.0); + assert_eq!(primary.window_minutes, Some(1440)); + assert_eq!(primary.resets_at, at("2026-01-15T00:00:00Z")); + assert_eq!(primary.reset_description, None); + let flash = model_specific.expect("flash window when pro is primary"); + assert_eq!(flash.used_percent, (1.0 - 0.9) * 100.0); + assert_eq!(flash.window_minutes, Some(1440)); + assert_eq!(flash.resets_at, at("2026-01-16T00:00:00.5Z")); + assert_eq!(email.as_deref(), Some("user@example.com")); +} + +#[test] +fn quota_falls_back_to_flash_then_any_model() { + let (primary, model_specific, email) = parse_buckets( + vec![ + bucket(Some("gemini-2.5-flash"), Some(0.3), Some("not a date")), + bucket(Some("other-model"), Some(0.1), None), + ], + None, + ) + .unwrap(); + assert_eq!(primary.used_percent, (1.0 - 0.3) * 100.0); + assert_eq!(primary.resets_at, None); + assert!(model_specific.is_none()); + assert_eq!(email, None); + + let (primary, model_specific, _) = parse_buckets( + vec![bucket( + Some("other-model"), + Some(0.25), + Some("2026-01-15T00:00:00+02:00"), + )], + None, + ) + .unwrap(); + assert_eq!(primary.used_percent, (1.0 - 0.25) * 100.0); + assert_eq!(primary.resets_at, at("2026-01-14T22:00:00Z")); + assert!(model_specific.is_none()); +} + +#[test] +fn quota_without_usable_model_buckets_reports_an_unused_window() { + // A fraction of 1.0 never beats the per-model starting value, so its + // reset time is dropped. + let (primary, model_specific, _) = parse_buckets( + vec![ + bucket( + Some("gemini-2.5-pro"), + Some(1.0), + Some("2026-01-15T00:00:00Z"), + ), + bucket(None, Some(0.2), Some("2026-01-15T00:00:00Z")), + ], + None, + ) + .unwrap(); + assert_eq!(primary.used_percent, 0.0); + assert_eq!(primary.resets_at, None); + assert!(model_specific.is_none()); + + let (primary, model_specific, _) = + parse_buckets(vec![bucket(None, Some(0.2), None)], None).unwrap(); + assert_eq!(primary.used_percent, 0.0); + assert_eq!(primary.window_minutes, Some(1440)); + assert_eq!(primary.resets_at, None); + assert!(model_specific.is_none()); +} + +#[test] +fn quota_without_buckets_is_a_parse_error() { + let empty = parse_buckets(Vec::new(), None).unwrap_err(); + assert!(matches!(empty, ProviderError::Parse(msg) if msg == "Empty quota buckets")); + let missing = GeminiApi::new() + .parse_quota_response(QuotaResponse { buckets: None }, None) + .unwrap_err(); + assert!(matches!(missing, ProviderError::Parse(msg) if msg == "No quota buckets in response")); +} + +#[test] +fn bundled_cli_layout_yields_oauth_client_credentials() { + // npm global layout on Windows: %APPDATA%\npm\gemini.cmd next to + // node_modules\@google\gemini-cli\bundle\chunk-*.js (no gemini-cli-core/dist). + let dir = tempfile::tempdir().unwrap(); + let bin_dir = dir.path(); + let bundle = bin_dir + .join("node_modules") + .join("@google") + .join("gemini-cli") + .join("bundle"); + std::fs::create_dir_all(&bundle).unwrap(); + std::fs::write(bundle.join("chunk-AAA.js"), "var x = 1;").unwrap(); + std::fs::write( + bundle.join("chunk-BBB.js"), + r#"var OAUTH_CLIENT_ID = "id-123.apps.googleusercontent.com"; var OAUTH_CLIENT_SECRET = 'secret-xyz';"#, + ) + .unwrap(); + + assert!( + GeminiApi::oauth_credentials_from_candidates(GeminiApi::binary_oauth_candidates(bin_dir)) + .is_none(), + "legacy dist layout must not match" + ); + let creds = + GeminiApi::bundled_cli_oauth_credentials(bin_dir).expect("bundle chunks should be scanned"); + assert_eq!(creds.client_id, "id-123.apps.googleusercontent.com"); + assert_eq!(creds.client_secret, "secret-xyz"); +} + +const BUNDLE_CHUNK_WITH_CONSTANTS: &str = r#"var OAUTH_CLIENT_ID = "id-456.apps.googleusercontent.com"; var OAUTH_CLIENT_SECRET = "secret-abc";"#; + +fn write_gemini_bundle(node_modules: &Path) -> PathBuf { + let bundle = node_modules + .join("@google") + .join("gemini-cli") + .join("bundle"); + std::fs::create_dir_all(&bundle).unwrap(); + std::fs::write(bundle.join("gemini.js"), "import './chunk-A.js';").unwrap(); + std::fs::write(bundle.join("chunk-A.js"), BUNDLE_CHUNK_WITH_CONSTANTS).unwrap(); + bundle +} + +#[test] +fn symlinked_binary_inside_bundle_yields_oauth_client_credentials() { + // Unix npm/Homebrew: bin/gemini canonicalizes to .../gemini-cli/bundle/gemini.js. + let dir = tempfile::tempdir().unwrap(); + let bundle = write_gemini_bundle(&dir.path().join("lib").join("node_modules")); + + let creds = GeminiApi::bundled_cli_oauth_credentials(&bundle) + .expect("the bundle that holds the binary should be scanned"); + assert_eq!(creds.client_id, "id-456.apps.googleusercontent.com"); + assert_eq!(creds.client_secret, "secret-abc"); +} + +#[test] +fn unrelated_bundle_directory_is_not_scanned() { + let dir = tempfile::tempdir().unwrap(); + let other = dir.path().join("other-tool").join("bundle"); + std::fs::create_dir_all(&other).unwrap(); + std::fs::write(other.join("chunk.js"), BUNDLE_CHUNK_WITH_CONSTANTS).unwrap(); + + assert!(GeminiApi::bundled_cli_oauth_credentials(&other).is_none()); +} + +#[test] +fn npm_global_node_modules_bundle_yields_oauth_client_credentials() { + // %APPDATA%\npm\node_modules fallback when `gemini` is not on PATH. + let dir = tempfile::tempdir().unwrap(); + let node_modules = dir.path().join("npm").join("node_modules"); + write_gemini_bundle(&node_modules); + + let creds = GeminiApi::node_modules_oauth_credentials(&node_modules) + .expect("bundle under the npm global node_modules should be scanned"); + assert_eq!(creds.client_secret, "secret-abc"); +} + +#[test] +fn fnm_windows_and_unix_layouts_yield_bundle_credentials() { + let dir = tempfile::tempdir().unwrap(); + let versions = dir.path().join("node-versions"); + // Windows fnm keeps global packages directly under installation\node_modules. + write_gemini_bundle( + &versions + .join("v22.0.0") + .join("installation") + .join("node_modules"), + ); + let creds = GeminiApi::fnm_oauth_credentials_from(&versions) + .expect("Windows fnm layout should be scanned"); + assert_eq!(creds.client_id, "id-456.apps.googleusercontent.com"); + + let unix_dir = tempfile::tempdir().unwrap(); + let unix_versions = unix_dir.path().join("node-versions"); + write_gemini_bundle( + &unix_versions + .join("v22.0.0") + .join("installation") + .join("lib") + .join("node_modules"), + ); + assert!(GeminiApi::fnm_oauth_credentials_from(&unix_versions).is_some()); +} + +#[test] +fn paid_tier_name_overrides_generic_tier_fallbacks() { + let status = parse_code_assist_status( + r#"{ + "currentTier": { "id": "free-tier" }, + "paidTier": { "name": "Gemini Code Assist in Google One AI Pro" } + }"#, + ); + + assert_eq!( + resolve_account_plan(&status, Some("example.com")), + Some("Gemini Code Assist in Google One AI Pro".to_string()) + ); + + let standard = parse_code_assist_status( + r#"{ + "currentTier": { "id": "standard-tier" }, + "paidTier": { "name": "Plus" } + }"#, + ); + + assert_eq!( + resolve_account_plan(&standard, None), + Some("Plus".to_string()) + ); +} + +#[test] +fn consumer_shutdown_signal_excludes_paid_and_workspace_accounts() { + let shutdown = parse_code_assist_status( + r#"{ + "ineligibleTiers": [ + {"tier":{"id":"free-tier"},"reason":"UNSUPPORTED_CLIENT"} + ] + }"#, + ); + assert!(is_consumer_client_unsupported(&shutdown, None)); + assert!(!is_consumer_client_unsupported( + &shutdown, + Some("example.com") + )); + + let paid = parse_code_assist_status( + r#"{ + "paidTier":{"name":"Gemini Code Assist Standard"}, + "ineligibleTiers":[ + {"tier":{"id":"free-tier"},"reason":"UNSUPPORTED_CLIENT"} + ] + }"#, + ); + assert!(!is_consumer_client_unsupported(&paid, None)); + + let standard = parse_code_assist_status( + r#"{ + "currentTier":{"id":"standard-tier"}, + "ineligibleTiers":[ + {"tier":{"id":"free-tier"},"reason":"UNSUPPORTED_CLIENT"} + ] + }"#, + ); + assert!(!is_consumer_client_unsupported(&standard, None)); +} + +#[test] +fn generic_tier_fallbacks_remain_when_paid_tier_is_absent() { + let free_tier = parse_code_assist_status(r#"{"currentTier":{"id":"free-tier"}}"#); + let paid = parse_code_assist_status(r#"{"currentTier":{"id":"standard-tier"}}"#); + + assert_eq!( + resolve_account_plan(&free_tier, Some("example.com")), + Some("Workspace".to_string()) + ); + assert_eq!( + resolve_account_plan(&free_tier, None), + Some("Free".to_string()) + ); + assert_eq!(resolve_account_plan(&paid, None), Some("Paid".to_string())); +} + +#[test] +fn invalid_code_assist_response_does_not_create_a_generic_plan() { + let status = parse_code_assist_status("not json"); + + assert_eq!(resolve_account_plan(&status, Some("example.com")), None); +} + +#[test] +fn malformed_paid_tier_preserves_current_tier_fallback() { + let status = parse_code_assist_status(r#"{"currentTier":{"id":"free-tier"},"paidTier":[]}"#); + + assert_eq!( + resolve_account_plan(&status, Some("example.com")), + Some("Workspace".to_string()) + ); +} diff --git a/rust/src/providers/gemini/mod.rs b/rust/src/providers/gemini/mod.rs index 0554652f8c..d17be8a53f 100755 --- a/rust/src/providers/gemini/mod.rs +++ b/rust/src/providers/gemini/mod.rs @@ -9,7 +9,6 @@ use async_trait::async_trait; use crate::core::{ FetchContext, Provider, ProviderError, ProviderFetchResult, ProviderId, SourceMode, - UsageSnapshot, }; pub use api::GeminiApi; @@ -43,18 +42,7 @@ impl Provider for GeminiProvider { tracing::debug!("Fetching Gemini usage via API"); match self.api.fetch_quota(ctx).await { - Ok((primary, model_specific, email, plan)) => { - let mut usage = UsageSnapshot::new(primary); - if let Some(ms) = model_specific { - usage = usage.with_model_specific(ms); - } - if let Some(e) = email { - usage = usage.with_email(e); - } - usage = usage.with_login_method(plan.unwrap_or_else(|| "Gemini CLI".to_string())); - - Ok(ProviderFetchResult::new(usage, "cli")) - } + Ok(usage) => Ok(ProviderFetchResult::new(usage, "cli")), Err(e) => { tracing::warn!("Gemini API fetch failed: {}", e); Err(e) diff --git a/rust/src/providers/infini.rs b/rust/src/providers/infini.rs index ce6d32cdc2..c723414ba5 100644 --- a/rust/src/providers/infini.rs +++ b/rust/src/providers/infini.rs @@ -279,6 +279,106 @@ mod tests { assert_eq!(usage.seven_day_percentage(), 0.0); assert_eq!(usage.thirty_day_percentage(), 0.0); } + + // Cases moved from the never-compiled rust/tests/providers/test_infini.rs. + + async fn usage_server(status: usize, body: &str) -> (mockito::ServerGuard, mockito::Mock) { + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("GET", "/maas/coding/usage") + .with_status(status) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + (server, mock) + } + + #[tokio::test] + async fn client_fetch_usage_success() { + let (server, mock) = usage_server( + 200, + r#"{ + "5_hour": {"quota": 5000, "used": 1000, "remain": 4000}, + "7_day": {"quota": 30000, "used": 5000, "remain": 25000}, + "30_day": {"quota": 60000, "used": 10000, "remain": 50000} + }"#, + ) + .await; + let client = InfiniClient::new("sk-cp-test-key".to_string()).with_base_url(server.url()); + + let usage = client.fetch_usage().await.unwrap(); + + mock.assert_async().await; + assert_eq!(usage.five_hour.quota, 5000); + assert_eq!(usage.seven_day.used, 5000); + } + + #[tokio::test] + async fn client_fetch_usage_unauthorized() { + let (server, mock) = usage_server(401, "").await; + let client = InfiniClient::new("invalid-key".to_string()).with_base_url(server.url()); + + let result = client.fetch_usage().await; + + mock.assert_async().await; + assert!(matches!(result, Err(InfiniError::Unauthorized))); + } + + /// One provider instance covers the four single-assert identity tests + /// (id, metadata, sources, supports_web) from the old file. + #[test] + fn provider_identity_and_sources() { + let provider = InfiniProvider::new("sk-cp-test".to_string()); + assert_eq!(provider.id(), ProviderId::Infini); + let meta = provider.metadata(); + assert_eq!(meta.id, ProviderId::Infini); + assert_eq!(meta.display_name, "Infini"); + let sources = provider.available_sources(); + assert!(sources.contains(&SourceMode::Auto)); + assert!(sources.contains(&SourceMode::Web)); + assert!(provider.supports_web()); + } + + #[tokio::test] + async fn provider_fetch_usage_success() { + let (server, mock) = usage_server( + 200, + r#"{ + "5_hour": {"quota": 5000, "used": 2500, "remain": 2500}, + "7_day": {"quota": 30000, "used": 15000, "remain": 15000}, + "30_day": {"quota": 60000, "used": 30000, "remain": 30000} + }"#, + ) + .await; + let provider = InfiniProvider::new(String::new()).with_base_url(server.url()); + let ctx = FetchContext { + api_key: Some("sk-cp-test-key".to_string()), + ..Default::default() + }; + + let result = provider.fetch_usage(&ctx).await.unwrap(); + + mock.assert_async().await; + assert_eq!(result.usage.primary.used_percent, 50.0); + let secondary = result.usage.secondary.expect("7-day window"); + assert_eq!(secondary.used_percent, 50.0); + } + + #[tokio::test] + async fn provider_fetch_usage_unauthorized() { + let (server, mock) = usage_server(401, "").await; + let provider = InfiniProvider::new(String::new()).with_base_url(server.url()); + let ctx = FetchContext { + api_key: Some("invalid-key".to_string()), + ..Default::default() + }; + + let result = provider.fetch_usage(&ctx).await; + + mock.assert_async().await; + assert!(result.is_err()); + } } // ==================== InfiniProvider ==================== diff --git a/rust/src/providers/kiro/cli_path.rs b/rust/src/providers/kiro/cli_path.rs new file mode 100755 index 0000000000..15e486b4e4 --- /dev/null +++ b/rust/src/providers/kiro/cli_path.rs @@ -0,0 +1,117 @@ +//! Kiro CLI binary detection. + +use std::path::{Path, PathBuf}; +use std::sync::OnceLock; + +/// Cached CLI path +static CLI_PATH: OnceLock> = OnceLock::new(); +const KIRO_CLI_PATH_ENV: &str = "CODEXBAR_KIRO_CLI_PATH"; + +fn is_allowed_kiro_binary(path: &Path) -> bool { + if !path.is_file() { + return false; + } + + let Some(file_name) = path.file_name().and_then(|name| name.to_str()) else { + return false; + }; + + #[cfg(target_os = "windows")] + { + file_name.eq_ignore_ascii_case("kiro-cli.exe") + } + + #[cfg(not(target_os = "windows"))] + { + file_name == "kiro-cli" || file_name == "kiro" + } +} + +fn env_override_cli_path() -> Option { + let raw = std::env::var(KIRO_CLI_PATH_ENV).ok()?; + let trimmed = raw.trim(); + if trimmed.is_empty() { + return None; + } + + let path = PathBuf::from(trimmed); + if is_allowed_kiro_binary(&path) { + return Some(path); + } + + None +} + +/// Find Kiro CLI binary path +pub fn find_kiro_cli() -> Option { + CLI_PATH + .get_or_init(|| { + // 1. Check explicit environment override first + if let Some(path) = env_override_cli_path() { + return Some(path); + } + + // 2. Hardened PATH lookup - use which but validate the result + // (avoids CWD hijacking by not executing bare command names) + if let Ok(path) = which::which("kiro-cli") + && is_allowed_kiro_binary(&path) + { + return Some(path); + } + if let Ok(path) = which::which("kiro") + && is_allowed_kiro_binary(&path) + { + return Some(path); + } + + // 3. Fall back to known install locations + #[cfg(target_os = "windows")] + { + let possible_paths = [ + dirs::data_local_dir() + .map(|p| p.join("Programs").join("Kiro").join("kiro-cli.exe")), + Some(PathBuf::from("C:\\Program Files\\Kiro\\kiro-cli.exe")), + ]; + for path in possible_paths.into_iter().flatten() { + if is_allowed_kiro_binary(&path) { + return Some(path); + } + } + } + + None + }) + .clone() +} + +#[cfg(test)] +mod tests { + use super::*; + use std::fs::File; + + fn temp_binary_path(name: &str) -> PathBuf { + let dir = + std::env::temp_dir().join(format!("codexbar-kiro-version-test-{}", std::process::id())); + std::fs::create_dir_all(&dir).expect("create temp binary directory"); + let path = dir.join(name); + File::create(&path).expect("create temp binary placeholder"); + path + } + + #[test] + #[cfg(target_os = "windows")] + fn windows_rejects_gui_kiro_binary_as_cli() { + let cli_path = temp_binary_path("kiro-cli.exe"); + let gui_path = temp_binary_path("kiro.exe"); + + assert!(is_allowed_kiro_binary(&cli_path)); + assert!( + !is_allowed_kiro_binary(&gui_path), + "kiro.exe is the Electron GUI app on Windows; running it as a CLI spawns the IDE" + ); + + // Best-effort cleanup of temp binaries; leftover files are harmless. + let _removed_cli = std::fs::remove_file(cli_path); + let _removed_gui = std::fs::remove_file(gui_path); + } +} diff --git a/rust/src/providers/kiro/mod.rs b/rust/src/providers/kiro/mod.rs index 6fbdb83ee9..6da0b8aa19 100755 --- a/rust/src/providers/kiro/mod.rs +++ b/rust/src/providers/kiro/mod.rs @@ -3,25 +3,18 @@ //! Fetches usage data from Kiro (Amazon's AI coding assistant) //! Uses kiro-cli for authentication and usage fetching +mod cli_path; #[cfg(test)] mod tests; mod usage_limits; -pub mod version; - -// Re-exports for version compatibility checking -#[allow( - unused_imports, - reason = "imports needed for future Kiro provider wiring" -)] -pub use version::{ - KiroVersion, detect_version, find_kiro_cli, get_version, is_compatible, is_installed, -}; + +pub use cli_path::find_kiro_cli; use async_trait::async_trait; use chrono::Datelike; use regex_lite::Regex; -use std::path::PathBuf; -use std::process::Stdio; +use std::path::Path; +use std::process::{Output, Stdio}; use tokio::process::Command; use crate::core::{ @@ -52,45 +45,31 @@ impl KiroProvider { Self } - /// Get Kiro config directory - fn get_kiro_config_path() -> Option { - #[cfg(target_os = "windows")] - { - dirs::config_dir().map(|p| p.join("Kiro")) - } - #[cfg(not(target_os = "windows"))] - { - dirs::home_dir().map(|p| p.join(".kiro")) - } - } - - /// Find Kiro CLI binary - fn which_kiro() -> Option { - version::find_kiro_cli() - } - - /// Check if user is logged in by running `kiro-cli whoami` - async fn ensure_logged_in(&self) -> Result<(), ProviderError> { - let cli_path = Self::which_kiro().ok_or_else(|| { - ProviderError::NotInstalled( - "kiro-cli not found. Install from https://kiro.dev".to_string(), - ) - })?; - + /// Runs `kiro-cli` with piped output and no console window. + async fn run_kiro( + cli_path: &Path, + args: &[&str], + env: &[(&str, &str)], + ) -> Result { #[cfg(windows)] const CREATE_NO_WINDOW: u32 = 0x08000000; - let mut cmd = Command::new(&cli_path); - cmd.arg("whoami") + let mut cmd = Command::new(cli_path); + cmd.args(args) + .envs(env.iter().copied()) .stdout(Stdio::piped()) .stderr(Stdio::piped()); #[cfg(windows)] cmd.creation_flags(CREATE_NO_WINDOW); - let output = cmd - .output() + cmd.output() .await - .map_err(|e| ProviderError::Other(format!("Failed to run kiro-cli: {}", e)))?; + .map_err(|e| ProviderError::Other(format!("Failed to run kiro-cli: {}", e))) + } + + /// Check if user is logged in by running `kiro-cli whoami` + async fn ensure_logged_in(cli_path: &Path) -> Result<(), ProviderError> { + let output = Self::run_kiro(cli_path, &["whoami"], &[]).await?; let stdout = String::from_utf8_lossy(&output.stdout).to_lowercase(); let stderr = String::from_utf8_lossy(&output.stderr).to_lowercase(); @@ -112,31 +91,22 @@ impl KiroProvider { /// Fetch usage via kiro-cli async fn fetch_via_cli(&self) -> Result { - // First ensure we're logged in - self.ensure_logged_in().await?; - - let cli_path = Self::which_kiro() - .ok_or_else(|| ProviderError::NotInstalled("kiro-cli not found".to_string()))?; + let cli_path = cli_path::find_kiro_cli().ok_or_else(|| { + ProviderError::NotInstalled( + "kiro-cli not found. Install from https://kiro.dev".to_string(), + ) + })?; + Self::ensure_logged_in(&cli_path).await?; - // Run the usage command. // Windows intentionally uses pipe-first (stdout/stderr Stdio::piped) rather than a dual // ConPTY path: kiro-cli `/usage` under --no-interactive emits parseable text on pipes, and // a second ConPTY probe would add flaky process-lifetime cost without better quota data. - #[cfg(windows)] - const CREATE_NO_WINDOW: u32 = 0x08000000; - - let mut cmd = Command::new(&cli_path); - cmd.args(["chat", "--no-interactive", "/usage"]) - .env("TERM", "xterm-256color") - .stdout(Stdio::piped()) - .stderr(Stdio::piped()); - #[cfg(windows)] - cmd.creation_flags(CREATE_NO_WINDOW); - - let output = cmd - .output() - .await - .map_err(|e| ProviderError::Other(format!("Failed to run kiro-cli: {}", e)))?; + let output = Self::run_kiro( + &cli_path, + &["chat", "--no-interactive", "/usage"], + &[("TERM", "xterm-256color")], + ) + .await?; let stdout = String::from_utf8_lossy(&output.stdout); let stderr = String::from_utf8_lossy(&output.stderr); diff --git a/rust/src/providers/kiro/usage_limits.rs b/rust/src/providers/kiro/usage_limits.rs index a90a4081ed..cb2fe9d50f 100644 --- a/rust/src/providers/kiro/usage_limits.rs +++ b/rust/src/providers/kiro/usage_limits.rs @@ -271,33 +271,22 @@ fn read_identity(path: &Path) -> Result { ) .optional() .map_err(|error| ProviderError::Other(format!("Kiro profile lookup: {error}")))?; - let access_token = - json_string(token_json.as_deref(), "access_token").ok_or(ProviderError::AuthRequired)?; + let access_token = json_string(token_json.as_deref(), "access_token", true) + .ok_or(ProviderError::AuthRequired)?; let profile_arn = - json_string_exact(profile_json.as_deref(), "arn").ok_or(ProviderError::AuthRequired)?; + json_string(profile_json.as_deref(), "arn", false).ok_or(ProviderError::AuthRequired)?; Ok(KiroIdentity { access_token, profile_arn, }) } -fn json_string(json: Option<&str>, key: &str) -> Option { - serde_json::from_str::(json?) - .ok()? - .get(key)? - .as_str() - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(str::to_string) -} - -fn json_string_exact(json: Option<&str>, key: &str) -> Option { - serde_json::from_str::(json?) - .ok()? - .get(key)? - .as_str() - .filter(|value| !value.is_empty()) - .map(str::to_string) +/// A non-empty string field of a JSON document; the ARN is matched untrimmed. +fn json_string(json: Option<&str>, key: &str, trim: bool) -> Option { + let value = serde_json::from_str::(json?).ok()?; + let text = value.get(key)?.as_str()?; + let text = if trim { text.trim() } else { text }; + (!text.is_empty()).then(|| text.to_string()) } fn endpoint_for_profile_arn(profile_arn: &str) -> Option<&'static str> { diff --git a/rust/src/providers/kiro/version.rs b/rust/src/providers/kiro/version.rs deleted file mode 100755 index 3bccb59ce7..0000000000 --- a/rust/src/providers/kiro/version.rs +++ /dev/null @@ -1,399 +0,0 @@ -//! Kiro CLI Version Detection -//! -//! Detect and parse Kiro CLI version for compatibility checks. - -#[cfg(windows)] -use std::os::windows::process::CommandExt; -use std::path::{Path, PathBuf}; -use std::process::Command; -use std::sync::OnceLock; - -/// Cached CLI path -static CLI_PATH: OnceLock> = OnceLock::new(); - -/// Cached CLI version -static CLI_VERSION: OnceLock> = OnceLock::new(); -const KIRO_CLI_PATH_ENV: &str = "CODEXBAR_KIRO_CLI_PATH"; - -fn is_allowed_kiro_binary(path: &Path) -> bool { - if !path.is_file() { - return false; - } - - let Some(file_name) = path.file_name().and_then(|name| name.to_str()) else { - return false; - }; - - #[cfg(target_os = "windows")] - { - file_name.eq_ignore_ascii_case("kiro-cli.exe") - } - - #[cfg(not(target_os = "windows"))] - { - file_name == "kiro-cli" || file_name == "kiro" - } -} - -fn env_override_cli_path() -> Option { - let raw = std::env::var(KIRO_CLI_PATH_ENV).ok()?; - let trimmed = raw.trim(); - if trimmed.is_empty() { - return None; - } - - let path = PathBuf::from(trimmed); - if is_allowed_kiro_binary(&path) { - return Some(path); - } - - None -} - -/// Kiro CLI version info -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct KiroVersion { - /// Major version number - pub major: u32, - /// Minor version number - pub minor: u32, - /// Patch version number - pub patch: u32, - /// Pre-release suffix (e.g., "beta.1") - pub prerelease: Option, - /// Build metadata - pub build: Option, - /// Raw version string - pub raw: String, -} - -impl KiroVersion { - /// Parse a version string - pub fn parse(version: &str) -> Option { - let trimmed = version.trim(); - if trimmed.is_empty() { - return None; - } - - // Handle "kiro-cli X.Y.Z" prefix - let version_part = if trimmed.to_lowercase().starts_with("kiro-cli ") { - &trimmed[9..] - } else if trimmed.to_lowercase().starts_with("kiro ") { - &trimmed[5..] - } else { - trimmed - } - .trim(); - - // Split off pre-release and build metadata - let (version_core, prerelease, build) = Self::split_version_parts(version_part); - - // Parse X.Y.Z - let parts: Vec<&str> = version_core.split('.').collect(); - if parts.is_empty() { - return None; - } - - let major = parts.first().and_then(|s| s.parse().ok()).unwrap_or(0); - let minor = parts.get(1).and_then(|s| s.parse().ok()).unwrap_or(0); - let patch = parts.get(2).and_then(|s| s.parse().ok()).unwrap_or(0); - - Some(KiroVersion { - major, - minor, - patch, - prerelease, - build, - raw: trimmed.to_string(), - }) - } - - /// Split version into core, prerelease, and build parts - fn split_version_parts(version: &str) -> (String, Option, Option) { - let mut core = version.to_string(); - let mut prerelease = None; - let mut build = None; - - // Extract build metadata first (after +) - if let Some(plus_idx) = core.find('+') { - build = Some(core[plus_idx + 1..].to_string()); - core = core[..plus_idx].to_string(); - } - - // Extract prerelease (after -) - if let Some(dash_idx) = core.find('-') { - prerelease = Some(core[dash_idx + 1..].to_string()); - core = core[..dash_idx].to_string(); - } - - (core, prerelease, build) - } - - /// Check if this version is at least the specified version - pub fn at_least(&self, major: u32, minor: u32, patch: u32) -> bool { - if self.major > major { - return true; - } - if self.major < major { - return false; - } - if self.minor > minor { - return true; - } - if self.minor < minor { - return false; - } - self.patch >= patch - } - - /// Check if this is a prerelease version - pub fn is_prerelease(&self) -> bool { - self.prerelease.is_some() - } - - /// Get display string - pub fn display(&self) -> String { - format!("{}.{}.{}", self.major, self.minor, self.patch) - } -} - -impl std::fmt::Display for KiroVersion { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!(f, "{}", self.raw) - } -} - -impl PartialOrd for KiroVersion { - fn partial_cmp(&self, other: &Self) -> Option { - Some(self.cmp(other)) - } -} - -impl Ord for KiroVersion { - fn cmp(&self, other: &Self) -> std::cmp::Ordering { - match self.major.cmp(&other.major) { - std::cmp::Ordering::Equal => {} - ord => return ord, - } - match self.minor.cmp(&other.minor) { - std::cmp::Ordering::Equal => {} - ord => return ord, - } - self.patch.cmp(&other.patch) - } -} - -/// Find Kiro CLI binary path -pub fn find_kiro_cli() -> Option { - CLI_PATH - .get_or_init(|| { - // 1. Check explicit environment override first - if let Some(path) = env_override_cli_path() { - return Some(path); - } - - // 2. Hardened PATH lookup - use which but validate the result - // (avoids CWD hijacking by not executing bare command names) - if let Ok(path) = which::which("kiro-cli") - && is_allowed_kiro_binary(&path) - { - return Some(path); - } - if let Ok(path) = which::which("kiro") - && is_allowed_kiro_binary(&path) - { - return Some(path); - } - - // 3. Fall back to known install locations - #[cfg(target_os = "windows")] - { - let possible_paths = [ - dirs::data_local_dir() - .map(|p| p.join("Programs").join("Kiro").join("kiro-cli.exe")), - Some(PathBuf::from("C:\\Program Files\\Kiro\\kiro-cli.exe")), - ]; - for path in possible_paths.into_iter().flatten() { - if is_allowed_kiro_binary(&path) { - return Some(path); - } - } - } - - None - }) - .clone() -} - -/// Detect Kiro CLI version -pub fn detect_version() -> Option { - CLI_VERSION - .get_or_init(|| { - let cli_path = find_kiro_cli()?; - - #[cfg(windows)] - const CREATE_NO_WINDOW: u32 = 0x08000000; - - let mut cmd = Command::new(&cli_path); - cmd.arg("--version"); - #[cfg(windows)] - cmd.creation_flags(CREATE_NO_WINDOW); - - let output = cmd.output().ok()?; - - if !output.status.success() { - return None; - } - - let stdout = String::from_utf8_lossy(&output.stdout); - let stderr = String::from_utf8_lossy(&output.stderr); - let combined = if stdout.trim().is_empty() { - stderr.to_string() - } else { - stdout.to_string() - }; - - let trimmed = combined.trim(); - if trimmed.is_empty() { - return None; - } - - // Output is like "kiro-cli 1.23.1" or just "1.23.1" - let version = if trimmed.to_lowercase().starts_with("kiro-cli ") { - trimmed[9..].trim().to_string() - } else if trimmed.to_lowercase().starts_with("kiro ") { - trimmed[5..].trim().to_string() - } else { - trimmed.to_string() - }; - - Some(version) - }) - .clone() -} - -/// Get parsed Kiro version -pub fn get_version() -> Option { - detect_version().and_then(|v| KiroVersion::parse(&v)) -} - -/// Check if Kiro CLI is installed -pub fn is_installed() -> bool { - find_kiro_cli().is_some() -} - -/// Check if the installed Kiro CLI version is compatible -pub fn is_compatible(min_major: u32, min_minor: u32, min_patch: u32) -> bool { - match get_version() { - Some(v) => v.at_least(min_major, min_minor, min_patch), - None => false, - } -} - -/// Reset cached values (for testing) -#[cfg(test)] -pub fn reset_cache() { - // Can't reset OnceLock in stable Rust, this is just for documentation -} - -#[cfg(test)] -mod tests { - use super::*; - use std::fs::File; - - fn temp_binary_path(name: &str) -> PathBuf { - let dir = - std::env::temp_dir().join(format!("codexbar-kiro-version-test-{}", std::process::id())); - std::fs::create_dir_all(&dir).expect("create temp binary directory"); - let path = dir.join(name); - File::create(&path).expect("create temp binary placeholder"); - path - } - - #[test] - fn test_parse_version_simple() { - let v = KiroVersion::parse("1.23.4").unwrap(); - assert_eq!(v.major, 1); - assert_eq!(v.minor, 23); - assert_eq!(v.patch, 4); - assert!(v.prerelease.is_none()); - } - - #[test] - fn test_parse_version_with_prefix() { - let v = KiroVersion::parse("kiro-cli 1.2.3").unwrap(); - assert_eq!(v.major, 1); - assert_eq!(v.minor, 2); - assert_eq!(v.patch, 3); - - let v = KiroVersion::parse("Kiro 2.0.0").unwrap(); - assert_eq!(v.major, 2); - } - - #[test] - fn test_parse_version_with_prerelease() { - let v = KiroVersion::parse("1.0.0-beta.1").unwrap(); - assert_eq!(v.major, 1); - assert_eq!(v.minor, 0); - assert_eq!(v.patch, 0); - assert_eq!(v.prerelease, Some("beta.1".to_string())); - assert!(v.is_prerelease()); - } - - #[test] - fn test_parse_version_with_build() { - let v = KiroVersion::parse("1.0.0+build123").unwrap(); - assert_eq!(v.build, Some("build123".to_string())); - } - - #[test] - fn test_version_comparison() { - let v1 = KiroVersion::parse("1.2.3").unwrap(); - let v2 = KiroVersion::parse("1.2.4").unwrap(); - let v3 = KiroVersion::parse("1.3.0").unwrap(); - let v4 = KiroVersion::parse("2.0.0").unwrap(); - - assert!(v1 < v2); - assert!(v2 < v3); - assert!(v3 < v4); - assert!(v1 < v4); - } - - #[test] - fn test_at_least() { - let v = KiroVersion::parse("1.5.2").unwrap(); - - assert!(v.at_least(1, 5, 2)); - assert!(v.at_least(1, 5, 0)); - assert!(v.at_least(1, 4, 0)); - assert!(v.at_least(0, 9, 0)); - - assert!(!v.at_least(1, 5, 3)); - assert!(!v.at_least(1, 6, 0)); - assert!(!v.at_least(2, 0, 0)); - } - - #[test] - fn test_display() { - let v = KiroVersion::parse("1.2.3-beta+build").unwrap(); - assert_eq!(v.display(), "1.2.3"); - assert_eq!(v.to_string(), "1.2.3-beta+build"); - } - - #[test] - #[cfg(target_os = "windows")] - fn windows_rejects_gui_kiro_binary_as_cli() { - let cli_path = temp_binary_path("kiro-cli.exe"); - let gui_path = temp_binary_path("kiro.exe"); - - assert!(is_allowed_kiro_binary(&cli_path)); - assert!( - !is_allowed_kiro_binary(&gui_path), - "kiro.exe is the Electron GUI app on Windows; running it as a CLI spawns the IDE" - ); - - // Best-effort cleanup of temp binaries; leftover files are harmless. - let _removed_cli = std::fs::remove_file(cli_path); - let _removed_gui = std::fs::remove_file(gui_path); - } -} diff --git a/rust/src/providers/qwencloud/fields.rs b/rust/src/providers/qwencloud/fields.rs new file mode 100644 index 0000000000..b970487dc7 --- /dev/null +++ b/rust/src/providers/qwencloud/fields.rs @@ -0,0 +1,372 @@ +//! Lenient lookups over the expanded Qwen Cloud console JSON. + +use chrono::{DateTime, NaiveDate, NaiveDateTime, TimeZone, Utc}; +use serde_json::Value; + +pub(super) fn expand_json_strings(value: Value) -> Value { + match value { + Value::Array(values) => Value::Array(values.into_iter().map(expand_json_strings).collect()), + Value::Object(map) => Value::Object( + map.into_iter() + .map(|(key, value)| (key, expand_json_strings(value))) + .collect(), + ), + Value::String(text) => serde_json::from_str::(&text) + .ok() + .filter(|nested| nested.is_object() || nested.is_array()) + .map(expand_json_strings) + .unwrap_or(Value::String(text)), + other => other, + } +} + +pub(super) fn percentage_points(ratio: Option) -> Option { + let ratio = ratio.filter(|v| v.is_finite())?; + Some((ratio.clamp(0.0, 1.0) * 100.0).clamp(0.0, 100.0)) +} + +pub(super) fn number_field(value: &Value, key: &str) -> Option { + value.as_object().and_then(|map| parse_f64(map.get(key))) +} + +pub(super) fn date_field(value: &Value, key: &str) -> Option> { + value.as_object().and_then(|map| parse_date(map.get(key))) +} + +/// Depth-first search: `probe` runs on each node before its children, and +/// the first `Some` wins. Probes return `None` for arrays and scalars. +fn deep_find<'a, T>(value: &'a Value, probe: &impl Fn(&'a Value) -> Option) -> Option { + if let Some(found) = probe(value) { + return Some(found); + } + match value { + Value::Object(map) => map.values().find_map(|nested| deep_find(nested, probe)), + Value::Array(values) => values.iter().find_map(|nested| deep_find(nested, probe)), + _ => None, + } +} + +pub(super) fn find_object_containing_any_of(value: &Value, keys: &[&str]) -> Option { + deep_find(value, &|node| { + let map = node.as_object()?; + keys.iter() + .any(|key| map.contains_key(*key)) + .then(|| node.clone()) + }) +} + +pub(super) fn find_first_value_for_key(value: &Value, key: &str) -> Option { + deep_find(value, &|node| node.as_object()?.get(key).cloned()) +} + +pub(super) const PLAN_NAME_KEYS: &[&str] = &[ + "planName", + "plan_name", + "packageName", + "package_name", + "commodityName", + "commodity_name", + "instanceName", + "instance_name", + "displayName", + "display_name", + "name", + "title", + "planType", + "plan_type", + "ProductName", + "productName", + "InstanceCode", +]; + +pub(super) const USED_QUOTA_KEYS: &[&str] = &[ + "usedQuota", + "used_quota", + "usedCredits", + "usedCredit", + "consumedCredits", + "usage", + "used", + "usedAmount", + "consumeAmount", + "usedValue", + "UsedValue", + "consumedValue", + "ConsumedValue", +]; + +pub(super) const TOTAL_QUOTA_KEYS: &[&str] = &[ + "totalQuota", + "total_quota", + "totalCredits", + "totalCredit", + "quota", + "creditLimit", + "creditsTotal", + "monthlyTotalQuota", + "amount", + "totalValue", + "TotalValue", + "CycleTotalValue", + "cycleTotalValue", + "subscriptionTotalNumber", + "SubscriptionTotalNumber", +]; + +pub(super) const REMAINING_QUOTA_KEYS: &[&str] = &[ + "remainingQuota", + "remainQuota", + "remainingCredits", + "remainingCredit", + "availableCredits", + "balance", + "remaining", + "availableAmount", + "remainAmount", + "totalSurplusValue", + "TotalSurplusValue", + "surplusValue", + "SurplusValue", + "CycleSurplusValue", + "cycleSurplusValue", +]; + +pub(super) const RESET_DATE_KEYS: &[&str] = &[ + "nextRefreshTime", + "resetTime", + "periodEndTime", + "billingCycleEnd", + "billCycleEndTime", + "expireTime", + "expirationTime", + "endTime", + "EndTime", + "validEndTime", + "instanceEndTime", + "nearestExpireDate", + "NearestExpireDate", +]; + +pub(super) fn find_token_plan_instance(value: &Value) -> Option { + find_first_object( + value, + &[ + "tokenPlanInstanceInfo", + "token_plan_instance_info", + "instanceInfo", + "instance_info", + ], + ) + .or_else(|| { + find_first_array( + value, + &[ + "tokenPlanInstanceInfos", + "token_plan_instance_infos", + "instanceInfos", + "instances", + "EquityList", + "Data", + "data", + "successResponse", + ], + ) + .and_then(|values| { + values + .into_iter() + .filter(Value::is_object) + .max_by_key(active_signal_score) + }) + }) +} + +pub(super) fn find_quota_info(value: &Value) -> Option { + find_first_object( + value, + &[ + "quotaInfo", + "quota_info", + "tokenPlanQuotaInfo", + "token_plan_quota_info", + ], + ) + .or_else(|| { + find_first_array(value, &["EquityList", "equityList"]).and_then(|values| { + values.into_iter().find(|item| { + item.as_object().is_some_and(|map| { + map.contains_key("CycleTotalValue") + || map.contains_key("cycleTotalValue") + || map.contains_key("CycleSurplusValue") + }) + }) + }) + }) + .or_else(|| { + let keys: Vec<&str> = USED_QUOTA_KEYS + .iter() + .chain(TOTAL_QUOTA_KEYS.iter()) + .chain(REMAINING_QUOTA_KEYS.iter()) + .copied() + .collect(); + find_object_containing_any_of(value, &keys) + }) +} + +fn find_first_object(value: &Value, keys: &[&str]) -> Option { + deep_find(value, &|node| { + let map = node.as_object()?; + keys.iter() + .find_map(|key| map.get(*key).filter(|v| v.is_object()).cloned()) + }) +} + +fn find_first_array(value: &Value, keys: &[&str]) -> Option> { + deep_find(value, &|node| { + let map = node.as_object()?; + keys.iter() + .find_map(|key| map.get(*key).and_then(Value::as_array).cloned()) + }) +} + +fn first_string(value: &Value, keys: &[&str]) -> Option { + value + .as_object() + .and_then(|map| keys.iter().find_map(|key| parse_string(map.get(*key)))) +} + +pub(super) fn find_first_string(value: &Value, keys: &[&str]) -> Option { + deep_find(value, &|node| first_string(node, keys)) +} + +pub(super) fn first_f64(value: &Value, keys: &[&str]) -> Option { + value + .as_object() + .and_then(|map| keys.iter().find_map(|key| parse_f64(map.get(*key)))) +} + +pub(super) fn find_first_f64(value: &Value, keys: &[&str]) -> Option { + deep_find(value, &|node| first_f64(node, keys)) +} + +pub(super) fn find_first_i64(value: &Value, keys: &[&str]) -> Option { + deep_find(value, &|node| { + let map = node.as_object()?; + keys.iter().find_map(|key| parse_i64(map.get(*key))) + }) +} + +pub(super) fn find_first_bool(value: &Value, keys: &[&str]) -> Option { + deep_find(value, &|node| { + let map = node.as_object()?; + keys.iter().find_map(|key| parse_bool(map.get(*key))) + }) +} + +pub(super) fn first_date(value: &Value, keys: &[&str]) -> Option> { + value + .as_object() + .and_then(|map| keys.iter().find_map(|key| parse_date(map.get(*key)))) +} + +pub(super) fn find_first_date(value: &Value, keys: &[&str]) -> Option> { + deep_find(value, &|node| first_date(node, keys)) +} + +fn parse_string(value: Option<&Value>) -> Option { + value? + .as_str() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) +} + +fn parse_f64(value: Option<&Value>) -> Option { + match value? { + Value::Number(number) => number.as_f64(), + Value::String(text) => text.trim().replace(',', "").parse().ok(), + _ => None, + } + .filter(|v| v.is_finite()) +} + +fn parse_i64(value: Option<&Value>) -> Option { + match value? { + Value::Number(number) => number.as_i64().or_else(|| { + number.as_f64().map(|v| { + // Quota numbers are small integers; any fractional part is discarded deliberately. + #[expect(clippy::cast_possible_truncation, reason = "quota counts fit i64")] + let v = v as i64; + v + }) + }), + Value::String(text) => text.trim().replace(',', "").parse().ok(), + _ => None, + } +} + +fn parse_bool(value: Option<&Value>) -> Option { + match value? { + Value::Bool(flag) => Some(*flag), + Value::Number(number) => number.as_i64().map(|v| v != 0), + Value::String(text) => match text.trim().to_lowercase().as_str() { + "true" | "1" | "yes" | "active" | "valid" | "normal" => Some(true), + "false" | "0" | "no" | "inactive" | "invalid" | "expired" => Some(false), + _ => None, + }, + _ => None, + } +} + +fn parse_date(value: Option<&Value>) -> Option> { + if let Some(raw) = parse_i64(value) { + if raw > 1_000_000_000_000 { + return Utc.timestamp_opt(raw / 1000, 0).single(); + } + if raw > 1_000_000_000 { + return Utc.timestamp_opt(raw, 0).single(); + } + } + let text = parse_string(value)?; + if let Ok(date) = DateTime::parse_from_rfc3339(&text) { + return Some(date.with_timezone(&Utc)); + } + if let Ok(date) = NaiveDate::parse_from_str(&text, "%Y-%m-%d") + && let Some(date_time) = date.and_hms_opt(0, 0, 0) + { + return Some(date_time.and_utc()); + } + for format in ["%Y-%m-%d %H:%M", "%Y-%m-%d %H:%M:%S"] { + if let Ok(date) = NaiveDateTime::parse_from_str(&text, format) { + return Some(date.and_utc()); + } + } + None +} + +fn active_signal_score(value: &Value) -> i32 { + let status = first_string(value, &["status", "instanceStatus", "state", "Status"]) + .unwrap_or_default() + .to_uppercase(); + if ["VALID", "ACTIVE", "NORMAL"].contains(&status.as_str()) { + return 3; + } + if [ + "EXPIRED", + "INVALID", + "INACTIVE", + "DISABLED", + "TERMINATED", + "STOPPED", + ] + .contains(&status.as_str()) + { + return -1; + } + parse_bool( + value + .as_object() + .and_then(|map| map.get("isActive").or_else(|| map.get("active"))), + ) + .map(|active| if active { 3 } else { -1 }) + .unwrap_or(0) +} diff --git a/rust/src/providers/qwencloud/mod.rs b/rust/src/providers/qwencloud/mod.rs index 702ce930db..1d8c4f45a4 100644 --- a/rust/src/providers/qwencloud/mod.rs +++ b/rust/src/providers/qwencloud/mod.rs @@ -4,11 +4,12 @@ //! `POST https://cs-data.qwencloud.com/data/api.json?...` //! with `IntlBroadScopeAspnGateway` / `sfm_bailian`. +mod fields; #[cfg(test)] mod monthly_tests; use async_trait::async_trait; -use chrono::{DateTime, NaiveDate, NaiveDateTime, TimeZone, Utc}; +use chrono::{DateTime, Utc}; use regex_lite::Regex; use serde_json::{Map, Value, json}; use uuid::Uuid; @@ -18,6 +19,13 @@ use crate::core::{ UsageSnapshot, }; use crate::providers::{browser_cookie_header, strip_cookie_prefix}; +use fields::{ + PLAN_NAME_KEYS, REMAINING_QUOTA_KEYS, RESET_DATE_KEYS, TOTAL_QUOTA_KEYS, USED_QUOTA_KEYS, + date_field, expand_json_strings, find_first_bool, find_first_date, find_first_f64, + find_first_i64, find_first_string, find_first_value_for_key, find_object_containing_any_of, + find_quota_info, find_token_plan_instance, first_date, first_f64, number_field, + percentage_points, +}; const GATEWAY_BASE_URL: &str = "https://home.qwencloud.com"; const DATA_GATEWAY_BASE_URL: &str = "https://cs-data.qwencloud.com"; @@ -105,42 +113,19 @@ impl QwenCloudProvider { .await .ok_or(ProviderError::AuthRequired)?; - let usage_body = Self::post_api( - &client, - USAGE_API, - Map::new(), - &sec_token, - &cookie_header, - ctx, - ) - .await?; - - let subscription_body = Self::post_api_optional( - &client, - SUBSCRIPTION_API, - { - let mut data = Map::new(); - data.insert( - "commodityCode".into(), - Value::String(PRODUCT_CODE.to_string()), - ); - data - }, - &sec_token, - &cookie_header, - ctx, - ) - .await; - - let quota_config_body = Self::post_api_optional( - &client, - QUOTA_CONFIG_API, - Map::new(), - &sec_token, - &cookie_header, - ctx, - ) - .await; + let (client, sec_token, cookie_header) = + (&client, sec_token.as_str(), cookie_header.as_str()); + let post = move |api: &'static str, data| { + Self::post_api(client, api, data, sec_token, cookie_header, ctx) + }; + let usage_body = post(USAGE_API, Map::new()).await?; + let mut subscription_params = Map::new(); + subscription_params.insert( + "commodityCode".into(), + Value::String(PRODUCT_CODE.to_string()), + ); + let subscription_body = post(SUBSCRIPTION_API, subscription_params).await.ok(); + let quota_config_body = post(QUOTA_CONFIG_API, Map::new()).await.ok(); let snapshot = Self::parse( &usage_body, @@ -284,19 +269,6 @@ impl QwenCloudProvider { Ok(body.to_vec()) } - async fn post_api_optional( - client: &reqwest::Client, - api: &str, - data_parameters: Map, - sec_token: &str, - cookie_header: &str, - ctx: &FetchContext, - ) -> Option> { - Self::post_api(client, api, data_parameters, sec_token, cookie_header, ctx) - .await - .ok() - } - fn parse( usage_data: &[u8], subscription_data: Option<&[u8]>, @@ -445,25 +417,23 @@ fn build_params_json( mut data_parameters: Map, cookie_header: &str, ) -> String { - let mut cornerstone = Map::new(); - cornerstone.insert( - "feTraceId".into(), - Value::String(Uuid::new_v4().to_string().to_lowercase()), - ); - cornerstone.insert("feURL".into(), Value::String(DASHBOARD_URL.to_string())); - cornerstone.insert("protocol".into(), Value::String("V2".into())); - cornerstone.insert("console".into(), Value::String("ONE_CONSOLE".into())); - cornerstone.insert("productCode".into(), Value::String("p_efm".into())); - cornerstone.insert("domain".into(), Value::String("home.qwencloud.com".into())); - cornerstone.insert("consoleSite".into(), Value::String("QWENCLOUD".into())); - cornerstone.insert("userNickName".into(), Value::String(String::new())); - cornerstone.insert("userPrincipalName".into(), Value::String(String::new())); - cornerstone.insert("xsp_lang".into(), Value::String(LANGUAGE.into())); + let mut cornerstone = json!({ + "feTraceId": Uuid::new_v4().to_string().to_lowercase(), + "feURL": DASHBOARD_URL, + "protocol": "V2", + "console": "ONE_CONSOLE", + "productCode": "p_efm", + "domain": "home.qwencloud.com", + "consoleSite": "QWENCLOUD", + "userNickName": "", + "userPrincipalName": "", + "xsp_lang": LANGUAGE, + }); if let Some(cna) = cookie_value("cna", cookie_header) { - cornerstone.insert("X-Anonymous-Id".into(), Value::String(cna)); + cornerstone["X-Anonymous-Id"] = Value::String(cna); } - data_parameters.insert("cornerstoneParam".into(), Value::Object(cornerstone)); + data_parameters.insert("cornerstoneParam".into(), cornerstone); json!({ "Api": api, @@ -716,459 +686,6 @@ fn is_likely_login_html(data: &[u8]) -> bool { && (text.contains("login") || text.contains("sign in") || text.contains("signin")) } -fn expand_json_strings(value: Value) -> Value { - match value { - Value::Array(values) => Value::Array(values.into_iter().map(expand_json_strings).collect()), - Value::Object(map) => Value::Object( - map.into_iter() - .map(|(key, value)| (key, expand_json_strings(value))) - .collect(), - ), - Value::String(text) => serde_json::from_str::(&text) - .ok() - .filter(|nested| nested.is_object() || nested.is_array()) - .map(expand_json_strings) - .unwrap_or(Value::String(text)), - other => other, - } -} - -fn percentage_points(ratio: Option) -> Option { - let ratio = ratio.filter(|v| v.is_finite())?; - Some((ratio.clamp(0.0, 1.0) * 100.0).clamp(0.0, 100.0)) -} - -fn number_field(value: &Value, key: &str) -> Option { - value.as_object().and_then(|map| parse_f64(map.get(key))) -} - -fn date_field(value: &Value, key: &str) -> Option> { - value.as_object().and_then(|map| parse_date(map.get(key))) -} - -fn find_object_containing_any_of(value: &Value, keys: &[&str]) -> Option { - match value { - Value::Object(map) => { - if keys.iter().any(|key| map.contains_key(*key)) { - return Some(Value::Object(map.clone())); - } - map.values() - .find_map(|nested| find_object_containing_any_of(nested, keys)) - } - Value::Array(values) => values - .iter() - .find_map(|nested| find_object_containing_any_of(nested, keys)), - _ => None, - } -} - -fn find_first_value_for_key(value: &Value, key: &str) -> Option { - match value { - Value::Object(map) => { - if let Some(nested) = map.get(key) { - return Some(nested.clone()); - } - map.values() - .find_map(|nested| find_first_value_for_key(nested, key)) - } - Value::Array(values) => values - .iter() - .find_map(|nested| find_first_value_for_key(nested, key)), - _ => None, - } -} - -const PLAN_NAME_KEYS: &[&str] = &[ - "planName", - "plan_name", - "packageName", - "package_name", - "commodityName", - "commodity_name", - "instanceName", - "instance_name", - "displayName", - "display_name", - "name", - "title", - "planType", - "plan_type", - "ProductName", - "productName", - "InstanceCode", -]; - -const USED_QUOTA_KEYS: &[&str] = &[ - "usedQuota", - "used_quota", - "usedCredits", - "usedCredit", - "consumedCredits", - "usage", - "used", - "usedAmount", - "consumeAmount", - "usedValue", - "UsedValue", - "consumedValue", - "ConsumedValue", -]; - -const TOTAL_QUOTA_KEYS: &[&str] = &[ - "totalQuota", - "total_quota", - "totalCredits", - "totalCredit", - "quota", - "creditLimit", - "creditsTotal", - "monthlyTotalQuota", - "amount", - "totalValue", - "TotalValue", - "CycleTotalValue", - "cycleTotalValue", - "subscriptionTotalNumber", - "SubscriptionTotalNumber", -]; - -const REMAINING_QUOTA_KEYS: &[&str] = &[ - "remainingQuota", - "remainQuota", - "remainingCredits", - "remainingCredit", - "availableCredits", - "balance", - "remaining", - "availableAmount", - "remainAmount", - "totalSurplusValue", - "TotalSurplusValue", - "surplusValue", - "SurplusValue", - "CycleSurplusValue", - "cycleSurplusValue", -]; - -const RESET_DATE_KEYS: &[&str] = &[ - "nextRefreshTime", - "resetTime", - "periodEndTime", - "billingCycleEnd", - "billCycleEndTime", - "expireTime", - "expirationTime", - "endTime", - "EndTime", - "validEndTime", - "instanceEndTime", - "nearestExpireDate", - "NearestExpireDate", -]; - -fn find_token_plan_instance(value: &Value) -> Option { - find_first_object( - value, - &[ - "tokenPlanInstanceInfo", - "token_plan_instance_info", - "instanceInfo", - "instance_info", - ], - ) - .or_else(|| { - find_first_array( - value, - &[ - "tokenPlanInstanceInfos", - "token_plan_instance_infos", - "instanceInfos", - "instances", - "EquityList", - "Data", - "data", - "successResponse", - ], - ) - .and_then(|values| { - values - .into_iter() - .filter(Value::is_object) - .max_by_key(active_signal_score) - }) - }) -} - -fn find_quota_info(value: &Value) -> Option { - find_first_object( - value, - &[ - "quotaInfo", - "quota_info", - "tokenPlanQuotaInfo", - "token_plan_quota_info", - ], - ) - .or_else(|| { - find_first_array(value, &["EquityList", "equityList"]).and_then(|values| { - values.into_iter().find(|item| { - item.as_object().is_some_and(|map| { - map.contains_key("CycleTotalValue") - || map.contains_key("cycleTotalValue") - || map.contains_key("CycleSurplusValue") - }) - }) - }) - }) - .or_else(|| { - let keys: Vec<&str> = USED_QUOTA_KEYS - .iter() - .chain(TOTAL_QUOTA_KEYS.iter()) - .chain(REMAINING_QUOTA_KEYS.iter()) - .copied() - .collect(); - find_first_object_with_any_key(value, &keys) - }) -} - -fn find_first_object(value: &Value, keys: &[&str]) -> Option { - match value { - Value::Object(map) => { - for key in keys { - if let Some(nested) = map.get(*key).filter(|v| v.is_object()) { - return Some(nested.clone()); - } - } - map.values() - .find_map(|nested| find_first_object(nested, keys)) - } - Value::Array(values) => values - .iter() - .find_map(|nested| find_first_object(nested, keys)), - _ => None, - } -} - -fn find_first_object_with_any_key(value: &Value, keys: &[&str]) -> Option { - match value { - Value::Object(map) => { - if keys.iter().any(|key| map.contains_key(*key)) { - return Some(Value::Object(map.clone())); - } - map.values() - .find_map(|nested| find_first_object_with_any_key(nested, keys)) - } - Value::Array(values) => values - .iter() - .find_map(|nested| find_first_object_with_any_key(nested, keys)), - _ => None, - } -} - -fn find_first_array(value: &Value, keys: &[&str]) -> Option> { - match value { - Value::Object(map) => { - for key in keys { - if let Some(Value::Array(values)) = map.get(*key) { - return Some(values.clone()); - } - } - map.values() - .find_map(|nested| find_first_array(nested, keys)) - } - Value::Array(values) => values - .iter() - .find_map(|nested| find_first_array(nested, keys)), - _ => None, - } -} - -fn first_string(value: &Value, keys: &[&str]) -> Option { - value - .as_object() - .and_then(|map| keys.iter().find_map(|key| parse_string(map.get(*key)))) -} - -fn find_first_string(value: &Value, keys: &[&str]) -> Option { - first_string(value, keys).or_else(|| match value { - Value::Object(map) => map - .values() - .find_map(|nested| find_first_string(nested, keys)), - Value::Array(values) => values - .iter() - .find_map(|nested| find_first_string(nested, keys)), - _ => None, - }) -} - -fn first_f64(value: &Value, keys: &[&str]) -> Option { - value - .as_object() - .and_then(|map| keys.iter().find_map(|key| parse_f64(map.get(*key)))) -} - -fn find_first_f64(value: &Value, keys: &[&str]) -> Option { - first_f64(value, keys).or_else(|| match value { - Value::Object(map) => map.values().find_map(|nested| find_first_f64(nested, keys)), - Value::Array(values) => values - .iter() - .find_map(|nested| find_first_f64(nested, keys)), - _ => None, - }) -} - -fn find_first_i64(value: &Value, keys: &[&str]) -> Option { - match value { - Value::Object(map) => { - for key in keys { - if let Some(parsed) = parse_i64(map.get(*key)) { - return Some(parsed); - } - } - map.values().find_map(|nested| find_first_i64(nested, keys)) - } - Value::Array(values) => values - .iter() - .find_map(|nested| find_first_i64(nested, keys)), - _ => None, - } -} - -fn find_first_bool(value: &Value, keys: &[&str]) -> Option { - match value { - Value::Object(map) => { - for key in keys { - if let Some(parsed) = parse_bool(map.get(*key)) { - return Some(parsed); - } - } - map.values() - .find_map(|nested| find_first_bool(nested, keys)) - } - Value::Array(values) => values - .iter() - .find_map(|nested| find_first_bool(nested, keys)), - _ => None, - } -} - -fn first_date(value: &Value, keys: &[&str]) -> Option> { - value - .as_object() - .and_then(|map| keys.iter().find_map(|key| parse_date(map.get(*key)))) -} - -fn find_first_date(value: &Value, keys: &[&str]) -> Option> { - first_date(value, keys).or_else(|| match value { - Value::Object(map) => map - .values() - .find_map(|nested| find_first_date(nested, keys)), - Value::Array(values) => values - .iter() - .find_map(|nested| find_first_date(nested, keys)), - _ => None, - }) -} - -fn parse_string(value: Option<&Value>) -> Option { - value? - .as_str() - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) -} - -fn parse_f64(value: Option<&Value>) -> Option { - match value? { - Value::Number(number) => number.as_f64(), - Value::String(text) => text.trim().replace(',', "").parse().ok(), - _ => None, - } - .filter(|v| v.is_finite()) -} - -fn parse_i64(value: Option<&Value>) -> Option { - match value? { - Value::Number(number) => number.as_i64().or_else(|| { - number.as_f64().map(|v| { - // Quota numbers are small integers; any fractional part is discarded deliberately. - #[expect(clippy::cast_possible_truncation, reason = "quota counts fit i64")] - let v = v as i64; - v - }) - }), - Value::String(text) => text.trim().replace(',', "").parse().ok(), - _ => None, - } -} - -fn parse_bool(value: Option<&Value>) -> Option { - match value? { - Value::Bool(flag) => Some(*flag), - Value::Number(number) => number.as_i64().map(|v| v != 0), - Value::String(text) => match text.trim().to_lowercase().as_str() { - "true" | "1" | "yes" | "active" | "valid" | "normal" => Some(true), - "false" | "0" | "no" | "inactive" | "invalid" | "expired" => Some(false), - _ => None, - }, - _ => None, - } -} - -fn parse_date(value: Option<&Value>) -> Option> { - if let Some(raw) = parse_i64(value) { - if raw > 1_000_000_000_000 { - return Utc.timestamp_opt(raw / 1000, 0).single(); - } - if raw > 1_000_000_000 { - return Utc.timestamp_opt(raw, 0).single(); - } - } - let text = parse_string(value)?; - if let Ok(date) = DateTime::parse_from_rfc3339(&text) { - return Some(date.with_timezone(&Utc)); - } - if let Ok(date) = NaiveDate::parse_from_str(&text, "%Y-%m-%d") - && let Some(date_time) = date.and_hms_opt(0, 0, 0) - { - return Some(date_time.and_utc()); - } - for format in ["%Y-%m-%d %H:%M", "%Y-%m-%d %H:%M:%S"] { - if let Ok(date) = NaiveDateTime::parse_from_str(&text, format) { - return Some(date.and_utc()); - } - } - None -} - -fn active_signal_score(value: &Value) -> i32 { - let status = first_string(value, &["status", "instanceStatus", "state", "Status"]) - .unwrap_or_default() - .to_uppercase(); - if ["VALID", "ACTIVE", "NORMAL"].contains(&status.as_str()) { - return 3; - } - if [ - "EXPIRED", - "INVALID", - "INACTIVE", - "DISABLED", - "TERMINATED", - "STOPPED", - ] - .contains(&status.as_str()) - { - return -1; - } - parse_bool( - value - .as_object() - .and_then(|map| map.get("isActive").or_else(|| map.get("active"))), - ) - .map(|active| if active { 3 } else { -1 }) - .unwrap_or(0) -} - fn used_percent(used: Option, total: Option, remaining: Option) -> Option { let total = total.filter(|total| *total > 0.0)?; let used = used.or_else(|| remaining.map(|remaining| total - remaining))?; @@ -1241,238 +758,4 @@ fn format_count_decimal(raw: &str) -> String { } #[cfg(test)] -mod tests { - use super::*; - - #[test] - fn parses_current_token_plan_5h_and_weekly() { - let inner = r#"{ - "code": 0, - "data": { - "per5HourPercentage": 0.03, - "per5HourResetTime": 1700003600000, - "per1WeekPercentage": 0.01, - "per1WeekResetTime": 1700086400000 - }, - "success": true - }"#; - let payload = serde_json::json!({ - "data": { - "DataV2": { - "data": inner, - }, - }, - "httpStatusCode": 200, - }); - - let snapshot = - QwenCloudProvider::parse(payload.to_string().as_bytes(), None, None).unwrap(); - assert_eq!(snapshot.five_hour_used_percent, Some(3.0)); - assert_eq!( - snapshot.five_hour_resets_at, - Some(Utc.timestamp_opt(1_700_003_600, 0).single().unwrap()) - ); - assert_eq!(snapshot.weekly_used_percent, Some(1.0)); - assert_eq!( - snapshot.weekly_resets_at, - Some(Utc.timestamp_opt(1_700_086_400, 0).single().unwrap()) - ); - - let usage = QwenCloudProvider::new() - .snapshot_to_usage(snapshot) - .unwrap(); - assert_eq!(usage.primary.used_percent, 3.0); - assert_eq!(usage.primary.window_minutes, Some(FIVE_HOUR_MINUTES)); - assert_eq!(usage.primary_label, None); - assert_eq!(usage.secondary.as_ref().map(|w| w.used_percent), Some(1.0)); - assert_eq!( - usage.secondary.as_ref().and_then(|w| w.window_minutes), - Some(WEEKLY_MINUTES) - ); - } - - #[test] - fn parses_personal_usage_fixture_shape() { - let payload = serde_json::json!({ - "code": "200", - "data": { - "DataV2": { - "data": { - "success": true, - "data": { - "per5HourPercentage": 0.0009973083333333333, - "per5HourResetTime": 1784813220000_i64, - "per1WeekPercentage": 0.0003014725, - "per1WeekResetTime": 1785234900000_i64 - } - }, - "success": true, - "httpStatus": 200 - } - }, - "successResponse": true - }); - let usage = QwenCloudProvider::new() - .snapshot_to_usage( - QwenCloudProvider::parse(payload.to_string().as_bytes(), None, None).unwrap(), - ) - .unwrap(); - assert!((usage.primary.used_percent - 0.09973083333333333).abs() < 1e-9); - assert_eq!(usage.primary.window_minutes, Some(300)); - assert!(usage.secondary.is_some()); - assert_eq!( - usage.secondary.as_ref().and_then(|w| w.window_minutes), - Some(10080) - ); - } - - #[test] - fn parses_nested_equity_list_legacy() { - let payload = serde_json::json!({ - "code": "200", - "successResponse": true, - "data": { - "TotalCount": 1, - "Data": [ - { - "InstanceCode": "qwen-token-plan", - "Status": "NORMAL", - "EndTime": 1_701_000_000_000_i64, - "EquityList": [ - { - "Type": "CREDITS", - "CycleTotalValue": "1000", - "CycleSurplusValue": "875" - } - ] - } - ] - } - }); - let snapshot = - QwenCloudProvider::parse(payload.to_string().as_bytes(), None, None).unwrap(); - assert_eq!(snapshot.total_quota, Some(1000.0)); - assert_eq!(snapshot.remaining_quota, Some(875.0)); - let usage = QwenCloudProvider::new() - .snapshot_to_usage(snapshot) - .unwrap(); - assert_eq!(usage.primary.used_percent, 12.5); - assert_eq!(usage.primary.window_minutes, Some(LEGACY_MINUTES)); - } - - #[test] - fn parses_flat_subscription_summary_legacy() { - let payload = serde_json::json!({ - "Success": true, - "Data": { - "TotalCount": 1, - "TotalValue": 2000, - "TotalSurplusValue": 1500 - } - }); - let snapshot = - QwenCloudProvider::parse(payload.to_string().as_bytes(), None, None).unwrap(); - assert_eq!(snapshot.total_quota, Some(2000.0)); - assert_eq!(snapshot.remaining_quota, Some(1500.0)); - assert_eq!( - used_percent( - snapshot.used_quota, - snapshot.total_quota, - snapshot.remaining_quota - ), - Some(25.0) - ); - } - - #[test] - fn login_payload_maps_to_auth_required() { - let err = QwenCloudProvider::parse( - br#"{"code":"ConsoleNeedLogin","message":"You need to log in.","successResponse":false}"#, - None, - None, - ) - .unwrap_err(); - assert!(matches!(err, ProviderError::AuthRequired)); - } - - #[test] - fn attaches_plan_name_from_subscription() { - let usage_payload = serde_json::json!({ - "data": { - "per5HourPercentage": 0.1, - "per5HourResetTime": 1700003600000_i64, - "per1WeekPercentage": 0.2, - "per1WeekResetTime": 1700086400000_i64 - } - }); - let subscription = serde_json::json!({ - "data": { "specCode": "pro" } - }); - let snapshot = QwenCloudProvider::parse( - usage_payload.to_string().as_bytes(), - Some(subscription.to_string().as_bytes()), - None, - ) - .unwrap(); - assert_eq!(snapshot.plan_name.as_deref(), Some("Pro")); - let usage = QwenCloudProvider::new() - .snapshot_to_usage(snapshot) - .unwrap(); - assert_eq!(usage.login_method.as_deref(), Some("Pro")); - } - - #[test] - fn extracts_sec_token_from_html() { - assert_eq!( - extract_sec_token(r#""#).as_deref(), - Some("qwen-html-token") - ); - assert_eq!( - cookie_value("login_aliyunid_csrf", "foo=bar; login_aliyunid_csrf=tok"), - Some("tok".to_string()) - ); - } - - #[test] - fn metadata_labels_match_upstream() { - let provider = QwenCloudProvider::new(); - assert_eq!(provider.metadata().session_label, "5-hour"); - assert_eq!(provider.metadata().weekly_label, "Weekly"); - assert!(!provider.metadata().default_enabled); - assert_eq!(provider.metadata().dashboard_url, Some(DASHBOARD_URL)); - } - - #[test] - fn weekly_only_shape_promotes_weekly_to_primary() { - // Real-world fixture: the account exposes only the weekly window, so - // there is no per5HourPercentage at all. Previously this failed with - // "Qwen Cloud usage windows missing"; now the weekly window becomes - // the primary window instead. - let payload = serde_json::json!({ - "code": "200", - "data": { - "DataV2": { - "data": { - "success": true, - "data": { - "per1WeekPercentage": 0.8439116574633999, - "per1WeekResetTime": 1785234900000_i64 - } - }, - "success": true, - "httpStatus": 200 - } - }, - "successResponse": true - }); - let usage = QwenCloudProvider::new() - .snapshot_to_usage( - QwenCloudProvider::parse(payload.to_string().as_bytes(), None, None).unwrap(), - ) - .unwrap(); - assert!((usage.primary.used_percent - 84.39116574634).abs() < 1e-9); - assert_eq!(usage.primary.window_minutes, Some(WEEKLY_MINUTES)); - assert_eq!(usage.primary_label.as_deref(), Some("Weekly")); - assert!(usage.secondary.is_none()); - } -} +mod tests; diff --git a/rust/src/providers/qwencloud/monthly_tests.rs b/rust/src/providers/qwencloud/monthly_tests.rs index 4e677716be..00a60add23 100644 --- a/rust/src/providers/qwencloud/monthly_tests.rs +++ b/rust/src/providers/qwencloud/monthly_tests.rs @@ -1,6 +1,7 @@ //! Monthly token-plan window (upstream 0.66.0, `TokenPlanMonthlyWindowTests`). use super::*; +use chrono::TimeZone; const MONTHLY: &str = r#"{"per1MonthPercentage":0.25,"per1MonthResetTime":1791043200000}"#; diff --git a/rust/src/providers/qwencloud/tests.rs b/rust/src/providers/qwencloud/tests.rs new file mode 100644 index 0000000000..b47e2b5226 --- /dev/null +++ b/rust/src/providers/qwencloud/tests.rs @@ -0,0 +1,269 @@ +use super::*; +use chrono::TimeZone; + +#[test] +fn parses_current_token_plan_5h_and_weekly() { + let inner = r#"{ + "code": 0, + "data": { + "per5HourPercentage": 0.03, + "per5HourResetTime": 1700003600000, + "per1WeekPercentage": 0.01, + "per1WeekResetTime": 1700086400000 + }, + "success": true + }"#; + let payload = serde_json::json!({ + "data": { + "DataV2": { + "data": inner, + }, + }, + "httpStatusCode": 200, + }); + + let snapshot = QwenCloudProvider::parse(payload.to_string().as_bytes(), None, None).unwrap(); + assert_eq!(snapshot.five_hour_used_percent, Some(3.0)); + assert_eq!( + snapshot.five_hour_resets_at, + Some(Utc.timestamp_opt(1_700_003_600, 0).single().unwrap()) + ); + assert_eq!(snapshot.weekly_used_percent, Some(1.0)); + assert_eq!( + snapshot.weekly_resets_at, + Some(Utc.timestamp_opt(1_700_086_400, 0).single().unwrap()) + ); + + let usage = QwenCloudProvider::new() + .snapshot_to_usage(snapshot) + .unwrap(); + assert_eq!(usage.primary.used_percent, 3.0); + assert_eq!(usage.primary.window_minutes, Some(FIVE_HOUR_MINUTES)); + assert_eq!(usage.primary_label, None); + assert_eq!(usage.secondary.as_ref().map(|w| w.used_percent), Some(1.0)); + assert_eq!( + usage.secondary.as_ref().and_then(|w| w.window_minutes), + Some(WEEKLY_MINUTES) + ); +} + +#[test] +fn parses_personal_usage_fixture_shape() { + let payload = serde_json::json!({ + "code": "200", + "data": { + "DataV2": { + "data": { + "success": true, + "data": { + "per5HourPercentage": 0.0009973083333333333, + "per5HourResetTime": 1784813220000_i64, + "per1WeekPercentage": 0.0003014725, + "per1WeekResetTime": 1785234900000_i64 + } + }, + "success": true, + "httpStatus": 200 + } + }, + "successResponse": true + }); + let usage = QwenCloudProvider::new() + .snapshot_to_usage( + QwenCloudProvider::parse(payload.to_string().as_bytes(), None, None).unwrap(), + ) + .unwrap(); + assert!((usage.primary.used_percent - 0.09973083333333333).abs() < 1e-9); + assert_eq!(usage.primary.window_minutes, Some(300)); + assert!(usage.secondary.is_some()); + assert_eq!( + usage.secondary.as_ref().and_then(|w| w.window_minutes), + Some(10080) + ); +} + +#[test] +fn parses_nested_equity_list_legacy() { + let payload = serde_json::json!({ + "code": "200", + "successResponse": true, + "data": { + "TotalCount": 1, + "Data": [ + { + "InstanceCode": "qwen-token-plan", + "Status": "NORMAL", + "EndTime": 1_701_000_000_000_i64, + "EquityList": [ + { + "Type": "CREDITS", + "CycleTotalValue": "1000", + "CycleSurplusValue": "875" + } + ] + } + ] + } + }); + let snapshot = QwenCloudProvider::parse(payload.to_string().as_bytes(), None, None).unwrap(); + assert_eq!(snapshot.total_quota, Some(1000.0)); + assert_eq!(snapshot.remaining_quota, Some(875.0)); + let usage = QwenCloudProvider::new() + .snapshot_to_usage(snapshot) + .unwrap(); + assert_eq!(usage.primary.used_percent, 12.5); + assert_eq!(usage.primary.window_minutes, Some(LEGACY_MINUTES)); +} + +#[test] +fn parses_flat_subscription_summary_legacy() { + let payload = serde_json::json!({ + "Success": true, + "Data": { + "TotalCount": 1, + "TotalValue": 2000, + "TotalSurplusValue": 1500 + } + }); + let snapshot = QwenCloudProvider::parse(payload.to_string().as_bytes(), None, None).unwrap(); + assert_eq!(snapshot.total_quota, Some(2000.0)); + assert_eq!(snapshot.remaining_quota, Some(1500.0)); + assert_eq!( + used_percent( + snapshot.used_quota, + snapshot.total_quota, + snapshot.remaining_quota + ), + Some(25.0) + ); +} + +#[test] +fn login_payload_maps_to_auth_required() { + let err = QwenCloudProvider::parse( + br#"{"code":"ConsoleNeedLogin","message":"You need to log in.","successResponse":false}"#, + None, + None, + ) + .unwrap_err(); + assert!(matches!(err, ProviderError::AuthRequired)); +} + +#[test] +fn attaches_plan_name_from_subscription() { + let usage_payload = serde_json::json!({ + "data": { + "per5HourPercentage": 0.1, + "per5HourResetTime": 1700003600000_i64, + "per1WeekPercentage": 0.2, + "per1WeekResetTime": 1700086400000_i64 + } + }); + let subscription = serde_json::json!({ + "data": { "specCode": "pro" } + }); + let snapshot = QwenCloudProvider::parse( + usage_payload.to_string().as_bytes(), + Some(subscription.to_string().as_bytes()), + None, + ) + .unwrap(); + assert_eq!(snapshot.plan_name.as_deref(), Some("Pro")); + let usage = QwenCloudProvider::new() + .snapshot_to_usage(snapshot) + .unwrap(); + assert_eq!(usage.login_method.as_deref(), Some("Pro")); +} + +#[test] +fn extracts_sec_token_from_html() { + assert_eq!( + extract_sec_token(r#""#).as_deref(), + Some("qwen-html-token") + ); + assert_eq!( + cookie_value("login_aliyunid_csrf", "foo=bar; login_aliyunid_csrf=tok"), + Some("tok".to_string()) + ); +} + +#[test] +fn metadata_labels_match_upstream() { + let provider = QwenCloudProvider::new(); + assert_eq!(provider.metadata().session_label, "5-hour"); + assert_eq!(provider.metadata().weekly_label, "Weekly"); + assert!(!provider.metadata().default_enabled); + assert_eq!(provider.metadata().dashboard_url, Some(DASHBOARD_URL)); +} + +#[test] +fn weekly_only_shape_promotes_weekly_to_primary() { + // Real-world fixture: the account exposes only the weekly window, so + // there is no per5HourPercentage at all. Previously this failed with + // "Qwen Cloud usage windows missing"; now the weekly window becomes + // the primary window instead. + let payload = serde_json::json!({ + "code": "200", + "data": { + "DataV2": { + "data": { + "success": true, + "data": { + "per1WeekPercentage": 0.8439116574633999, + "per1WeekResetTime": 1785234900000_i64 + } + }, + "success": true, + "httpStatus": 200 + } + }, + "successResponse": true + }); + let usage = QwenCloudProvider::new() + .snapshot_to_usage( + QwenCloudProvider::parse(payload.to_string().as_bytes(), None, None).unwrap(), + ) + .unwrap(); + assert!((usage.primary.used_percent - 84.39116574634).abs() < 1e-9); + assert_eq!(usage.primary.window_minutes, Some(WEEKLY_MINUTES)); + assert_eq!(usage.primary_label.as_deref(), Some("Weekly")); + assert!(usage.secondary.is_none()); +} + +/// Pins the console `params` payload: the cornerstone fields, the optional +/// `cna` anonymous id, and the caller's data parameters. +#[test] +fn params_json_wraps_data_with_cornerstone_fields() { + for (cookie, anonymous_id) in [("cna=anon-1; other=x", Some("anon-1")), ("other=x", None)] { + let mut data = Map::new(); + data.insert("commodityCode".into(), json!("sfm_tokenplan_public_cn")); + let params = build_params_json("zeldaEasy.test.api", data, cookie); + let value: Value = serde_json::from_str(¶ms).unwrap(); + let trace = value["Data"]["cornerstoneParam"]["feTraceId"] + .as_str() + .unwrap(); + assert_eq!(trace.len(), 36); + assert_eq!(trace, trace.to_lowercase()); + let mut cornerstone = json!({ + "feTraceId": trace, + "feURL": DASHBOARD_URL, + "protocol": "V2", + "console": "ONE_CONSOLE", + "productCode": "p_efm", + "domain": "home.qwencloud.com", + "consoleSite": "QWENCLOUD", + "userNickName": "", + "userPrincipalName": "", + "xsp_lang": "en-US", + }); + if let Some(id) = anonymous_id { + cornerstone["X-Anonymous-Id"] = json!(id); + } + let expected = json!({ + "Api": "zeldaEasy.test.api", + "V": "1.0", + "Data": {"commodityCode": "sfm_tokenplan_public_cn", "cornerstoneParam": cornerstone}, + }); + assert_eq!(params, expected.to_string()); + } +} diff --git a/rust/src/providers/sub2api/mod.rs b/rust/src/providers/sub2api/mod.rs index 4016f0f625..c2aeec26b9 100644 --- a/rust/src/providers/sub2api/mod.rs +++ b/rust/src/providers/sub2api/mod.rs @@ -85,11 +85,10 @@ struct UsageResponse { unit: Option, balance: Option, quota: Option, - #[serde(default, rename = "rate_limits")] + #[serde(default)] rate_limits: Option>, subscription: Option, usage: Option, - #[serde(rename = "expires_at")] expires_at: Option, } @@ -107,25 +106,17 @@ struct RateLimitResponse { limit: f64, used: f64, remaining: f64, - #[serde(rename = "reset_at")] reset_at: Option, } #[derive(Debug, Deserialize)] struct SubscriptionResponse { - #[serde(rename = "daily_usage_usd")] daily_usage_usd: Option, - #[serde(rename = "weekly_usage_usd")] weekly_usage_usd: Option, - #[serde(rename = "monthly_usage_usd")] monthly_usage_usd: Option, - #[serde(rename = "daily_limit_usd")] daily_limit_usd: Option, - #[serde(rename = "weekly_limit_usd")] weekly_limit_usd: Option, - #[serde(rename = "monthly_limit_usd")] monthly_limit_usd: Option, - #[serde(rename = "expires_at")] expires_at: Option, } @@ -138,9 +129,7 @@ struct UsageBlockResponse { #[derive(Debug, Deserialize)] struct TotalsResponse { requests: Option, - #[serde(rename = "total_tokens")] total_tokens: Option, - #[serde(rename = "actual_cost")] actual_cost: Option, } @@ -515,34 +504,17 @@ fn snapshot_from_parsed(parsed: ParsedUsage) -> ProviderFetchResult { )) } else if let Some(first) = parsed.rate_limits.first() { // Rate-limit-only payload: promote the first window to primary. - UsageSnapshot::new(RateWindow::with_details( - used_percent(first.used, first.limit), - window_minutes(&first.window), - first.reset_at, - Some(amount_description(first.used, first.limit, &parsed.unit)), - )) + UsageSnapshot::new(rate_limit_window(first, &parsed.unit)) } else { UsageSnapshot::new(RateWindow::informational("Key quota")) }; let skip_first_rate_limit = parsed.quota.is_none() && !parsed.rate_limits.is_empty(); - for (idx, rate_limit) in parsed.rate_limits.iter().enumerate() { - if skip_first_rate_limit && idx == 0 { - continue; - } - snap = snap.with_extra_rate_window( - rate_limit.window.clone(), - rate_limit_title(&rate_limit.window), - RateWindow::with_details( - used_percent(rate_limit.used, rate_limit.limit), - window_minutes(&rate_limit.window), - rate_limit.reset_at, - Some(amount_description( - rate_limit.used, - rate_limit.limit, - &parsed.unit, - )), - ), - ); + for rate_limit in parsed + .rate_limits + .iter() + .skip(usize::from(skip_first_rate_limit)) + { + snap = with_rate_limit_extra(snap, rate_limit, &parsed.unit); } snap } @@ -560,39 +532,18 @@ fn snapshot_from_parsed(parsed: ParsedUsage) -> ProviderFetchResult { } Kind::Unknown => { // Prefer any totals over a blank "No quota data" row. - if let Some(today) = &parsed.today { - UsageSnapshot::new(RateWindow::informational(totals_description( - today, - &parsed.unit, - ))) - } else if let Some(total) = &parsed.total { - UsageSnapshot::new(RateWindow::informational(totals_description( - total, - &parsed.unit, - ))) - } else { - UsageSnapshot::new(RateWindow::informational("No quota data")) - } + let description = parsed.today.as_ref().or(parsed.total.as_ref()).map_or_else( + || "No quota data".to_string(), + |t| totals_description(t, &parsed.unit), + ); + UsageSnapshot::new(RateWindow::informational(description)) } }; // Rate-limit extras for subscription/wallet kinds that also ship them. if kind != Kind::KeyQuota { for rate_limit in &parsed.rate_limits { - snapshot = snapshot.with_extra_rate_window( - rate_limit.window.clone(), - rate_limit_title(&rate_limit.window), - RateWindow::with_details( - used_percent(rate_limit.used, rate_limit.limit), - window_minutes(&rate_limit.window), - rate_limit.reset_at, - Some(amount_description( - rate_limit.used, - rate_limit.limit, - &parsed.unit, - )), - ), - ); + snapshot = with_rate_limit_extra(snapshot, rate_limit, &parsed.unit); } } @@ -659,6 +610,27 @@ fn snapshot_from_parsed(parsed: ParsedUsage) -> ProviderFetchResult { result } +fn rate_limit_window(rate_limit: &ParsedRateLimit, unit: &str) -> RateWindow { + RateWindow::with_details( + used_percent(rate_limit.used, rate_limit.limit), + window_minutes(&rate_limit.window), + rate_limit.reset_at, + Some(amount_description(rate_limit.used, rate_limit.limit, unit)), + ) +} + +fn with_rate_limit_extra( + snapshot: UsageSnapshot, + rate_limit: &ParsedRateLimit, + unit: &str, +) -> UsageSnapshot { + snapshot.with_extra_rate_window( + rate_limit.window.clone(), + rate_limit_title(&rate_limit.window), + rate_limit_window(rate_limit, unit), + ) +} + fn classify_usage_kind(parsed: &ParsedUsage) -> Kind { if parsed.subscription.is_some() { Kind::Subscription @@ -773,284 +745,4 @@ fn totals_description(totals: &ParsedTotals, unit: &str) -> String { } #[cfg(test)] -mod tests { - use super::*; - - #[test] - fn parses_quota_limited_key_usage() { - let json = r#" - { - "mode": "quota_limited", - "isValid": true, - "status": "active", - "remaining": 75, - "unit": "USD", - "quota": { - "limit": 100, - "used": 25, - "remaining": 75, - "unit": "USD" - }, - "rate_limits": [ - { - "window": "5h", - "limit": 20, - "used": 5, - "remaining": 15, - "reset_at": "2026-07-11T12:30:00Z" - }, - { - "window": "7d", - "limit": 200, - "used": 40, - "remaining": 160 - } - ], - "expires_at": "2026-08-01T00:00:00Z", - "usage": { - "today": { - "requests": 4, - "total_tokens": 1200, - "actual_cost": 1.25 - }, - "total": { - "requests": 40, - "total_tokens": 12000, - "actual_cost": 25 - } - } - } - "#; - - let parsed = parse_usage_body(json).unwrap(); - assert_eq!(parsed.mode, "quota_limited"); - assert_eq!(parsed.quota.as_ref().unwrap().remaining, 75.0); - assert_eq!(parsed.rate_limits.len(), 2); - assert_eq!(parsed.today.as_ref().unwrap().total_tokens, 1200); - - let result = snapshot_from_parsed(parsed); - assert!((result.usage.primary.used_percent - 25.0).abs() < f64::EPSILON); - assert!(result.cost.is_none()); - assert!( - result - .usage - .extra_rate_windows - .iter() - .any(|w| w.id == "5h" && w.window.window_minutes == Some(300)) - ); - assert!(result.usage.extra_rate_windows.iter().any(|w| { - w.id == "today" - && w.window - .reset_description - .as_deref() - .is_some_and(|d| d.contains("1200 tokens")) - })); - assert!( - result - .usage - .extra_rate_windows - .iter() - .any(|w| w.id == "expires") - ); - } - - #[test] - fn parses_subscription_usage_windows() { - let json = r#" - { - "mode": "unrestricted", - "isValid": true, - "planName": "Claude Team", - "remaining": 8, - "unit": "USD", - "subscription": { - "daily_usage_usd": 2, - "weekly_usage_usd": 10, - "monthly_usage_usd": 30, - "daily_limit_usd": 10, - "weekly_limit_usd": 40, - "monthly_limit_usd": 100, - "expires_at": "2026-08-15T00:00:00.123Z" - } - } - "#; - - let result = snapshot_from_parsed(parse_usage_body(json).unwrap()); - assert!((result.usage.primary.used_percent - 20.0).abs() < f64::EPSILON); - assert!( - (result.usage.secondary.as_ref().unwrap().used_percent - 25.0).abs() < f64::EPSILON - ); - assert!((result.usage.tertiary.as_ref().unwrap().used_percent - 30.0).abs() < f64::EPSILON); - assert!(result.usage.account_organization.is_none()); - assert_eq!( - result.usage.login_method.as_deref(), - Some("Claude Team (unrestricted)") - ); - assert!( - result - .usage - .extra_rate_windows - .iter() - .any(|w| w.id == "expires") - ); - } - - #[test] - fn subscription_without_limits_shows_usage_not_empty() { - let json = r#" - { - "mode": "unrestricted", - "planName": "Soft plan", - "unit": "USD", - "subscription": { - "daily_usage_usd": 2.5, - "weekly_usage_usd": 10, - "monthly_usage_usd": 30 - }, - "balance": 12.0 - } - "#; - let result = snapshot_from_parsed(parse_usage_body(json).unwrap()); - assert!(result.usage.primary.is_informational); - assert!( - result - .usage - .primary - .reset_description - .as_deref() - .is_some_and(|d| d.contains("$2.50") && d.contains("day")) - ); - assert!(result.cost.is_some()); - assert_eq!(result.cost.as_ref().unwrap().limit, Some(12.0)); - assert!(result.usage.account_organization.is_none()); - assert!( - result - .usage - .login_method - .as_deref() - .is_some_and(|m| m.starts_with("Soft plan")) - ); - } - - #[test] - fn preserves_authoritative_subscription_windows() { - let json = r#" - { - "mode": "unrestricted", - "subscription": { - "daily_usage_usd": 120.23, - "weekly_usage_usd": 229.20, - "monthly_usage_usd": 1296.23, - "daily_limit_usd": 120, - "weekly_limit_usd": 700, - "monthly_limit_usd": 2800 - } - } - "#; - - let result = snapshot_from_parsed(parse_usage_body(json).unwrap()); - assert!((result.usage.primary.used_percent - 100.0).abs() < f64::EPSILON); - assert!( - (result.usage.secondary.as_ref().unwrap().used_percent - (229.20 / 700.0 * 100.0)) - .abs() - < 0.001 - ); - assert_eq!( - result.usage.primary.reset_description.as_deref(), - Some("$120.23 / $120.00") - ); - assert_eq!( - result - .usage - .secondary - .as_ref() - .and_then(|w| w.reset_description.as_deref()), - Some("$229.20 / $700.00") - ); - } - - #[test] - fn parses_wallet_balance_only() { - let json = r#" - { - "mode": "unrestricted", - "isValid": true, - "planName": "Wallet plan", - "remaining": 42.5, - "unit": "USD", - "balance": 42.5 - } - "#; - - let result = snapshot_from_parsed(parse_usage_body(json).unwrap()); - assert!(result.usage.primary.is_informational); - assert_eq!( - result.usage.primary.reset_description.as_deref(), - Some("$42.50 balance") - ); - assert_eq!( - result.usage.login_method.as_deref(), - Some("Wallet plan (unrestricted)") - ); - let cost = result.cost.unwrap(); - assert_eq!(cost.limit, Some(42.5)); - assert_eq!(cost.period, "balance"); - } - - #[test] - fn invalid_credentials_flag_is_detected() { - let json = r#"{"mode":"unrestricted","isValid":false}"#; - let parsed = parse_usage_body(json).unwrap(); - assert!(!parsed.is_valid); - } - - #[test] - fn usage_url_accepts_root_versioned_and_complete_urls() { - let root = Url::parse("https://api.example.com").unwrap(); - assert_eq!( - usage_url(&root).unwrap().as_str(), - "https://api.example.com/v1/usage" - ); - - let versioned = Url::parse("https://api.example.com/v1").unwrap(); - assert_eq!( - usage_url(&versioned).unwrap().as_str(), - "https://api.example.com/v1/usage" - ); - - let complete = Url::parse("https://api.example.com/v1/usage").unwrap(); - assert_eq!( - usage_url(&complete).unwrap().as_str(), - "https://api.example.com/v1/usage" - ); - } - - #[test] - fn settings_allow_https_and_loopback_http_only() { - assert!(validated_sub2api_base_url("https://api.example.com").is_ok()); - assert!(validated_sub2api_base_url("http://127.0.0.1:8080").is_ok()); - assert!(validated_sub2api_base_url("http://api.example.com").is_err()); - assert!(validated_sub2api_base_url("https://user:pass@api.example.com").is_err()); - assert!(validated_sub2api_base_url("https://api.example.com?token=secret").is_err()); - assert!(validated_sub2api_base_url("https://api.example.com#fragment").is_err()); - } - - #[test] - fn cleans_quoted_env_values() { - assert_eq!( - clean_env_value(" \"sk-test\" ").as_deref(), - Some("sk-test") - ); - assert_eq!(clean_env_value("''"), None); - } - - #[test] - fn usage_request_includes_days_and_timezone() { - let base = Url::parse("https://api.example.com").unwrap(); - let url = usage_request_url(&base).unwrap(); - let query: std::collections::HashMap<_, _> = url.query_pairs().into_owned().collect(); - assert_eq!(query.get("days").map(String::as_str), Some("30")); - assert!(query.contains_key("timezone")); - assert!(!query.get("timezone").unwrap().is_empty()); - } -} +mod tests; diff --git a/rust/src/providers/sub2api/tests.rs b/rust/src/providers/sub2api/tests.rs new file mode 100644 index 0000000000..7f5e4dec20 --- /dev/null +++ b/rust/src/providers/sub2api/tests.rs @@ -0,0 +1,276 @@ +use super::*; + +#[test] +fn parses_quota_limited_key_usage() { + let json = r#" + { + "mode": "quota_limited", + "isValid": true, + "status": "active", + "remaining": 75, + "unit": "USD", + "quota": { + "limit": 100, + "used": 25, + "remaining": 75, + "unit": "USD" + }, + "rate_limits": [ + { + "window": "5h", + "limit": 20, + "used": 5, + "remaining": 15, + "reset_at": "2026-07-11T12:30:00Z" + }, + { + "window": "7d", + "limit": 200, + "used": 40, + "remaining": 160 + } + ], + "expires_at": "2026-08-01T00:00:00Z", + "usage": { + "today": { + "requests": 4, + "total_tokens": 1200, + "actual_cost": 1.25 + }, + "total": { + "requests": 40, + "total_tokens": 12000, + "actual_cost": 25 + } + } + } + "#; + + let parsed = parse_usage_body(json).unwrap(); + assert_eq!(parsed.mode, "quota_limited"); + assert_eq!(parsed.quota.as_ref().unwrap().remaining, 75.0); + assert_eq!(parsed.rate_limits.len(), 2); + assert_eq!(parsed.today.as_ref().unwrap().total_tokens, 1200); + + let result = snapshot_from_parsed(parsed); + assert!((result.usage.primary.used_percent - 25.0).abs() < f64::EPSILON); + assert!(result.cost.is_none()); + assert!( + result + .usage + .extra_rate_windows + .iter() + .any(|w| w.id == "5h" && w.window.window_minutes == Some(300)) + ); + assert!(result.usage.extra_rate_windows.iter().any(|w| { + w.id == "today" + && w.window + .reset_description + .as_deref() + .is_some_and(|d| d.contains("1200 tokens")) + })); + assert!( + result + .usage + .extra_rate_windows + .iter() + .any(|w| w.id == "expires") + ); +} + +#[test] +fn parses_subscription_usage_windows() { + let json = r#" + { + "mode": "unrestricted", + "isValid": true, + "planName": "Claude Team", + "remaining": 8, + "unit": "USD", + "subscription": { + "daily_usage_usd": 2, + "weekly_usage_usd": 10, + "monthly_usage_usd": 30, + "daily_limit_usd": 10, + "weekly_limit_usd": 40, + "monthly_limit_usd": 100, + "expires_at": "2026-08-15T00:00:00.123Z" + } + } + "#; + + let result = snapshot_from_parsed(parse_usage_body(json).unwrap()); + assert!((result.usage.primary.used_percent - 20.0).abs() < f64::EPSILON); + assert!((result.usage.secondary.as_ref().unwrap().used_percent - 25.0).abs() < f64::EPSILON); + assert!((result.usage.tertiary.as_ref().unwrap().used_percent - 30.0).abs() < f64::EPSILON); + assert!(result.usage.account_organization.is_none()); + assert_eq!( + result.usage.login_method.as_deref(), + Some("Claude Team (unrestricted)") + ); + assert!( + result + .usage + .extra_rate_windows + .iter() + .any(|w| w.id == "expires") + ); +} + +#[test] +fn subscription_without_limits_shows_usage_not_empty() { + let json = r#" + { + "mode": "unrestricted", + "planName": "Soft plan", + "unit": "USD", + "subscription": { + "daily_usage_usd": 2.5, + "weekly_usage_usd": 10, + "monthly_usage_usd": 30 + }, + "balance": 12.0 + } + "#; + let result = snapshot_from_parsed(parse_usage_body(json).unwrap()); + assert!(result.usage.primary.is_informational); + assert!( + result + .usage + .primary + .reset_description + .as_deref() + .is_some_and(|d| d.contains("$2.50") && d.contains("day")) + ); + assert!(result.cost.is_some()); + assert_eq!(result.cost.as_ref().unwrap().limit, Some(12.0)); + assert!(result.usage.account_organization.is_none()); + assert!( + result + .usage + .login_method + .as_deref() + .is_some_and(|m| m.starts_with("Soft plan")) + ); +} + +#[test] +fn preserves_authoritative_subscription_windows() { + let json = r#" + { + "mode": "unrestricted", + "subscription": { + "daily_usage_usd": 120.23, + "weekly_usage_usd": 229.20, + "monthly_usage_usd": 1296.23, + "daily_limit_usd": 120, + "weekly_limit_usd": 700, + "monthly_limit_usd": 2800 + } + } + "#; + + let result = snapshot_from_parsed(parse_usage_body(json).unwrap()); + assert!((result.usage.primary.used_percent - 100.0).abs() < f64::EPSILON); + assert!( + (result.usage.secondary.as_ref().unwrap().used_percent - (229.20 / 700.0 * 100.0)).abs() + < 0.001 + ); + assert_eq!( + result.usage.primary.reset_description.as_deref(), + Some("$120.23 / $120.00") + ); + assert_eq!( + result + .usage + .secondary + .as_ref() + .and_then(|w| w.reset_description.as_deref()), + Some("$229.20 / $700.00") + ); +} + +#[test] +fn parses_wallet_balance_only() { + let json = r#" + { + "mode": "unrestricted", + "isValid": true, + "planName": "Wallet plan", + "remaining": 42.5, + "unit": "USD", + "balance": 42.5 + } + "#; + + let result = snapshot_from_parsed(parse_usage_body(json).unwrap()); + assert!(result.usage.primary.is_informational); + assert_eq!( + result.usage.primary.reset_description.as_deref(), + Some("$42.50 balance") + ); + assert_eq!( + result.usage.login_method.as_deref(), + Some("Wallet plan (unrestricted)") + ); + let cost = result.cost.unwrap(); + assert_eq!(cost.limit, Some(42.5)); + assert_eq!(cost.period, "balance"); +} + +#[test] +fn invalid_credentials_flag_is_detected() { + let json = r#"{"mode":"unrestricted","isValid":false}"#; + let parsed = parse_usage_body(json).unwrap(); + assert!(!parsed.is_valid); +} + +#[test] +fn usage_url_accepts_root_versioned_and_complete_urls() { + let root = Url::parse("https://api.example.com").unwrap(); + assert_eq!( + usage_url(&root).unwrap().as_str(), + "https://api.example.com/v1/usage" + ); + + let versioned = Url::parse("https://api.example.com/v1").unwrap(); + assert_eq!( + usage_url(&versioned).unwrap().as_str(), + "https://api.example.com/v1/usage" + ); + + let complete = Url::parse("https://api.example.com/v1/usage").unwrap(); + assert_eq!( + usage_url(&complete).unwrap().as_str(), + "https://api.example.com/v1/usage" + ); +} + +#[test] +fn settings_allow_https_and_loopback_http_only() { + assert!(validated_sub2api_base_url("https://api.example.com").is_ok()); + assert!(validated_sub2api_base_url("http://127.0.0.1:8080").is_ok()); + assert!(validated_sub2api_base_url("http://api.example.com").is_err()); + assert!(validated_sub2api_base_url("https://user:pass@api.example.com").is_err()); + assert!(validated_sub2api_base_url("https://api.example.com?token=secret").is_err()); + assert!(validated_sub2api_base_url("https://api.example.com#fragment").is_err()); +} + +#[test] +fn cleans_quoted_env_values() { + assert_eq!( + clean_env_value(" \"sk-test\" ").as_deref(), + Some("sk-test") + ); + assert_eq!(clean_env_value("''"), None); +} + +#[test] +fn usage_request_includes_days_and_timezone() { + let base = Url::parse("https://api.example.com").unwrap(); + let url = usage_request_url(&base).unwrap(); + let query: std::collections::HashMap<_, _> = url.query_pairs().into_owned().collect(); + assert_eq!(query.get("days").map(String::as_str), Some("30")); + assert!(query.contains_key("timezone")); + assert!(!query.get("timezone").unwrap().is_empty()); +} diff --git a/rust/src/providers/zai/mcp_details.rs b/rust/src/providers/zai/mcp_details.rs deleted file mode 100755 index d5e9e88f60..0000000000 --- a/rust/src/providers/zai/mcp_details.rs +++ /dev/null @@ -1,451 +0,0 @@ -//! Zai (z.ai) Usage Statistics and MCP Details -//! -//! Provides detailed usage tracking for Zai provider including: -//! - Token limits -//! - Time limits -//! - Per-model usage details for MCP (Model Context Protocol) - -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; - -/// Z.ai limit types -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] -pub enum ZaiLimitType { - /// Token-based limit - TokensLimit, - /// Credit-based limit (credit Coding Plans, upstream 0.49.0 #2724) - CreditLimit, - /// Time-based limit - TimeLimit, -} - -impl ZaiLimitType { - pub fn from_string(s: &str) -> Option { - match s { - "TOKENS_LIMIT" => Some(ZaiLimitType::TokensLimit), - "CREDIT_LIMIT" => Some(ZaiLimitType::CreditLimit), - "TIME_LIMIT" => Some(ZaiLimitType::TimeLimit), - _ => None, - } - } -} - -/// Z.ai limit time unit -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] -pub enum ZaiLimitUnit { - Unknown, - Days, - Hours, - Minutes, -} - -impl ZaiLimitUnit { - pub fn from_int(n: i32) -> Self { - match n { - 1 => ZaiLimitUnit::Days, - 3 => ZaiLimitUnit::Hours, - 5 => ZaiLimitUnit::Minutes, - _ => ZaiLimitUnit::Unknown, - } - } - - pub fn label(&self) -> &'static str { - match self { - ZaiLimitUnit::Days => "day", - ZaiLimitUnit::Hours => "hour", - ZaiLimitUnit::Minutes => "minute", - ZaiLimitUnit::Unknown => "unit", - } - } -} - -/// Per-model usage detail for MCP tools -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ZaiUsageDetail { - /// Model code (e.g., "claude-3-opus", "claude-3-sonnet") - pub model_code: String, - /// Token usage count - pub usage: i64, -} - -impl ZaiUsageDetail { - pub fn new(model_code: impl Into, usage: i64) -> Self { - Self { - model_code: model_code.into(), - usage, - } - } - - /// Format usage as human-readable string - pub fn format_usage(&self) -> String { - if self.usage >= 1_000_000_000 { - format!("{:.1}B tokens", self.usage as f64 / 1_000_000_000.0) - } else if self.usage >= 1_000_000 { - format!("{:.1}M tokens", self.usage as f64 / 1_000_000.0) - } else if self.usage >= 1_000 { - format!("{:.1}K tokens", self.usage as f64 / 1_000.0) - } else { - format!("{} tokens", self.usage) - } - } -} - -/// A single limit entry from Z.ai -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ZaiLimitEntry { - /// Type of limit - pub limit_type: ZaiLimitType, - /// Time unit for the window - pub unit: ZaiLimitUnit, - /// Number of units in the window - pub number: i32, - /// Total usage allowed - pub usage: i64, - /// Current value used - pub current_value: i64, - /// Remaining allocation - pub remaining: i64, - /// Usage percentage (0-100) - pub percentage: f64, - /// Per-model usage details - pub usage_details: Vec, - /// When the limit resets - pub next_reset_time: Option>, -} - -impl ZaiLimitEntry { - /// Calculate used percentage from values - pub fn used_percent(&self) -> f64 { - if self.usage <= 0 { - return self.percentage; - } - - let limit = self.usage.max(0); - if limit == 0 { - return 0.0; - } - - let used_from_remaining = limit - self.remaining; - let used = used_from_remaining - .max(self.current_value) - .max(0) - .min(limit); - let percent = (used as f64 / limit as f64) * 100.0; - percent.clamp(0.0, 100.0) - } - - /// Get window duration in minutes - pub fn window_minutes(&self) -> Option { - if self.number <= 0 { - return None; - } - match self.unit { - ZaiLimitUnit::Minutes => Some(self.number), - ZaiLimitUnit::Hours => Some(self.number * 60), - ZaiLimitUnit::Days => Some(self.number * 24 * 60), - ZaiLimitUnit::Unknown => None, - } - } - - /// Get window description (e.g., "1 hour", "7 days") - pub fn window_description(&self) -> Option { - if self.number <= 0 { - return None; - } - if self.unit == ZaiLimitUnit::Unknown { - return None; - } - - let unit_label = self.unit.label(); - if self.number == 1 { - Some(format!("{} {}", self.number, unit_label)) - } else { - Some(format!("{} {}s", self.number, unit_label)) - } - } - - /// Get window label for display (e.g., "1 hour window") - pub fn window_label(&self) -> Option { - self.window_description().map(|d| format!("{} window", d)) - } -} - -/// Complete Z.ai usage snapshot -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ZaiUsageSnapshot { - /// Token limit entry - pub token_limit: Option, - /// Time limit entry - pub time_limit: Option, - /// User's plan name - pub plan_name: Option, - /// When this snapshot was captured - pub updated_at: DateTime, -} - -impl ZaiUsageSnapshot { - /// Check if this snapshot contains valid data - pub fn is_valid(&self) -> bool { - self.token_limit.is_some() || self.time_limit.is_some() - } - - /// Get the primary limit (tokens preferred over time) - pub fn primary_limit(&self) -> Option<&ZaiLimitEntry> { - self.token_limit.as_ref().or(self.time_limit.as_ref()) - } - - /// Get the secondary limit - pub fn secondary_limit(&self) -> Option<&ZaiLimitEntry> { - if self.token_limit.is_some() && self.time_limit.is_some() { - self.time_limit.as_ref() - } else { - None - } - } -} - -/// MCP Details menu data for UI -#[derive(Debug, Clone)] -pub struct McpDetailsMenu { - /// Window label (e.g., "1 hour window") - pub window_label: Option, - /// Reset time description - pub reset_description: Option, - /// Sorted per-model usage details - pub usage_details: Vec, -} - -impl McpDetailsMenu { - /// Build menu data from a Z.ai snapshot - pub fn from_snapshot(snapshot: &ZaiUsageSnapshot) -> Option { - let time_limit = snapshot.time_limit.as_ref()?; - if time_limit.usage_details.is_empty() { - return None; - } - - let window_label = time_limit.window_label(); - - let reset_description = time_limit.next_reset_time.map(|reset| { - let now = Utc::now(); - if reset <= now { - "now".to_string() - } else { - let duration = reset - now; - let hours = duration.num_hours(); - let minutes = duration.num_minutes() % 60; - - if hours > 24 { - let days = hours / 24; - format!("{}d {}h", days, hours % 24) - } else if hours > 0 { - format!("{}h {}m", hours, minutes) - } else { - format!("{}m", minutes) - } - } - }); - - // Sort by model code - let mut usage_details = time_limit.usage_details.clone(); - usage_details.sort_by(|a, b| { - a.model_code - .to_lowercase() - .cmp(&b.model_code.to_lowercase()) - }); - - Some(Self { - window_label, - reset_description, - usage_details, - }) - } - - /// Generate menu items as (label, value) pairs - pub fn menu_items(&self) -> Vec<(String, String)> { - let mut items = Vec::new(); - - if let Some(window) = &self.window_label { - items.push(("Window".to_string(), window.clone())); - } - - if let Some(reset) = &self.reset_description { - items.push(("Resets".to_string(), reset.clone())); - } - - for detail in &self.usage_details { - items.push((detail.model_code.clone(), detail.format_usage())); - } - - items - } -} - -/// API response structures for parsing -#[derive(Debug, Deserialize)] -pub(crate) struct ZaiQuotaLimitResponse { - pub code: i32, - pub msg: String, - pub data: Option, - pub success: bool, -} - -impl ZaiQuotaLimitResponse { - pub fn is_success(&self) -> bool { - self.success && self.code == 200 - } -} - -#[derive(Debug, Deserialize)] -pub(crate) struct ZaiQuotaLimitData { - pub limits: Vec, - #[serde(alias = "plan")] - #[serde(alias = "plan_type")] - #[serde(alias = "packageName")] - pub plan_name: Option, -} - -#[derive(Debug, Deserialize)] -pub(crate) struct ZaiLimitRaw { - #[serde(rename = "type")] - pub limit_type: String, - pub unit: i32, - pub number: i32, - pub usage: i64, - #[serde(rename = "currentValue")] - pub current_value: i64, - pub remaining: i64, - pub percentage: i32, - #[serde(rename = "usageDetails")] - pub usage_details: Option>, - #[serde(rename = "nextResetTime")] - pub next_reset_time: Option, -} - -#[derive(Debug, Deserialize)] -pub(crate) struct ZaiUsageDetailRaw { - #[serde(rename = "modelCode")] - pub model_code: String, - pub usage: i64, -} - -impl ZaiLimitRaw { - pub fn to_limit_entry(&self) -> Option { - let limit_type = ZaiLimitType::from_string(&self.limit_type)?; - let unit = ZaiLimitUnit::from_int(self.unit); - - let next_reset = self.next_reset_time.map(|ms| { - let secs = ms / 1000; - // ms % 1000 is at most 999, so nanoseconds stay below u32::MAX. - #[expect( - clippy::cast_possible_truncation, - reason = "remainder < 1000 keeps nanoseconds within u32" - )] - let nsecs = ((ms % 1000) * 1_000_000) as u32; - DateTime::from_timestamp(secs, nsecs).unwrap_or_else(Utc::now) - }); - - let usage_details = self - .usage_details - .as_ref() - .map(|details| { - details - .iter() - .map(|d| ZaiUsageDetail::new(&d.model_code, d.usage)) - .collect() - }) - .unwrap_or_default(); - - Some(ZaiLimitEntry { - limit_type, - unit, - number: self.number, - usage: self.usage, - current_value: self.current_value, - remaining: self.remaining, - percentage: self.percentage as f64, - usage_details, - next_reset_time: next_reset, - }) - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_limit_type_from_string() { - assert_eq!( - ZaiLimitType::from_string("TOKENS_LIMIT"), - Some(ZaiLimitType::TokensLimit) - ); - assert_eq!( - ZaiLimitType::from_string("TIME_LIMIT"), - Some(ZaiLimitType::TimeLimit) - ); - assert_eq!(ZaiLimitType::from_string("INVALID"), None); - } - - #[test] - fn test_limit_unit_from_int() { - assert_eq!(ZaiLimitUnit::from_int(1), ZaiLimitUnit::Days); - assert_eq!(ZaiLimitUnit::from_int(3), ZaiLimitUnit::Hours); - assert_eq!(ZaiLimitUnit::from_int(5), ZaiLimitUnit::Minutes); - assert_eq!(ZaiLimitUnit::from_int(99), ZaiLimitUnit::Unknown); - } - - #[test] - fn test_usage_detail_format() { - let detail = ZaiUsageDetail::new("glm-5", 1_500_000_000); - assert_eq!(detail.usage, 1_500_000_000); - assert_eq!(detail.format_usage(), "1.5B tokens"); - - let detail = ZaiUsageDetail::new("claude-3-opus", 1_500_000); - assert_eq!(detail.format_usage(), "1.5M tokens"); - - let detail = ZaiUsageDetail::new("claude-3-sonnet", 5_000); - assert_eq!(detail.format_usage(), "5.0K tokens"); - - let detail = ZaiUsageDetail::new("gpt-4", 500); - assert_eq!(detail.format_usage(), "500 tokens"); - } - - #[test] - fn test_window_description() { - let entry = ZaiLimitEntry { - limit_type: ZaiLimitType::TimeLimit, - unit: ZaiLimitUnit::Hours, - number: 1, - usage: 1000, - current_value: 500, - remaining: 500, - percentage: 50.0, - usage_details: vec![], - next_reset_time: None, - }; - - assert_eq!(entry.window_description(), Some("1 hour".to_string())); - assert_eq!(entry.window_label(), Some("1 hour window".to_string())); - } - - #[test] - fn test_window_minutes() { - let mut entry = ZaiLimitEntry { - limit_type: ZaiLimitType::TimeLimit, - unit: ZaiLimitUnit::Hours, - number: 2, - usage: 1000, - current_value: 0, - remaining: 1000, - percentage: 0.0, - usage_details: vec![], - next_reset_time: None, - }; - - assert_eq!(entry.window_minutes(), Some(120)); - - entry.unit = ZaiLimitUnit::Days; - entry.number = 7; - assert_eq!(entry.window_minutes(), Some(7 * 24 * 60)); - } -} diff --git a/rust/src/providers/zai/mod.rs b/rust/src/providers/zai/mod.rs index c3b5d22d32..04caa64067 100755 --- a/rust/src/providers/zai/mod.rs +++ b/rust/src/providers/zai/mod.rs @@ -4,19 +4,10 @@ //! Uses API token stored in Windows Credential Manager mod balance; -pub mod mcp_details; pub mod region; mod reset_plausibility; pub mod settings; -// Re-exports for MCP details menu -#[allow( - unused_imports, - reason = "imports needed for future ZAI provider wiring" -)] -pub use mcp_details::{ - McpDetailsMenu, ZaiLimitEntry, ZaiLimitType, ZaiLimitUnit, ZaiUsageDetail, ZaiUsageSnapshot, -}; pub use region::ZaiRegion; pub use settings::ZaiSettingsError; @@ -253,26 +244,11 @@ impl ZaiProvider { }) } - /// Validate endpoint overrides against the region *before* any bearer - /// request (upstream #2621/#2623: canonical cross-region overrides are - /// rejected pre-auth; custom relay hosts stay legal). - fn validate_endpoint_overrides( - env: &settings::EnvMap, - region: ZaiRegion, - ) -> Result<(), ProviderError> { - ZaiSettingsReader::validate_endpoint_overrides(env, region) - .map_err(|err| ProviderError::Other(err.to_string())) - } - /// Quota URL: `Z_AI_QUOTA_URL` full override → `Z_AI_API_HOST` host /// override → the selected region's canonical endpoint. fn quota_url(env: &settings::EnvMap, region: ZaiRegion) -> Result { let provider_err = |err: ZaiSettingsError| ProviderError::Other(err.to_string()); - let env_get = |key: &str| env.get(key).and_then(|raw| settings::cleaned(raw)); - if env_get(settings::ZAI_QUOTA_URL_ENV).is_some() { - let url = ZaiSettingsReader::quota_url_override(env) - .map_err(provider_err)? - .expect("override present"); + if let Some(url) = ZaiSettingsReader::quota_url_override(env).map_err(provider_err)? { return Ok(url); } if let Some(url) = ZaiSettingsReader::quota_url_from_api_host(env).map_err(provider_err)? { @@ -321,8 +297,9 @@ impl ZaiProvider { let env = settings::process_env(); let region = Self::effective_region(ctx, &env); // Canonical cross-region overrides are rejected before any bearer - // token is sent (upstream #2621/#2623). - Self::validate_endpoint_overrides(&env, region)?; + // token is sent; custom relay hosts stay legal (upstream #2621/#2623). + ZaiSettingsReader::validate_endpoint_overrides(&env, region) + .map_err(|err| ProviderError::Other(err.to_string()))?; let api_token = Self::get_api_token(ctx.api_key.as_deref(), region, &env)?; let client = crate::core::credentialed_http_client_builder() @@ -551,14 +528,7 @@ impl ZaiProvider { /// Returns `None` when number ≤ 0 or unit is unknown (upstream windowMinutes). fn window_minutes(l: &ZaiLimit) -> Option { let number = l.number.filter(|&n| n > 0)? as u32; - let unit = l.unit?; - let minutes_per_unit = match unit { - 1 => 1440, // days - 3 => 60, // hours - 5 => 1, // minutes - 6 => 10080, // weeks - _ => return None, - }; + let (minutes_per_unit, _, _) = unit_spec(l.unit?)?; Some(number * minutes_per_unit) } } @@ -577,41 +547,22 @@ fn rate_window_reset_description(l: &ZaiLimit, window_mins: Option) -> Opti fn window_label(l: &ZaiLimit) -> Option { let number = l.number.filter(|&n| n > 0)?; - let unit = l.unit?; - let unit_label = match unit { - 1 => { - if number == 1 { - "day" - } else { - "days" - } - } - 3 => { - if number == 1 { - "hour" - } else { - "hours" - } - } - 5 => { - if number == 1 { - "minute" - } else { - "minutes" - } - } - 6 => { - if number == 1 { - "week" - } else { - "weeks" - } - } - _ => return None, - }; + let (_, one, many) = unit_spec(l.unit?)?; + let unit_label = if number == 1 { one } else { many }; Some(format!("{number} {unit_label} window")) } +/// z.ai limit unit code → (minutes per unit, singular label, plural label). +fn unit_spec(unit: i32) -> Option<(u32, &'static str, &'static str)> { + Some(match unit { + 1 => (1440, "day", "days"), + 3 => (60, "hour", "hours"), + 5 => (1, "minute", "minutes"), + 6 => (10080, "week", "weeks"), + _ => return None, + }) +} + impl ZaiTeamContext { fn from_env() -> Option { let organization_id = std::env::var(ZAI_BIGMODEL_ORG_ENV) diff --git a/rust/src/providers/zai/region.rs b/rust/src/providers/zai/region.rs index 838ef62dce..696257727c 100644 --- a/rust/src/providers/zai/region.rs +++ b/rust/src/providers/zai/region.rs @@ -8,8 +8,6 @@ use reqwest::Url; /// Canonical quota API path shared by both regions. const QUOTA_PATH: &str = "/api/monitor/usage/quota/limit"; -/// Per-model usage API path shared by both regions. -const MODEL_USAGE_PATH: &str = "/api/monitor/usage/model-usage"; /// Which z.ai API plane a credential belongs to. #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] @@ -44,13 +42,6 @@ impl ZaiRegion { .expect("region quota URL is a valid constant") } - /// Model-usage endpoint for this region. - pub fn model_usage_url(self) -> Url { - self.base_url() - .join(MODEL_USAGE_PATH) - .expect("region model-usage URL is a valid constant") - } - /// Canonical host for this region's quota endpoint. pub fn canonical_host(self) -> &'static str { match self { @@ -59,24 +50,6 @@ impl ZaiRegion { } } - /// Personal-plan dashboard for this region. - pub fn dashboard_url(self) -> Url { - Url::parse(match self { - ZaiRegion::Global => "https://z.ai/manage-apikey/coding-plan/personal/my-plan", - ZaiRegion::BigModelCn => "https://bigmodel.cn/coding-plan/personal/usage", - }) - .expect("region dashboard URL is a valid constant") - } - - /// Team dashboard for this region (global reuses the personal dashboard). - pub fn team_dashboard_url(self) -> Url { - match self { - ZaiRegion::Global => self.dashboard_url(), - ZaiRegion::BigModelCn => Url::parse("https://bigmodel.cn/coding-plan/team/usage-stats") - .expect("region team dashboard URL is a valid constant"), - } - } - /// Parse the persisted settings `api_region` value into a region. /// /// Accepted values cover the upstream region IDs (`global`, @@ -132,48 +105,22 @@ mod tests { ZaiRegion::BigModelCn.quota_limit_url().as_str(), "https://open.bigmodel.cn/api/monitor/usage/quota/limit" ); - assert_eq!( - ZaiRegion::Global.model_usage_url().as_str(), - "https://api.z.ai/api/monitor/usage/model-usage" - ); - assert_eq!( - ZaiRegion::BigModelCn.model_usage_url().as_str(), - "https://open.bigmodel.cn/api/monitor/usage/model-usage" - ); - assert_eq!( - ZaiRegion::BigModelCn.team_dashboard_url().as_str(), - "https://bigmodel.cn/coding-plan/team/usage-stats" - ); } #[test] fn settings_aliases_map_to_regions() { - assert_eq!( - ZaiRegion::from_settings_value(Some("cn")), - ZaiRegion::BigModelCn - ); - assert_eq!( - ZaiRegion::from_settings_value(Some(" bigmodel ")), - ZaiRegion::BigModelCn - ); - assert_eq!( - ZaiRegion::from_settings_value(Some("bigmodel-cn")), - ZaiRegion::BigModelCn - ); - assert_eq!( - ZaiRegion::from_settings_value(Some("bigmodel_cn")), - ZaiRegion::BigModelCn - ); - assert_eq!( - ZaiRegion::from_settings_value(Some("global")), - ZaiRegion::Global - ); - assert_eq!( - ZaiRegion::from_settings_value(Some("intl")), - ZaiRegion::Global - ); - assert_eq!(ZaiRegion::from_settings_value(Some("")), ZaiRegion::Global); - assert_eq!(ZaiRegion::from_settings_value(None), ZaiRegion::Global); + for (raw, region) in [ + (Some("cn"), ZaiRegion::BigModelCn), + (Some(" bigmodel "), ZaiRegion::BigModelCn), + (Some("bigmodel-cn"), ZaiRegion::BigModelCn), + (Some("bigmodel_cn"), ZaiRegion::BigModelCn), + (Some("global"), ZaiRegion::Global), + (Some("intl"), ZaiRegion::Global), + (Some(""), ZaiRegion::Global), + (None, ZaiRegion::Global), + ] { + assert_eq!(ZaiRegion::from_settings_value(raw), region, "{raw:?}"); + } } #[test] diff --git a/rust/src/providers/zai/settings.rs b/rust/src/providers/zai/settings.rs index 71320ca2bd..29e76b8e3b 100644 --- a/rust/src/providers/zai/settings.rs +++ b/rust/src/providers/zai/settings.rs @@ -46,10 +46,6 @@ pub const BIGMODEL_API_KEY_RELATIVE_PATHS: [&str; 3] = [ /// Errors mirroring upstream `ZaiSettingsError`. #[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] pub enum ZaiSettingsError { - #[error( - "z.ai API token not found. Set apiKey in CodexBar settings, Z_AI_API_KEY, or a BigModel CN credential." - )] - MissingToken, #[error("z.ai endpoint override {0} must use HTTPS or a bare host.")] InvalidEndpointOverride(&'static str), #[error("z.ai endpoint override {0} does not match the selected {1} region.")] @@ -118,25 +114,16 @@ impl ZaiSettingsReader { /// `InvalidEndpointOverride` for non-HTTPS values so a broken override /// never silently downgrades the transfer. pub fn quota_url_override(env: &EnvMap) -> Result, ZaiSettingsError> { - env_get(env, ZAI_QUOTA_URL_ENV) - .and_then(cleaned) - .map(|raw| { - normalized_https_url(&raw) - .ok_or(ZaiSettingsError::InvalidEndpointOverride(ZAI_QUOTA_URL_ENV)) - }) - .transpose() + override_url(env, ZAI_QUOTA_URL_ENV) } /// `Z_AI_API_HOST` override expanded to the quota endpoint. pub fn quota_url_from_api_host(env: &EnvMap) -> Result, ZaiSettingsError> { - let Some(raw) = env_get(env, ZAI_API_HOST_ENV).and_then(cleaned) else { - return Ok(None); - }; - let mut url = normalized_https_url(&raw) - .ok_or(ZaiSettingsError::InvalidEndpointOverride(ZAI_API_HOST_ENV))?; - url.set_path("api/monitor/usage/quota/limit"); - url.set_query(None); - Ok(Some(url)) + Ok(override_url(env, ZAI_API_HOST_ENV)?.map(|mut url| { + url.set_path("api/monitor/usage/quota/limit"); + url.set_query(None); + url + })) } /// Validate all endpoint overrides against the selected region *before* @@ -153,8 +140,7 @@ impl ZaiSettingsReader { env: &EnvMap, region: ZaiRegion, ) -> Result<(), ZaiSettingsError> { - if env_get(env, ZAI_QUOTA_URL_ENV).and_then(cleaned).is_some() { - let url = Self::quota_url_override(env)?.expect("override present"); + if let Some(url) = Self::quota_url_override(env)? { return validate_known_host(&url, region, ZAI_QUOTA_URL_ENV); } Self::validate_api_host_endpoint_override(env, region) @@ -164,15 +150,22 @@ impl ZaiSettingsReader { env: &EnvMap, region: ZaiRegion, ) -> Result<(), ZaiSettingsError> { - let Some(raw) = env_get(env, ZAI_API_HOST_ENV).and_then(cleaned) else { - return Ok(()); - }; - let url = normalized_https_url(&raw) - .ok_or(ZaiSettingsError::InvalidEndpointOverride(ZAI_API_HOST_ENV))?; - validate_known_host(&url, region, ZAI_API_HOST_ENV) + match override_url(env, ZAI_API_HOST_ENV)? { + Some(url) => validate_known_host(&url, region, ZAI_API_HOST_ENV), + None => Ok(()), + } } } +/// A set endpoint-override env var as an HTTPS URL; `InvalidEndpointOverride` +/// for non-HTTPS values so a broken override never silently downgrades. +fn override_url(env: &EnvMap, key: &'static str) -> Result, ZaiSettingsError> { + env_get(env, key) + .and_then(cleaned) + .map(|raw| normalized_https_url(&raw).ok_or(ZaiSettingsError::InvalidEndpointOverride(key))) + .transpose() +} + /// Canonical cross-region override rejection (upstream `validateKnownHost`). /// /// Only canonical plane hosts are pinned: `api.z.ai` overrides under BigModel @@ -253,17 +246,22 @@ mod tests { .collect() } + fn token(map: &EnvMap, region: ZaiRegion) -> Option { + ZaiSettingsReader::api_token(map, Path::new("/nonexistent"), region) + } + + fn write_relay(home: &Path, dir: &str, name: &str, contents: &str) { + let dir = home.join(dir); + std::fs::create_dir_all(&dir).expect("mkdir relay"); + std::fs::write(dir.join(name), contents).expect("write relay key"); + } + #[test] fn api_token_reads_from_environment() { let map = env(&[("Z_AI_API_KEY", "abc123")]); + assert_eq!(token(&map, ZaiRegion::Global).as_deref(), Some("abc123")); assert_eq!( - ZaiSettingsReader::api_token(&map, Path::new("/nonexistent"), ZaiRegion::Global) - .as_deref(), - Some("abc123") - ); - assert_eq!( - ZaiSettingsReader::api_token(&map, Path::new("/nonexistent"), ZaiRegion::BigModelCn) - .as_deref(), + token(&map, ZaiRegion::BigModelCn).as_deref(), Some("abc123") ); } @@ -272,10 +270,7 @@ mod tests { fn legacy_alias_feeds_both_regions() { let map = env(&[("ZAI_API_TOKEN", "legacy-token")]); for region in [ZaiRegion::Global, ZaiRegion::BigModelCn] { - assert_eq!( - ZaiSettingsReader::api_token(&map, Path::new("/nonexistent"), region).as_deref(), - Some("legacy-token") - ); + assert_eq!(token(&map, region).as_deref(), Some("legacy-token")); } } @@ -283,14 +278,10 @@ mod tests { fn bigmodel_aliases_are_available_only_to_china_region() { let map = env(&[("BIGMODEL_API_KEY", "china-token")]); assert_eq!( - ZaiSettingsReader::api_token(&map, Path::new("/nonexistent"), ZaiRegion::BigModelCn) - .as_deref(), + token(&map, ZaiRegion::BigModelCn).as_deref(), Some("china-token") ); - assert_eq!( - ZaiSettingsReader::api_token(&map, Path::new("/nonexistent"), ZaiRegion::Global), - None - ); + assert_eq!(token(&map, ZaiRegion::Global), None); } #[test] @@ -302,14 +293,12 @@ mod tests { ("BIGMODEL_API_KEY", "bigmodel"), ]); assert_eq!( - ZaiSettingsReader::api_token(&map, Path::new("/nonexistent"), ZaiRegion::BigModelCn) - .as_deref(), + token(&map, ZaiRegion::BigModelCn).as_deref(), Some("bigmodel") ); let map = env(&[("GLM_API_KEY", "glm"), ("ZHIPUAI_API_KEY", "zhipuai")]); assert_eq!( - ZaiSettingsReader::api_token(&map, Path::new("/nonexistent"), ZaiRegion::BigModelCn) - .as_deref(), + token(&map, ZaiRegion::BigModelCn).as_deref(), Some("zhipuai") ); } @@ -317,13 +306,12 @@ mod tests { #[test] fn glm_relay_file_is_available_only_to_china_region() { let home = tempfile::tempdir().expect("tempdir"); - let relay_dir = home.path().join(".coding-relay"); - std::fs::create_dir_all(&relay_dir).expect("mkdir relay"); - std::fs::write( - relay_dir.join("glm-api-key"), + write_relay( + home.path(), + ".coding-relay", + "glm-api-key", " relay-china-token\nignored-second-line", - ) - .expect("write relay key"); + ); let map = env(&[]); assert_eq!( @@ -346,9 +334,7 @@ mod tests { (".config/zhipu", "api_key", "relay-zhipu"), ] { let home = tempfile::tempdir().expect("tempdir"); - let dir = home.path().join(dir); - std::fs::create_dir_all(&dir).expect("mkdir"); - std::fs::write(dir.join(name), format!("{token}\n")).expect("write"); + write_relay(home.path(), dir, name, &format!("{token}\n")); let map = env(&[]); assert_eq!( ZaiSettingsReader::api_token(&map, home.path(), ZaiRegion::BigModelCn).as_deref(), @@ -363,9 +349,7 @@ mod tests { (".coding-relay", "glm-api-key", "relay-glm"), (".config/zhipu", "api_key", "relay-zhipu"), ] { - let dir = home.path().join(dir); - std::fs::create_dir_all(&dir).expect("mkdir"); - std::fs::write(dir.join(name), format!("{token}\n")).expect("write"); + write_relay(home.path(), dir, name, &format!("{token}\n")); } let map = env(&[]); assert_eq!( @@ -377,9 +361,7 @@ mod tests { #[test] fn unreadable_or_empty_relay_files_are_skipped() { let home = tempfile::tempdir().expect("tempdir"); - let relay_dir = home.path().join(".coding-relay"); - std::fs::create_dir_all(&relay_dir).expect("mkdir relay"); - std::fs::write(relay_dir.join("glm-api-key"), "\n \n").expect("write empty"); + write_relay(home.path(), ".coding-relay", "glm-api-key", "\n \n"); let map = env(&[]); assert_eq!( ZaiSettingsReader::api_token(&map, home.path(), ZaiRegion::BigModelCn), diff --git a/rust/src/providers/zai/tests.rs b/rust/src/providers/zai/tests.rs index b02ecd2781..c89b045ca8 100644 --- a/rust/src/providers/zai/tests.rs +++ b/rust/src/providers/zai/tests.rs @@ -65,23 +65,18 @@ fn parses_workspace_pair_as_team_context() { #[test] fn parses_successful_response_without_message() { - let provider = ZaiProvider::new(); - let quota: ZaiQuotaResponse = serde_json::from_value(serde_json::json!({ - "code": 200, - "data": { - "planName": "BigModel CN", - "limits": [{ - "type": "TOKENS_LIMIT", - "used": 10, - "limit": 100, - "unit": 3, - "number": 5 - }] - } - })) - .unwrap(); - - let usage = provider.parse_quota_response("a).unwrap().usage; + let data = serde_json::json!({ + "planName": "BigModel CN", + "limits": [{ + "type": "TOKENS_LIMIT", + "used": 10, + "limit": 100, + "unit": 3, + "number": 5 + }] + }); + + let usage = parse_data(data).unwrap().usage; assert_eq!(usage.login_method.as_deref(), Some("BigModel CN")); assert_eq!(usage.primary.used_percent, 10.0); @@ -89,25 +84,20 @@ fn parses_successful_response_without_message() { #[test] fn parses_current_api_percentage_and_reset_time() { - let provider = ZaiProvider::new(); - let quota: ZaiQuotaResponse = serde_json::from_value(serde_json::json!({ - "code": 200, - "data": { - "limits": [{ - "type": "TOKENS_LIMIT", - "unit": 3, - "number": 5, - "usage": 800000000, - "currentValue": 600000000, - "remaining": 200000000, - "percentage": 75, - "nextResetTime": 1770648402389_i64 - }] - } - })) - .unwrap(); - - let usage = provider.parse_quota_response("a).unwrap().usage; + let data = serde_json::json!({ + "limits": [{ + "type": "TOKENS_LIMIT", + "unit": 3, + "number": 5, + "usage": 800000000, + "currentValue": 600000000, + "remaining": 200000000, + "percentage": 75, + "nextResetTime": 1770648402389_i64 + }] + }); + + let usage = parse_data(data).unwrap().usage; assert_eq!(usage.primary.used_percent, 75.0); assert_eq!(usage.primary.window_minutes, Some(300)); @@ -117,38 +107,31 @@ fn parses_current_api_percentage_and_reset_time() { #[test] fn five_hour_reset_plausibility_drops_impossible_timestamp() { let now = Utc::now(); - let quota: ZaiQuotaResponse = serde_json::from_value(serde_json::json!({ - "code": 200, - "data": {"limits": [ - { - "type": "TOKENS_LIMIT", - "unit": 3, - "number": 5, - "percentage": 25, - "nextResetTime": (now + chrono::Duration::hours(10)).timestamp_millis() - }, - { - "type": "TOKENS_LIMIT", - "unit": 6, - "number": 1, - "percentage": 9, - "nextResetTime": (now + chrono::Duration::days(6)).timestamp_millis() - }, - { - "type": "TIME_LIMIT", - "unit": 5, - "number": 1, - "percentage": 22, - "nextResetTime": (now + chrono::Duration::days(20)).timestamp_millis() - } - ]} - })) - .unwrap(); + let data = serde_json::json!({"limits": [ + { + "type": "TOKENS_LIMIT", + "unit": 3, + "number": 5, + "percentage": 25, + "nextResetTime": (now + chrono::Duration::hours(10)).timestamp_millis() + }, + { + "type": "TOKENS_LIMIT", + "unit": 6, + "number": 1, + "percentage": 9, + "nextResetTime": (now + chrono::Duration::days(6)).timestamp_millis() + }, + { + "type": "TIME_LIMIT", + "unit": 5, + "number": 1, + "percentage": 22, + "nextResetTime": (now + chrono::Duration::days(20)).timestamp_millis() + } + ]}); - let usage = ZaiProvider::new() - .parse_quota_response("a) - .unwrap() - .usage; + let usage = parse_data(data).unwrap().usage; assert_eq!(usage.primary.used_percent, 25.0); assert_eq!(usage.primary.window_minutes, Some(300)); assert_eq!(usage.primary.reset_description.as_deref(), Some("5-hour")); @@ -183,37 +166,32 @@ fn credit_limit_plan_drives_primary_and_weekly_windows() { // Upstream 0.49.0 #2724/#2712: credit-based Coding Plans report // CREDIT_LIMIT rows shaped like TOKENS_LIMIT. Without this, usage // sticks at 0% used / 100% remaining. - let provider = ZaiProvider::new(); - let quota: ZaiQuotaResponse = serde_json::from_value(serde_json::json!({ - "code": 200, - "data": { - "planName": "GLM Coding Lite", - "limits": [ - { - "type": "CREDIT_LIMIT", - "unit": 3, - "number": 5, - "usage": 500, - "currentValue": 475, - "remaining": 25, - "percentage": 95, - "nextResetTime": 1770648402389_i64 - }, - { - "type": "CREDIT_LIMIT", - "unit": 6, - "number": 1, - "usage": 3000, - "currentValue": 1200, - "remaining": 1800, - "percentage": 40 - } - ] - } - })) - .unwrap(); + let data = serde_json::json!({ + "planName": "GLM Coding Lite", + "limits": [ + { + "type": "CREDIT_LIMIT", + "unit": 3, + "number": 5, + "usage": 500, + "currentValue": 475, + "remaining": 25, + "percentage": 95, + "nextResetTime": 1770648402389_i64 + }, + { + "type": "CREDIT_LIMIT", + "unit": 6, + "number": 1, + "usage": 3000, + "currentValue": 1200, + "remaining": 1800, + "percentage": 40 + } + ] + }); - let usage = provider.parse_quota_response("a).unwrap().usage; + let usage = parse_data(data).unwrap().usage; // Shortest window (5h credits) is the primary; longest (weekly) secondary. assert!((usage.primary.used_percent - 95.0).abs() < f64::EPSILON); @@ -230,24 +208,19 @@ fn usage_signal_overrides_stale_percentage() { // Upstream 0.49.0 `parseLimit`: a positive `usage` total makes the // absolute used signal authoritative; the API's `percentage` is only // trusted without it. - let provider = ZaiProvider::new(); - let quota: ZaiQuotaResponse = serde_json::from_value(serde_json::json!({ - "code": 200, - "data": { - "limits": [{ - "type": "CREDIT_LIMIT", - "unit": 3, - "number": 5, - "usage": 500, - "currentValue": 25, - "remaining": 475, - "percentage": 95 - }] - } - })) - .unwrap(); - - let usage = provider.parse_quota_response("a).unwrap().usage; + let data = serde_json::json!({ + "limits": [{ + "type": "CREDIT_LIMIT", + "unit": 3, + "number": 5, + "usage": 500, + "currentValue": 25, + "remaining": 475, + "percentage": 95 + }] + }); + + let usage = parse_data(data).unwrap().usage; assert!((usage.primary.used_percent - 5.0).abs() < f64::EPSILON); } @@ -256,24 +229,19 @@ fn usage_signal_overrides_stale_percentage() { fn time_limit_primary_carries_mcp_label_without_duration() { // Upstream 0.48.0: TIME_LIMIT (MCP) windows no longer keep explicit // duration minutes and label as "MCP", not the old monthly sentinel. - let provider = ZaiProvider::new(); - let quota: ZaiQuotaResponse = serde_json::from_value(serde_json::json!({ - "code": 200, - "data": { - "limits": [{ - "type": "TIME_LIMIT", - "unit": 3, - "number": 5, - "usage": 100, - "currentValue": 20, - "remaining": 80, - "percentage": 25, - "nextResetTime": 123000_i64 - }] - } - })) - .unwrap(); - let usage = provider.parse_quota_response("a).unwrap().usage; + let data = serde_json::json!({ + "limits": [{ + "type": "TIME_LIMIT", + "unit": 3, + "number": 5, + "usage": 100, + "currentValue": 20, + "remaining": 80, + "percentage": 25, + "nextResetTime": 123000_i64 + }] + }); + let usage = parse_data(data).unwrap().usage; assert_eq!(usage.primary.window_minutes, None); assert_eq!(usage.primary.reset_description.as_deref(), Some("MCP")); assert!(usage.primary.resets_at.is_some()); @@ -281,24 +249,19 @@ fn time_limit_primary_carries_mcp_label_without_duration() { #[test] fn bare_time_limit_primary_has_no_window_duration() { - let provider = ZaiProvider::new(); - let quota: ZaiQuotaResponse = serde_json::from_value(serde_json::json!({ - "code": 200, - "data": { - "limits": [{ - "type": "TIME_LIMIT", - "unit": 1, - "number": 0, - "usage": 100, - "currentValue": 20, - "remaining": 80, - "percentage": 25, - "nextResetTime": 123000_i64 - }] - } - })) - .unwrap(); - let usage = provider.parse_quota_response("a).unwrap().usage; + let data = serde_json::json!({ + "limits": [{ + "type": "TIME_LIMIT", + "unit": 1, + "number": 0, + "usage": 100, + "currentValue": 20, + "remaining": 80, + "percentage": 25, + "nextResetTime": 123000_i64 + }] + }); + let usage = parse_data(data).unwrap().usage; assert_eq!(usage.primary.window_minutes, None); assert_eq!(usage.primary.reset_description.as_deref(), Some("MCP")); } @@ -308,28 +271,23 @@ fn mcp_limit_renders_separate_named_window() { // Upstream 0.48.0 GLM Coding Plan layout: coding-limit primary + // MCP as a named extra window; MCP 1-minute marker no longer maps // to a monthly sentinel secondary. - let provider = ZaiProvider::new(); - let quota: ZaiQuotaResponse = serde_json::from_value(serde_json::json!({ - "code": 200, - "data": { - "limits": [ - { - "type": "TOKENS_LIMIT", - "unit": 6, - "number": 1, - "percentage": 34 - }, - { - "type": "TIME_LIMIT", - "unit": 5, - "number": 1, - "percentage": 10 - } - ] - } - })) - .unwrap(); - let usage = provider.parse_quota_response("a).unwrap().usage; + let data = serde_json::json!({ + "limits": [ + { + "type": "TOKENS_LIMIT", + "unit": 6, + "number": 1, + "percentage": 34 + }, + { + "type": "TIME_LIMIT", + "unit": 5, + "number": 1, + "percentage": 10 + } + ] + }); + let usage = parse_data(data).unwrap().usage; assert_eq!(usage.primary.window_minutes, Some(10080)); assert_eq!( @@ -352,29 +310,24 @@ fn mcp_limit_renders_separate_named_window() { fn session_five_hour_window_becomes_primary_over_weekly() { // Upstream 0.48.0 GLM Coding Plan: 2+ TOKENS_LIMIT entries → // shortest (5-hour) window primary, longest (weekly) secondary. - let provider = ZaiProvider::new(); - let quota: ZaiQuotaResponse = serde_json::from_value(serde_json::json!({ - "code": 200, - "data": { - "limits": [ - { - "type": "TOKENS_LIMIT", - "unit": 3, - "number": 5, - "percentage": 55, - "nextResetTime": 1770648402389_i64 - }, - { - "type": "TOKENS_LIMIT", - "unit": 6, - "number": 1, - "percentage": 34 - } - ] - } - })) - .unwrap(); - let usage = provider.parse_quota_response("a).unwrap().usage; + let data = serde_json::json!({ + "limits": [ + { + "type": "TOKENS_LIMIT", + "unit": 3, + "number": 5, + "percentage": 55, + "nextResetTime": 1770648402389_i64 + }, + { + "type": "TOKENS_LIMIT", + "unit": 6, + "number": 1, + "percentage": 34 + } + ] + }); + let usage = parse_data(data).unwrap().usage; assert_eq!(usage.primary.used_percent, 55.0); assert_eq!(usage.primary.window_minutes, Some(300)); @@ -388,51 +341,36 @@ fn session_five_hour_window_becomes_primary_over_weekly() { #[test] fn plan_name_falls_back_to_level_key() { - let provider = ZaiProvider::new(); - let quota: ZaiQuotaResponse = serde_json::from_value(serde_json::json!({ - "code": 200, - "data": { - "level": "GLM Coding Plan", - "limits": [] - } - })) - .unwrap(); - let usage = provider.parse_quota_response("a).unwrap().usage; + let data = serde_json::json!({ + "level": "GLM Coding Plan", + "limits": [] + }); + let usage = parse_data(data).unwrap().usage; assert_eq!(usage.login_method.as_deref(), Some("GLM Coding Plan")); for key in ["plan", "plan_type", "packageName"] { - let quota: ZaiQuotaResponse = serde_json::from_value(serde_json::json!({ - "code": 200, - "data": { key: "Coding Plan", "limits": [] } - })) - .unwrap(); - let usage = provider.parse_quota_response("a).unwrap().usage; + let data = serde_json::json!({ key: "Coding Plan", "limits": [] }); + let usage = parse_data(data).unwrap().usage; assert_eq!(usage.login_method.as_deref(), Some("Coding Plan"), "{key}"); } } #[test] fn empty_plan_fields_fall_back_to_default() { - let provider = ZaiProvider::new(); - let quota: ZaiQuotaResponse = serde_json::from_value(serde_json::json!({ - "code": 200, - "data": { "planName": " ", "level": "", "limits": [] } - })) - .unwrap(); - let usage = provider.parse_quota_response("a).unwrap().usage; + let data = serde_json::json!({ "planName": " ", "level": "", "limits": [] }); + let usage = parse_data(data).unwrap().usage; assert_eq!(usage.login_method.as_deref(), Some("z.ai")); } #[test] fn preserves_api_code_error_message() { - let provider = ZaiProvider::new(); let quota: ZaiQuotaResponse = serde_json::from_value(serde_json::json!({ "code": 401, "message": "invalid token" })) .unwrap(); - let error = provider.parse_quota_response("a).unwrap_err(); + let error = ZaiProvider::new().parse_quota_response("a).unwrap_err(); assert!(error.to_string().contains("invalid token")); } diff --git a/rust/src/providers/zoommate/mod.rs b/rust/src/providers/zoommate/mod.rs index d5093112fe..cbd8ff6bcd 100644 --- a/rust/src/providers/zoommate/mod.rs +++ b/rust/src/providers/zoommate/mod.rs @@ -185,19 +185,15 @@ impl ZoomMateProvider { async fn fetch_via_web(&self, ctx: &FetchContext) -> Result { let request_context = self.resolve_request_context(ctx).await?; - match self + let result = self .fetch_credits_status(&request_context, ctx.web_timeout) - .await + .await; + if let (Err(ProviderError::AuthRequired), Some(key)) = + (&result, request_context.cache_key.as_deref()) { - Ok(snap) => Ok(snap), - Err(ProviderError::AuthRequired) => { - if let Some(key) = request_context.cache_key.as_deref() { - cache_invalidate(key); - } - Err(ProviderError::AuthRequired) - } - Err(e) => Err(e), + cache_invalidate(key); } + result } async fn resolve_request_context( @@ -344,28 +340,9 @@ impl ZoomMateProvider { ctx: &RequestContext, timeout_secs: u64, ) -> Result { - let preferred = ctx.preferred_host.clone(); - let authorization = ctx.authorization.clone(); - let headers = ctx.headers.clone(); - let cookie_by_host = ctx.cookie_by_host.clone(); - let account_email = ctx.account_email.clone(); - - with_api_host_failover(preferred.as_deref(), |host| { - let authorization = authorization.clone(); - let headers = headers.clone(); - let cookie_by_host = cookie_by_host.clone(); - let account_email = account_email.clone(); - async move { - self.fetch_credits_status_on_host( - host, - &authorization, - &headers, - &cookie_by_host, - account_email.as_deref(), - timeout_secs, - ) + with_api_host_failover(ctx.preferred_host.as_deref(), |host| async move { + self.fetch_credits_status_on_host(host, ctx, timeout_secs) .await - } }) .await } @@ -373,10 +350,7 @@ impl ZoomMateProvider { async fn fetch_credits_status_on_host( &self, host: &str, - authorization: &str, - headers: &HashMap, - cookie_by_host: &HashMap, - account_email: Option<&str>, + ctx: &RequestContext, timeout_secs: u64, ) -> Result { let url = format!("https://{host}{CREDITS_STATUS_PATH}"); @@ -392,7 +366,7 @@ impl ZoomMateProvider { .header("Sec-Fetch-Site", "same-site"); // Captured headers first (except Origin/Referer/Authorization which we pin). - for (name, value) in headers { + for (name, value) in &ctx.headers { if name.eq_ignore_ascii_case("origin") || name.eq_ignore_ascii_case("referer") || name.eq_ignore_ascii_case("authorization") @@ -403,12 +377,12 @@ impl ZoomMateProvider { req = req.header(name.as_str(), value.as_str()); } // Upstream #2627: only the header scoped to THIS destination host is sent. - if let Some(cookie) = cookie_by_host.get(host) { + if let Some(cookie) = ctx.cookie_by_host.get(host) { req = req.header("Cookie", cookie); } // Fixed Origin/Referer so captured values never widen the first-party boundary. req = req - .header("Authorization", authorization) + .header("Authorization", ctx.authorization.as_str()) .header("Origin", ORIGIN_REFERER) .header("Referer", ORIGIN_REFERER); @@ -434,7 +408,7 @@ impl ZoomMateProvider { Ok(snapshot_from_credit_status( &credit_status, - account_email, + ctx.account_email.as_deref(), Utc::now(), )) } @@ -517,25 +491,14 @@ fn date_from_millis(raw: Option) -> Option> { } fn hosts_preferred(preferred: Option<&str>) -> Vec<&'static str> { - let preferred = preferred.map(|h| h.to_ascii_lowercase()); - match preferred.as_deref() { - Some(h) if API_HOSTS.iter().any(|x| x.eq_ignore_ascii_case(h)) => { - let mut out = vec![ - API_HOSTS - .iter() - .copied() - .find(|x| x.eq_ignore_ascii_case(h)) - .unwrap_or(API_HOSTS[0]), - ]; - for host in API_HOSTS { - if !host.eq_ignore_ascii_case(h) { - out.push(*host); - } - } - out - } - _ => API_HOSTS.to_vec(), + let mut hosts = API_HOSTS.to_vec(); + if let Some(index) = + preferred.and_then(|h| hosts.iter().position(|x| x.eq_ignore_ascii_case(h))) + { + let host = hosts.remove(index); + hosts.insert(0, host); } + hosts } /// Runs `operation` per host. Auth (401/403 → AuthRequired) and parse errors do @@ -553,8 +516,7 @@ where for (index, host) in hosts.iter().enumerate() { match operation(host).await { Ok(v) => return Ok(v), - Err(ProviderError::AuthRequired) => return Err(ProviderError::AuthRequired), - Err(ProviderError::Parse(msg)) => return Err(ProviderError::Parse(msg)), + Err(e) if !should_failover(&e) => return Err(e), Err(e) => { last_err = Some(e); if index + 1 < hosts.len() { @@ -768,271 +730,9 @@ pub(crate) fn snapshot_from_credit_status( snap } -// Classify failover errors for pure unit tests. -#[cfg(test)] -pub(crate) fn should_failover(err: &ProviderError) -> bool { +fn should_failover(err: &ProviderError) -> bool { !matches!(err, ProviderError::AuthRequired | ProviderError::Parse(_)) } #[cfg(test)] -mod tests { - use super::*; - - fn sample_status_json() -> &'static str { - r#"{ - "data": { - "credit_status": { - "budget_cap": 1000.0, - "used_credit": 250.0, - "remaining_credit": 750.0, - "overage_credit": 0.0, - "allow_overage": false, - "cycle_start_date": 1722470400000, - "cycle_end_date": 1725148800000, - "is_quota_available": true, - "is_unlimited": false - } - }, - "status_code": 200 - }"# - } - - #[test] - fn parses_credits_status_fixture() { - let envelope: CreditsStatusEnvelope = - serde_json::from_str(sample_status_json()).expect("fixture parses"); - let status = envelope.data.unwrap().credit_status.unwrap(); - let snap = snapshot_from_credit_status(&status, Some("user@zoom.us"), Utc::now()); - assert!((snap.primary.used_percent - 25.0).abs() < 0.01); - assert_eq!(snap.primary.reset_description.as_deref(), Some("Credits")); - assert!(snap.primary.resets_at.is_some()); - // 31 days ≈ 44640 minutes (2024-08-01 → 2024-09-01) - assert_eq!(snap.primary.window_minutes, Some(44640)); - assert_eq!(snap.account_email.as_deref(), Some("user@zoom.us")); - assert_eq!(snap.login_method.as_deref(), Some("Cookie")); - } - - #[test] - fn unlimited_or_zero_cap_yields_zero_percent() { - let unlimited = CreditStatus { - budget_cap: Some(100.0), - used_credit: Some(50.0), - is_unlimited: Some(true), - ..Default::default() - }; - let snap = snapshot_from_credit_status(&unlimited, None, Utc::now()); - assert_eq!(snap.primary.used_percent, 0.0); - assert!(snap.primary.resets_at.is_none()); - - let zero_cap = CreditStatus { - budget_cap: Some(0.0), - used_credit: Some(10.0), - is_unlimited: Some(false), - ..Default::default() - }; - let snap = snapshot_from_credit_status(&zero_cap, None, Utc::now()); - assert_eq!(snap.primary.used_percent, 0.0); - } - - #[test] - fn clamps_used_percent_to_100() { - let status = CreditStatus { - budget_cap: Some(100.0), - used_credit: Some(150.0), - is_unlimited: Some(false), - cycle_end_date: Some(1_900_000_000_000), - ..Default::default() - }; - let snap = snapshot_from_credit_status(&status, None, Utc::now()); - assert_eq!(snap.primary.used_percent, 100.0); - } - - #[test] - fn manual_curl_capture_requires_allowed_url_and_authorization() { - let good = "curl 'https://ai.zoom.us/ai-computer/api/v1/credits/status' \ - -H 'Authorization: Bearer tok-abc' -H 'Cookie: session=xyz'"; - let ctx = request_context_from_manual(good).expect("valid capture"); - assert_eq!(ctx.authorization, "Bearer tok-abc"); - assert_eq!(ctx.preferred_host.as_deref(), Some("ai.zoom.us")); - assert_eq!( - ctx.cookie_by_host.get("ai.zoom.us").map(String::as_str), - Some("session=xyz") - ); - - // Wrong path - assert!( - request_context_from_manual( - "curl 'https://ai.zoom.us/ai-computer/api/v1/other' -H 'Authorization: Bearer x'" - ) - .is_none() - ); - // Query rejected - assert!( - request_context_from_manual( - "curl 'https://ai.zoom.us/ai-computer/api/v1/credits/status?x=1' \ - -H 'Authorization: Bearer x'" - ) - .is_none() - ); - // Missing auth - assert!( - request_context_from_manual( - "curl 'https://ai.zoom.us/ai-computer/api/v1/credits/status' -H 'Cookie: a=b'" - ) - .is_none() - ); - // Bad host - assert!( - request_context_from_manual( - "curl 'https://evil.example/ai-computer/api/v1/credits/status' \ - -H 'Authorization: Bearer x'" - ) - .is_none() - ); - } - - #[test] - fn hosts_preferred_promotes_capture_host() { - assert_eq!( - hosts_preferred(Some("zoommate.zoom.us")), - vec!["zoommate.zoom.us", "ai.zoom.us"] - ); - assert_eq!( - hosts_preferred(None), - vec!["ai.zoom.us", "zoommate.zoom.us"] - ); - } - - #[test] - fn failover_skips_auth_and_parse() { - assert!(!should_failover(&ProviderError::AuthRequired)); - assert!(!should_failover(&ProviderError::Parse("x".into()))); - assert!(should_failover(&ProviderError::Other("HTTP 500".into()))); - assert!(should_failover(&ProviderError::NoCookies)); - } - - #[test] - fn bearer_header_normalizes_prefix() { - assert_eq!(bearer_header_value("tok"), "Bearer tok"); - assert_eq!(bearer_header_value("Bearer tok"), "Bearer tok"); - assert_eq!(bearer_header_value("bearer tok"), "bearer tok"); - } - - #[test] - fn jwt_exp_reads_payload() { - // header.payload.sig — payload = {"exp": 2000000000} - let payload = base64::Engine::encode( - &base64::engine::general_purpose::URL_SAFE_NO_PAD, - br#"{"exp":2000000000}"#, - ); - let token = format!("aaa.{payload}.sig"); - assert_eq!(jwt_exp_unix(&token), Some(2_000_000_000)); - assert_eq!(jwt_exp_unix("not-a-jwt"), None); - } - - #[test] - fn cookie_fingerprint_is_stable() { - let mut a = HashMap::new(); - a.insert("ai.zoom.us".into(), "c=1".into()); - a.insert("zoommate.zoom.us".into(), "c=2".into()); - let mut b = HashMap::new(); - b.insert("zoommate.zoom.us".into(), "c=2".into()); - b.insert("ai.zoom.us".into(), "c=1".into()); - assert_eq!(cookie_fingerprint(&a), cookie_fingerprint(&b)); - assert_eq!(cookie_fingerprint(&a).len(), 64); - } - - // ── F16: browser cookie scope preservation (upstream #2627) ─── - - #[derive(Deserialize)] - struct CookieScopeFixture { - records: Vec, - } - - #[derive(Deserialize)] - #[serde(rename_all = "camelCase")] - struct FixtureRecord { - source_domain: String, - domain: String, - scope: String, - name: String, - value: String, - } - - /// Upstream fixture `issue-2507-cookie-scope.json`, copied verbatim: the raw - /// browser host key lives in `sourceDomain`; our `Cookie.domain` carries the - /// same raw form, and the leading dot encodes the fixture's explicit scope. - fn issue_2507_records() -> Vec { - let fixture: CookieScopeFixture = serde_json::from_str(include_str!( - "../fixtures/zoommate/issue-2507-cookie-scope.json" - )) - .unwrap(); - fixture - .records - .into_iter() - .map(|record| { - // Guard the dot↔scope derivation against fixture drift. - let derived = if record.source_domain.starts_with('.') { - "domain" - } else { - "hostOnly" - }; - assert_eq!(derived, record.scope, "{}", record.name); - assert_eq!( - record.source_domain.trim_start_matches('.'), - record.domain, - "{}", - record.name - ); - Cookie { - name: record.name, - value: record.value, - domain: record.source_domain, - path: "/".into(), - expires: None, - is_secure: true, - is_http_only: true, - } - }) - .collect() - } - - #[test] - fn issue_2507_fixture_routes_parent_cookie_to_both_hosts_without_leaks() { - let records = issue_2507_records(); - assert_eq!( - cookie_header_for_host(&records, "ai.zoom.us").as_deref(), - Some("parent=fake; ai-only=fake") - ); - assert_eq!( - cookie_header_for_host(&records, "zoommate.zoom.us").as_deref(), - Some("parent=fake; mate-only=fake") - ); - } - - #[test] - fn cookie_scope_filter_follows_rfc_6265_scope() { - assert!(cookie_is_sendable_to_host("ai.zoom.us", "ai.zoom.us")); - assert!(!cookie_is_sendable_to_host( - "ai.zoom.us", - "zoommate.zoom.us" - )); - assert!(cookie_is_sendable_to_host(".zoom.us", "ai.zoom.us")); - assert!(cookie_is_sendable_to_host(".zoom.us", "zoommate.zoom.us")); - // Plain zoom.us is host-only: never sent to leaf API hosts. - assert!(!cookie_is_sendable_to_host("zoom.us", "ai.zoom.us")); - // Sibling subdomains are not destinations, and the host-only - // marketing cookie doesn't roam either. - assert!(!cookie_is_sendable_to_host( - "marketing.zoom.us", - "ai.zoom.us" - )); - assert!(cookie_header_for_host(&issue_2507_records(), "marketing.zoom.us").is_none()); - // Suffix-lookalike attackers and empty domains never match. - assert!(!cookie_is_sendable_to_host( - "zoom.us.attacker.com", - "ai.zoom.us" - )); - assert!(!cookie_is_sendable_to_host("", "ai.zoom.us")); - } -} +mod tests; diff --git a/rust/src/providers/zoommate/tests.rs b/rust/src/providers/zoommate/tests.rs new file mode 100644 index 0000000000..1b0f5e6b66 --- /dev/null +++ b/rust/src/providers/zoommate/tests.rs @@ -0,0 +1,259 @@ +use super::*; + +fn sample_status_json() -> &'static str { + r#"{ + "data": { + "credit_status": { + "budget_cap": 1000.0, + "used_credit": 250.0, + "remaining_credit": 750.0, + "overage_credit": 0.0, + "allow_overage": false, + "cycle_start_date": 1722470400000, + "cycle_end_date": 1725148800000, + "is_quota_available": true, + "is_unlimited": false + } + }, + "status_code": 200 + }"# +} + +#[test] +fn parses_credits_status_fixture() { + let envelope: CreditsStatusEnvelope = + serde_json::from_str(sample_status_json()).expect("fixture parses"); + let status = envelope.data.unwrap().credit_status.unwrap(); + let snap = snapshot_from_credit_status(&status, Some("user@zoom.us"), Utc::now()); + assert!((snap.primary.used_percent - 25.0).abs() < 0.01); + assert_eq!(snap.primary.reset_description.as_deref(), Some("Credits")); + assert!(snap.primary.resets_at.is_some()); + // 31 days ≈ 44640 minutes (2024-08-01 → 2024-09-01) + assert_eq!(snap.primary.window_minutes, Some(44640)); + assert_eq!(snap.account_email.as_deref(), Some("user@zoom.us")); + assert_eq!(snap.login_method.as_deref(), Some("Cookie")); +} + +#[test] +fn unlimited_or_zero_cap_yields_zero_percent() { + let unlimited = CreditStatus { + budget_cap: Some(100.0), + used_credit: Some(50.0), + is_unlimited: Some(true), + ..Default::default() + }; + let snap = snapshot_from_credit_status(&unlimited, None, Utc::now()); + assert_eq!(snap.primary.used_percent, 0.0); + assert!(snap.primary.resets_at.is_none()); + + let zero_cap = CreditStatus { + budget_cap: Some(0.0), + used_credit: Some(10.0), + is_unlimited: Some(false), + ..Default::default() + }; + let snap = snapshot_from_credit_status(&zero_cap, None, Utc::now()); + assert_eq!(snap.primary.used_percent, 0.0); +} + +#[test] +fn clamps_used_percent_to_100() { + let status = CreditStatus { + budget_cap: Some(100.0), + used_credit: Some(150.0), + is_unlimited: Some(false), + cycle_end_date: Some(1_900_000_000_000), + ..Default::default() + }; + let snap = snapshot_from_credit_status(&status, None, Utc::now()); + assert_eq!(snap.primary.used_percent, 100.0); +} + +#[test] +fn manual_curl_capture_requires_allowed_url_and_authorization() { + let good = "curl 'https://ai.zoom.us/ai-computer/api/v1/credits/status' \ + -H 'Authorization: Bearer tok-abc' -H 'Cookie: session=xyz'"; + let ctx = request_context_from_manual(good).expect("valid capture"); + assert_eq!(ctx.authorization, "Bearer tok-abc"); + assert_eq!(ctx.preferred_host.as_deref(), Some("ai.zoom.us")); + assert_eq!( + ctx.cookie_by_host.get("ai.zoom.us").map(String::as_str), + Some("session=xyz") + ); + + // Wrong path + assert!( + request_context_from_manual( + "curl 'https://ai.zoom.us/ai-computer/api/v1/other' -H 'Authorization: Bearer x'" + ) + .is_none() + ); + // Query rejected + assert!( + request_context_from_manual( + "curl 'https://ai.zoom.us/ai-computer/api/v1/credits/status?x=1' \ + -H 'Authorization: Bearer x'" + ) + .is_none() + ); + // Missing auth + assert!( + request_context_from_manual( + "curl 'https://ai.zoom.us/ai-computer/api/v1/credits/status' -H 'Cookie: a=b'" + ) + .is_none() + ); + // Bad host + assert!( + request_context_from_manual( + "curl 'https://evil.example/ai-computer/api/v1/credits/status' \ + -H 'Authorization: Bearer x'" + ) + .is_none() + ); +} + +#[test] +fn hosts_preferred_promotes_capture_host() { + assert_eq!( + hosts_preferred(Some("zoommate.zoom.us")), + vec!["zoommate.zoom.us", "ai.zoom.us"] + ); + assert_eq!( + hosts_preferred(None), + vec!["ai.zoom.us", "zoommate.zoom.us"] + ); +} + +#[test] +fn failover_skips_auth_and_parse() { + assert!(!should_failover(&ProviderError::AuthRequired)); + assert!(!should_failover(&ProviderError::Parse("x".into()))); + assert!(should_failover(&ProviderError::Other("HTTP 500".into()))); + assert!(should_failover(&ProviderError::NoCookies)); +} + +#[test] +fn bearer_header_normalizes_prefix() { + assert_eq!(bearer_header_value("tok"), "Bearer tok"); + assert_eq!(bearer_header_value("Bearer tok"), "Bearer tok"); + assert_eq!(bearer_header_value("bearer tok"), "bearer tok"); +} + +#[test] +fn jwt_exp_reads_payload() { + // header.payload.sig — payload = {"exp": 2000000000} + let payload = base64::Engine::encode( + &base64::engine::general_purpose::URL_SAFE_NO_PAD, + br#"{"exp":2000000000}"#, + ); + let token = format!("aaa.{payload}.sig"); + assert_eq!(jwt_exp_unix(&token), Some(2_000_000_000)); + assert_eq!(jwt_exp_unix("not-a-jwt"), None); +} + +#[test] +fn cookie_fingerprint_is_stable() { + let mut a = HashMap::new(); + a.insert("ai.zoom.us".into(), "c=1".into()); + a.insert("zoommate.zoom.us".into(), "c=2".into()); + let mut b = HashMap::new(); + b.insert("zoommate.zoom.us".into(), "c=2".into()); + b.insert("ai.zoom.us".into(), "c=1".into()); + assert_eq!(cookie_fingerprint(&a), cookie_fingerprint(&b)); + assert_eq!(cookie_fingerprint(&a).len(), 64); +} + +// ── F16: browser cookie scope preservation (upstream #2627) ─── + +#[derive(Deserialize)] +struct CookieScopeFixture { + records: Vec, +} + +#[derive(Deserialize)] +#[serde(rename_all = "camelCase")] +struct FixtureRecord { + source_domain: String, + domain: String, + scope: String, + name: String, + value: String, +} + +/// Upstream fixture `issue-2507-cookie-scope.json`, copied verbatim: the raw +/// browser host key lives in `sourceDomain`; our `Cookie.domain` carries the +/// same raw form, and the leading dot encodes the fixture's explicit scope. +fn issue_2507_records() -> Vec { + let fixture: CookieScopeFixture = serde_json::from_str(include_str!( + "../fixtures/zoommate/issue-2507-cookie-scope.json" + )) + .unwrap(); + fixture + .records + .into_iter() + .map(|record| { + // Guard the dot↔scope derivation against fixture drift. + let derived = if record.source_domain.starts_with('.') { + "domain" + } else { + "hostOnly" + }; + assert_eq!(derived, record.scope, "{}", record.name); + assert_eq!( + record.source_domain.trim_start_matches('.'), + record.domain, + "{}", + record.name + ); + Cookie { + name: record.name, + value: record.value, + domain: record.source_domain, + path: "/".into(), + expires: None, + is_secure: true, + is_http_only: true, + } + }) + .collect() +} + +#[test] +fn issue_2507_fixture_routes_parent_cookie_to_both_hosts_without_leaks() { + let records = issue_2507_records(); + assert_eq!( + cookie_header_for_host(&records, "ai.zoom.us").as_deref(), + Some("parent=fake; ai-only=fake") + ); + assert_eq!( + cookie_header_for_host(&records, "zoommate.zoom.us").as_deref(), + Some("parent=fake; mate-only=fake") + ); +} + +#[test] +fn cookie_scope_filter_follows_rfc_6265_scope() { + assert!(cookie_is_sendable_to_host("ai.zoom.us", "ai.zoom.us")); + assert!(!cookie_is_sendable_to_host( + "ai.zoom.us", + "zoommate.zoom.us" + )); + assert!(cookie_is_sendable_to_host(".zoom.us", "ai.zoom.us")); + assert!(cookie_is_sendable_to_host(".zoom.us", "zoommate.zoom.us")); + // Plain zoom.us is host-only: never sent to leaf API hosts. + assert!(!cookie_is_sendable_to_host("zoom.us", "ai.zoom.us")); + // Sibling subdomains are not destinations, and the host-only + // marketing cookie doesn't roam either. + assert!(!cookie_is_sendable_to_host( + "marketing.zoom.us", + "ai.zoom.us" + )); + assert!(cookie_header_for_host(&issue_2507_records(), "marketing.zoom.us").is_none()); + // Suffix-lookalike attackers and empty domains never match. + assert!(!cookie_is_sendable_to_host( + "zoom.us.attacker.com", + "ai.zoom.us" + )); + assert!(!cookie_is_sendable_to_host("", "ai.zoom.us")); +} diff --git a/rust/tests/providers/test_infini.rs b/rust/tests/providers/test_infini.rs deleted file mode 100644 index fb333ab3e9..0000000000 --- a/rust/tests/providers/test_infini.rs +++ /dev/null @@ -1,119 +0,0 @@ -//! InfiniClient API 测试 - -use codexbar::providers::infini::{InfiniClient, InfiniError, InfiniUsage, UsagePeriod}; -use codexbar::providers::InfiniProvider; -use codexbar::core::{Provider, ProviderId, FetchContext, SourceMode}; - -#[tokio::test] -async fn test_fetch_usage_success() { - let mut server = mockito::Server::new(); - let mock = server - .mock("GET", "/maas/coding/usage") - .with_status(200) - .with_header("content-type", "application/json") - .with_body( - r#"{ - "5_hour": {"quota": 5000, "used": 1000, "remain": 4000}, - "7_day": {"quota": 30000, "used": 5000, "remain": 25000}, - "30_day": {"quota": 60000, "used": 10000, "remain": 50000} - }"#, - ) - .create(); - - let client = InfiniClient::new("sk-cp-test-key".to_string()).with_base_url(server.url()); - - let usage = client.fetch_usage().await.unwrap(); - - mock.assert(); - assert_eq!(usage.five_hour.quota, 5000); - assert_eq!(usage.seven_day.used, 5000); -} - -#[tokio::test] -async fn test_fetch_usage_unauthorized() { - let mut server = mockito::Server::new(); - let mock = server.mock("GET", "/maas/coding/usage").with_status(401).create(); - - let client = InfiniClient::new("invalid-key".to_string()).with_base_url(server.url()); - - let result = client.fetch_usage().await; - - mock.assert(); - assert!(matches!(result, Err(InfiniError::Unauthorized))); -} - -// ==================== InfiniProvider Tests ==================== - -#[test] -fn test_infini_provider_id() { - let provider = InfiniProvider::new("sk-cp-test".to_string()); - assert_eq!(provider.id(), ProviderId::Infini); -} - -#[test] -fn test_infini_provider_metadata() { - let provider = InfiniProvider::new("sk-cp-test".to_string()); - let meta = provider.metadata(); - assert_eq!(meta.id, ProviderId::Infini); - assert_eq!(meta.display_name, "Infini"); -} - -#[test] -fn test_infini_provider_available_sources() { - let provider = InfiniProvider::new("sk-cp-test".to_string()); - let sources = provider.available_sources(); - assert!(sources.contains(&SourceMode::Auto)); - assert!(sources.contains(&SourceMode::Web)); -} - -#[test] -fn test_infini_provider_supports_web() { - let provider = InfiniProvider::new("sk-cp-test".to_string()); - assert!(provider.supports_web()); -} - -#[tokio::test] -async fn test_infini_provider_fetch_usage_success() { - let mut server = mockito::Server::new(); - let mock = server - .mock("GET", "/maas/coding/usage") - .with_status(200) - .with_header("content-type", "application/json") - .with_body( - r#"{ - "5_hour": {"quota": 5000, "used": 2500, "remain": 2500}, - "7_day": {"quota": 30000, "used": 15000, "remain": 15000}, - "30_day": {"quota": 60000, "used": 30000, "remain": 30000} - }"#, - ) - .create(); - - let provider = InfiniProvider::new(String::new()).with_base_url(server.url()); - let ctx = FetchContext { - api_key: Some("sk-cp-test-key".to_string()), - ..Default::default() - }; - let result = provider.fetch_usage(&ctx).await.unwrap(); - - mock.assert(); - assert_eq!(result.usage.primary.used_percent, 50.0); - assert!(result.usage.secondary.is_some()); - let secondary = result.usage.secondary.unwrap(); - assert_eq!(secondary.used_percent, 50.0); -} - -#[tokio::test] -async fn test_infini_provider_fetch_usage_unauthorized() { - let mut server = mockito::Server::new(); - let mock = server.mock("GET", "/maas/coding/usage").with_status(401).create(); - - let provider = InfiniProvider::new(String::new()).with_base_url(server.url()); - let ctx = FetchContext { - api_key: Some("invalid-key".to_string()), - ..Default::default() - }; - let result = provider.fetch_usage(&ctx).await; - - mock.assert(); - assert!(result.is_err()); -}