From 9668faa8f669357b565b9c2f92f7dc397a9c3fa6 Mon Sep 17 00:00:00 2001 From: RCD <90105158+Finesssee@users.noreply.github.com> Date: Thu, 1 Oct 2026 11:15:13 +0700 Subject: [PATCH 1/2] Cache Hugging Face identity for 12 h, explain 403 billing access, keep API-only off cookies --- .../providers/huggingface/identity_cache.rs | 240 ++++++++++++++++++ rust/src/providers/huggingface/mod.rs | 138 ++++++++-- rust/src/providers/huggingface/wallet.rs | 136 +++++++++- 3 files changed, 488 insertions(+), 26 deletions(-) create mode 100644 rust/src/providers/huggingface/identity_cache.rs diff --git a/rust/src/providers/huggingface/identity_cache.rs b/rust/src/providers/huggingface/identity_cache.rs new file mode 100644 index 0000000000..13b8b85adf --- /dev/null +++ b/rust/src/providers/huggingface/identity_cache.rs @@ -0,0 +1,240 @@ +use std::collections::HashMap; +use std::future::Future; +use std::sync::{Mutex, OnceLock}; + +use chrono::{DateTime, Duration, Utc}; +use sha2::{Digest, Sha256}; + +use super::IdentitySnapshot; + +const IDENTITY_TTL: Duration = Duration::seconds(12 * 60 * 60); +const MAX_IDENTITY_ENTRIES: usize = 32; + +#[derive(Debug, Clone)] +struct CachedIdentity { + identity: IdentitySnapshot, + fetched_at: DateTime, +} + +#[derive(Debug, Default)] +pub(super) struct IdentityCache { + entries: HashMap, +} + +impl IdentityCache { + pub(super) fn get(&mut self, token: &str, now: DateTime) -> Option { + let key = token_key(token); + let fresh = self + .entries + .get(&key) + .is_some_and(|entry| is_fresh(entry, &now)); + if !fresh { + drop(self.entries.remove(&key)); + return None; + } + self.entries.get(&key).map(|entry| entry.identity.clone()) + } + + pub(super) fn insert(&mut self, token: &str, identity: IdentitySnapshot, now: DateTime) { + self.remove_expired(&now); + let key = token_key(token); + drop(self.entries.remove(&key)); + if self.entries.len() >= MAX_IDENTITY_ENTRIES { + let oldest_key = self + .entries + .iter() + .min_by_key(|(_, entry)| entry.fetched_at) + .map(|(key, _)| key.clone()); + if let Some(oldest_key) = oldest_key { + drop(self.entries.remove(&oldest_key)); + } + } + self.entries.insert( + key, + CachedIdentity { + identity, + fetched_at: now, + }, + ); + } + + fn remove_expired(&mut self, now: &DateTime) { + self.entries.retain(|_, entry| is_fresh(entry, now)); + } +} + +fn is_fresh(entry: &CachedIdentity, now: &DateTime) -> bool { + let age = now.signed_duration_since(entry.fetched_at); + age >= Duration::zero() && age < IDENTITY_TTL +} + +fn token_key(token: &str) -> String { + const HEX: &[u8; 16] = b"0123456789abcdef"; + let digest = Sha256::digest(token.as_bytes()); + let mut key = String::with_capacity(64); + for byte in digest { + key.push(HEX[(byte >> 4) as usize] as char); + key.push(HEX[(byte & 0x0f) as usize] as char); + } + key +} + +static PROCESS_IDENTITY_CACHE: OnceLock> = OnceLock::new(); + +pub(super) fn process_identity_cache() -> &'static Mutex { + PROCESS_IDENTITY_CACHE.get_or_init(|| Mutex::new(IdentityCache::default())) +} + +pub(super) async fn get_or_fetch_identity( + cache: &Mutex, + token: &str, + now: DateTime, + fetch: F, +) -> Option +where + F: FnOnce() -> Fut, + Fut: Future>, +{ + if let Some(identity) = cache + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .get(token, now) + { + return Some(identity); + } + + let identity = fetch().await?; + cache + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .insert(token, identity.clone(), now); + Some(identity) +} + +#[cfg(test)] +mod tests { + use std::sync::atomic::{AtomicUsize, Ordering}; + + use super::*; + + fn at(seconds: i64) -> DateTime { + DateTime::from_timestamp(seconds, 0).unwrap() + } + + fn identity(name: &str) -> IdentitySnapshot { + IdentitySnapshot { + user_id: None, + name: Some(name.to_string()), + email: None, + plan: None, + } + } + + #[tokio::test] + async fn identity_cache_is_isolated_per_token_and_expires_with_fetch_clock() { + let cache = Mutex::new(IdentityCache::default()); + let calls = AtomicUsize::new(0); + let t = at(1_800_000_000); + + let first = get_or_fetch_identity(&cache, "fixture-a", t, || async { + calls.fetch_add(1, Ordering::SeqCst); + Some(identity("fixture-a")) + }) + .await + .unwrap(); + let second = get_or_fetch_identity(&cache, "fixture-b", t, || async { + calls.fetch_add(1, Ordering::SeqCst); + Some(identity("fixture-b")) + }) + .await + .unwrap(); + let cached = + get_or_fetch_identity(&cache, "fixture-a", t + Duration::seconds(60), || async { + calls.fetch_add(1, Ordering::SeqCst); + Some(identity("unexpected")) + }) + .await + .unwrap(); + let refreshed = + get_or_fetch_identity(&cache, "fixture-a", t + Duration::hours(13), || async { + calls.fetch_add(1, Ordering::SeqCst); + Some(identity("fixture-a-refreshed")) + }) + .await + .unwrap(); + + assert_eq!(first.name.as_deref(), Some("fixture-a")); + assert_eq!(second.name.as_deref(), Some("fixture-b")); + assert_eq!(cached, first); + assert_eq!(refreshed.name.as_deref(), Some("fixture-a-refreshed")); + assert_eq!(calls.load(Ordering::SeqCst), 3); + let cache = cache.lock().unwrap(); + assert!(cache.entries.keys().all(|key| !key.contains("fixture-a"))); + assert!(cache.entries.keys().all(|key| !key.contains("fixture-b"))); + assert!(cache.entries.keys().all(|key| key.len() == 64)); + } + + #[test] + fn identity_cache_expires_at_twelve_hours_and_when_clock_moves_backwards() { + let mut cache = IdentityCache::default(); + let t = at(1_800_000_000); + cache.insert("fixture", identity("cached"), t); + assert_eq!(cache.get("fixture", t + Duration::hours(12)), None); + + cache.insert("fixture", identity("cached"), t); + assert_eq!(cache.get("fixture", t - Duration::seconds(1)), None); + assert!(cache.entries.is_empty()); + } + + #[tokio::test] + async fn empty_identity_is_not_cached_but_name_only_identity_is() { + let cache = Mutex::new(IdentityCache::default()); + let calls = AtomicUsize::new(0); + let t = at(1_800_000_000); + cache + .lock() + .unwrap() + .insert("fixture", identity("stale"), t - IDENTITY_TTL); + + assert_eq!( + get_or_fetch_identity(&cache, "fixture", t, || async { + calls.fetch_add(1, Ordering::SeqCst); + None + }) + .await, + None + ); + let named = get_or_fetch_identity(&cache, "fixture", t, || async { + calls.fetch_add(1, Ordering::SeqCst); + Some(identity("only-name")) + }) + .await + .unwrap(); + let cached = get_or_fetch_identity(&cache, "fixture", t + Duration::seconds(1), || async { + calls.fetch_add(1, Ordering::SeqCst); + Some(identity("unexpected")) + }) + .await + .unwrap(); + + assert_eq!(named.name.as_deref(), Some("only-name")); + assert_eq!(cached, named); + assert_eq!(calls.load(Ordering::SeqCst), 2); + } + + #[test] + fn cache_evicts_old_entries_at_its_size_limit() { + let mut cache = IdentityCache::default(); + let t = at(1_800_000_000); + for index in 0..=MAX_IDENTITY_ENTRIES { + cache.insert( + &format!("token-{index}"), + identity("name"), + t + Duration::seconds(i64::try_from(index).unwrap()), + ); + } + assert_eq!(cache.entries.len(), MAX_IDENTITY_ENTRIES); + assert_eq!(cache.get("token-0", t + Duration::seconds(40)), None); + assert!(cache.get("token-32", t + Duration::seconds(40)).is_some()); + } +} diff --git a/rust/src/providers/huggingface/mod.rs b/rust/src/providers/huggingface/mod.rs index 102cbd3d35..474dd7f889 100644 --- a/rust/src/providers/huggingface/mod.rs +++ b/rust/src/providers/huggingface/mod.rs @@ -1,9 +1,9 @@ //! Hugging Face billing provider. //! //! Hugging Face exposes inference billing and optional ZeroGPU usage through -//! authenticated JSON endpoints. The billing data is presented as cost and -//! transient detail rows; it is deliberately not converted into a quota -//! window or a persisted identity record. +//! authenticated JSON endpoints, with a prepaid wallet balance available from +//! the browser session in Auto mode. Billing data is presented as cost and +//! transient detail rows; identity is kept only in a short-lived process cache. use async_trait::async_trait; use chrono::{DateTime, Datelike, TimeZone, Utc}; @@ -18,9 +18,11 @@ use crate::core::{ ProviderFetchResult, ProviderId, ProviderMetadata, RateWindow, SourceMode, UsageSnapshot, }; +mod identity_cache; mod wallet; -use wallet::{WalletCandidate, matching_wallet_balance, parse_wallet_balance}; +use identity_cache::{get_or_fetch_identity, process_identity_cache}; +use wallet::{WalletCandidate, fetch_matching_wallet_balance, parse_wallet_balance}; const BILLING_URL: &str = "https://huggingface.co/api/settings/billing/usage-v2"; const WHOAMI_URL: &str = "https://huggingface.co/api/whoami-v2"; @@ -170,16 +172,25 @@ impl HuggingFaceProvider { let zerogpu_url = Url::parse(ZEROGPU_URL) .map_err(|_| ProviderError::Other("Invalid Hugging Face ZeroGPU URL.".to_string()))?; - let (billing, identity, zerogpu, wallet_candidate) = tokio::join!( - self.fetch_json(billing_url, &token, PRIMARY_TIMEOUT), - self.fetch_optional_json(whoami_url, &token), + let (billing, identity, zerogpu) = tokio::join!( + self.fetch_json( + billing_url, + &token, + PRIMARY_TIMEOUT, + classify_billing_status + ), + get_or_fetch_identity(process_identity_cache(), &token, now, || async { + let profile = self.fetch_optional_json(whoami_url, &token).await?; + parse_identity(&profile) + }), self.fetch_optional_json(zerogpu_url, &token), - self.fetch_optional_wallet_candidate(), ); let billing = parse_billing(billing?)?; - let identity = identity.and_then(|value| parse_identity(&value)); let zerogpu = zerogpu.and_then(|value| parse_zerogpu(&value)); - let balance = matching_wallet_balance(identity.as_ref(), wallet_candidate); + let balance = fetch_matching_wallet_balance(ctx.source_mode, identity.as_ref(), || { + self.fetch_optional_wallet_candidate() + }) + .await; Ok(build_result(billing, identity, zerogpu, balance)) } @@ -217,7 +228,7 @@ impl HuggingFaceProvider { .send() .await?; if !response.status().is_success() { - return Err(classify_status(response.status())); + return Err(classify_optional_status(response.status())); } let body = read_bounded_body(response, "wallet response").await?; String::from_utf8(body).map_err(|_| { @@ -229,7 +240,9 @@ impl HuggingFaceProvider { } async fn fetch_optional_json(&self, url: Url, token: &str) -> Option { - self.fetch_json(url, token, OPTIONAL_TIMEOUT).await.ok() + self.fetch_json(url, token, OPTIONAL_TIMEOUT, classify_optional_status) + .await + .ok() } async fn fetch_json( @@ -237,6 +250,7 @@ impl HuggingFaceProvider { url: Url, token: &str, timeout: Duration, + classify: fn(StatusCode) -> ProviderError, ) -> Result { // Wrapper timeout, not just the client's PRIMARY_TIMEOUT: the // optional-fetch path (fetch_optional_json) overrides this with @@ -252,7 +266,7 @@ impl HuggingFaceProvider { .await?; let status = response.status(); if !status.is_success() { - return Err(classify_status(status)); + return Err(classify(status)); } let body = read_bounded_body(response, "JSON body").await?; @@ -462,14 +476,12 @@ fn parse_identity(value: &Value) -> Option { .get("isPro") .and_then(Value::as_bool) .map(|is_pro| if is_pro { "Pro" } else { "Free" }.to_string()); - (user_id.is_some() || name.is_some() || email.is_some() || plan.is_some()).then_some( - IdentitySnapshot { - user_id, - name, - email, - plan, - }, - ) + (user_id.is_some() || name.is_some() || email.is_some()).then_some(IdentitySnapshot { + user_id, + name, + email, + plan, + }) } fn safe_text(value: Option<&str>) -> Option { @@ -577,7 +589,7 @@ fn format_usd(value: f64) -> String { format!("${value:.2}") } -fn classify_status(status: StatusCode) -> ProviderError { +fn classify_optional_status(status: StatusCode) -> ProviderError { match status { StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN => ProviderError::AuthRequired, StatusCode::TOO_MANY_REQUESTS => { @@ -590,6 +602,17 @@ fn classify_status(status: StatusCode) -> ProviderError { } } +fn classify_billing_status(status: StatusCode) -> ProviderError { + if status == StatusCode::FORBIDDEN { + ProviderError::Other( + "The Hugging Face token lacks billing access. Use a classic read token or enable Billing read on a fine-grained token." + .to_string(), + ) + } else { + classify_optional_status(status) + } +} + #[cfg(test)] mod tests { use super::*; @@ -775,6 +798,7 @@ mod tests { assert_eq!(identity.email.as_deref(), Some("n@example.test")); assert_eq!(identity.plan.as_deref(), Some("Pro")); assert!(parse_identity(&json!({"email": "bad\nemail"})).is_none()); + assert!(parse_identity(&json!({"isPro": true})).is_none()); } #[test] @@ -786,20 +810,84 @@ mod tests { StatusCode::INTERNAL_SERVER_ERROR, StatusCode::BAD_REQUEST, ] { - let error = classify_status(status).to_string(); + let error = classify_optional_status(status).to_string(); assert!(!error.contains(token)); assert!(!error.contains("response body")); } assert!(matches!( - classify_status(StatusCode::UNAUTHORIZED), + classify_optional_status(StatusCode::UNAUTHORIZED), ProviderError::AuthRequired )); assert!(matches!( - classify_status(StatusCode::FORBIDDEN), + classify_optional_status(StatusCode::FORBIDDEN), ProviderError::AuthRequired )); } + #[test] + fn billing_status_errors_retain_actionable_classification() { + let token = "hf_secret_fixture"; + assert!(matches!( + classify_billing_status(StatusCode::UNAUTHORIZED), + ProviderError::AuthRequired + )); + + let forbidden = classify_billing_status(StatusCode::FORBIDDEN).to_string(); + assert_eq!( + forbidden, + "The Hugging Face token lacks billing access. Use a classic read token or enable Billing read on a fine-grained token." + ); + assert!(!forbidden.contains(token)); + + assert_eq!( + classify_billing_status(StatusCode::TOO_MANY_REQUESTS).to_string(), + "Hugging Face API rate limited (HTTP 429)." + ); + assert_eq!( + classify_billing_status(StatusCode::SERVICE_UNAVAILABLE).to_string(), + "Hugging Face service unavailable (HTTP 5xx)." + ); + } + + #[tokio::test] + async fn html_billing_403_body_still_returns_the_permission_message() { + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.unwrap(); + let mut request = [0; 1024]; + assert!(stream.read(&mut request).await.unwrap() > 0); + let body = "Billing access denied"; + let response = format!( + "HTTP/1.1 403 Forbidden\r\nContent-Type: text/html\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}", + body.len(), + body + ); + stream.write_all(response.as_bytes()).await.unwrap(); + }); + + let mut provider = HuggingFaceProvider::new(); + provider.client = Client::builder().no_proxy().build().unwrap(); + let url = Url::parse(&format!("http://{address}/billing")).unwrap(); + let error = provider + .fetch_json( + url, + "hf_secret_fixture", + PRIMARY_TIMEOUT, + classify_billing_status, + ) + .await + .unwrap_err(); + server.await.unwrap(); + assert_eq!( + error.to_string(), + "The Hugging Face token lacks billing access. Use a classic read token or enable Billing read on a fine-grained token." + ); + assert!(!error.to_string().contains("hf_secret_fixture")); + } + #[test] fn result_uses_cost_and_transient_details_without_quota_windows() { let result = build_result( diff --git a/rust/src/providers/huggingface/wallet.rs b/rust/src/providers/huggingface/wallet.rs index 9958c6891f..88a8614e9b 100644 --- a/rust/src/providers/huggingface/wallet.rs +++ b/rust/src/providers/huggingface/wallet.rs @@ -1,7 +1,9 @@ +use std::future::Future; + use serde_json::Value; use super::IdentitySnapshot; -use crate::core::ProviderError; +use crate::core::{ProviderError, SourceMode}; #[derive(Debug, Clone, PartialEq)] pub(super) struct WalletCandidate { @@ -18,6 +20,25 @@ pub(super) fn matching_wallet_balance( (candidate.user_id == expected_user_id).then_some(candidate.balance) } +pub(super) async fn fetch_matching_wallet_balance( + source_mode: SourceMode, + identity: Option<&IdentitySnapshot>, + load_candidate: F, +) -> Option +where + F: FnOnce() -> Fut, + Fut: Future>, +{ + if source_mode != SourceMode::Auto + || identity + .and_then(|identity| identity.user_id.as_deref()) + .is_none() + { + return None; + } + matching_wallet_balance(identity, load_candidate().await) +} + pub(super) fn parse_wallet_balance(html: &str) -> Result { let mut current = Vec::new(); let mut legacy = Vec::new(); @@ -125,7 +146,32 @@ fn decode_html_entities(raw: &str) -> Result { #[cfg(test)] mod tests { + use std::sync::atomic::{AtomicUsize, Ordering}; + + use chrono::{DateTime, Duration, Utc}; + use super::*; + use crate::providers::huggingface::identity_cache::{IdentityCache, get_or_fetch_identity}; + + fn identity(user_id: Option<&str>) -> IdentitySnapshot { + IdentitySnapshot { + user_id: user_id.map(str::to_string), + name: Some("fixture".to_string()), + email: None, + plan: None, + } + } + + fn candidate(user_id: &str, balance: f64) -> WalletCandidate { + WalletCandidate { + user_id: user_id.to_string(), + balance, + } + } + + fn at(seconds: i64) -> DateTime { + DateTime::from_timestamp(seconds, 0).unwrap() + } #[test] fn candidate_is_attached_only_to_the_matching_token_identity() { @@ -155,6 +201,94 @@ mod tests { assert_eq!(matching_wallet_balance(None, Some(candidate)), None); } + #[tokio::test] + async fn wallet_loader_runs_only_for_auto_with_a_user_id() { + let calls = AtomicUsize::new(0); + let api_only_identity = identity(Some("user-a")); + assert_eq!( + fetch_matching_wallet_balance(SourceMode::OAuth, Some(&api_only_identity), || { + calls.fetch_add(1, Ordering::SeqCst); + std::future::ready(Some(candidate("user-a", 12.5))) + }) + .await, + None + ); + assert_eq!(calls.load(Ordering::SeqCst), 0); + + let auto_identity = identity(Some("user-a")); + assert_eq!( + fetch_matching_wallet_balance(SourceMode::Auto, Some(&auto_identity), || { + calls.fetch_add(1, Ordering::SeqCst); + std::future::ready(Some(candidate("user-a", 12.5))) + }) + .await, + Some(12.5) + ); + assert_eq!(calls.load(Ordering::SeqCst), 1); + + let no_id_identity = identity(None); + assert_eq!( + fetch_matching_wallet_balance(SourceMode::Auto, Some(&no_id_identity), || { + calls.fetch_add(1, Ordering::SeqCst); + std::future::ready(Some(candidate("user-a", 12.5))) + }) + .await, + None + ); + assert_eq!( + fetch_matching_wallet_balance(SourceMode::Auto, None, || { + calls.fetch_add(1, Ordering::SeqCst); + std::future::ready(Some(candidate("user-a", 12.5))) + }) + .await, + None + ); + assert_eq!(calls.load(Ordering::SeqCst), 1); + } + + #[tokio::test] + async fn wallet_loader_rejects_mismatched_browser_identity() { + let identity = identity(Some("api-user")); + let calls = AtomicUsize::new(0); + assert_eq!( + fetch_matching_wallet_balance(SourceMode::Auto, Some(&identity), || { + calls.fetch_add(1, Ordering::SeqCst); + std::future::ready(Some(candidate("browser-user", 40.0))) + }) + .await, + None + ); + assert_eq!(calls.load(Ordering::SeqCst), 1); + } + + #[tokio::test] + async fn cached_identity_cannot_keep_wallet_after_browser_account_changes() { + let token = "fixture-token"; + let now = at(1_800_000_000); + let cache = std::sync::Mutex::new(IdentityCache::default()); + cache + .lock() + .unwrap() + .insert(token, identity(Some("api-user")), now); + let cached_identity = + get_or_fetch_identity(&cache, token, now + Duration::minutes(1), || async { + panic!("fresh identity should be served from cache") + }) + .await + .unwrap(); + let calls = AtomicUsize::new(0); + + let balance = + fetch_matching_wallet_balance(SourceMode::Auto, Some(&cached_identity), || { + calls.fetch_add(1, Ordering::SeqCst); + std::future::ready(Some(candidate("new-browser-user", 40.0))) + }) + .await; + + assert_eq!(balance, None); + assert_eq!(calls.load(Ordering::SeqCst), 1); + } + #[test] fn parser_prefers_unique_current_balance_and_supports_legacy_cents() { let current = r#"
"#; From 78d8932da7d2a383e0cbe6536fd2eddc328d4476 Mon Sep 17 00:00:00 2001 From: RCD <90105158+Finesssee@users.noreply.github.com> Date: Thu, 1 Oct 2026 16:05:12 +0700 Subject: [PATCH 2/2] Keep Hugging Face wallet for an Auto source with a selected token account The shell forced the OAuth (API-only) source for any token account with an environment override, which skipped the prepaid wallet. A new Provider::token_account_preserves_auto_source hook lets Hugging Face keep Auto, matching upstream's base-source resolver; an explicit API source still stays API-only. --- .../src-tauri/src/commands/mod.rs | 3 + .../src-tauri/src/commands/providers.rs | 4 +- .../commands/token_account_source_tests.rs | 74 +++++++++++++++++++ rust/src/core/provider.rs | 11 +++ rust/src/providers/huggingface/mod.rs | 6 ++ 5 files changed, 97 insertions(+), 1 deletion(-) create mode 100644 apps/desktop-tauri/src-tauri/src/commands/token_account_source_tests.rs diff --git a/apps/desktop-tauri/src-tauri/src/commands/mod.rs b/apps/desktop-tauri/src-tauri/src/commands/mod.rs index fdfaeda967..40b6085bbd 100644 --- a/apps/desktop-tauri/src-tauri/src/commands/mod.rs +++ b/apps/desktop-tauri/src-tauri/src/commands/mod.rs @@ -78,6 +78,9 @@ pub(crate) use usage_items::*; #[cfg(test)] mod tests; +#[cfg(test)] +mod token_account_source_tests; + pub use chart::*; pub use spend_contract::*; pub use tokens::*; diff --git a/apps/desktop-tauri/src-tauri/src/commands/providers.rs b/apps/desktop-tauri/src-tauri/src/commands/providers.rs index 103d47b869..9c0f673f40 100644 --- a/apps/desktop-tauri/src-tauri/src/commands/providers.rs +++ b/apps/desktop-tauri/src-tauri/src/commands/providers.rs @@ -92,7 +92,9 @@ pub(crate) fn build_fetch_context( let (mut source_mode, mut cookie_header, fails_closed_without_cookie) = if id.cookie_domain().is_none() { - let source_mode = if active_token_env.is_some() { + let keeps_auto = + usage_source == SourceMode::Auto && provider.token_account_preserves_auto_source(); + let source_mode = if active_token_env.is_some() && !keeps_auto { SourceMode::OAuth } else { usage_source diff --git a/apps/desktop-tauri/src-tauri/src/commands/token_account_source_tests.rs b/apps/desktop-tauri/src-tauri/src/commands/token_account_source_tests.rs new file mode 100644 index 0000000000..ee1cee114e --- /dev/null +++ b/apps/desktop-tauri/src-tauri/src/commands/token_account_source_tests.rs @@ -0,0 +1,74 @@ +//! Source-mode resolution for selected token accounts. + +use std::collections::HashMap; + +use codexbar::core::{ + ProviderAccountData, ProviderId, SourceMode, TokenAccount, instantiate_provider, +}; +use codexbar::settings::{ApiKeys, ManualCookies, Settings}; + +fn token_accounts(id: ProviderId, token: &str) -> HashMap { + let mut data = ProviderAccountData::new(); + data.add_account(TokenAccount::new("Work", token)); + HashMap::from([(id, data)]) +} + +fn context_with_usage_source( + id: ProviderId, + usage_source: &str, + accounts: &HashMap, +) -> codexbar::core::FetchContext { + let mut settings = Settings::default(); + settings.set_usage_source(id, usage_source); + super::build_fetch_context( + id, + &settings, + &ManualCookies::default(), + &ApiKeys::default(), + accounts, + ) +} + +#[test] +fn huggingface_token_account_keeps_auto_so_the_wallet_is_still_read() { + let accounts = token_accounts(ProviderId::HuggingFace, "hf_account_token"); + let ctx = context_with_usage_source(ProviderId::HuggingFace, "auto", &accounts); + + assert_eq!(ctx.source_mode, SourceMode::Auto); + assert_eq!(ctx.api_key.as_deref(), Some("hf_account_token")); +} + +#[test] +fn huggingface_token_account_with_explicit_api_source_stays_api_only() { + let accounts = token_accounts(ProviderId::HuggingFace, "hf_account_token"); + let ctx = context_with_usage_source(ProviderId::HuggingFace, "oauth", &accounts); + + assert_eq!(ctx.source_mode, SourceMode::OAuth); + assert_eq!(ctx.api_key.as_deref(), Some("hf_account_token")); +} + +#[test] +fn huggingface_token_account_with_an_unsupported_stored_source_falls_back_to_api() { + let accounts = token_accounts(ProviderId::HuggingFace, "hf_account_token"); + for stale in ["web", "cli"] { + let ctx = context_with_usage_source(ProviderId::HuggingFace, stale, &accounts); + assert_eq!(ctx.source_mode, SourceMode::OAuth, "{stale}"); + } +} + +#[test] +fn huggingface_without_a_token_account_follows_the_usage_source() { + let ctx = context_with_usage_source(ProviderId::HuggingFace, "auto", &HashMap::new()); + assert_eq!(ctx.source_mode, SourceMode::Auto); +} + +#[test] +fn only_huggingface_keeps_auto_for_token_accounts() { + for id in ProviderId::all() { + assert_eq!( + instantiate_provider(*id).token_account_preserves_auto_source(), + *id == ProviderId::HuggingFace, + "{id:?}" + ); + } +} diff --git a/rust/src/core/provider.rs b/rust/src/core/provider.rs index cc954839f7..a96a36a73b 100755 --- a/rust/src/core/provider.rs +++ b/rust/src/core/provider.rs @@ -860,6 +860,17 @@ pub trait Provider: Send + Sync { false } + /// Whether a selected token account leaves an `Auto` usage source as `Auto`. + /// + /// The shell normally maps a token account with an environment override to + /// the OAuth (API-only) lane. A provider whose `Auto` source adds an + /// optional extra on top of the API credential, such as the Hugging Face + /// prepaid wallet, opts in so the account token does not silently drop it. + /// An explicitly chosen non-Auto usage source still maps to OAuth. + fn token_account_preserves_auto_source(&self) -> bool { + false + } + /// How the shell treats a manual cookie source with no cookie present. /// /// `Fallback` lets the shell remap to its generic browser-cookie attempt. diff --git a/rust/src/providers/huggingface/mod.rs b/rust/src/providers/huggingface/mod.rs index 474dd7f889..beb959eb04 100644 --- a/rust/src/providers/huggingface/mod.rs +++ b/rust/src/providers/huggingface/mod.rs @@ -325,6 +325,12 @@ impl Provider for HuggingFaceProvider { fn available_sources(&self) -> Vec { vec![SourceMode::Auto, SourceMode::OAuth] } + + /// Upstream keeps the user's source for a selected API token account, so an + /// `Auto` fetch still reads the prepaid wallet from the browser session. + fn token_account_preserves_auto_source(&self) -> bool { + true + } } fn resolve_token(ctx: &FetchContext) -> Result {