1use 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#[derive(Debug, Clone)]
20pub struct FetchResult {
21 pub document: Value,
23 pub expires_at: Option<f64>,
25}
26
27pub type DocumentFetcherFn = Arc<
29 dyn Fn() -> Pin<Box<dyn Future<Output = Result<FetchResult, AuthplaneError>> + Send>>
30 + Send
31 + Sync,
32>;
33
34pub 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 expired: bool,
47}
48
49#[derive(Debug, Default)]
50struct State {
51 cached: Option<CachedDocument>,
52 refresh_task: Option<JoinHandle<()>>,
53}
54
55pub 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 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 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 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 let _ = cache.self_handle.set(Arc::downgrade(&cache));
116 cache
117 }
118
119 pub fn document_type(&self) -> &str {
121 &self.document_type
122 }
123
124 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 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 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 pub async fn get(&self, force_refresh: bool) -> Result<Value, AuthplaneError> {
158 self.get_inner(force_refresh, true).await
159 }
160
161 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 {
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 let _guard = self.fetch_lock.lock().await;
208
209 {
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 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 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
316fn 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 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 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 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 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 assert_eq!(cache.get(false).await.expect("stale fallback")["v"], 1);
473 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 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}