Skip to main content

authplane_sdk/cache/
document_cache.rs

1//! Generic JSON document cache: TTL with HTTP cache-header awareness,
2//! background refresh at 80 % of effective TTL, stale fallback on fetch
3//! errors, lock-coordinated fetching, and an optional change callback.
4
5use std::future::Future;
6use std::pin::Pin;
7use std::sync::{Arc, Weak};
8use std::time::{Duration, Instant};
9
10use serde_json::Value;
11use tokio::sync::{Mutex, OnceCell};
12use tokio::task::JoinHandle;
13
14#[cfg(test)]
15use crate::AuthError;
16use crate::AuthplaneError;
17
18/// Result returned by a [`DocumentFetcherFn`].
19#[derive(Debug, Clone)]
20pub struct FetchResult {
21    /// The parsed JSON body.
22    pub document: Value,
23    /// Absolute Unix expiry timestamp derived from cache headers, if any.
24    pub expires_at: Option<f64>,
25}
26
27/// Type alias for the boxed async fetcher closure used by [`DocumentCache`].
28pub type DocumentFetcherFn = Arc<
29    dyn Fn() -> Pin<Box<dyn Future<Output = Result<FetchResult, AuthplaneError>> + Send>>
30        + Send
31        + Sync,
32>;
33
34/// Type alias for the on_change callback used by [`DocumentCache`].
35pub type DocumentChangeCallback =
36    Arc<dyn Fn(Value, Value) -> Pin<Box<dyn Future<Output = ()> + Send>> + Send + Sync>;
37
38#[derive(Debug)]
39struct CachedDocument {
40    body: Value,
41    cache_time: Instant,
42    server_expires_at_unix: Option<f64>,
43    /// Set by [`DocumentCache::expire`]: the next `get` must re-fetch no
44    /// matter how warm the TTL says this entry is, but the body stays
45    /// available as the stale fallback should that fetch fail.
46    expired: bool,
47}
48
49#[derive(Debug, Default)]
50struct State {
51    cached: Option<CachedDocument>,
52    refresh_task: Option<JoinHandle<()>>,
53}
54
55/// Generic JSON document cache shared by metadata and JWKS layers.
56pub struct DocumentCache {
57    fetcher: DocumentFetcherFn,
58    refresh_seconds: u64,
59    document_type: String,
60    on_change: Option<DocumentChangeCallback>,
61    error_factory: Box<dyn Fn(&str) -> AuthplaneError + Send + Sync>,
62    state: Arc<Mutex<State>>,
63    fetch_lock: Arc<Mutex<()>>,
64    /// Weak self-pointer used to spawn a background-refresh task without
65    /// creating an `Arc<Self>` cycle that would leak the cache.
66    self_handle: OnceCell<Weak<DocumentCache>>,
67}
68
69impl std::fmt::Debug for DocumentCache {
70    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
71        f.debug_struct("DocumentCache")
72            .field("document_type", &self.document_type)
73            .field("refresh_seconds", &self.refresh_seconds)
74            .finish()
75    }
76}
77
78impl DocumentCache {
79    /// Build a new cache with default error factory (`transport_error`).
80    pub fn new(
81        fetcher: DocumentFetcherFn,
82        refresh_seconds: u64,
83        document_type: impl Into<String>,
84        on_change: Option<DocumentChangeCallback>,
85    ) -> Arc<Self> {
86        Self::with_error_factory(
87            fetcher,
88            refresh_seconds,
89            document_type,
90            on_change,
91            Box::new(default_error_factory),
92        )
93    }
94
95    /// Build a cache with a custom error factory (e.g. `JwksFetchError`).
96    pub fn with_error_factory(
97        fetcher: DocumentFetcherFn,
98        refresh_seconds: u64,
99        document_type: impl Into<String>,
100        on_change: Option<DocumentChangeCallback>,
101        error_factory: Box<dyn Fn(&str) -> AuthplaneError + Send + Sync>,
102    ) -> Arc<Self> {
103        let cache = Arc::new(Self {
104            fetcher,
105            refresh_seconds: refresh_seconds.max(1),
106            document_type: document_type.into(),
107            on_change,
108            error_factory,
109            state: Arc::new(Mutex::new(State::default())),
110            fetch_lock: Arc::new(Mutex::new(())),
111            self_handle: OnceCell::new(),
112        });
113        // Store a Weak handle so background tasks do not keep the cache
114        // alive past the user's last Arc drop.
115        let _ = cache.self_handle.set(Arc::downgrade(&cache));
116        cache
117    }
118
119    /// Document type label.
120    pub fn document_type(&self) -> &str {
121        &self.document_type
122    }
123
124    /// Configured refresh interval in seconds.
125    pub fn refresh_seconds(&self) -> u64 {
126        self.refresh_seconds
127    }
128
129    fn effective_expires_at(&self, doc: &CachedDocument, now: Instant) -> Instant {
130        if doc.expired {
131            // `cache_time` is never later than any `now` a caller compares
132            // against, so an expired entry always misses the warm path.
133            return doc.cache_time;
134        }
135        let configured = doc.cache_time + Duration::from_secs(self.refresh_seconds);
136        if let Some(server_unix) = doc.server_expires_at_unix {
137            // Translate the server's absolute Unix expiry into a monotonic
138            // Instant relative to "now" so it can be compared cleanly.
139            let now_unix = now_unix_seconds();
140            if server_unix <= now_unix {
141                return doc.cache_time;
142            }
143            let delta = Duration::from_secs_f64((server_unix - now_unix).max(0.0));
144            let server_instant = now + delta;
145            return std::cmp::min(configured, server_instant);
146        }
147        configured
148    }
149
150    /// Return the cached document, fetching or refreshing as needed.
151    ///
152    /// Concurrent callers serialise behind an internal mutex so only one
153    /// fetch is in flight at a time. On fetch errors a stale cache is
154    /// served, silently — the crate carries no logging facility, so the
155    /// failed refresh leaves no trace; if no cache existed yet, the error
156    /// factory is invoked.
157    pub async fn get(&self, force_refresh: bool) -> Result<Value, AuthplaneError> {
158        self.get_inner(force_refresh, true).await
159    }
160
161    /// Force a refresh and report a failed fetch instead of falling back to
162    /// the cached document.
163    ///
164    /// [`Self::get`] answers a failed refresh with the last good body, which
165    /// is the right posture for a caller that only wants *a* document. A
166    /// caller whose next decision depends on the document being current —
167    /// rebinding a rotated `jwks_uri`, say — needs to know the refresh did
168    /// not happen, otherwise it treats the previous interval's body as a
169    /// fresh answer and schedules its next attempt as if it had one. The
170    /// cached body is left in place either way, so other readers keep being
171    /// served through the failure.
172    pub(crate) async fn refresh_strict(&self) -> Result<Value, AuthplaneError> {
173        self.get_inner(true, false).await
174    }
175
176    async fn get_inner(
177        &self,
178        force_refresh: bool,
179        stale_fallback: bool,
180    ) -> Result<Value, AuthplaneError> {
181        let now = Instant::now();
182        // Fast path: cached value still warm.
183        {
184            let mut state = self.state.lock().await;
185            if !force_refresh && let Some(doc) = state.cached.as_ref() {
186                let expires = self.effective_expires_at(doc, now);
187                if now < expires {
188                    let body = doc.body.clone();
189                    let ttl = expires.saturating_duration_since(doc.cache_time);
190                    let elapsed = now.saturating_duration_since(doc.cache_time);
191                    let needs_bg_refresh = !ttl.is_zero() && elapsed >= ttl.mul_f64(0.8);
192                    let task_idle = state
193                        .refresh_task
194                        .as_ref()
195                        .map(|task| task.is_finished())
196                        .unwrap_or(true);
197                    if needs_bg_refresh && task_idle {
198                        self.do_spawn_background_refresh(&mut state);
199                    }
200                    drop(state);
201                    return Ok(body);
202                }
203            }
204        }
205
206        // Slow path: fetch under fetch_lock to coordinate with concurrent callers.
207        let _guard = self.fetch_lock.lock().await;
208
209        // Re-check after acquiring the lock — another caller may have refreshed.
210        {
211            let state = self.state.lock().await;
212            if !force_refresh && let Some(doc) = state.cached.as_ref() {
213                let expires = self.effective_expires_at(doc, Instant::now());
214                if Instant::now() < expires {
215                    return Ok(doc.body.clone());
216                }
217            }
218        }
219
220        match (self.fetcher)().await {
221            Ok(result) => {
222                let mut state = self.state.lock().await;
223                let old_body = state.cached.as_ref().map(|doc| doc.body.clone());
224                let new_body = result.document.clone();
225                state.cached = Some(CachedDocument {
226                    body: new_body.clone(),
227                    cache_time: Instant::now(),
228                    server_expires_at_unix: result.expires_at,
229                    expired: false,
230                });
231                drop(state);
232                if let (Some(callback), Some(old)) = (self.on_change.as_ref(), old_body)
233                    && old != new_body
234                {
235                    let cb = callback.clone();
236                    let old_clone = old;
237                    let new_clone = new_body.clone();
238                    tokio::spawn(async move {
239                        (cb)(old_clone, new_clone).await;
240                    });
241                }
242                Ok(new_body)
243            }
244            Err(error) => {
245                if stale_fallback {
246                    let state = self.state.lock().await;
247                    if let Some(doc) = state.cached.as_ref() {
248                        return Ok(doc.body.clone());
249                    }
250                }
251                Err((self.error_factory)(&format!(
252                    "Failed to fetch {}: {error}",
253                    self.document_type
254                )))
255            }
256        }
257    }
258
259    /// Expire the cached document so the next [`Self::get`] re-fetches,
260    /// while keeping the body available as the stale fallback.
261    ///
262    /// Needed when the fetch target itself changes rather than expiring —
263    /// an RFC 8414 `jwks_uri` rotation, for instance. The cached body came
264    /// from the withdrawn URL, so the TTL that would normally govern it no
265    /// longer says anything about its freshness.
266    ///
267    /// Expired, not dropped: dropping would leave the fallback branch with
268    /// nothing to serve, so one transient failure at the new target would
269    /// fail every caller — including ones the retired document was
270    /// answering a moment earlier. Keeping the body preserves the cache's
271    /// last-known-good posture; retired content is served only while a
272    /// fetch of the new target is failing.
273    ///
274    /// Takes `fetch_lock` first, so a fetch already in flight against the
275    /// old target commits before the expiry lands and cannot resurrect a
276    /// warm entry for a full TTL afterwards. Safe to await from callers
277    /// holding no cache locks; nothing here is held across the fetcher.
278    pub(crate) async fn expire(&self) {
279        let _guard = self.fetch_lock.lock().await;
280        let mut state = self.state.lock().await;
281        if let Some(doc) = state.cached.as_mut() {
282            doc.expired = true;
283        }
284    }
285
286    /// Cancel the background refresh task (if any). Safe to call multiple times.
287    pub async fn aclose(&self) {
288        let mut state = self.state.lock().await;
289        if let Some(task) = state.refresh_task.take() {
290            task.abort();
291        }
292    }
293
294    fn do_spawn_background_refresh(&self, state: &mut State) {
295        if state
296            .refresh_task
297            .as_ref()
298            .map(|task| !task.is_finished())
299            .unwrap_or(false)
300        {
301            return;
302        }
303        let weak = match self.self_handle.get() {
304            Some(handle) => handle.clone(),
305            None => return,
306        };
307        let task = tokio::spawn(async move {
308            if let Some(cache) = weak.upgrade() {
309                let _ = cache.get(true).await;
310            }
311        });
312        state.refresh_task = Some(task);
313    }
314}
315
316// Thin wrapper around the shared `errors::transport_error` helper, kept
317// as a function so it can be passed as a `Box<dyn Fn(&str) -> AuthplaneError>`
318// default. New transport-error sites should call the helper directly
319// rather than going through this thunk.
320fn default_error_factory(message: &str) -> AuthplaneError {
321    crate::errors::transport_error(message)
322}
323
324use crate::time_utils::unix_now_secs_f64 as now_unix_seconds;
325
326#[cfg(test)]
327mod tests {
328    use super::*;
329    use std::sync::atomic::{AtomicUsize, Ordering};
330    use tokio::sync::Mutex as TokioMutex;
331
332    fn make_fetcher(
333        responses: Vec<Result<FetchResult, AuthplaneError>>,
334    ) -> (DocumentFetcherFn, Arc<AtomicUsize>) {
335        let counter = Arc::new(AtomicUsize::new(0));
336        let counter_clone = counter.clone();
337        let queue = Arc::new(TokioMutex::new(responses));
338        let fetcher: DocumentFetcherFn = Arc::new(move || {
339            let counter = counter_clone.clone();
340            let queue = queue.clone();
341            Box::pin(async move {
342                counter.fetch_add(1, Ordering::SeqCst);
343                let mut q = queue.lock().await;
344                if q.is_empty() {
345                    return Err(AuthplaneError::Auth(AuthError {
346                        message: "no more responses".to_string(),
347                        code: "test".to_string(),
348                        status_code: None,
349                    }));
350                }
351                q.remove(0)
352            })
353        });
354        (fetcher, counter)
355    }
356
357    #[tokio::test]
358    async fn first_get_invokes_fetcher_and_caches_result() {
359        let (fetcher, counter) = make_fetcher(vec![Ok(FetchResult {
360            document: serde_json::json!({"a": 1}),
361            expires_at: None,
362        })]);
363        let cache = DocumentCache::new(fetcher, 60, "test", None);
364        let body = cache.get(false).await.expect("ok");
365        assert_eq!(body, serde_json::json!({"a": 1}));
366        // Subsequent get returns cache without re-fetching.
367        let body2 = cache.get(false).await.expect("ok");
368        assert_eq!(body2, serde_json::json!({"a": 1}));
369        assert_eq!(counter.load(Ordering::SeqCst), 1);
370    }
371
372    #[tokio::test]
373    async fn force_refresh_bypasses_cache() {
374        let (fetcher, counter) = make_fetcher(vec![
375            Ok(FetchResult {
376                document: serde_json::json!({"v": 1}),
377                expires_at: None,
378            }),
379            Ok(FetchResult {
380                document: serde_json::json!({"v": 2}),
381                expires_at: None,
382            }),
383        ]);
384        let cache = DocumentCache::new(fetcher, 600, "test", None);
385        let v1 = cache.get(false).await.expect("ok");
386        assert_eq!(v1["v"], 1);
387        let v2 = cache.get(true).await.expect("ok");
388        assert_eq!(v2["v"], 2);
389        assert_eq!(counter.load(Ordering::SeqCst), 2);
390    }
391
392    #[tokio::test]
393    async fn fetch_failure_falls_back_to_stale() {
394        let (fetcher, _counter) = make_fetcher(vec![
395            Ok(FetchResult {
396                document: serde_json::json!({"k": "first"}),
397                expires_at: None,
398            }),
399            Err(AuthplaneError::Auth(AuthError {
400                message: "boom".to_string(),
401                code: "transport_error".to_string(),
402                status_code: None,
403            })),
404        ]);
405        let cache = DocumentCache::new(fetcher, 600, "test", None);
406        let _ = cache.get(false).await.expect("first ok");
407        // Force refresh fails; stale cache must be returned.
408        let stale = cache.get(true).await.expect("stale");
409        assert_eq!(stale["k"], "first");
410    }
411
412    #[tokio::test]
413    async fn first_fetch_failure_propagates_error() {
414        let (fetcher, _counter) = make_fetcher(vec![Err(AuthplaneError::Auth(AuthError {
415            message: "boom".to_string(),
416            code: "transport_error".to_string(),
417            status_code: None,
418        }))]);
419        let cache = DocumentCache::new(fetcher, 60, "test", None);
420        let result = cache.get(false).await;
421        assert!(result.is_err());
422    }
423
424    #[tokio::test]
425    async fn expire_forces_the_next_get_to_refetch() {
426        let (fetcher, counter) = make_fetcher(vec![
427            Ok(FetchResult {
428                document: serde_json::json!({"v": 1}),
429                expires_at: None,
430            }),
431            Ok(FetchResult {
432                document: serde_json::json!({"v": 2}),
433                expires_at: None,
434            }),
435        ]);
436        // A long TTL would normally keep serving the first document.
437        let cache = DocumentCache::new(fetcher, 3600, "test", None);
438        assert_eq!(cache.get(false).await.expect("first")["v"], 1);
439        cache.expire().await;
440        assert_eq!(cache.get(false).await.expect("second")["v"], 2);
441        assert_eq!(counter.load(Ordering::SeqCst), 2);
442    }
443
444    #[tokio::test]
445    async fn expire_keeps_the_stale_body_as_fallback_when_the_refetch_fails() {
446        // The expiry must not strip the cache of its last-known-good
447        // posture: if the re-fetch it forces fails, callers are served the
448        // retired body rather than an error. This is the difference between
449        // expiring and dropping — a dropped entry would fail every caller
450        // for as long as the new target is unreachable.
451        let (fetcher, counter) = make_fetcher(vec![
452            Ok(FetchResult {
453                document: serde_json::json!({"v": 1}),
454                expires_at: None,
455            }),
456            Err(AuthplaneError::Auth(AuthError {
457                message: "new target unreachable".to_string(),
458                code: "transport_error".to_string(),
459                status_code: None,
460            })),
461            Ok(FetchResult {
462                document: serde_json::json!({"v": 2}),
463                expires_at: None,
464            }),
465        ]);
466        let cache = DocumentCache::new(fetcher, 3600, "test", None);
467        assert_eq!(cache.get(false).await.expect("first")["v"], 1);
468
469        cache.expire().await;
470
471        // The forced re-fetch fails; the retired body must still answer.
472        assert_eq!(cache.get(false).await.expect("stale fallback")["v"], 1);
473        // The entry stays expired, so the next get retries the fetch
474        // rather than settling back onto the retired body.
475        assert_eq!(cache.get(false).await.expect("recovered")["v"], 2);
476        assert_eq!(counter.load(Ordering::SeqCst), 3);
477    }
478
479    #[tokio::test]
480    async fn on_change_callback_fires_when_document_changes() {
481        let (fetcher, _counter) = make_fetcher(vec![
482            Ok(FetchResult {
483                document: serde_json::json!({"v": 1}),
484                expires_at: None,
485            }),
486            Ok(FetchResult {
487                document: serde_json::json!({"v": 2}),
488                expires_at: None,
489            }),
490        ]);
491        let invoked = Arc::new(AtomicUsize::new(0));
492        let invoked_cb = invoked.clone();
493        let on_change: DocumentChangeCallback = Arc::new(move |_old, _new| {
494            let invoked_cb = invoked_cb.clone();
495            Box::pin(async move {
496                invoked_cb.fetch_add(1, Ordering::SeqCst);
497            })
498        });
499        let cache = DocumentCache::new(fetcher, 600, "test", Some(on_change));
500        cache.get(false).await.expect("first");
501        cache.get(true).await.expect("second");
502        // The callback runs on a tokio task; give it a moment.
503        tokio::task::yield_now().await;
504        tokio::time::sleep(std::time::Duration::from_millis(20)).await;
505        assert_eq!(invoked.load(Ordering::SeqCst), 1);
506    }
507}