use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use sha2::{Digest, Sha256};
use tokio::sync::Mutex;
use tokio::time::Instant;
use tracing::{debug, warn};
use crate::default_token_factory::{FetchedToken, MintReason};
use crate::ZerobusResult;
pub(crate) const DEFAULT_REFRESH_BUFFER: Duration = Duration::from_secs(300);
struct CachedToken {
value: String,
expires_at: Instant,
}
impl CachedToken {
fn is_expired(&self) -> bool {
Instant::now() >= self.expires_at
}
}
#[derive(Clone, PartialEq, Eq, Hash)]
struct TokenKey {
client_id: String,
secret_digest: [u8; 32],
table_name: String,
}
impl TokenKey {
fn new(client_id: &str, client_secret: &str, table_name: &str) -> Self {
let secret_digest = Sha256::digest(client_secret.as_bytes()).into();
Self {
client_id: client_id.to_string(),
secret_digest,
table_name: table_name.to_string(),
}
}
}
type Slot = Arc<Mutex<Option<CachedToken>>>;
pub(crate) struct TokenCache {
entries: Mutex<HashMap<TokenKey, Slot>>,
refresh_buffer: Duration,
enabled: bool,
}
impl TokenCache {
pub(crate) fn new(enabled: bool, refresh_buffer: Duration) -> Self {
Self {
entries: Mutex::new(HashMap::new()),
refresh_buffer,
enabled,
}
}
pub(crate) async fn get_or_fetch<F, Fut>(
&self,
client_id: &str,
client_secret: &str,
table_name: &str,
fetch: F,
) -> ZerobusResult<String>
where
F: FnOnce(MintReason) -> Fut,
Fut: std::future::Future<Output = ZerobusResult<FetchedToken>>,
{
if !self.enabled {
return fetch(MintReason::CacheDisabled)
.await
.map(|fetched| fetched.token);
}
let key = TokenKey::new(client_id, client_secret, table_name);
let slot = {
let mut entries = self.entries.lock().await;
if !entries.contains_key(&key) {
Self::prune_expired(&mut entries);
}
Arc::clone(entries.entry(key).or_default())
};
let mut guard = slot.lock().await;
if let Some(cached) = guard.as_ref() {
if !self.needs_refresh(cached) {
debug!(table = %table_name, "token cache hit, reusing cached token");
return Ok(cached.value.clone());
}
}
let reason = if guard.is_some() {
MintReason::Refresh
} else {
MintReason::ColdMiss
};
let fetched = match fetch(reason).await {
Ok(fetched) => fetched,
Err(err) => {
if err.is_retryable() {
if let Some(cached) = guard.as_ref() {
if !cached.is_expired() {
warn!(table = %table_name, "token refresh failed (retryable); serving still-valid cached token");
return Ok(cached.value.clone());
}
}
}
return Err(err);
}
};
let token = fetched.token.clone();
let expires_at = fetched
.expires_in
.and_then(|ttl| Instant::now().checked_add(ttl));
match expires_at {
Some(expires_at) => {
*guard = Some(CachedToken {
value: fetched.token,
expires_at,
});
}
None => {
let keep_existing = guard.as_ref().is_some_and(|cached| !cached.is_expired());
if !keep_existing {
*guard = None;
}
}
}
Ok(token)
}
pub(crate) async fn invalidate(&self, client_id: &str, client_secret: &str, table_name: &str) {
if !self.enabled {
return;
}
let key = TokenKey::new(client_id, client_secret, table_name);
if self.entries.lock().await.remove(&key).is_some() {
debug!(table = %table_name, "token cache entry invalidated after auth rejection");
}
}
fn needs_refresh(&self, cached: &CachedToken) -> bool {
match Instant::now().checked_add(self.refresh_buffer) {
Some(deadline) => deadline >= cached.expires_at,
None => true,
}
}
fn prune_expired(entries: &mut HashMap<TokenKey, Slot>) {
entries.retain(|_, slot| match slot.try_lock() {
Ok(guard) => match guard.as_ref() {
Some(cached) => !cached.is_expired(),
None => true,
},
Err(_) => true,
});
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
fn fetched(token: &str, ttl_secs: Option<u64>) -> FetchedToken {
FetchedToken {
token: token.to_string(),
expires_in: ttl_secs.map(Duration::from_secs),
}
}
#[tokio::test]
async fn caches_token_across_calls() {
let cache = TokenCache::new(true, Duration::from_secs(60));
let calls = AtomicUsize::new(0);
let make = |_reason| async {
calls.fetch_add(1, Ordering::SeqCst);
Ok(fetched("tok", Some(3600)))
};
let a = cache
.get_or_fetch("id", "secret", "c.s.t", make)
.await
.unwrap();
let b = cache
.get_or_fetch("id", "secret", "c.s.t", make)
.await
.unwrap();
assert_eq!(a, "tok");
assert_eq!(b, "tok");
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"second call should hit cache"
);
}
#[tokio::test]
async fn refetches_when_within_refresh_buffer() {
let cache = TokenCache::new(true, Duration::from_secs(60));
let calls = AtomicUsize::new(0);
let make = |_reason| async {
let n = calls.fetch_add(1, Ordering::SeqCst);
Ok(fetched(&format!("tok{n}"), Some(1)))
};
let a = cache
.get_or_fetch("id", "secret", "c.s.t", make)
.await
.unwrap();
let b = cache
.get_or_fetch("id", "secret", "c.s.t", make)
.await
.unwrap();
assert_eq!(a, "tok0");
assert_eq!(b, "tok1");
assert_eq!(calls.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn separate_tables_get_separate_entries() {
let cache = TokenCache::new(true, Duration::from_secs(60));
let calls = AtomicUsize::new(0);
let make = |_reason| async {
let n = calls.fetch_add(1, Ordering::SeqCst);
Ok(fetched(&format!("tok{n}"), Some(3600)))
};
let a = cache
.get_or_fetch("id", "secret", "c.s.t1", make)
.await
.unwrap();
let b = cache
.get_or_fetch("id", "secret", "c.s.t2", make)
.await
.unwrap();
assert_ne!(a, b);
assert_eq!(calls.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn rotated_secret_gets_new_entry() {
let cache = TokenCache::new(true, Duration::from_secs(60));
let calls = AtomicUsize::new(0);
let make = |_reason| async {
let n = calls.fetch_add(1, Ordering::SeqCst);
Ok(fetched(&format!("tok{n}"), Some(3600)))
};
cache
.get_or_fetch("id", "secret-v1", "c.s.t", make)
.await
.unwrap();
cache
.get_or_fetch("id", "secret-v2", "c.s.t", make)
.await
.unwrap();
assert_eq!(calls.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn token_without_ttl_is_not_cached() {
let cache = TokenCache::new(true, Duration::from_secs(60));
let calls = AtomicUsize::new(0);
let make = |_reason| async {
calls.fetch_add(1, Ordering::SeqCst);
Ok(fetched("tok", None))
};
cache
.get_or_fetch("id", "secret", "c.s.t", make)
.await
.unwrap();
cache
.get_or_fetch("id", "secret", "c.s.t", make)
.await
.unwrap();
assert_eq!(calls.load(Ordering::SeqCst), 2, "no TTL means no caching");
}
#[tokio::test]
async fn invalidate_forces_remint_on_next_call() {
let cache = TokenCache::new(true, Duration::from_secs(60));
let calls = AtomicUsize::new(0);
let make = |_reason| async {
calls.fetch_add(1, Ordering::SeqCst);
Ok(fetched("tok", Some(3600)))
};
cache
.get_or_fetch("id", "secret", "c.s.t", make)
.await
.unwrap();
cache.invalidate("id", "secret", "c.s.t").await;
cache
.get_or_fetch("id", "secret", "c.s.t", make)
.await
.unwrap();
assert_eq!(calls.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn disabled_cache_always_fetches() {
let cache = TokenCache::new(false, Duration::from_secs(60));
let calls = AtomicUsize::new(0);
let make = |_reason| async {
calls.fetch_add(1, Ordering::SeqCst);
Ok(fetched("tok", Some(3600)))
};
cache
.get_or_fetch("id", "secret", "c.s.t", make)
.await
.unwrap();
cache
.get_or_fetch("id", "secret", "c.s.t", make)
.await
.unwrap();
assert_eq!(calls.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn fetch_error_leaves_no_cached_entry() {
let cache = TokenCache::new(true, Duration::from_secs(60));
let err = cache
.get_or_fetch("id", "secret", "c.s.t", |_reason| async {
Err(crate::ZerobusError::TokenFetchError("boom".to_string()))
})
.await;
assert!(err.is_err());
let ok = cache
.get_or_fetch("id", "secret", "c.s.t", |_reason| async {
Ok(fetched("tok", Some(3600)))
})
.await
.unwrap();
assert_eq!(ok, "tok");
}
#[tokio::test]
async fn refresh_failure_serves_still_valid_token() {
let cache = TokenCache::new(true, Duration::from_secs(60));
let seeded = cache
.get_or_fetch("id", "secret", "c.s.t", |_reason| async {
Ok(fetched("valid", Some(30)))
})
.await
.unwrap();
assert_eq!(seeded, "valid");
let served = cache
.get_or_fetch("id", "secret", "c.s.t", |_reason| async {
Err(crate::ZerobusError::TokenFetchError("blip".to_string()))
})
.await
.unwrap();
assert_eq!(served, "valid");
}
#[tokio::test]
async fn refresh_failure_propagates_non_retryable_error() {
let cache = TokenCache::new(true, Duration::from_secs(60));
cache
.get_or_fetch("id", "secret", "c.s.t", |_reason| async {
Ok(fetched("valid", Some(30)))
})
.await
.unwrap();
let result = cache
.get_or_fetch("id", "secret", "c.s.t", |_reason| async {
Err(crate::ZerobusError::InvalidUCTokenError(
"revoked".to_string(),
))
})
.await;
assert!(matches!(
result,
Err(crate::ZerobusError::InvalidUCTokenError(_))
));
}
#[tokio::test]
async fn no_ttl_response_does_not_evict_valid_token() {
let cache = TokenCache::new(true, Duration::from_secs(60));
cache
.get_or_fetch("id", "secret", "c.s.t", |_reason| async {
Ok(fetched("valid", Some(30)))
})
.await
.unwrap();
let fresh = cache
.get_or_fetch("id", "secret", "c.s.t", |_reason| async {
Ok(fetched("nottl", None))
})
.await
.unwrap();
assert_eq!(fresh, "nottl");
let served = cache
.get_or_fetch("id", "secret", "c.s.t", |_reason| async {
Err(crate::ZerobusError::TokenFetchError("blip".to_string()))
})
.await
.unwrap();
assert_eq!(served, "valid");
}
#[tokio::test]
async fn single_flight_mints_once_for_concurrent_callers() {
let cache = Arc::new(TokenCache::new(true, Duration::from_secs(60)));
let calls = Arc::new(AtomicUsize::new(0));
let mut handles = Vec::new();
for _ in 0..16 {
let cache = Arc::clone(&cache);
let calls = Arc::clone(&calls);
handles.push(tokio::spawn(async move {
cache
.get_or_fetch("id", "secret", "c.s.t", |_reason| async {
calls.fetch_add(1, Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(20)).await;
Ok(fetched("tok", Some(3600)))
})
.await
.unwrap()
}));
}
for handle in handles {
assert_eq!(handle.await.unwrap(), "tok");
}
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"single-flight must mint exactly once for concurrent same-key callers"
);
}
}