Skip to main content

authplane_sdk/cache/
jwks_cache.rs

1//! JWKS cache with key-ID lookup, force-refresh on `kid` miss, stale
2//! fallback on fetch errors, and background refresh at 80 % of TTL.
3
4use std::sync::Arc;
5use std::time::{Duration, Instant};
6
7use serde_json::Value;
8use tokio::sync::Mutex;
9
10#[cfg(test)]
11use crate::AuthError;
12use crate::AuthplaneError;
13use crate::cache::document_cache::{DocumentCache, DocumentFetcherFn};
14use crate::constants::jwk_params;
15
16/// JWKS cache.
17#[derive(Clone)]
18pub struct JwksCache {
19    inner: Arc<DocumentCache>,
20    last_force_refresh: Arc<Mutex<Option<Instant>>>,
21    min_force_refresh_interval: Duration,
22}
23
24impl std::fmt::Debug for JwksCache {
25    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
26        f.debug_struct("JwksCache")
27            .field("inner", &self.inner)
28            .field(
29                "min_force_refresh_interval",
30                &self.min_force_refresh_interval,
31            )
32            .finish()
33    }
34}
35
36impl JwksCache {
37    /// Default minimum interval between force-refreshes triggered by `kid` misses.
38    pub const DEFAULT_MIN_FORCE_REFRESH_SECONDS: u64 = 30;
39
40    /// Build a JWKS cache.
41    pub fn new(fetcher: DocumentFetcherFn, refresh_seconds: u64) -> Self {
42        let inner = DocumentCache::with_error_factory(
43            fetcher,
44            refresh_seconds,
45            "jwks",
46            None,
47            Box::new(jwks_error_factory),
48        );
49        Self {
50            inner,
51            last_force_refresh: Arc::new(Mutex::new(None)),
52            min_force_refresh_interval: Duration::from_secs(
53                Self::DEFAULT_MIN_FORCE_REFRESH_SECONDS,
54            ),
55        }
56    }
57
58    /// Override the minimum force-refresh interval (test hook).
59    pub fn with_min_force_refresh_interval(mut self, interval: Duration) -> Self {
60        self.min_force_refresh_interval = interval;
61        self
62    }
63
64    /// Underlying [`DocumentCache`] (for `aclose()` plumbing).
65    pub fn document_cache(&self) -> Arc<DocumentCache> {
66        self.inner.clone()
67    }
68
69    /// Cancel any background refresh task.
70    pub async fn aclose(&self) {
71        self.inner.aclose().await;
72    }
73
74    /// Expire the cached JWKS and clear the force-refresh rate limiter.
75    ///
76    /// Called when the fetcher is rebound to a new `jwks_uri` (RFC 8414
77    /// rotation): keys retrieved from the withdrawn URL must stop being
78    /// served from the warm cache — the next lookup fetches from the
79    /// rebound URI — and the next `kid` miss must be free to force a
80    /// refresh instead of being throttled by a force-refresh that ran
81    /// against the old URL.
82    ///
83    /// Expired, not dropped. The retired key set stays available as the
84    /// stale fallback, so an AS that publishes rotated metadata before the
85    /// new endpoint is live — or a transient failure there — degrades to
86    /// serving the last good keys instead of failing every verification.
87    pub(crate) async fn expire(&self) {
88        self.inner.expire().await;
89        *self.last_force_refresh.lock().await = None;
90    }
91
92    /// Return the full JWKS document (forces a fetch if no cached value yet).
93    pub async fn get(&self, force_refresh: bool) -> Result<Value, AuthplaneError> {
94        self.inner.get(force_refresh).await
95    }
96
97    /// Look up a JWK by `kid`, optionally restricting to a specific algorithm.
98    ///
99    /// On a miss, the cache force-refreshes (rate-limited via
100    /// `min_force_refresh_interval`) and tries again, so a key rotation at
101    /// the same `jwks_uri` is followed without a caller-visible refresh.
102    pub async fn get_key_by_kid(
103        &self,
104        kid: &str,
105        algorithm: Option<&str>,
106    ) -> Result<Option<Value>, AuthplaneError> {
107        if let Some(jwk) = self.find_key(self.inner.get(false).await?, kid, algorithm) {
108            return Ok(Some(jwk));
109        }
110        // Miss: maybe rotate. Rate-limit force-refresh attempts.
111        if !self.try_record_force_refresh().await {
112            return Ok(None);
113        }
114        let document = self.inner.get(true).await?;
115        Ok(self.find_key(document, kid, algorithm))
116    }
117
118    fn find_key(&self, document: Value, kid: &str, algorithm: Option<&str>) -> Option<Value> {
119        let keys = document.get("keys").and_then(Value::as_array)?;
120        for entry in keys {
121            if !entry.is_object() {
122                continue;
123            }
124            if entry.get(jwk_params::KID).and_then(Value::as_str) != Some(kid) {
125                continue;
126            }
127            if let Some(use_value) = entry.get(jwk_params::USE).and_then(Value::as_str)
128                && use_value != jwk_params::USE_SIG
129            {
130                continue;
131            }
132            if let Some(ops) = entry.get(jwk_params::KEY_OPS).and_then(Value::as_array)
133                && !ops.iter().any(|op| {
134                    op.as_str()
135                        .map(|s| s == jwk_params::KEY_OPS_VERIFY)
136                        .unwrap_or(false)
137                })
138            {
139                continue;
140            }
141            if let (Some(expected), Some(jwk_alg)) = (
142                algorithm,
143                entry.get(jwk_params::ALG).and_then(Value::as_str),
144            ) && jwk_alg != expected
145            {
146                continue;
147            }
148            return Some(entry.clone());
149        }
150        None
151    }
152
153    async fn try_record_force_refresh(&self) -> bool {
154        let mut last = self.last_force_refresh.lock().await;
155        let now = Instant::now();
156        match *last {
157            Some(prev) if now.saturating_duration_since(prev) < self.min_force_refresh_interval => {
158                false
159            }
160            _ => {
161                *last = Some(now);
162                true
163            }
164        }
165    }
166}
167
168// Thin wrapper around the shared `errors::auth_error` helper. Kept as a
169// function so it can be handed to `DocumentCache::with_error_factory`
170// as a `Box<dyn Fn(&str) -> AuthplaneError>`; the JWKS-specific error
171// code (`jwks_fetch_error`) is bound here.
172fn jwks_error_factory(message: &str) -> AuthplaneError {
173    crate::errors::auth_error("jwks_fetch_error", message)
174}
175
176#[cfg(test)]
177mod tests {
178    use super::*;
179    use crate::cache::document_cache::FetchResult;
180    use std::pin::Pin;
181    use std::sync::Arc as StdArc;
182    use std::sync::atomic::{AtomicUsize, Ordering};
183    use tokio::sync::Mutex as TokioMutex;
184
185    fn fixed_fetcher(responses: Vec<Value>) -> (DocumentFetcherFn, StdArc<AtomicUsize>) {
186        let counter = StdArc::new(AtomicUsize::new(0));
187        let counter_clone = counter.clone();
188        let queue = StdArc::new(TokioMutex::new(responses));
189        let fetcher: DocumentFetcherFn = StdArc::new(move || {
190            let counter = counter_clone.clone();
191            let queue = queue.clone();
192            Box::pin(async move {
193                counter.fetch_add(1, Ordering::SeqCst);
194                let mut q = queue.lock().await;
195                if q.is_empty() {
196                    return Err(AuthplaneError::Auth(AuthError {
197                        message: "exhausted".to_string(),
198                        code: "transport_error".to_string(),
199                        status_code: None,
200                    }));
201                }
202                Ok(FetchResult {
203                    document: q.remove(0),
204                    expires_at: None,
205                })
206            }) as Pin<Box<_>>
207        });
208        (fetcher, counter)
209    }
210
211    fn jwks_with(kids: &[&str]) -> Value {
212        let keys: Vec<Value> = kids
213            .iter()
214            .map(|kid| {
215                serde_json::json!({
216                    "kid": kid,
217                    "kty": "RSA",
218                    "alg": "RS256",
219                    "use": "sig",
220                    "n": "abc",
221                    "e": "AQAB",
222                })
223            })
224            .collect();
225        serde_json::json!({"keys": keys})
226    }
227
228    #[tokio::test]
229    async fn lookup_finds_existing_kid() {
230        let (fetcher, counter) = fixed_fetcher(vec![jwks_with(&["k1", "k2"])]);
231        let cache = JwksCache::new(fetcher, 600);
232        let jwk = cache
233            .get_key_by_kid("k1", Some("RS256"))
234            .await
235            .expect("ok")
236            .expect("found");
237        assert_eq!(jwk["kid"], "k1");
238        assert_eq!(counter.load(Ordering::SeqCst), 1);
239    }
240
241    #[tokio::test]
242    async fn missing_kid_triggers_force_refresh() {
243        let (fetcher, counter) = fixed_fetcher(vec![jwks_with(&["k1"]), jwks_with(&["k1", "k2"])]);
244        let cache = JwksCache::new(fetcher, 600);
245        let jwk = cache.get_key_by_kid("k2", None).await.expect("ok");
246        assert!(jwk.is_some(), "second fetch should expose k2");
247        assert_eq!(counter.load(Ordering::SeqCst), 2);
248    }
249
250    #[tokio::test]
251    async fn force_refresh_is_rate_limited() {
252        let (fetcher, counter) = fixed_fetcher(vec![
253            jwks_with(&["k1"]),
254            jwks_with(&["k1"]),
255            jwks_with(&["k1"]),
256        ]);
257        let cache =
258            JwksCache::new(fetcher, 600).with_min_force_refresh_interval(Duration::from_secs(60));
259        // First miss: triggers force refresh.
260        cache.get_key_by_kid("missing", None).await.expect("ok");
261        // Second miss within the rate-limit window: no extra fetch.
262        cache.get_key_by_kid("missing", None).await.expect("ok");
263        assert_eq!(counter.load(Ordering::SeqCst), 2); // initial get + 1 force refresh
264    }
265
266    #[tokio::test]
267    async fn algorithm_mismatch_is_skipped() {
268        let mismatched = serde_json::json!({"keys": [{
269            "kid": "k1", "kty": "RSA", "alg": "ES256", "use": "sig",
270            "n": "abc", "e": "AQAB",
271        }]});
272        let (fetcher, _counter) = fixed_fetcher(vec![mismatched.clone(), mismatched]);
273        let cache = JwksCache::new(fetcher, 600);
274        let jwk = cache.get_key_by_kid("k1", Some("RS256")).await.expect("ok");
275        assert!(jwk.is_none());
276    }
277
278    #[tokio::test]
279    async fn expire_bypasses_cached_keys_and_clears_the_force_refresh_limiter() {
280        // Models a `jwks_uri` rotation: the retired document is cached,
281        // then the fetcher is rebound and the cache expired. The next
282        // lookup must go back to the fetcher rather than answering from
283        // keys retrieved at the withdrawn URL.
284        let (fetcher, counter) = fixed_fetcher(vec![jwks_with(&["old"]), jwks_with(&["new"])]);
285        // A long TTL, so only `expire` can force the second fetch.
286        let cache = JwksCache::new(fetcher, 3600);
287        assert!(
288            cache
289                .get_key_by_kid("old", None)
290                .await
291                .expect("ok")
292                .is_some()
293        );
294        // The miss above consumed the one force-refresh the rate limiter
295        // allows; `expire` must clear it along with the document's TTL.
296        assert_eq!(counter.load(Ordering::SeqCst), 1);
297
298        cache.expire().await;
299
300        assert!(
301            cache
302                .get_key_by_kid("new", None)
303                .await
304                .expect("ok")
305                .is_some()
306        );
307        assert_eq!(counter.load(Ordering::SeqCst), 2);
308    }
309
310    #[tokio::test]
311    async fn enc_use_keys_are_skipped() {
312        let mixed = serde_json::json!({"keys": [{
313            "kid": "k1", "kty": "RSA", "alg": "RS256", "use": "enc",
314            "n": "abc", "e": "AQAB",
315        }]});
316        let (fetcher, _counter) = fixed_fetcher(vec![mixed.clone(), mixed]);
317        let cache = JwksCache::new(fetcher, 600);
318        let jwk = cache.get_key_by_kid("k1", None).await.expect("ok");
319        assert!(jwk.is_none());
320    }
321
322    #[tokio::test]
323    async fn key_ops_without_verify_is_skipped() {
324        let restricted = serde_json::json!({"keys": [{
325            "kid": "k1", "kty": "RSA", "alg": "RS256",
326            "key_ops": ["sign"],
327            "n": "abc", "e": "AQAB",
328        }]});
329        let (fetcher, _counter) = fixed_fetcher(vec![restricted.clone(), restricted]);
330        let cache = JwksCache::new(fetcher, 600);
331        let jwk = cache.get_key_by_kid("k1", None).await.expect("ok");
332        assert!(jwk.is_none());
333    }
334}