1use 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#[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 pub const DEFAULT_MIN_FORCE_REFRESH_SECONDS: u64 = 30;
39
40 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 pub fn with_min_force_refresh_interval(mut self, interval: Duration) -> Self {
60 self.min_force_refresh_interval = interval;
61 self
62 }
63
64 pub fn document_cache(&self) -> Arc<DocumentCache> {
66 self.inner.clone()
67 }
68
69 pub async fn aclose(&self) {
71 self.inner.aclose().await;
72 }
73
74 pub(crate) async fn expire(&self) {
88 self.inner.expire().await;
89 *self.last_force_refresh.lock().await = None;
90 }
91
92 pub async fn get(&self, force_refresh: bool) -> Result<Value, AuthplaneError> {
94 self.inner.get(force_refresh).await
95 }
96
97 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 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
168fn 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 cache.get_key_by_kid("missing", None).await.expect("ok");
261 cache.get_key_by_kid("missing", None).await.expect("ok");
263 assert_eq!(counter.load(Ordering::SeqCst), 2); }
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 let (fetcher, counter) = fixed_fetcher(vec![jwks_with(&["old"]), jwks_with(&["new"])]);
285 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 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}