Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions apps/desktop-tauri/src-tauri/src/commands/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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::*;
Expand Down
4 changes: 3 additions & 1 deletion apps/desktop-tauri/src-tauri/src/commands/providers.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
@@ -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<ProviderId, ProviderAccountData> {
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<ProviderId, ProviderAccountData>,
) -> 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:?}"
);
}
}
11 changes: 11 additions & 0 deletions rust/src/core/provider.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
240 changes: 240 additions & 0 deletions rust/src/providers/huggingface/identity_cache.rs
Original file line number Diff line number Diff line change
@@ -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<Utc>,
}

#[derive(Debug, Default)]
pub(super) struct IdentityCache {
entries: HashMap<String, CachedIdentity>,
}

impl IdentityCache {
pub(super) fn get(&mut self, token: &str, now: DateTime<Utc>) -> Option<IdentitySnapshot> {
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<Utc>) {
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<Utc>) {
self.entries.retain(|_, entry| is_fresh(entry, now));
}
}

fn is_fresh(entry: &CachedIdentity, now: &DateTime<Utc>) -> 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<Mutex<IdentityCache>> = OnceLock::new();

pub(super) fn process_identity_cache() -> &'static Mutex<IdentityCache> {
PROCESS_IDENTITY_CACHE.get_or_init(|| Mutex::new(IdentityCache::default()))
}

pub(super) async fn get_or_fetch_identity<F, Fut>(
cache: &Mutex<IdentityCache>,
token: &str,
now: DateTime<Utc>,
fetch: F,
) -> Option<IdentitySnapshot>
where
F: FnOnce() -> Fut,
Fut: Future<Output = Option<IdentitySnapshot>>,
{
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<Utc> {
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());
}
}
Loading