use std::collections::{BTreeSet, HashMap};
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use base64::Engine as _;
use parking_lot::Mutex;
use tokio::sync::Mutex as AsyncMutex;
use crate::internal_auth::{InternalAuthNError, InternalAuthenticator, PlatformIdentity};
pub const DEFAULT_TOKEN_REVIEW_CACHE_TTL: Duration = Duration::from_secs(30);
pub const MAX_TOKEN_REVIEW_CACHE_TTL: Duration = Duration::from_mins(5);
const NEGATIVE_CACHE_TTL: Duration = Duration::from_secs(1);
pub const DEFAULT_AUTHENTICATION_TIMEOUT: Duration = Duration::from_secs(10);
const SWEEP_INTERVAL: u32 = 32;
pub const MAX_CACHE_ENTRIES: usize = 10_000;
#[derive(Debug, thiserror::Error)]
#[error(
"internal-auth cache TTL must be > 0 and <= {MAX_TOKEN_REVIEW_CACHE_TTL:?}, got {actual:?}"
)]
pub struct InvalidCacheTtl {
actual: Duration,
}
#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
struct TokenKey([u8; 32]);
impl TokenKey {
fn of(token: &str) -> Self {
use sha2::{Digest as _, Sha256};
Self(Sha256::digest(token.as_bytes()).into())
}
}
impl std::fmt::Debug for TokenKey {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "TokenKey({:02x}{:02x}..)", self.0[0], self.0[1])
}
}
enum CacheEntry {
Valid {
identity: PlatformIdentity,
expires_at: Instant,
},
Rejected {
expires_at: Instant,
},
}
impl CacheEntry {
fn expires_at(&self) -> Instant {
match self {
Self::Valid { expires_at, .. } | Self::Rejected { expires_at } => *expires_at,
}
}
}
enum CacheLookup {
Valid(PlatformIdentity),
Rejected,
Miss,
}
struct ExpiringCache {
entries: HashMap<TokenKey, CacheEntry>,
by_expiry: BTreeSet<(Instant, TokenKey)>,
}
impl ExpiringCache {
fn new() -> Self {
Self {
entries: HashMap::new(),
by_expiry: BTreeSet::new(),
}
}
fn len(&self) -> usize {
self.entries.len()
}
fn contains_key(&self, key: &TokenKey) -> bool {
self.entries.contains_key(key)
}
fn get(&self, key: &TokenKey) -> Option<&CacheEntry> {
self.entries.get(key)
}
fn insert(&mut self, key: TokenKey, entry: CacheEntry) {
let expires_at = entry.expires_at();
if let Some(previous) = self.entries.insert(key, entry) {
self.by_expiry.remove(&(previous.expires_at(), key));
}
self.by_expiry.insert((expires_at, key));
}
fn remove(&mut self, key: &TokenKey) {
if let Some(entry) = self.entries.remove(key) {
self.by_expiry.remove(&(entry.expires_at(), *key));
}
}
fn sweep_expired(&mut self, now: Instant) {
let live = self.by_expiry.split_off(&(now, TokenKey([0; 32])));
let expired = std::mem::replace(&mut self.by_expiry, live);
for (_, key) in expired {
self.entries.remove(&key);
}
}
fn evict_soonest(&mut self) {
if let Some((expires_at, key)) = self.by_expiry.pop_first() {
self.entries.remove(&key);
debug_assert!(
!self.entries.contains_key(&key),
"index and map disagreed about {key:?} at {expires_at:?}"
);
}
}
}
enum ExpClaim {
NotJwt,
Expires(u64),
Unreadable,
}
fn jwt_exp_claim(token: &str) -> ExpClaim {
let mut parts = token.split('.');
let (Some(_header), Some(payload_b64), Some(_signature), None) =
(parts.next(), parts.next(), parts.next(), parts.next())
else {
return ExpClaim::NotJwt;
};
let Ok(payload) = base64::engine::general_purpose::URL_SAFE_NO_PAD.decode(payload_b64) else {
return ExpClaim::Unreadable;
};
let Ok(value) = serde_json::from_slice::<serde_json::Value>(&payload) else {
return ExpClaim::Unreadable;
};
match value.get("exp").and_then(serde_json::Value::as_u64) {
Some(exp) => ExpClaim::Expires(exp),
None => ExpClaim::Unreadable,
}
}
fn clamped_expiry(token: &str, now: Instant, ttl: Duration) -> Instant {
let ttl_expiry = now.checked_add(ttl).unwrap_or(now);
let exp_secs = match jwt_exp_claim(token) {
ExpClaim::NotJwt => return ttl_expiry,
ExpClaim::Unreadable => return now,
ExpClaim::Expires(exp) => exp,
};
let Some(exp_at) = UNIX_EPOCH.checked_add(Duration::from_secs(exp_secs)) else {
return ttl_expiry;
};
let Ok(remaining) = exp_at.duration_since(SystemTime::now()) else {
return now;
};
let exp_expiry = now.checked_add(remaining).unwrap_or(ttl_expiry);
ttl_expiry.min(exp_expiry)
}
pub struct CachingInternalAuthenticator<A> {
inner: A,
ttl: Duration,
authentication_timeout: Duration,
cache: Mutex<ExpiringCache>,
inflight: Mutex<HashMap<TokenKey, Arc<AsyncMutex<()>>>>,
sweep_counter: AtomicU32,
}
impl<A> std::fmt::Debug for CachingInternalAuthenticator<A> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CachingInternalAuthenticator")
.field("ttl", &self.ttl)
.finish_non_exhaustive()
}
}
impl<A> CachingInternalAuthenticator<A> {
pub fn new(inner: A, ttl: Duration) -> Result<Self, InvalidCacheTtl> {
if ttl == Duration::ZERO || ttl > MAX_TOKEN_REVIEW_CACHE_TTL {
return Err(InvalidCacheTtl { actual: ttl });
}
Ok(Self {
inner,
ttl,
authentication_timeout: DEFAULT_AUTHENTICATION_TIMEOUT,
cache: Mutex::new(ExpiringCache::new()),
inflight: Mutex::new(HashMap::new()),
sweep_counter: AtomicU32::new(0),
})
}
#[must_use]
pub fn with_default_ttl(inner: A) -> Self {
Self {
inner,
ttl: DEFAULT_TOKEN_REVIEW_CACHE_TTL,
authentication_timeout: DEFAULT_AUTHENTICATION_TIMEOUT,
cache: Mutex::new(ExpiringCache::new()),
inflight: Mutex::new(HashMap::new()),
sweep_counter: AtomicU32::new(0),
}
}
#[must_use]
pub fn with_authentication_timeout(mut self, timeout: Duration) -> Self {
self.authentication_timeout = timeout;
self
}
fn lookup(&self, key: &TokenKey, now: Instant) -> CacheLookup {
let mut cache = self.cache.lock();
let Some(entry) = cache.get(key) else {
return CacheLookup::Miss;
};
if entry.expires_at() <= now {
cache.remove(key);
return CacheLookup::Miss;
}
match entry {
CacheEntry::Valid { identity, .. } => CacheLookup::Valid(identity.clone()),
CacheEntry::Rejected { .. } => CacheLookup::Rejected,
}
}
fn insert(&self, key: TokenKey, entry: CacheEntry, now: Instant) {
let mut cache = self.cache.lock();
let count = self.sweep_counter.fetch_add(1, Ordering::Relaxed) + 1;
if count.is_multiple_of(SWEEP_INTERVAL) {
cache.sweep_expired(now);
}
if cache.len() >= MAX_CACHE_ENTRIES && !cache.contains_key(&key) {
cache.sweep_expired(now);
if cache.len() >= MAX_CACHE_ENTRIES {
tracing::warn!(
entries = cache.len(),
max_entries = MAX_CACHE_ENTRIES,
"internal-auth cache is full of unexpired entries; evicting a live entry"
);
cache.evict_soonest();
}
}
cache.insert(key, entry);
}
fn token_lock(&self, key: TokenKey) -> Arc<AsyncMutex<()>> {
let mut inflight = self.inflight.lock();
Arc::clone(
inflight
.entry(key)
.or_insert_with(|| Arc::new(AsyncMutex::new(()))),
)
}
fn release_token_lock(&self, key: &TokenKey, lock: &Arc<AsyncMutex<()>>) {
let mut inflight = self.inflight.lock();
if Arc::strong_count(lock) <= 2 {
inflight.remove(key);
}
}
}
struct ReleaseTokenLockOnDrop<'a, A> {
owner: &'a CachingInternalAuthenticator<A>,
key: TokenKey,
lock: &'a Arc<AsyncMutex<()>>,
}
impl<A> Drop for ReleaseTokenLockOnDrop<'_, A> {
fn drop(&mut self) {
self.owner.release_token_lock(&self.key, self.lock);
}
}
impl<A: InternalAuthenticator> CachingInternalAuthenticator<A> {
async fn authenticate_at(
&self,
token: &str,
now: Instant,
) -> Result<PlatformIdentity, InternalAuthNError> {
let key = TokenKey::of(token);
match self.lookup(&key, now) {
CacheLookup::Valid(identity) => return Ok(identity),
CacheLookup::Rejected => return Err(InternalAuthNError::InvalidToken),
CacheLookup::Miss => {}
}
let lock = self.token_lock(key);
let _release = ReleaseTokenLockOnDrop {
owner: self,
key,
lock: &lock,
};
let _guard = lock.lock().await;
let now = now.max(Instant::now());
match self.lookup(&key, now) {
CacheLookup::Valid(identity) => return Ok(identity),
CacheLookup::Rejected => return Err(InternalAuthNError::InvalidToken),
CacheLookup::Miss => {}
}
let result = self.inner.authenticate(token).await;
let stored_at = Instant::now();
match result {
Ok(identity) => {
let expires_at = clamped_expiry(token, stored_at, self.ttl);
self.insert(
key,
CacheEntry::Valid {
identity: identity.clone(),
expires_at,
},
stored_at,
);
Ok(identity)
}
Err(InternalAuthNError::InvalidToken) => {
self.insert(
key,
CacheEntry::Rejected {
expires_at: stored_at + NEGATIVE_CACHE_TTL,
},
stored_at,
);
Err(InternalAuthNError::InvalidToken)
}
Err(err) => Err(err),
}
}
}
impl<A: InternalAuthenticator> InternalAuthenticator for CachingInternalAuthenticator<A> {
async fn authenticate(&self, token: &str) -> Result<PlatformIdentity, InternalAuthNError> {
match tokio::time::timeout(
self.authentication_timeout,
self.authenticate_at(token, Instant::now()),
)
.await
{
Ok(result) => result,
Err(_elapsed) => Err(InternalAuthNError::Unavailable),
}
}
}
#[cfg(test)]
#[cfg_attr(coverage_nightly, coverage(off))]
mod tests {
use super::*;
use std::sync::atomic::AtomicUsize;
struct CountingAuth {
calls: AtomicUsize,
mode: Mutex<Mode>,
delay: Mutex<Duration>,
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum Mode {
Succeed,
Unavailable,
Invalid,
}
impl CountingAuth {
fn new() -> Self {
Self {
calls: AtomicUsize::new(0),
mode: Mutex::new(Mode::Succeed),
delay: Mutex::new(Duration::ZERO),
}
}
fn calls(&self) -> usize {
self.calls.load(Ordering::SeqCst)
}
fn set_mode(&self, mode: Mode) {
*self.mode.lock() = mode;
}
fn set_delay(&self, delay: Duration) {
*self.delay.lock() = delay;
}
}
impl InternalAuthenticator for CountingAuth {
async fn authenticate(&self, token: &str) -> Result<PlatformIdentity, InternalAuthNError> {
self.calls.fetch_add(1, Ordering::SeqCst);
let delay = *self.delay.lock();
if !delay.is_zero() {
tokio::time::sleep(delay).await;
}
match *self.mode.lock() {
Mode::Succeed => Ok(PlatformIdentity::Shared {
name: token.to_owned(),
}),
Mode::Unavailable => Err(InternalAuthNError::Unavailable),
Mode::Invalid => Err(InternalAuthNError::InvalidToken),
}
}
}
#[tokio::test]
async fn a_very_long_token_is_cached_like_any_other() {
let cached =
CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
let huge = "x".repeat(64 * 1024);
cached.authenticate(&huge).await.unwrap();
cached.authenticate(&huge).await.unwrap();
assert_eq!(cached.cache.lock().len(), 1);
assert_eq!(
cached.inner.calls(),
1,
"the second call must be served from cache"
);
}
#[test]
fn the_cache_key_is_a_fixed_size_digest_of_the_token() {
let short = TokenKey::of("t");
let long = TokenKey::of(&"x".repeat(64 * 1024));
assert_eq!(short.0.len(), long.0.len());
assert_ne!(short, long);
assert_eq!(
short,
TokenKey::of("t"),
"the same token must map to itself"
);
assert!(!format!("{:?}", TokenKey::of("super-secret")).contains("super-secret"));
}
#[test]
fn the_expiry_index_tracks_the_map_through_overwrites_and_removals() {
let mut cache = ExpiringCache::new();
let now = Instant::now();
cache.insert(
TokenKey::of("a"),
valid_entry("a", now + Duration::from_secs(30)),
);
cache.insert(
TokenKey::of("b"),
valid_entry("b", now + Duration::from_secs(10)),
);
cache.insert(
TokenKey::of("a"),
valid_entry("a", now + Duration::from_secs(5)),
);
assert_eq!(cache.len(), 2);
assert_eq!(
cache.by_expiry.len(),
2,
"index must not keep a stale entry"
);
cache.evict_soonest();
assert!(
!cache.contains_key(&TokenKey::of("a")),
"`a` now expires soonest"
);
assert!(cache.contains_key(&TokenKey::of("b")));
assert_eq!(cache.by_expiry.len(), 1);
cache.remove(&TokenKey::of("b"));
assert_eq!(cache.len(), 0);
assert_eq!(cache.by_expiry.len(), 0);
}
#[test]
fn sweeping_removes_exactly_the_expired() {
let mut cache = ExpiringCache::new();
let now = Instant::now();
cache.insert(
TokenKey::of("expired"),
valid_entry("expired", now.checked_sub(Duration::from_secs(1)).unwrap()),
);
cache.insert(
TokenKey::of("live"),
valid_entry("live", now + Duration::from_mins(1)),
);
cache.sweep_expired(now);
assert!(!cache.contains_key(&TokenKey::of("expired")));
assert!(cache.contains_key(&TokenKey::of("live")));
assert_eq!(
cache.by_expiry.len(),
1,
"the index must shrink with the map"
);
}
#[tokio::test(start_paused = true)]
async fn with_authentication_timeout_overrides_the_default() {
let short = Duration::from_millis(50);
let cached = CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1))
.unwrap()
.with_authentication_timeout(short);
cached.inner.set_delay(short * 2);
let err = cached
.authenticate("tok")
.await
.expect_err("a backend slower than the overridden deadline must not resolve");
assert!(matches!(err, InternalAuthNError::Unavailable));
}
#[tokio::test(start_paused = true)]
async fn a_hung_backend_times_out_rather_than_parking_every_caller() {
let cached =
CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
cached.inner.set_delay(DEFAULT_AUTHENTICATION_TIMEOUT * 2);
let err = cached
.authenticate("tok")
.await
.expect_err("a backend that outlasts the deadline must not resolve");
assert!(
matches!(err, InternalAuthNError::Unavailable),
"a timeout is a backend availability problem, not a verdict on the credential"
);
}
#[tokio::test(start_paused = true)]
async fn concurrent_callers_share_one_deadline_rather_than_queueing() {
let cached = Arc::new(
CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap(),
);
cached.inner.set_delay(DEFAULT_AUTHENTICATION_TIMEOUT * 10);
let started = tokio::time::Instant::now();
let mut callers = Vec::new();
for _ in 0..3 {
let cached = Arc::clone(&cached);
callers.push(tokio::spawn(async move {
let outcome = cached.authenticate("same-token").await;
(outcome, started.elapsed())
}));
}
for caller in callers {
let (outcome, elapsed) = caller.await.expect("task must not panic");
assert!(
matches!(outcome, Err(InternalAuthNError::Unavailable)),
"a hung backend must surface as Unavailable, got {outcome:?}"
);
assert!(
elapsed < DEFAULT_AUTHENTICATION_TIMEOUT * 2,
"caller waited {elapsed:?}, i.e. it queued behind another \
caller's deadline instead of holding its own"
);
}
}
#[tokio::test(start_paused = true)]
async fn a_very_long_token_obeys_the_deadline_too() {
let cached =
CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
cached.inner.set_delay(DEFAULT_AUTHENTICATION_TIMEOUT * 2);
let huge = "x".repeat(64 * 1024);
let err = cached
.authenticate(&huge)
.await
.expect_err("a long token must be bounded as well");
assert!(matches!(err, InternalAuthNError::Unavailable));
}
#[tokio::test(start_paused = true)]
async fn a_timed_out_validation_is_not_cached() {
let cached =
CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
cached.inner.set_delay(DEFAULT_AUTHENTICATION_TIMEOUT * 2);
drop(cached.authenticate("tok").await);
cached.inner.set_delay(Duration::ZERO);
let identity = cached
.authenticate("tok")
.await
.expect("a recovered backend must be consulted again");
assert_eq!(identity.peer_name(), "tok");
assert_eq!(
cached.inner.calls(),
2,
"the timed-out attempt must not have been cached"
);
}
#[tokio::test]
async fn second_call_within_ttl_hits_cache() {
let cached =
CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
let a = cached.authenticate("tok").await.unwrap();
let b = cached.authenticate("tok").await.unwrap();
assert_eq!(a, b);
assert_eq!(
cached.inner.calls(),
1,
"second call must be served from cache"
);
assert_eq!(
a.peer_name(),
"tok",
"cached identity must match the token it was issued for"
);
}
#[tokio::test]
async fn distinct_tokens_are_cached_independently() {
let cached =
CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
let a = cached.authenticate("a").await.unwrap();
let b = cached.authenticate("b").await.unwrap();
let a2 = cached.authenticate("a").await.unwrap();
assert_eq!(
cached.inner.calls(),
2,
"each distinct token validated once"
);
assert_eq!(a.peer_name(), "a");
assert_eq!(b.peer_name(), "b");
assert_eq!(
a2.peer_name(),
"a",
"cache must not confuse token identities"
);
}
#[tokio::test]
async fn entry_expires_after_ttl() {
let cached =
CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_millis(20))
.unwrap();
let t0 = Instant::now();
cached.authenticate_at("tok", t0).await.unwrap();
assert_eq!(cached.inner.calls(), 1);
cached
.authenticate_at("tok", t0 + Duration::from_millis(21))
.await
.unwrap();
assert_eq!(
cached.inner.calls(),
2,
"expired entry must be re-validated"
);
}
#[tokio::test]
async fn unavailable_errors_are_not_cached() {
let cached =
CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
cached.inner.set_mode(Mode::Unavailable);
assert!(cached.authenticate("tok").await.is_err());
assert!(cached.authenticate("tok").await.is_err());
assert_eq!(
cached.inner.calls(),
2,
"backend outages must not be cached"
);
cached.inner.set_mode(Mode::Succeed);
cached.authenticate("tok").await.unwrap();
cached.authenticate("tok").await.unwrap();
assert_eq!(
cached.inner.calls(),
3,
"recovery validated once, then cached"
);
}
#[tokio::test]
async fn invalid_token_rejections_are_briefly_negative_cached() {
let cached =
CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
cached.inner.set_mode(Mode::Invalid);
let t0 = Instant::now();
let err = cached.authenticate("bad").await.unwrap_err();
assert!(matches!(err, InternalAuthNError::InvalidToken));
let err = cached.authenticate("bad").await.unwrap_err();
assert!(matches!(err, InternalAuthNError::InvalidToken));
assert_eq!(
cached.inner.calls(),
1,
"a cached rejection must not re-hit the backend"
);
cached.inner.set_mode(Mode::Succeed);
let identity = cached
.authenticate_at("bad", t0 + NEGATIVE_CACHE_TTL + Duration::from_millis(1))
.await
.unwrap();
assert_eq!(
identity.peer_name(),
"bad",
"the token validates once the backend accepts it"
);
assert_eq!(cached.inner.calls(), 2);
}
#[test]
fn new_rejects_zero_and_over_max_ttl() {
assert!(CachingInternalAuthenticator::new(CountingAuth::new(), Duration::ZERO).is_err());
assert!(
CachingInternalAuthenticator::new(
CountingAuth::new(),
MAX_TOKEN_REVIEW_CACHE_TTL + Duration::from_secs(1)
)
.is_err()
);
assert!(
CachingInternalAuthenticator::new(CountingAuth::new(), MAX_TOKEN_REVIEW_CACHE_TTL)
.is_ok()
);
}
fn valid_entry(name: &str, expires_at: Instant) -> CacheEntry {
CacheEntry::Valid {
identity: PlatformIdentity::Shared {
name: name.to_owned(),
},
expires_at,
}
}
#[test]
fn capacity_bound_evicts_soonest_to_expire_when_full() {
let cached =
CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
let now = Instant::now();
for i in 0..MAX_CACHE_ENTRIES {
let name = format!("tok-{i}");
let expires_at = now + Duration::from_mins(1) + Duration::from_micros(i as u64);
cached.insert(TokenKey::of(&name), valid_entry(&name, expires_at), now);
}
assert_eq!(cached.cache.lock().len(), MAX_CACHE_ENTRIES);
cached.insert(
TokenKey::of("overflow"),
valid_entry("overflow", now + Duration::from_mins(2)),
now,
);
let cache = cached.cache.lock();
assert_eq!(
cache.len(),
MAX_CACHE_ENTRIES,
"cache must never grow past MAX_CACHE_ENTRIES"
);
assert!(
cache.contains_key(&TokenKey::of("overflow")),
"the newly inserted token must be present"
);
assert!(
!cache.contains_key(&TokenKey::of("tok-0")),
"the soonest-to-expire entry must be evicted to make room"
);
}
#[test]
fn full_cache_reclaims_expired_before_scanning_for_a_victim() {
let cached =
CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
let now = Instant::now();
for i in 1..MAX_CACHE_ENTRIES {
let name = format!("tok-{i}");
let expires_at = now + Duration::from_mins(5) + Duration::from_micros(i as u64);
cached.insert(TokenKey::of(&name), valid_entry(&name, expires_at), now);
}
cached.insert(
TokenKey::of("expired"),
valid_entry("expired", now.checked_sub(Duration::from_secs(1)).unwrap()),
now,
);
assert_eq!(cached.cache.lock().len(), MAX_CACHE_ENTRIES);
cached.insert(
TokenKey::of("overflow"),
valid_entry("overflow", now + Duration::from_mins(10)),
now,
);
let cache = cached.cache.lock();
assert_eq!(cache.len(), MAX_CACHE_ENTRIES);
assert!(cache.contains_key(&TokenKey::of("overflow")));
assert!(
!cache.contains_key(&TokenKey::of("expired")),
"the expired entry must be reclaimed by the sweep, making room"
);
assert!(
cache.contains_key(&TokenKey::of("tok-1")),
"a still-valid entry must not be evicted while an expired one exists"
);
}
#[tokio::test]
async fn concurrent_misses_for_the_same_token_single_flight() {
let cached = Arc::new(
CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap(),
);
let mut handles = Vec::new();
for _ in 0..8 {
let cached = Arc::clone(&cached);
handles.push(tokio::spawn(
async move { cached.authenticate("burst").await },
));
}
for handle in handles {
handle.await.unwrap().unwrap();
}
assert_eq!(
cached.inner.calls(),
1,
"a burst of concurrent misses for the same token must collapse to one backend call"
);
}
#[tokio::test]
async fn aborting_an_inflight_call_still_releases_the_single_flight_slot() {
let auth = CountingAuth::new();
auth.set_delay(Duration::from_millis(200));
let cached =
Arc::new(CachingInternalAuthenticator::new(auth, Duration::from_mins(1)).unwrap());
let handle = {
let cached = Arc::clone(&cached);
tokio::spawn(async move { cached.authenticate("burst").await })
};
tokio::time::sleep(Duration::from_millis(20)).await;
handle.abort();
let result = handle.await;
assert!(
result.unwrap_err().is_cancelled(),
"the task must actually have been aborted mid-flight for this test to be meaningful"
);
assert!(
cached.inflight.lock().is_empty(),
"aborting an in-flight call must not leak its single-flight slot"
);
let identity = tokio::time::timeout(Duration::from_secs(1), cached.authenticate("burst"))
.await
.expect("post-abort call must not hang")
.unwrap();
assert_eq!(identity.peer_name(), "burst");
}
#[test]
fn clamped_expiry_never_extends_trust_past_an_expired_jwt() {
let now = Instant::now();
let payload = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(br#"{"exp":1}"#);
let expired_token = format!("h.{payload}.s");
assert_eq!(
clamped_expiry(&expired_token, now, Duration::from_mins(1)),
now,
"a JWT already expired per its own claim must not be trusted at all"
);
}
#[test]
fn clamped_expiry_uses_jwt_remaining_life_when_shorter_than_ttl() {
let now = Instant::now();
let exp = (SystemTime::now() + Duration::from_secs(2))
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs();
let payload =
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(format!(r#"{{"exp":{exp}}}"#));
let token = format!("h.{payload}.s");
let expiry = clamped_expiry(&token, now, Duration::from_mins(5));
assert!(
expiry > now,
"a token with life left must produce a non-zero cache lifetime"
);
assert!(
expiry <= now + Duration::from_secs(2),
"a JWT with less remaining life than the configured ttl must clamp to the JWT's expiry"
);
}
#[tokio::test]
async fn second_waiter_reuses_result_populated_while_it_waited() {
let auth = CountingAuth::new();
auth.set_delay(Duration::from_millis(50));
let cached =
Arc::new(CachingInternalAuthenticator::new(auth, Duration::from_mins(1)).unwrap());
let first = {
let cached = Arc::clone(&cached);
tokio::spawn(async move { cached.authenticate("burst").await })
};
tokio::time::sleep(Duration::from_millis(10)).await;
let second = cached.authenticate("burst").await.unwrap();
let first = first.await.unwrap().unwrap();
assert_eq!(first, second);
assert_eq!(
cached.inner.calls(),
1,
"a waiter that blocked on the single-flight lock must reuse the result the winner stored"
);
}
#[tokio::test]
async fn cancelling_a_waiter_still_releases_its_inflight_registration() {
let cached = Arc::new(
CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap(),
);
cached.inner.set_delay(Duration::from_secs(30));
for i in 0..32 {
let token = format!("tok-{i}");
let a = tokio::spawn({
let cached = Arc::clone(&cached);
let token = token.clone();
async move { cached.authenticate(&token).await }
});
tokio::time::sleep(Duration::from_millis(5)).await;
let b = tokio::spawn({
let cached = Arc::clone(&cached);
let token = token.clone();
async move { cached.authenticate(&token).await }
});
tokio::time::sleep(Duration::from_millis(5)).await;
a.abort();
b.abort();
drop(a.await);
drop(b.await);
}
assert_eq!(
cached.inflight.lock().len(),
0,
"every cancelled waiter -- holding the lock or still queued for it -- \
must release its `inflight` registration"
);
}
#[tokio::test]
async fn second_waiter_reuses_rejection_populated_while_it_waited() {
let auth = CountingAuth::new();
auth.set_delay(Duration::from_millis(50));
auth.set_mode(Mode::Invalid);
let cached =
Arc::new(CachingInternalAuthenticator::new(auth, Duration::from_mins(1)).unwrap());
let first = {
let cached = Arc::clone(&cached);
tokio::spawn(async move { cached.authenticate("burst").await })
};
tokio::time::sleep(Duration::from_millis(10)).await;
let second = cached.authenticate("burst").await;
let first = first.await.unwrap();
assert!(matches!(first, Err(InternalAuthNError::InvalidToken)));
assert!(matches!(second, Err(InternalAuthNError::InvalidToken)));
assert_eq!(
cached.inner.calls(),
1,
"a waiter that blocked on the single-flight lock must reuse the cached rejection"
);
}
#[tokio::test]
async fn lock_waiter_revalidates_an_entry_that_expired_during_the_wait() {
let auth = CountingAuth::new();
auth.set_delay(Duration::from_millis(50));
let cached =
Arc::new(CachingInternalAuthenticator::new(auth, Duration::from_mins(1)).unwrap());
let payload = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(br#"{"exp":1}"#);
let token = format!("h.{payload}.s");
let first = {
let cached = Arc::clone(&cached);
let token = token.clone();
tokio::spawn(async move { cached.authenticate(&token).await })
};
tokio::time::sleep(Duration::from_millis(10)).await;
cached.authenticate(&token).await.unwrap();
first.await.unwrap().unwrap();
assert_eq!(
cached.inner.calls(),
2,
"a waiter must re-validate an entry that expired during its wait, not serve it stale"
);
}
fn jwt_with_payload(payload: &[u8]) -> String {
let encoded = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(payload);
format!("h.{encoded}.s")
}
#[test]
fn jwt_exp_claim_extracts_and_ignores_non_jwt() {
let token = jwt_with_payload(br#"{"exp":123}"#);
assert!(matches!(jwt_exp_claim(&token), ExpClaim::Expires(123)));
assert!(matches!(jwt_exp_claim("not-a-jwt"), ExpClaim::NotJwt));
assert!(matches!(
jwt_exp_claim("shared-secret-token"),
ExpClaim::NotJwt
));
}
#[test]
fn jwt_exp_claim_separates_unreadable_from_absent() {
assert!(matches!(
jwt_exp_claim(&jwt_with_payload(br"{}")),
ExpClaim::Unreadable
));
assert!(matches!(
jwt_exp_claim(&jwt_with_payload(br#"{"exp":"soon"}"#)),
ExpClaim::Unreadable
));
assert!(matches!(
jwt_exp_claim(&jwt_with_payload(br#"{"exp":1.5}"#)),
ExpClaim::Unreadable
));
assert!(matches!(
jwt_exp_claim(&jwt_with_payload(br"not json")),
ExpClaim::Unreadable
));
assert!(matches!(
jwt_exp_claim("h.!!!not-base64!!!.s"),
ExpClaim::Unreadable
));
}
#[test]
fn unreadable_exp_is_not_cached_for_the_full_ttl() {
let now = Instant::now();
let ttl = Duration::from_mins(5);
assert_eq!(
clamped_expiry("shared-secret-token", now, ttl),
now + ttl,
"an opaque credential has no expiry to clamp against"
);
let unreadable = jwt_with_payload(br#"{"exp":"soon"}"#);
assert_eq!(
clamped_expiry(&unreadable, now, ttl),
now,
"a declared-but-unreadable expiry must not be cached for the full TTL"
);
}
#[test]
fn absurd_exp_claims_do_not_panic() {
let now = Instant::now();
let ttl = Duration::from_secs(30);
for payload in [
format!(r#"{{"exp":{}}}"#, u64::MAX),
format!(r#"{{"exp":{}}}"#, i64::MAX),
r#"{"exp":0}"#.to_owned(),
] {
let token = jwt_with_payload(payload.as_bytes());
let expiry = clamped_expiry(&token, now, ttl);
assert!(
expiry <= now + ttl,
"the clamp must never extend trust beyond the configured TTL"
);
}
}
}