use std::collections::HashMap;
use std::future::Future;
use std::sync::{Arc, OnceLock, PoisonError, RwLock};
use std::time::{Duration, SystemTime};
use super::error::CredentialError;
const EXPIRY_SKEW: Duration = Duration::from_secs(120);
const MIN_TTL: Duration = Duration::from_secs(60);
const MAX_TTL: Duration = Duration::from_secs(3600);
#[derive(Debug, Clone)]
struct CachedToken {
token: String,
expires_at: SystemTime,
}
#[derive(Default)]
struct TokenCache {
entries: RwLock<HashMap<String, CachedToken>>,
mints: tokio::sync::Mutex<HashMap<String, Arc<tokio::sync::Mutex<()>>>>,
}
fn cache() -> &'static TokenCache {
static CACHE: OnceLock<TokenCache> = OnceLock::new();
CACHE.get_or_init(TokenCache::default)
}
#[must_use]
pub fn clamp_ttl(declared: Option<u64>) -> Duration {
match declared {
None | Some(0) => MAX_TTL,
Some(secs) => Duration::from_secs(secs).clamp(MIN_TTL, MAX_TTL),
}
}
fn cached(key: &str) -> Option<String> {
let guard = cache()
.entries
.read()
.unwrap_or_else(PoisonError::into_inner);
let token = guard.get(key).and_then(|entry| {
(entry.expires_at > SystemTime::now() + EXPIRY_SKEW).then(|| entry.token.clone())
});
drop(guard);
token
}
fn store(key: &str, token: &str, ttl: Duration) {
cache()
.entries
.write()
.unwrap_or_else(PoisonError::into_inner)
.insert(
key.to_owned(),
CachedToken {
token: token.to_owned(),
expires_at: SystemTime::now() + ttl,
},
);
}
pub async fn token_for<F, Fut>(key: &str, mint: F) -> Result<String, CredentialError>
where
F: FnOnce() -> Fut,
Fut: Future<Output = Result<(String, Duration), CredentialError>>,
{
if let Some(token) = cached(key) {
return Ok(token);
}
let lock = {
let mut mints = cache().mints.lock().await;
Arc::clone(mints.entry(key.to_owned()).or_default())
};
let _minting = lock.lock().await;
if let Some(token) = cached(key) {
return Ok(token);
}
let (token, ttl) = mint().await?;
store(key, &token, ttl);
Ok(token)
}