Skip to main content

nio_client/
session.rs

1//! Token -> principal resolution with an L1 cache tier — the port of the
2//! normative resolver in nio's `check_client/src/session.rs` (issue #243/#245).
3//!
4//! An opaque session token is hashed in-process (`sha256`, hex) — the raw
5//! token never leaves the process — and resolved to
6//! `{principal, tenant_id, expires_at}` over `am.SessionService` on
7//! nio-client. The cache is LRU-bounded with a positive TTL carrying
8//! *downward-only* jitter (so the TTL is a hard staleness/revocation cap),
9//! negative tombstones for unknown tokens, single-flight coalescing of
10//! concurrent misses, refresh-ahead for hot entries, and an opt-in
11//! stale-if-error window.
12
13use crate::pb::session_service_client::SessionServiceClient;
14use crate::pb::{resolve_response, ResolveRequest};
15use crate::UserId;
16use chrono::{DateTime, Utc};
17use futures::future::{BoxFuture, FutureExt, Shared};
18use sha2::{Digest, Sha256};
19use std::collections::{BTreeMap, HashMap};
20use std::future::Future;
21use std::sync::atomic::{AtomicU64, Ordering};
22use std::sync::{Arc, Mutex};
23use std::time::{Duration, Instant};
24use tonic::transport::Channel;
25
26/// Bounds a single fill's fetch (Go: resolveTimeout). The fill runs on a
27/// detached task so one caller cancelling does not poison coalesced waiters;
28/// an elapsed timeout classifies as a transport error (stale-if-error
29/// eligible).
30const RESOLVE_TIMEOUT: Duration = Duration::from_secs(5);
31
32/// `hex(sha256(raw_token))` — the 64-char lowercase cache/wire key. The raw
33/// token is never stored or transmitted; only this hash is.
34pub fn token_hash(raw_token: &str) -> String {
35    let digest = Sha256::digest(raw_token.as_bytes());
36    hex::encode(digest)
37}
38
39/// Newtype wrapper around [`token_hash`] for call sites that prefer a type.
40pub struct TokenHash(pub String);
41
42impl TokenHash {
43    pub fn from_raw(raw_token: &str) -> Self {
44        TokenHash(token_hash(raw_token))
45    }
46    pub fn as_str(&self) -> &str {
47        &self.0
48    }
49}
50
51/// A resolved session: the principal, its tenant, and the wall-clock
52/// instant the session stops being valid.
53#[derive(Clone, Debug)]
54pub struct ResolvedSession {
55    pub principal: UserId,
56    pub tenant_id: String,
57    pub expires_at: DateTime<Utc>,
58}
59
60/// A resolution failure. `not_found` is *not* an error — it is `Ok(None)`.
61/// Only genuine faults (transport, backend) are errors; `Transport` marks the
62/// class eligible for stale-if-error fallback.
63#[derive(Clone, Debug, thiserror::Error)]
64pub enum ResolveError {
65    #[error("session resolve transport error: {0}")]
66    Transport(String),
67    #[error("session resolve backend error: {0}")]
68    Backend(String),
69}
70
71impl ResolveError {
72    fn is_transport(&self) -> bool {
73        matches!(self, ResolveError::Transport(_))
74    }
75}
76
77pub type ResolveFuture<'a> = BoxFuture<'a, Result<Option<ResolvedSession>, ResolveError>>;
78
79/// The backend fill for a cache miss. Implementations do the actual point
80/// read; `Ok(None)` means the token is unknown.
81pub trait SessionFetcher: Send + Sync + 'static {
82    fn fetch<'a>(&'a self, token_hash: &'a str) -> ResolveFuture<'a>;
83}
84
85/// What callers use: an already-cache-tiered resolve. Object-safe so state
86/// structs can hold `Arc<dyn SessionResolver>` without a type parameter.
87pub trait SessionResolver: Send + Sync + 'static {
88    fn resolve<'a>(&'a self, token_hash: &'a str) -> ResolveFuture<'a>;
89    /// Drop any cached entry for this hash (sign-out / revoke on the local node).
90    fn evict(&self, token_hash: &str);
91}
92
93/// Session-resolution cache tunables (issue #243). The library does not read
94/// environment variables; the process supplies the config.
95#[derive(Clone, Debug)]
96pub struct ResolverConfig {
97    /// L1 LRU capacity.
98    pub capacity: usize,
99    /// Positive entry TTL (hard staleness/revocation cap).
100    pub l1_ttl: Duration,
101    /// Negative tombstone TTL for unknown tokens.
102    pub neg_ttl: Duration,
103    /// Serve stale on transport error for this window; zero = off.
104    pub stale_if_error: Duration,
105}
106
107impl Default for ResolverConfig {
108    /// The #243 defaults: capacity 10000, L1 TTL 30s, neg TTL 2s,
109    /// stale-if-error off.
110    fn default() -> Self {
111        ResolverConfig {
112            capacity: 10_000,
113            l1_ttl: Duration::from_secs(30),
114            neg_ttl: Duration::from_secs(2),
115            stale_if_error: Duration::ZERO,
116        }
117    }
118}
119
120/// One cache slot. `outcome == None` is a negative tombstone.
121#[derive(Clone)]
122struct Entry {
123    outcome: Option<ResolvedSession>,
124    fetched_at: Instant,
125    fresh_until: Instant,
126    stale_until: Instant,
127    effective_ttl: Duration,
128}
129
130/// Bounded LRU keyed by token hash. A monotonically increasing recency
131/// generation orders entries in a `BTreeMap`, so eviction is `O(log n)` and
132/// `get` promotes to most-recently-used.
133struct Lru {
134    map: HashMap<String, (Entry, u64)>,
135    recency: BTreeMap<u64, String>,
136    next_gen: u64,
137    capacity: usize,
138}
139
140impl Lru {
141    fn new(capacity: usize) -> Self {
142        Lru {
143            map: HashMap::new(),
144            recency: BTreeMap::new(),
145            next_gen: 0,
146            capacity,
147        }
148    }
149
150    fn get(&mut self, key: &str) -> Option<Entry> {
151        if self.capacity == 0 {
152            return None;
153        }
154        let (entry, old_gen) = self.map.get(key)?;
155        let entry = entry.clone();
156        let old_gen = *old_gen;
157        self.recency.remove(&old_gen);
158        let gen = self.next_gen;
159        self.next_gen += 1;
160        self.recency.insert(gen, key.to_string());
161        if let Some(e) = self.map.get_mut(key) {
162            e.1 = gen;
163        }
164        Some(entry)
165    }
166
167    fn peek(&self, key: &str) -> Option<Entry> {
168        self.map.get(key).map(|(e, _)| e.clone())
169    }
170
171    fn put(&mut self, key: String, entry: Entry) {
172        if self.capacity == 0 {
173            return;
174        }
175        if let Some((_, old_gen)) = self.map.get(&key) {
176            let old_gen = *old_gen;
177            self.recency.remove(&old_gen);
178        }
179        let gen = self.next_gen;
180        self.next_gen += 1;
181        self.recency.insert(gen, key.clone());
182        self.map.insert(key, (entry, gen));
183        while self.map.len() > self.capacity {
184            let Some((&lru_gen, lru_key)) = self.recency.iter().next() else {
185                break;
186            };
187            let lru_key = lru_key.clone();
188            self.recency.remove(&lru_gen);
189            self.map.remove(&lru_key);
190        }
191    }
192
193    fn remove(&mut self, key: &str) {
194        if let Some((_, gen)) = self.map.remove(key) {
195            self.recency.remove(&gen);
196        }
197    }
198}
199
200type FillResult = Result<Option<ResolvedSession>, Arc<ResolveError>>;
201type SharedFill = Shared<BoxFuture<'static, FillResult>>;
202
203/// Single-flight de-duplication of identical concurrent fills (an SPA firing
204/// 20 parallel XHRs on a cold token triggers one fill, 19 waiters share it).
205#[derive(Clone, Default)]
206struct SingleFlight {
207    inflight: Arc<Mutex<HashMap<String, (u64, SharedFill)>>>,
208    next_id: Arc<AtomicU64>,
209}
210
211impl SingleFlight {
212    async fn run<F, Fut>(&self, key: &str, make: F) -> Result<Option<ResolvedSession>, ResolveError>
213    where
214        F: FnOnce() -> Fut,
215        Fut: Future<Output = Result<Option<ResolvedSession>, ResolveError>> + Send + 'static,
216    {
217        let shared = {
218            let mut map = self.inflight.lock().expect("singleflight mutex poisoned");
219            if let Some((_, existing)) = map.get(key) {
220                existing.clone()
221            } else {
222                let id = self.next_id.fetch_add(1, Ordering::Relaxed);
223                let inflight = Arc::clone(&self.inflight);
224                let owned_key = key.to_string();
225                let inner = make().map(|r| r.map_err(Arc::new));
226                // The map entry is removed inside the shared future so cleanup
227                // runs no matter which task drives the fill to completion — a
228                // leader dropped mid-await (client disconnect) must not leak
229                // the entry and freeze this hash on a memoized result.
230                let fut: SharedFill = async move {
231                    let result = inner.await;
232                    let mut map = inflight.lock().expect("singleflight mutex poisoned");
233                    if map
234                        .get(&owned_key)
235                        .map(|(gid, _)| *gid == id)
236                        .unwrap_or(false)
237                    {
238                        map.remove(&owned_key);
239                    }
240                    result
241                }
242                .boxed()
243                .shared();
244                map.insert(key.to_string(), (id, fut.clone()));
245                // Detach the fill (Go: detached context): it completes and
246                // populates the cache even if every waiter is cancelled.
247                tokio::spawn(fut.clone());
248                fut
249            }
250        };
251
252        shared.await.map_err(|e| (*e).clone())
253    }
254}
255
256struct ResolverInner {
257    fetcher: Arc<dyn SessionFetcher>,
258    cache: Mutex<Lru>,
259    flight: SingleFlight,
260    cfg: ResolverConfig,
261}
262
263impl ResolverInner {
264    fn effective_ttl(&self) -> Duration {
265        // Downward-only jitter U(0.8, 1.0): l1_ttl is a hard cap.
266        let jitter = 0.8 + 0.2 * rand::random::<f64>();
267        Duration::from_secs_f64(self.cfg.l1_ttl.as_secs_f64() * jitter)
268    }
269
270    async fn resolve(
271        self: Arc<Self>,
272        hash: String,
273    ) -> Result<Option<ResolvedSession>, ResolveError> {
274        let now = Instant::now();
275        let now_wall = Utc::now();
276
277        // 1. L1 lookup (guard dropped before any await).
278        let cached = {
279            let mut guard = self.cache.lock().expect("session cache mutex poisoned");
280            guard.get(&hash)
281        };
282        if let Some(entry) = cached {
283            let wall_valid = match &entry.outcome {
284                Some(s) => s.expires_at > now_wall,
285                None => true,
286            };
287            if !wall_valid {
288                // Positive entry past its session expiry: forced miss.
289                self.cache
290                    .lock()
291                    .expect("session cache mutex poisoned")
292                    .remove(&hash);
293            } else if now < entry.fresh_until {
294                // Hit. Refresh-ahead for hot positive entries.
295                if entry.outcome.is_some() {
296                    let remaining = entry.fresh_until.saturating_duration_since(now);
297                    if remaining.as_secs_f64() < 0.10 * entry.effective_ttl.as_secs_f64() {
298                        self.clone().spawn_refresh(hash.clone());
299                    }
300                }
301                return Ok(entry.outcome);
302            }
303        }
304
305        // 2. Miss (or stale): capture a stale candidate, then single-flight fill.
306        let stale = self.stale_candidate(&hash, now, now_wall);
307        let this = self.clone();
308        let key = hash.clone();
309        let outcome = self
310            .flight
311            .run(&hash, move || {
312                let this = this.clone();
313                async move { this.fill(key).await }
314            })
315            .await;
316
317        match outcome {
318            Ok(v) => Ok(v),
319            Err(e) => {
320                if e.is_transport() {
321                    if let Some((s, _fetched_at)) = stale {
322                        log::warn!("session resolver: serving stale entry on transport error: {e}");
323                        return Ok(Some(s));
324                    }
325                }
326                Err(e)
327            }
328        }
329    }
330
331    async fn fill(self: Arc<Self>, hash: String) -> Result<Option<ResolvedSession>, ResolveError> {
332        let fetched = tokio::time::timeout(RESOLVE_TIMEOUT, self.fetcher.fetch(&hash))
333            .await
334            .map_err(|_| {
335                ResolveError::Transport("session resolve timed out after 5s".to_string())
336            })??;
337        let now = Instant::now();
338        let entry = match &fetched {
339            Some(s) => {
340                let eff = self.effective_ttl();
341                let wall_remaining = (s.expires_at - Utc::now())
342                    .to_std()
343                    .unwrap_or(Duration::ZERO);
344                let ttl = eff.min(wall_remaining);
345                Entry {
346                    outcome: Some(s.clone()),
347                    fetched_at: now,
348                    fresh_until: now + ttl,
349                    stale_until: now + ttl + self.cfg.stale_if_error,
350                    effective_ttl: eff,
351                }
352            }
353            None => Entry {
354                outcome: None,
355                fetched_at: now,
356                fresh_until: now + self.cfg.neg_ttl,
357                stale_until: now + self.cfg.neg_ttl,
358                effective_ttl: self.cfg.neg_ttl,
359            },
360        };
361        {
362            let mut guard = self.cache.lock().expect("session cache mutex poisoned");
363            guard.put(hash, entry);
364        }
365        Ok(fetched)
366    }
367
368    fn stale_candidate(
369        &self,
370        hash: &str,
371        now: Instant,
372        now_wall: DateTime<Utc>,
373    ) -> Option<(ResolvedSession, Instant)> {
374        if self.cfg.stale_if_error.is_zero() {
375            return None;
376        }
377        let guard = self.cache.lock().expect("session cache mutex poisoned");
378        let entry = guard.peek(hash)?;
379        match entry.outcome {
380            Some(s) if s.expires_at > now_wall && now < entry.stale_until => {
381                Some((s, entry.fetched_at))
382            }
383            _ => None,
384        }
385    }
386
387    fn spawn_refresh(self: Arc<Self>, hash: String) {
388        tokio::spawn(async move {
389            let this = self.clone();
390            let key = hash.clone();
391            let _ = self
392                .flight
393                .run(&hash, move || {
394                    let this = this.clone();
395                    async move { this.fill(key).await }
396                })
397                .await;
398        });
399    }
400}
401
402/// The cache-tiered resolver. Constructed over any [`SessionFetcher`].
403#[derive(Clone)]
404pub struct CachedResolver {
405    shared: Arc<ResolverInner>,
406}
407
408impl CachedResolver {
409    pub fn new(fetcher: Arc<dyn SessionFetcher>, cfg: ResolverConfig) -> Self {
410        CachedResolver {
411            shared: Arc::new(ResolverInner {
412                cache: Mutex::new(Lru::new(cfg.capacity)),
413                flight: SingleFlight::default(),
414                fetcher,
415                cfg,
416            }),
417        }
418    }
419
420    /// Erase to the object-safe trait for storage in state structs.
421    pub fn into_dyn(self) -> Arc<dyn SessionResolver> {
422        Arc::new(self)
423    }
424}
425
426impl SessionResolver for CachedResolver {
427    fn resolve<'a>(&'a self, token_hash: &'a str) -> ResolveFuture<'a> {
428        let shared = self.shared.clone();
429        let hash = token_hash.to_string();
430        Box::pin(async move { shared.resolve(hash).await })
431    }
432
433    fn evict(&self, token_hash: &str) {
434        let mut guard = self
435            .shared
436            .cache
437            .lock()
438            .expect("session cache mutex poisoned");
439        guard.remove(token_hash);
440    }
441}
442
443/// Fleet path: resolve over `am.SessionService` on nio-client. The relying
444/// party supplies a connected [`Channel`] (see [`crate::connect_channel`];
445/// check and session are always distinct endpoints).
446pub struct GrpcSessionResolver;
447
448impl GrpcSessionResolver {
449    // Factory returning the object-safe trait; not a `Self` ctor.
450    #[allow(clippy::new_ret_no_self)]
451    pub fn new(channel: Channel, cfg: ResolverConfig) -> Arc<dyn SessionResolver> {
452        let client = SessionServiceClient::new(channel);
453        let fetcher: Arc<dyn SessionFetcher> = Arc::new(GrpcFetcher { client });
454        CachedResolver::new(fetcher, cfg).into_dyn()
455    }
456}
457
458struct GrpcFetcher {
459    client: SessionServiceClient<Channel>,
460}
461
462impl SessionFetcher for GrpcFetcher {
463    fn fetch<'a>(&'a self, token_hash: &'a str) -> ResolveFuture<'a> {
464        let mut client = self.client.clone();
465        let token_hash = token_hash.to_string();
466        Box::pin(async move {
467            let resp = client
468                .resolve(ResolveRequest { token_hash })
469                .await
470                .map_err(classify_status)?;
471            match resp.into_inner().outcome {
472                Some(resolve_response::Outcome::Session(s)) => Ok(Some(ResolvedSession {
473                    principal: UserId::try_from(s.principal)
474                        .map_err(|e| ResolveError::Backend(e.to_string()))?,
475                    tenant_id: s.tenant_id,
476                    expires_at: DateTime::from_timestamp(s.expires_at_unix_seconds, 0)
477                        .unwrap_or_else(Utc::now),
478                })),
479                // unknown / expired / revoked — deliberately indistinguishable.
480                Some(resolve_response::Outcome::NotFound(_)) | None => Ok(None),
481            }
482        })
483    }
484}
485
486fn classify_status(status: tonic::Status) -> ResolveError {
487    match status.code() {
488        tonic::Code::Unavailable | tonic::Code::DeadlineExceeded => {
489            ResolveError::Transport(status.message().to_string())
490        }
491        _ => ResolveError::Backend(status.message().to_string()),
492    }
493}
494
495#[cfg(test)]
496mod tests {
497    use super::*;
498    use std::sync::atomic::AtomicUsize;
499
500    struct CountingFetcher {
501        calls: Arc<AtomicUsize>,
502        session: Option<ResolvedSession>,
503    }
504
505    impl SessionFetcher for CountingFetcher {
506        fn fetch<'a>(&'a self, _token_hash: &'a str) -> ResolveFuture<'a> {
507            self.calls.fetch_add(1, Ordering::Relaxed);
508            let session = self.session.clone();
509            Box::pin(async move { Ok(session) })
510        }
511    }
512
513    struct FailingFetcher {
514        calls: Arc<AtomicUsize>,
515    }
516
517    impl SessionFetcher for FailingFetcher {
518        fn fetch<'a>(&'a self, _token_hash: &'a str) -> ResolveFuture<'a> {
519            self.calls.fetch_add(1, Ordering::Relaxed);
520            Box::pin(async move { Err(ResolveError::Transport("down".into())) })
521        }
522    }
523
524    fn session_valid_for(mins: i64) -> ResolvedSession {
525        ResolvedSession {
526            principal: UserId::try_from(42).unwrap(),
527            tenant_id: String::new(),
528            expires_at: Utc::now() + chrono::TimeDelta::minutes(mins),
529        }
530    }
531
532    fn cfg() -> ResolverConfig {
533        ResolverConfig {
534            capacity: 100,
535            l1_ttl: Duration::from_secs(30),
536            neg_ttl: Duration::from_secs(2),
537            stale_if_error: Duration::ZERO,
538        }
539    }
540
541    #[test]
542    fn token_hash_is_hex_sha256() {
543        // sha256("") — a known vector.
544        assert_eq!(
545            token_hash(""),
546            "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855"
547        );
548        assert_eq!(token_hash("abc").len(), 64);
549    }
550
551    #[test]
552    fn default_config_matches_243() {
553        let cfg = ResolverConfig::default();
554        assert_eq!(cfg.capacity, 10_000);
555        assert_eq!(cfg.l1_ttl, Duration::from_secs(30));
556        assert_eq!(cfg.neg_ttl, Duration::from_secs(2));
557        assert_eq!(cfg.stale_if_error, Duration::ZERO);
558    }
559
560    #[tokio::test]
561    async fn hit_does_not_refetch() {
562        let calls = Arc::new(AtomicUsize::new(0));
563        let fetcher = Arc::new(CountingFetcher {
564            calls: calls.clone(),
565            session: Some(session_valid_for(120)),
566        });
567        let r = CachedResolver::new(fetcher, cfg());
568        let first = r.resolve("deadbeef").await.unwrap();
569        assert!(first.is_some());
570        let _ = r.resolve("deadbeef").await.unwrap();
571        let _ = r.resolve("deadbeef").await.unwrap();
572        assert_eq!(calls.load(Ordering::Relaxed), 1, "L1 hit must not refetch");
573    }
574
575    #[tokio::test]
576    async fn unknown_token_is_tombstoned() {
577        let calls = Arc::new(AtomicUsize::new(0));
578        let fetcher = Arc::new(CountingFetcher {
579            calls: calls.clone(),
580            session: None,
581        });
582        let r = CachedResolver::new(fetcher, cfg());
583        assert!(r.resolve("nope").await.unwrap().is_none());
584        assert!(r.resolve("nope").await.unwrap().is_none());
585        assert_eq!(
586            calls.load(Ordering::Relaxed),
587            1,
588            "tombstone must not refetch"
589        );
590    }
591
592    #[tokio::test]
593    async fn evict_forces_refetch() {
594        let calls = Arc::new(AtomicUsize::new(0));
595        let fetcher = Arc::new(CountingFetcher {
596            calls: calls.clone(),
597            session: Some(session_valid_for(120)),
598        });
599        let r = CachedResolver::new(fetcher, cfg());
600        let _ = r.resolve("k").await.unwrap();
601        r.evict("k");
602        let _ = r.resolve("k").await.unwrap();
603        assert_eq!(
604            calls.load(Ordering::Relaxed),
605            2,
606            "evict must force a refetch"
607        );
608    }
609
610    #[tokio::test]
611    async fn expired_session_is_forced_miss() {
612        let calls = Arc::new(AtomicUsize::new(0));
613        // Session already expired: every lookup must be a forced miss.
614        let fetcher = Arc::new(CountingFetcher {
615            calls: calls.clone(),
616            session: Some(session_valid_for(-1)),
617        });
618        let r = CachedResolver::new(fetcher, cfg());
619        let _ = r.resolve("k").await.unwrap();
620        let _ = r.resolve("k").await.unwrap();
621        assert_eq!(
622            calls.load(Ordering::Relaxed),
623            2,
624            "expired entry must not serve from L1"
625        );
626    }
627
628    #[tokio::test]
629    async fn transport_error_propagates_without_stale_window() {
630        let calls = Arc::new(AtomicUsize::new(0));
631        let fetcher = Arc::new(FailingFetcher {
632            calls: calls.clone(),
633        });
634        let r = CachedResolver::new(fetcher, cfg());
635        let err = r.resolve("k").await.expect_err("must fail");
636        assert!(matches!(err, ResolveError::Transport(_)));
637    }
638
639    #[tokio::test]
640    async fn zero_capacity_disables_cache() {
641        let calls = Arc::new(AtomicUsize::new(0));
642        let fetcher = Arc::new(CountingFetcher {
643            calls: calls.clone(),
644            session: Some(session_valid_for(120)),
645        });
646        let mut c = cfg();
647        c.capacity = 0;
648        let r = CachedResolver::new(fetcher, c);
649        let _ = r.resolve("k").await.unwrap();
650        let _ = r.resolve("k").await.unwrap();
651        assert_eq!(
652            calls.load(Ordering::Relaxed),
653            2,
654            "capacity 0 must not cache"
655        );
656    }
657
658    #[tokio::test]
659    async fn lru_evicts_oldest() {
660        let calls = Arc::new(AtomicUsize::new(0));
661        let fetcher = Arc::new(CountingFetcher {
662            calls: calls.clone(),
663            session: Some(session_valid_for(120)),
664        });
665        let mut c = cfg();
666        c.capacity = 2;
667        let r = CachedResolver::new(fetcher, c);
668        let _ = r.resolve("a").await.unwrap();
669        let _ = r.resolve("b").await.unwrap();
670        let _ = r.resolve("c").await.unwrap(); // evicts "a"
671        let _ = r.resolve("a").await.unwrap(); // refetch
672        assert_eq!(calls.load(Ordering::Relaxed), 4, "lru must evict oldest");
673    }
674
675    struct GatedFirstFetcher {
676        calls: Arc<AtomicUsize>,
677        gate: Arc<tokio::sync::Semaphore>,
678        session: Option<ResolvedSession>,
679    }
680
681    impl SessionFetcher for GatedFirstFetcher {
682        fn fetch<'a>(&'a self, _token_hash: &'a str) -> ResolveFuture<'a> {
683            let n = self.calls.fetch_add(1, Ordering::Relaxed);
684            let gate = self.gate.clone();
685            let session = self.session.clone();
686            Box::pin(async move {
687                if n == 0 {
688                    let _permit = gate.acquire_owned().await.expect("gate closed");
689                }
690                Ok(session)
691            })
692        }
693    }
694
695    // Regression: a leader future dropped mid-fill (client disconnect) must
696    // not leak the inflight entry and freeze this hash on a memoized result —
697    // the fill runs detached and cleans up the flight map itself.
698    #[tokio::test]
699    async fn cancelled_leader_does_not_poison_future_resolves() {
700        let calls = Arc::new(AtomicUsize::new(0));
701        let gate = Arc::new(tokio::sync::Semaphore::new(0));
702        let fetcher = Arc::new(GatedFirstFetcher {
703            calls: calls.clone(),
704            gate: gate.clone(),
705            session: Some(session_valid_for(120)),
706        });
707        let r = CachedResolver::new(fetcher, cfg());
708
709        let leader = {
710            let r = r.clone();
711            tokio::spawn(async move { r.resolve("k").await })
712        };
713        while calls.load(Ordering::Relaxed) == 0 {
714            tokio::task::yield_now().await;
715        }
716        leader.abort();
717        let _ = leader.await;
718
719        gate.add_permits(1);
720        let got = r.resolve("k").await.unwrap();
721        assert!(got.is_some());
722        assert_eq!(calls.load(Ordering::Relaxed), 1);
723
724        r.evict("k");
725        let got = r.resolve("k").await.unwrap();
726        assert!(got.is_some());
727        assert_eq!(
728            calls.load(Ordering::Relaxed),
729            2,
730            "resolve after evict must refetch even after a cancelled leader"
731        );
732    }
733
734    struct HangingFetcher;
735
736    impl SessionFetcher for HangingFetcher {
737        fn fetch<'a>(&'a self, _token_hash: &'a str) -> ResolveFuture<'a> {
738            Box::pin(futures::future::pending())
739        }
740    }
741
742    #[tokio::test(start_paused = true)]
743    async fn fill_timeout_is_transport_error() {
744        let r = CachedResolver::new(Arc::new(HangingFetcher), cfg());
745        let err = r
746            .resolve("k")
747            .await
748            .expect_err("hung backend must time out");
749        assert!(matches!(err, ResolveError::Transport(_)));
750    }
751
752    struct SwitchableFetcher {
753        session: ResolvedSession,
754        err: Mutex<Option<ResolveError>>,
755    }
756
757    impl SessionFetcher for SwitchableFetcher {
758        fn fetch<'a>(&'a self, _token_hash: &'a str) -> ResolveFuture<'a> {
759            let err = self.err.lock().unwrap().clone();
760            let session = self.session.clone();
761            Box::pin(async move {
762                match err {
763                    Some(e) => Err(e),
764                    None => Ok(Some(session)),
765                }
766            })
767        }
768    }
769
770    fn stale_cfg() -> ResolverConfig {
771        ResolverConfig {
772            capacity: 100,
773            l1_ttl: Duration::from_millis(1),
774            neg_ttl: Duration::from_secs(2),
775            stale_if_error: Duration::from_secs(60),
776        }
777    }
778
779    #[tokio::test]
780    async fn transport_error_serves_stale_within_window() {
781        let fetcher = Arc::new(SwitchableFetcher {
782            session: session_valid_for(120),
783            err: Mutex::new(None),
784        });
785        let r = CachedResolver::new(fetcher.clone(), stale_cfg());
786        assert!(r.resolve("k").await.unwrap().is_some());
787        tokio::time::sleep(Duration::from_millis(20)).await; // past fresh_until
788        *fetcher.err.lock().unwrap() = Some(ResolveError::Transport("unavailable".into()));
789        let got = r
790            .resolve("k")
791            .await
792            .expect("transport error must serve stale");
793        assert!(got.is_some());
794    }
795
796    #[tokio::test]
797    async fn backend_error_propagates_despite_stale_window() {
798        let fetcher = Arc::new(SwitchableFetcher {
799            session: session_valid_for(120),
800            err: Mutex::new(None),
801        });
802        let r = CachedResolver::new(fetcher.clone(), stale_cfg());
803        assert!(r.resolve("k").await.unwrap().is_some());
804        tokio::time::sleep(Duration::from_millis(20)).await;
805        *fetcher.err.lock().unwrap() = Some(ResolveError::Backend("boom".into()));
806        r.resolve("k")
807            .await
808            .expect_err("backend error must propagate, not serve stale");
809    }
810
811    #[test]
812    fn classify_status_transport_vs_backend() {
813        assert!(matches!(
814            classify_status(tonic::Status::unavailable("x")),
815            ResolveError::Transport(_)
816        ));
817        assert!(matches!(
818            classify_status(tonic::Status::deadline_exceeded("x")),
819            ResolveError::Transport(_)
820        ));
821        assert!(matches!(
822            classify_status(tonic::Status::internal("x")),
823            ResolveError::Backend(_)
824        ));
825        assert!(matches!(
826            classify_status(tonic::Status::permission_denied("x")),
827            ResolveError::Backend(_)
828        ));
829    }
830
831    #[test]
832    fn token_hash_newtype_matches_fn() {
833        assert_eq!(TokenHash::from_raw("tok").as_str(), token_hash("tok"));
834    }
835}