1use std::collections::{BTreeSet, HashMap};
37use std::sync::Arc;
38use std::sync::atomic::{AtomicU32, Ordering};
39use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
40
41use base64::Engine as _;
42use parking_lot::Mutex;
43use tokio::sync::Mutex as AsyncMutex;
44
45use crate::internal_auth::{InternalAuthNError, InternalAuthenticator, PlatformIdentity};
46
47pub const DEFAULT_TOKEN_REVIEW_CACHE_TTL: Duration = Duration::from_secs(30);
53
54pub const MAX_TOKEN_REVIEW_CACHE_TTL: Duration = Duration::from_mins(5);
58
59const NEGATIVE_CACHE_TTL: Duration = Duration::from_secs(1);
63
64pub const DEFAULT_AUTHENTICATION_TIMEOUT: Duration = Duration::from_secs(10);
82
83const SWEEP_INTERVAL: u32 = 32;
86
87pub const MAX_CACHE_ENTRIES: usize = 10_000;
93
94#[derive(Debug, thiserror::Error)]
97#[error(
98 "internal-auth cache TTL must be > 0 and <= {MAX_TOKEN_REVIEW_CACHE_TTL:?}, got {actual:?}"
99)]
100pub struct InvalidCacheTtl {
101 actual: Duration,
102}
103
104#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
116struct TokenKey([u8; 32]);
117
118impl TokenKey {
119 fn of(token: &str) -> Self {
120 use sha2::{Digest as _, Sha256};
121 Self(Sha256::digest(token.as_bytes()).into())
122 }
123}
124
125impl std::fmt::Debug for TokenKey {
126 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
127 write!(f, "TokenKey({:02x}{:02x}..)", self.0[0], self.0[1])
130 }
131}
132
133enum CacheEntry {
135 Valid {
136 identity: PlatformIdentity,
137 expires_at: Instant,
138 },
139 Rejected {
140 expires_at: Instant,
141 },
142}
143
144impl CacheEntry {
145 fn expires_at(&self) -> Instant {
146 match self {
147 Self::Valid { expires_at, .. } | Self::Rejected { expires_at } => *expires_at,
148 }
149 }
150}
151
152enum CacheLookup {
154 Valid(PlatformIdentity),
155 Rejected,
156 Miss,
157}
158
159struct ExpiringCache {
171 entries: HashMap<TokenKey, CacheEntry>,
172 by_expiry: BTreeSet<(Instant, TokenKey)>,
175}
176
177impl ExpiringCache {
178 fn new() -> Self {
179 Self {
180 entries: HashMap::new(),
181 by_expiry: BTreeSet::new(),
182 }
183 }
184
185 fn len(&self) -> usize {
186 self.entries.len()
187 }
188
189 fn contains_key(&self, key: &TokenKey) -> bool {
190 self.entries.contains_key(key)
191 }
192
193 fn get(&self, key: &TokenKey) -> Option<&CacheEntry> {
194 self.entries.get(key)
195 }
196
197 fn insert(&mut self, key: TokenKey, entry: CacheEntry) {
198 let expires_at = entry.expires_at();
199 if let Some(previous) = self.entries.insert(key, entry) {
200 self.by_expiry.remove(&(previous.expires_at(), key));
201 }
202 self.by_expiry.insert((expires_at, key));
203 }
204
205 fn remove(&mut self, key: &TokenKey) {
206 if let Some(entry) = self.entries.remove(key) {
207 self.by_expiry.remove(&(entry.expires_at(), *key));
208 }
209 }
210
211 fn sweep_expired(&mut self, now: Instant) {
214 let live = self.by_expiry.split_off(&(now, TokenKey([0; 32])));
217 let expired = std::mem::replace(&mut self.by_expiry, live);
218 for (_, key) in expired {
219 self.entries.remove(&key);
220 }
221 }
222
223 fn evict_soonest(&mut self) {
226 if let Some((expires_at, key)) = self.by_expiry.pop_first() {
227 self.entries.remove(&key);
228 debug_assert!(
229 !self.entries.contains_key(&key),
230 "index and map disagreed about {key:?} at {expires_at:?}"
231 );
232 }
233 }
234}
235
236enum ExpClaim {
243 NotJwt,
245 Expires(u64),
247 Unreadable,
250}
251
252fn jwt_exp_claim(token: &str) -> ExpClaim {
260 let mut parts = token.split('.');
263 let (Some(_header), Some(payload_b64), Some(_signature), None) =
264 (parts.next(), parts.next(), parts.next(), parts.next())
265 else {
266 return ExpClaim::NotJwt;
267 };
268
269 let Ok(payload) = base64::engine::general_purpose::URL_SAFE_NO_PAD.decode(payload_b64) else {
270 return ExpClaim::Unreadable;
271 };
272 let Ok(value) = serde_json::from_slice::<serde_json::Value>(&payload) else {
273 return ExpClaim::Unreadable;
274 };
275 match value.get("exp").and_then(serde_json::Value::as_u64) {
276 Some(exp) => ExpClaim::Expires(exp),
277 None => ExpClaim::Unreadable,
278 }
279}
280
281fn clamped_expiry(token: &str, now: Instant, ttl: Duration) -> Instant {
290 let ttl_expiry = now.checked_add(ttl).unwrap_or(now);
291 let exp_secs = match jwt_exp_claim(token) {
292 ExpClaim::NotJwt => return ttl_expiry,
293 ExpClaim::Unreadable => return now,
296 ExpClaim::Expires(exp) => exp,
297 };
298
299 let Some(exp_at) = UNIX_EPOCH.checked_add(Duration::from_secs(exp_secs)) else {
300 return ttl_expiry;
303 };
304 let Ok(remaining) = exp_at.duration_since(SystemTime::now()) else {
305 return now;
308 };
309 let exp_expiry = now.checked_add(remaining).unwrap_or(ttl_expiry);
310 ttl_expiry.min(exp_expiry)
311}
312
313pub struct CachingInternalAuthenticator<A> {
337 inner: A,
338 ttl: Duration,
339 authentication_timeout: Duration,
343 cache: Mutex<ExpiringCache>,
344 inflight: Mutex<HashMap<TokenKey, Arc<AsyncMutex<()>>>>,
347 sweep_counter: AtomicU32,
348}
349
350impl<A> std::fmt::Debug for CachingInternalAuthenticator<A> {
351 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
352 f.debug_struct("CachingInternalAuthenticator")
353 .field("ttl", &self.ttl)
354 .finish_non_exhaustive()
355 }
356}
357
358impl<A> CachingInternalAuthenticator<A> {
359 pub fn new(inner: A, ttl: Duration) -> Result<Self, InvalidCacheTtl> {
367 if ttl == Duration::ZERO || ttl > MAX_TOKEN_REVIEW_CACHE_TTL {
368 return Err(InvalidCacheTtl { actual: ttl });
369 }
370 Ok(Self {
371 inner,
372 ttl,
373 authentication_timeout: DEFAULT_AUTHENTICATION_TIMEOUT,
374 cache: Mutex::new(ExpiringCache::new()),
375 inflight: Mutex::new(HashMap::new()),
376 sweep_counter: AtomicU32::new(0),
377 })
378 }
379
380 #[must_use]
382 pub fn with_default_ttl(inner: A) -> Self {
383 Self {
384 inner,
385 ttl: DEFAULT_TOKEN_REVIEW_CACHE_TTL,
386 authentication_timeout: DEFAULT_AUTHENTICATION_TIMEOUT,
387 cache: Mutex::new(ExpiringCache::new()),
388 inflight: Mutex::new(HashMap::new()),
389 sweep_counter: AtomicU32::new(0),
390 }
391 }
392
393 #[must_use]
397 pub fn with_authentication_timeout(mut self, timeout: Duration) -> Self {
398 self.authentication_timeout = timeout;
399 self
400 }
401
402 fn lookup(&self, key: &TokenKey, now: Instant) -> CacheLookup {
408 let mut cache = self.cache.lock();
409 let Some(entry) = cache.get(key) else {
410 return CacheLookup::Miss;
411 };
412 if entry.expires_at() <= now {
413 cache.remove(key);
414 return CacheLookup::Miss;
415 }
416 match entry {
417 CacheEntry::Valid { identity, .. } => CacheLookup::Valid(identity.clone()),
418 CacheEntry::Rejected { .. } => CacheLookup::Rejected,
419 }
420 }
421
422 fn insert(&self, key: TokenKey, entry: CacheEntry, now: Instant) {
433 let mut cache = self.cache.lock();
434 let count = self.sweep_counter.fetch_add(1, Ordering::Relaxed) + 1;
435 if count.is_multiple_of(SWEEP_INTERVAL) {
436 cache.sweep_expired(now);
437 }
438 if cache.len() >= MAX_CACHE_ENTRIES && !cache.contains_key(&key) {
439 cache.sweep_expired(now);
441 if cache.len() >= MAX_CACHE_ENTRIES {
442 tracing::warn!(
448 entries = cache.len(),
449 max_entries = MAX_CACHE_ENTRIES,
450 "internal-auth cache is full of unexpired entries; evicting a live entry"
451 );
452 cache.evict_soonest();
453 }
454 }
455 cache.insert(key, entry);
456 }
457
458 fn token_lock(&self, key: TokenKey) -> Arc<AsyncMutex<()>> {
460 let mut inflight = self.inflight.lock();
461 Arc::clone(
462 inflight
463 .entry(key)
464 .or_insert_with(|| Arc::new(AsyncMutex::new(()))),
465 )
466 }
467
468 fn release_token_lock(&self, key: &TokenKey, lock: &Arc<AsyncMutex<()>>) {
471 let mut inflight = self.inflight.lock();
472 if Arc::strong_count(lock) <= 2 {
475 inflight.remove(key);
476 }
477 }
478}
479
480struct ReleaseTokenLockOnDrop<'a, A> {
484 owner: &'a CachingInternalAuthenticator<A>,
485 key: TokenKey,
486 lock: &'a Arc<AsyncMutex<()>>,
487}
488
489impl<A> Drop for ReleaseTokenLockOnDrop<'_, A> {
490 fn drop(&mut self) {
491 self.owner.release_token_lock(&self.key, self.lock);
492 }
493}
494
495impl<A: InternalAuthenticator> CachingInternalAuthenticator<A> {
496 async fn authenticate_at(
514 &self,
515 token: &str,
516 now: Instant,
517 ) -> Result<PlatformIdentity, InternalAuthNError> {
518 let key = TokenKey::of(token);
521
522 match self.lookup(&key, now) {
523 CacheLookup::Valid(identity) => return Ok(identity),
524 CacheLookup::Rejected => return Err(InternalAuthNError::InvalidToken),
525 CacheLookup::Miss => {}
526 }
527
528 let lock = self.token_lock(key);
531 let _release = ReleaseTokenLockOnDrop {
539 owner: self,
540 key,
541 lock: &lock,
542 };
543 let _guard = lock.lock().await;
544
545 let now = now.max(Instant::now());
550 match self.lookup(&key, now) {
551 CacheLookup::Valid(identity) => return Ok(identity),
552 CacheLookup::Rejected => return Err(InternalAuthNError::InvalidToken),
553 CacheLookup::Miss => {}
554 }
555
556 let result = self.inner.authenticate(token).await;
566 let stored_at = Instant::now();
569
570 match result {
571 Ok(identity) => {
572 let expires_at = clamped_expiry(token, stored_at, self.ttl);
573 self.insert(
574 key,
575 CacheEntry::Valid {
576 identity: identity.clone(),
577 expires_at,
578 },
579 stored_at,
580 );
581 Ok(identity)
582 }
583 Err(InternalAuthNError::InvalidToken) => {
584 self.insert(
585 key,
586 CacheEntry::Rejected {
587 expires_at: stored_at + NEGATIVE_CACHE_TTL,
588 },
589 stored_at,
590 );
591 Err(InternalAuthNError::InvalidToken)
592 }
593 Err(err) => Err(err),
596 }
597 }
598}
599
600impl<A: InternalAuthenticator> InternalAuthenticator for CachingInternalAuthenticator<A> {
601 async fn authenticate(&self, token: &str) -> Result<PlatformIdentity, InternalAuthNError> {
602 match tokio::time::timeout(
616 self.authentication_timeout,
617 self.authenticate_at(token, Instant::now()),
618 )
619 .await
620 {
621 Ok(result) => result,
622 Err(_elapsed) => Err(InternalAuthNError::Unavailable),
623 }
624 }
625}
626
627#[cfg(test)]
628#[cfg_attr(coverage_nightly, coverage(off))]
629mod tests {
630 use super::*;
631 use std::sync::atomic::AtomicUsize;
632
633 struct CountingAuth {
637 calls: AtomicUsize,
638 mode: Mutex<Mode>,
639 delay: Mutex<Duration>,
644 }
645
646 #[derive(Clone, Copy, PartialEq, Eq)]
647 enum Mode {
648 Succeed,
649 Unavailable,
650 Invalid,
651 }
652
653 impl CountingAuth {
654 fn new() -> Self {
655 Self {
656 calls: AtomicUsize::new(0),
657 mode: Mutex::new(Mode::Succeed),
658 delay: Mutex::new(Duration::ZERO),
659 }
660 }
661 fn calls(&self) -> usize {
662 self.calls.load(Ordering::SeqCst)
663 }
664 fn set_mode(&self, mode: Mode) {
665 *self.mode.lock() = mode;
666 }
667 fn set_delay(&self, delay: Duration) {
668 *self.delay.lock() = delay;
669 }
670 }
671
672 impl InternalAuthenticator for CountingAuth {
673 async fn authenticate(&self, token: &str) -> Result<PlatformIdentity, InternalAuthNError> {
674 self.calls.fetch_add(1, Ordering::SeqCst);
675 let delay = *self.delay.lock();
676 if !delay.is_zero() {
677 tokio::time::sleep(delay).await;
678 }
679 match *self.mode.lock() {
680 Mode::Succeed => Ok(PlatformIdentity::Shared {
681 name: token.to_owned(),
682 }),
683 Mode::Unavailable => Err(InternalAuthNError::Unavailable),
684 Mode::Invalid => Err(InternalAuthNError::InvalidToken),
685 }
686 }
687 }
688
689 #[tokio::test]
690 async fn a_very_long_token_is_cached_like_any_other() {
691 let cached =
695 CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
696 let huge = "x".repeat(64 * 1024);
697
698 cached.authenticate(&huge).await.unwrap();
699 cached.authenticate(&huge).await.unwrap();
700
701 assert_eq!(cached.cache.lock().len(), 1);
702 assert_eq!(
703 cached.inner.calls(),
704 1,
705 "the second call must be served from cache"
706 );
707 }
708
709 #[test]
710 fn the_cache_key_is_a_fixed_size_digest_of_the_token() {
711 let short = TokenKey::of("t");
715 let long = TokenKey::of(&"x".repeat(64 * 1024));
716
717 assert_eq!(short.0.len(), long.0.len());
718 assert_ne!(short, long);
719 assert_eq!(
720 short,
721 TokenKey::of("t"),
722 "the same token must map to itself"
723 );
724
725 assert!(!format!("{:?}", TokenKey::of("super-secret")).contains("super-secret"));
727 }
728
729 #[test]
730 fn the_expiry_index_tracks_the_map_through_overwrites_and_removals() {
731 let mut cache = ExpiringCache::new();
734 let now = Instant::now();
735
736 cache.insert(
737 TokenKey::of("a"),
738 valid_entry("a", now + Duration::from_secs(30)),
739 );
740 cache.insert(
741 TokenKey::of("b"),
742 valid_entry("b", now + Duration::from_secs(10)),
743 );
744 cache.insert(
747 TokenKey::of("a"),
748 valid_entry("a", now + Duration::from_secs(5)),
749 );
750 assert_eq!(cache.len(), 2);
751 assert_eq!(
752 cache.by_expiry.len(),
753 2,
754 "index must not keep a stale entry"
755 );
756
757 cache.evict_soonest();
758 assert!(
759 !cache.contains_key(&TokenKey::of("a")),
760 "`a` now expires soonest"
761 );
762 assert!(cache.contains_key(&TokenKey::of("b")));
763 assert_eq!(cache.by_expiry.len(), 1);
764
765 cache.remove(&TokenKey::of("b"));
766 assert_eq!(cache.len(), 0);
767 assert_eq!(cache.by_expiry.len(), 0);
768 }
769
770 #[test]
771 fn sweeping_removes_exactly_the_expired() {
772 let mut cache = ExpiringCache::new();
773 let now = Instant::now();
774
775 cache.insert(
776 TokenKey::of("expired"),
777 valid_entry("expired", now.checked_sub(Duration::from_secs(1)).unwrap()),
778 );
779 cache.insert(
780 TokenKey::of("live"),
781 valid_entry("live", now + Duration::from_mins(1)),
782 );
783
784 cache.sweep_expired(now);
785
786 assert!(!cache.contains_key(&TokenKey::of("expired")));
787 assert!(cache.contains_key(&TokenKey::of("live")));
788 assert_eq!(
789 cache.by_expiry.len(),
790 1,
791 "the index must shrink with the map"
792 );
793 }
794
795 #[tokio::test(start_paused = true)]
796 async fn with_authentication_timeout_overrides_the_default() {
797 let short = Duration::from_millis(50);
800 let cached = CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1))
801 .unwrap()
802 .with_authentication_timeout(short);
803 cached.inner.set_delay(short * 2);
804
805 let err = cached
806 .authenticate("tok")
807 .await
808 .expect_err("a backend slower than the overridden deadline must not resolve");
809 assert!(matches!(err, InternalAuthNError::Unavailable));
810 }
811
812 #[tokio::test(start_paused = true)]
813 async fn a_hung_backend_times_out_rather_than_parking_every_caller() {
814 let cached =
815 CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
816 cached.inner.set_delay(DEFAULT_AUTHENTICATION_TIMEOUT * 2);
818
819 let err = cached
820 .authenticate("tok")
821 .await
822 .expect_err("a backend that outlasts the deadline must not resolve");
823 assert!(
824 matches!(err, InternalAuthNError::Unavailable),
825 "a timeout is a backend availability problem, not a verdict on the credential"
826 );
827 }
828
829 #[tokio::test(start_paused = true)]
830 async fn concurrent_callers_share_one_deadline_rather_than_queueing() {
831 let cached = Arc::new(
837 CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap(),
838 );
839 cached.inner.set_delay(DEFAULT_AUTHENTICATION_TIMEOUT * 10);
840
841 let started = tokio::time::Instant::now();
847 let mut callers = Vec::new();
848 for _ in 0..3 {
849 let cached = Arc::clone(&cached);
850 callers.push(tokio::spawn(async move {
851 let outcome = cached.authenticate("same-token").await;
852 (outcome, started.elapsed())
853 }));
854 }
855
856 for caller in callers {
857 let (outcome, elapsed) = caller.await.expect("task must not panic");
858 assert!(
859 matches!(outcome, Err(InternalAuthNError::Unavailable)),
860 "a hung backend must surface as Unavailable, got {outcome:?}"
861 );
862 assert!(
863 elapsed < DEFAULT_AUTHENTICATION_TIMEOUT * 2,
864 "caller waited {elapsed:?}, i.e. it queued behind another \
865 caller's deadline instead of holding its own"
866 );
867 }
868 }
869
870 #[tokio::test(start_paused = true)]
871 async fn a_very_long_token_obeys_the_deadline_too() {
872 let cached =
875 CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
876 cached.inner.set_delay(DEFAULT_AUTHENTICATION_TIMEOUT * 2);
877
878 let huge = "x".repeat(64 * 1024);
879 let err = cached
880 .authenticate(&huge)
881 .await
882 .expect_err("a long token must be bounded as well");
883 assert!(matches!(err, InternalAuthNError::Unavailable));
884 }
885
886 #[tokio::test(start_paused = true)]
887 async fn a_timed_out_validation_is_not_cached() {
888 let cached =
889 CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
890 cached.inner.set_delay(DEFAULT_AUTHENTICATION_TIMEOUT * 2);
891 drop(cached.authenticate("tok").await);
892
893 cached.inner.set_delay(Duration::ZERO);
896 let identity = cached
897 .authenticate("tok")
898 .await
899 .expect("a recovered backend must be consulted again");
900 assert_eq!(identity.peer_name(), "tok");
901 assert_eq!(
902 cached.inner.calls(),
903 2,
904 "the timed-out attempt must not have been cached"
905 );
906 }
907
908 #[tokio::test]
909 async fn second_call_within_ttl_hits_cache() {
910 let cached =
911 CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
912
913 let a = cached.authenticate("tok").await.unwrap();
914 let b = cached.authenticate("tok").await.unwrap();
915 assert_eq!(a, b);
916 assert_eq!(
917 cached.inner.calls(),
918 1,
919 "second call must be served from cache"
920 );
921 assert_eq!(
922 a.peer_name(),
923 "tok",
924 "cached identity must match the token it was issued for"
925 );
926 }
927
928 #[tokio::test]
929 async fn distinct_tokens_are_cached_independently() {
930 let cached =
931 CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
932
933 let a = cached.authenticate("a").await.unwrap();
934 let b = cached.authenticate("b").await.unwrap();
935 let a2 = cached.authenticate("a").await.unwrap();
936 assert_eq!(
937 cached.inner.calls(),
938 2,
939 "each distinct token validated once"
940 );
941 assert_eq!(a.peer_name(), "a");
942 assert_eq!(b.peer_name(), "b");
943 assert_eq!(
944 a2.peer_name(),
945 "a",
946 "cache must not confuse token identities"
947 );
948 }
949
950 #[tokio::test]
951 async fn entry_expires_after_ttl() {
952 let cached =
957 CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_millis(20))
958 .unwrap();
959 let t0 = Instant::now();
960
961 cached.authenticate_at("tok", t0).await.unwrap();
962 assert_eq!(cached.inner.calls(), 1);
963
964 cached
965 .authenticate_at("tok", t0 + Duration::from_millis(21))
966 .await
967 .unwrap();
968 assert_eq!(
969 cached.inner.calls(),
970 2,
971 "expired entry must be re-validated"
972 );
973 }
974
975 #[tokio::test]
976 async fn unavailable_errors_are_not_cached() {
977 let cached =
978 CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
979 cached.inner.set_mode(Mode::Unavailable);
980
981 assert!(cached.authenticate("tok").await.is_err());
982 assert!(cached.authenticate("tok").await.is_err());
983 assert_eq!(
984 cached.inner.calls(),
985 2,
986 "backend outages must not be cached"
987 );
988
989 cached.inner.set_mode(Mode::Succeed);
991 cached.authenticate("tok").await.unwrap();
992 cached.authenticate("tok").await.unwrap();
993 assert_eq!(
994 cached.inner.calls(),
995 3,
996 "recovery validated once, then cached"
997 );
998 }
999
1000 #[tokio::test]
1001 async fn invalid_token_rejections_are_briefly_negative_cached() {
1002 let cached =
1003 CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
1004 cached.inner.set_mode(Mode::Invalid);
1005 let t0 = Instant::now();
1006
1007 let err = cached.authenticate("bad").await.unwrap_err();
1008 assert!(matches!(err, InternalAuthNError::InvalidToken));
1009 let err = cached.authenticate("bad").await.unwrap_err();
1012 assert!(matches!(err, InternalAuthNError::InvalidToken));
1013 assert_eq!(
1014 cached.inner.calls(),
1015 1,
1016 "a cached rejection must not re-hit the backend"
1017 );
1018
1019 cached.inner.set_mode(Mode::Succeed);
1024 let identity = cached
1025 .authenticate_at("bad", t0 + NEGATIVE_CACHE_TTL + Duration::from_millis(1))
1026 .await
1027 .unwrap();
1028 assert_eq!(
1029 identity.peer_name(),
1030 "bad",
1031 "the token validates once the backend accepts it"
1032 );
1033 assert_eq!(cached.inner.calls(), 2);
1034 }
1035
1036 #[test]
1037 fn new_rejects_zero_and_over_max_ttl() {
1038 assert!(CachingInternalAuthenticator::new(CountingAuth::new(), Duration::ZERO).is_err());
1039 assert!(
1040 CachingInternalAuthenticator::new(
1041 CountingAuth::new(),
1042 MAX_TOKEN_REVIEW_CACHE_TTL + Duration::from_secs(1)
1043 )
1044 .is_err()
1045 );
1046 assert!(
1047 CachingInternalAuthenticator::new(CountingAuth::new(), MAX_TOKEN_REVIEW_CACHE_TTL)
1048 .is_ok()
1049 );
1050 }
1051
1052 fn valid_entry(name: &str, expires_at: Instant) -> CacheEntry {
1053 CacheEntry::Valid {
1054 identity: PlatformIdentity::Shared {
1055 name: name.to_owned(),
1056 },
1057 expires_at,
1058 }
1059 }
1060
1061 #[test]
1062 fn capacity_bound_evicts_soonest_to_expire_when_full() {
1063 let cached =
1064 CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
1065 let now = Instant::now();
1066
1067 for i in 0..MAX_CACHE_ENTRIES {
1068 let name = format!("tok-{i}");
1069 let expires_at = now + Duration::from_mins(1) + Duration::from_micros(i as u64);
1071 cached.insert(TokenKey::of(&name), valid_entry(&name, expires_at), now);
1072 }
1073 assert_eq!(cached.cache.lock().len(), MAX_CACHE_ENTRIES);
1074
1075 cached.insert(
1076 TokenKey::of("overflow"),
1077 valid_entry("overflow", now + Duration::from_mins(2)),
1078 now,
1079 );
1080
1081 let cache = cached.cache.lock();
1082 assert_eq!(
1083 cache.len(),
1084 MAX_CACHE_ENTRIES,
1085 "cache must never grow past MAX_CACHE_ENTRIES"
1086 );
1087 assert!(
1088 cache.contains_key(&TokenKey::of("overflow")),
1089 "the newly inserted token must be present"
1090 );
1091 assert!(
1092 !cache.contains_key(&TokenKey::of("tok-0")),
1093 "the soonest-to-expire entry must be evicted to make room"
1094 );
1095 }
1096
1097 #[test]
1098 fn full_cache_reclaims_expired_before_scanning_for_a_victim() {
1099 let cached =
1100 CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap();
1101 let now = Instant::now();
1102
1103 for i in 1..MAX_CACHE_ENTRIES {
1107 let name = format!("tok-{i}");
1108 let expires_at = now + Duration::from_mins(5) + Duration::from_micros(i as u64);
1109 cached.insert(TokenKey::of(&name), valid_entry(&name, expires_at), now);
1110 }
1111 cached.insert(
1114 TokenKey::of("expired"),
1115 valid_entry("expired", now.checked_sub(Duration::from_secs(1)).unwrap()),
1116 now,
1117 );
1118 assert_eq!(cached.cache.lock().len(), MAX_CACHE_ENTRIES);
1119
1120 cached.insert(
1121 TokenKey::of("overflow"),
1122 valid_entry("overflow", now + Duration::from_mins(10)),
1123 now,
1124 );
1125
1126 let cache = cached.cache.lock();
1127 assert_eq!(cache.len(), MAX_CACHE_ENTRIES);
1128 assert!(cache.contains_key(&TokenKey::of("overflow")));
1129 assert!(
1130 !cache.contains_key(&TokenKey::of("expired")),
1131 "the expired entry must be reclaimed by the sweep, making room"
1132 );
1133 assert!(
1134 cache.contains_key(&TokenKey::of("tok-1")),
1135 "a still-valid entry must not be evicted while an expired one exists"
1136 );
1137 }
1138
1139 #[tokio::test]
1140 async fn concurrent_misses_for_the_same_token_single_flight() {
1141 let cached = Arc::new(
1142 CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap(),
1143 );
1144
1145 let mut handles = Vec::new();
1146 for _ in 0..8 {
1147 let cached = Arc::clone(&cached);
1148 handles.push(tokio::spawn(
1149 async move { cached.authenticate("burst").await },
1150 ));
1151 }
1152 for handle in handles {
1153 handle.await.unwrap().unwrap();
1154 }
1155
1156 assert_eq!(
1157 cached.inner.calls(),
1158 1,
1159 "a burst of concurrent misses for the same token must collapse to one backend call"
1160 );
1161 }
1162
1163 #[tokio::test]
1164 async fn aborting_an_inflight_call_still_releases_the_single_flight_slot() {
1165 let auth = CountingAuth::new();
1166 auth.set_delay(Duration::from_millis(200));
1167 let cached =
1168 Arc::new(CachingInternalAuthenticator::new(auth, Duration::from_mins(1)).unwrap());
1169
1170 let handle = {
1171 let cached = Arc::clone(&cached);
1172 tokio::spawn(async move { cached.authenticate("burst").await })
1173 };
1174 tokio::time::sleep(Duration::from_millis(20)).await;
1177 handle.abort();
1178 let result = handle.await;
1179 assert!(
1180 result.unwrap_err().is_cancelled(),
1181 "the task must actually have been aborted mid-flight for this test to be meaningful"
1182 );
1183
1184 assert!(
1185 cached.inflight.lock().is_empty(),
1186 "aborting an in-flight call must not leak its single-flight slot"
1187 );
1188
1189 let identity = tokio::time::timeout(Duration::from_secs(1), cached.authenticate("burst"))
1192 .await
1193 .expect("post-abort call must not hang")
1194 .unwrap();
1195 assert_eq!(identity.peer_name(), "burst");
1196 }
1197
1198 #[test]
1199 fn clamped_expiry_never_extends_trust_past_an_expired_jwt() {
1200 let now = Instant::now();
1201 let payload = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(br#"{"exp":1}"#);
1202 let expired_token = format!("h.{payload}.s");
1203 assert_eq!(
1204 clamped_expiry(&expired_token, now, Duration::from_mins(1)),
1205 now,
1206 "a JWT already expired per its own claim must not be trusted at all"
1207 );
1208 }
1209
1210 #[test]
1211 fn clamped_expiry_uses_jwt_remaining_life_when_shorter_than_ttl() {
1212 let now = Instant::now();
1213 let exp = (SystemTime::now() + Duration::from_secs(2))
1214 .duration_since(UNIX_EPOCH)
1215 .unwrap()
1216 .as_secs();
1217 let payload =
1218 base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(format!(r#"{{"exp":{exp}}}"#));
1219 let token = format!("h.{payload}.s");
1220 let expiry = clamped_expiry(&token, now, Duration::from_mins(5));
1221
1222 assert!(
1230 expiry > now,
1231 "a token with life left must produce a non-zero cache lifetime"
1232 );
1233 assert!(
1234 expiry <= now + Duration::from_secs(2),
1235 "a JWT with less remaining life than the configured ttl must clamp to the JWT's expiry"
1236 );
1237 }
1238
1239 #[tokio::test]
1240 async fn second_waiter_reuses_result_populated_while_it_waited() {
1241 let auth = CountingAuth::new();
1242 auth.set_delay(Duration::from_millis(50));
1243 let cached =
1244 Arc::new(CachingInternalAuthenticator::new(auth, Duration::from_mins(1)).unwrap());
1245
1246 let first = {
1247 let cached = Arc::clone(&cached);
1248 tokio::spawn(async move { cached.authenticate("burst").await })
1249 };
1250 tokio::time::sleep(Duration::from_millis(10)).await;
1254 let second = cached.authenticate("burst").await.unwrap();
1255 let first = first.await.unwrap().unwrap();
1256
1257 assert_eq!(first, second);
1258 assert_eq!(
1259 cached.inner.calls(),
1260 1,
1261 "a waiter that blocked on the single-flight lock must reuse the result the winner stored"
1262 );
1263 }
1264
1265 #[tokio::test]
1266 async fn cancelling_a_waiter_still_releases_its_inflight_registration() {
1267 let cached = Arc::new(
1275 CachingInternalAuthenticator::new(CountingAuth::new(), Duration::from_mins(1)).unwrap(),
1276 );
1277 cached.inner.set_delay(Duration::from_secs(30));
1278
1279 for i in 0..32 {
1280 let token = format!("tok-{i}");
1281
1282 let a = tokio::spawn({
1283 let cached = Arc::clone(&cached);
1284 let token = token.clone();
1285 async move { cached.authenticate(&token).await }
1286 });
1287 tokio::time::sleep(Duration::from_millis(5)).await;
1289
1290 let b = tokio::spawn({
1291 let cached = Arc::clone(&cached);
1292 let token = token.clone();
1293 async move { cached.authenticate(&token).await }
1294 });
1295 tokio::time::sleep(Duration::from_millis(5)).await;
1298
1299 a.abort();
1305 b.abort();
1306 drop(a.await);
1307 drop(b.await);
1308 }
1309
1310 assert_eq!(
1311 cached.inflight.lock().len(),
1312 0,
1313 "every cancelled waiter -- holding the lock or still queued for it -- \
1314 must release its `inflight` registration"
1315 );
1316 }
1317
1318 #[tokio::test]
1319 async fn second_waiter_reuses_rejection_populated_while_it_waited() {
1320 let auth = CountingAuth::new();
1321 auth.set_delay(Duration::from_millis(50));
1322 auth.set_mode(Mode::Invalid);
1323 let cached =
1324 Arc::new(CachingInternalAuthenticator::new(auth, Duration::from_mins(1)).unwrap());
1325
1326 let first = {
1327 let cached = Arc::clone(&cached);
1328 tokio::spawn(async move { cached.authenticate("burst").await })
1329 };
1330 tokio::time::sleep(Duration::from_millis(10)).await;
1331 let second = cached.authenticate("burst").await;
1332 let first = first.await.unwrap();
1333
1334 assert!(matches!(first, Err(InternalAuthNError::InvalidToken)));
1335 assert!(matches!(second, Err(InternalAuthNError::InvalidToken)));
1336 assert_eq!(
1337 cached.inner.calls(),
1338 1,
1339 "a waiter that blocked on the single-flight lock must reuse the cached rejection"
1340 );
1341 }
1342
1343 #[tokio::test]
1344 async fn lock_waiter_revalidates_an_entry_that_expired_during_the_wait() {
1345 let auth = CountingAuth::new();
1346 auth.set_delay(Duration::from_millis(50));
1349 let cached =
1350 Arc::new(CachingInternalAuthenticator::new(auth, Duration::from_mins(1)).unwrap());
1351
1352 let payload = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(br#"{"exp":1}"#);
1358 let token = format!("h.{payload}.s");
1359
1360 let first = {
1361 let cached = Arc::clone(&cached);
1362 let token = token.clone();
1363 tokio::spawn(async move { cached.authenticate(&token).await })
1364 };
1365 tokio::time::sleep(Duration::from_millis(10)).await;
1368 cached.authenticate(&token).await.unwrap();
1369 first.await.unwrap().unwrap();
1370
1371 assert_eq!(
1372 cached.inner.calls(),
1373 2,
1374 "a waiter must re-validate an entry that expired during its wait, not serve it stale"
1375 );
1376 }
1377
1378 fn jwt_with_payload(payload: &[u8]) -> String {
1380 let encoded = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(payload);
1381 format!("h.{encoded}.s")
1382 }
1383
1384 #[test]
1385 fn jwt_exp_claim_extracts_and_ignores_non_jwt() {
1386 let token = jwt_with_payload(br#"{"exp":123}"#);
1387 assert!(matches!(jwt_exp_claim(&token), ExpClaim::Expires(123)));
1388
1389 assert!(matches!(jwt_exp_claim("not-a-jwt"), ExpClaim::NotJwt));
1390 assert!(matches!(
1391 jwt_exp_claim("shared-secret-token"),
1392 ExpClaim::NotJwt
1393 ));
1394 }
1395
1396 #[test]
1397 fn jwt_exp_claim_separates_unreadable_from_absent() {
1398 assert!(matches!(
1402 jwt_exp_claim(&jwt_with_payload(br"{}")),
1403 ExpClaim::Unreadable
1404 ));
1405 assert!(matches!(
1406 jwt_exp_claim(&jwt_with_payload(br#"{"exp":"soon"}"#)),
1407 ExpClaim::Unreadable
1408 ));
1409 assert!(matches!(
1410 jwt_exp_claim(&jwt_with_payload(br#"{"exp":1.5}"#)),
1411 ExpClaim::Unreadable
1412 ));
1413 assert!(matches!(
1414 jwt_exp_claim(&jwt_with_payload(br"not json")),
1415 ExpClaim::Unreadable
1416 ));
1417 assert!(matches!(
1418 jwt_exp_claim("h.!!!not-base64!!!.s"),
1419 ExpClaim::Unreadable
1420 ));
1421 }
1422
1423 #[test]
1424 fn unreadable_exp_is_not_cached_for_the_full_ttl() {
1425 let now = Instant::now();
1426 let ttl = Duration::from_mins(5);
1427
1428 assert_eq!(
1430 clamped_expiry("shared-secret-token", now, ttl),
1431 now + ttl,
1432 "an opaque credential has no expiry to clamp against"
1433 );
1434
1435 let unreadable = jwt_with_payload(br#"{"exp":"soon"}"#);
1438 assert_eq!(
1439 clamped_expiry(&unreadable, now, ttl),
1440 now,
1441 "a declared-but-unreadable expiry must not be cached for the full TTL"
1442 );
1443 }
1444
1445 #[test]
1446 fn absurd_exp_claims_do_not_panic() {
1447 let now = Instant::now();
1448 let ttl = Duration::from_secs(30);
1449
1450 for payload in [
1453 format!(r#"{{"exp":{}}}"#, u64::MAX),
1454 format!(r#"{{"exp":{}}}"#, i64::MAX),
1455 r#"{"exp":0}"#.to_owned(),
1456 ] {
1457 let token = jwt_with_payload(payload.as_bytes());
1458 let expiry = clamped_expiry(&token, now, ttl);
1459 assert!(
1460 expiry <= now + ttl,
1461 "the clamp must never extend trust beyond the configured TTL"
1462 );
1463 }
1464 }
1465}