1use crate::claims::AuthContext;
2use crate::error::VerifyError;
3use crate::keys::VerifyingKey;
4use crate::token::TokenVerifier;
5use std::collections::HashMap;
6use std::sync::Arc;
7use std::time::{Duration, Instant};
8use tokio::sync::RwLock;
9
10#[derive(Clone)]
12pub struct RotationKey {
13 pub key: VerifyingKey,
15 pub key_id: String,
17 pub added_at: Instant,
19 pub expires_at: Option<Instant>,
21 pub is_primary: bool,
23}
24
25impl RotationKey {
26 pub fn primary(key: VerifyingKey, key_id: impl Into<String>) -> Self {
28 Self {
29 key,
30 key_id: key_id.into(),
31 added_at: Instant::now(),
32 expires_at: None,
33 is_primary: true,
34 }
35 }
36
37 pub fn secondary(key: VerifyingKey, key_id: impl Into<String>, grace_period: Duration) -> Self {
39 Self {
40 key,
41 key_id: key_id.into(),
42 added_at: Instant::now(),
43 expires_at: Some(Instant::now() + grace_period),
44 is_primary: false,
45 }
46 }
47
48 pub fn is_expired(&self) -> bool {
50 self.expires_at
51 .map(|exp| Instant::now() > exp)
52 .unwrap_or(false)
53 }
54}
55
56fn build_inner_verifier(
87 key: crate::keys::VerifyingKey,
88 issuer: &str,
89 audiences: &crate::AudienceSet,
90 require_origin: bool,
91) -> TokenVerifier {
92 let verifier = match audiences.as_single() {
93 Some(single) => TokenVerifier::new(key, issuer, single),
94 None => TokenVerifier::with_audiences(key, issuer, audiences.iter())
95 .expect("an AudienceSet is never empty"),
96 };
97 if require_origin {
98 verifier.with_origin_validation()
99 } else {
100 verifier
101 }
102}
103
104pub struct MultiKeyVerifier {
105 keys: Arc<RwLock<HashMap<String, RotationKey>>>,
106 issuer: String,
107 audiences: crate::AudienceSet,
108 require_origin: bool,
109 cleanup_interval: Duration,
110 last_cleanup: Arc<RwLock<Instant>>,
111}
112
113impl MultiKeyVerifier {
114 pub fn new(
116 keys: Vec<RotationKey>,
117 issuer: impl Into<String>,
118 audience: impl Into<String>,
119 ) -> Self {
120 let key_map: HashMap<String, RotationKey> =
121 keys.into_iter().map(|k| (k.key_id.clone(), k)).collect();
122
123 Self {
124 keys: Arc::new(RwLock::new(key_map)),
125 issuer: issuer.into(),
126 audiences: crate::AudienceSet::single(audience),
127 require_origin: false,
128 cleanup_interval: Duration::from_secs(3600), last_cleanup: Arc::new(RwLock::new(Instant::now())),
130 }
131 }
132
133 pub fn from_single_key(
135 key: VerifyingKey,
136 key_id: impl Into<String>,
137 issuer: impl Into<String>,
138 audience: impl Into<String>,
139 ) -> Self {
140 Self::new(vec![RotationKey::primary(key, key_id)], issuer, audience)
141 }
142
143 pub fn with_origin_validation(mut self) -> Self {
145 self.require_origin = true;
146 self
147 }
148
149 pub fn with_cleanup_interval(mut self, interval: Duration) -> Self {
151 self.cleanup_interval = interval;
152 self
153 }
154
155 pub async fn add_key(&self, key: RotationKey) {
157 let mut keys = self.keys.write().await;
158
159 if key.is_primary {
161 for existing in keys.values_mut() {
162 if existing.is_primary {
163 existing.is_primary = false;
164 existing.expires_at = Some(Instant::now() + Duration::from_secs(86400));
166 }
168 }
169 }
170
171 keys.insert(key.key_id.clone(), key);
172 }
173
174 pub async fn remove_key(&self, key_id: &str) {
176 let mut keys = self.keys.write().await;
177 keys.remove(key_id);
178 }
179
180 pub async fn key_ids(&self) -> Vec<String> {
182 let keys = self.keys.read().await;
183 keys.keys().cloned().collect()
184 }
185
186 pub async fn primary_key_id(&self) -> Option<String> {
188 let keys = self.keys.read().await;
189 keys.values()
190 .find(|k| k.is_primary)
191 .map(|k| k.key_id.clone())
192 }
193
194 async fn cleanup_expired_keys(&self) {
196 let should_cleanup = {
197 let last = self.last_cleanup.read().await;
198 last.elapsed() >= self.cleanup_interval
199 };
200
201 if !should_cleanup {
202 return;
203 }
204
205 let mut keys = self.keys.write().await;
206 let expired: Vec<String> = keys
207 .iter()
208 .filter(|(_, k)| k.is_expired())
209 .map(|(id, _)| id.clone())
210 .collect();
211
212 for key_id in expired {
213 keys.remove(&key_id);
214 }
215
216 let mut last = self.last_cleanup.write().await;
218 *last = Instant::now();
219 }
220
221 pub async fn verify(
223 &self,
224 token: &str,
225 expected_origin: Option<&str>,
226 expected_client_ip: Option<&str>,
227 ) -> Result<AuthContext, VerifyError> {
228 self.cleanup_expired_keys().await;
230
231 let keys = self.keys.read().await;
232
233 if keys.is_empty() {
234 return Err(VerifyError::KeyNotFound("no keys configured".to_string()));
235 }
236
237 let mut last_error = None;
238
239 let mut key_order: Vec<&RotationKey> = keys.values().collect();
241 key_order.sort_by_key(|k| !k.is_primary); for key_entry in key_order {
244 if key_entry.is_expired() {
245 continue;
246 }
247
248 let verifier = build_inner_verifier(
249 key_entry.key.clone(),
250 &self.issuer,
251 &self.audiences,
252 self.require_origin,
253 );
254
255 match verifier.verify(token, expected_origin, expected_client_ip) {
256 Ok(ctx) => {
257 return Ok(ctx);
258 }
259 Err(VerifyError::InvalidSignature) => {
260 last_error = Some(VerifyError::InvalidSignature);
262 continue;
263 }
264 Err(e) => {
265 return Err(e);
267 }
268 }
269 }
270
271 Err(last_error.unwrap_or(VerifyError::InvalidSignature))
273 }
274
275 pub async fn verify_fast(
277 &self,
278 token: &str,
279 expected_origin: Option<&str>,
280 expected_client_ip: Option<&str>,
281 ) -> Result<AuthContext, VerifyError> {
282 let keys = self.keys.read().await;
283
284 if keys.is_empty() {
285 return Err(VerifyError::KeyNotFound("no keys configured".to_string()));
286 }
287
288 let mut last_error = None;
289
290 let mut key_order: Vec<&RotationKey> = keys.values().collect();
292 key_order.sort_by_key(|k| !k.is_primary);
293
294 for key_entry in key_order {
295 if key_entry.is_expired() {
296 continue;
297 }
298
299 let verifier = build_inner_verifier(
300 key_entry.key.clone(),
301 &self.issuer,
302 &self.audiences,
303 self.require_origin,
304 );
305
306 match verifier.verify(token, expected_origin, expected_client_ip) {
307 Ok(ctx) => return Ok(ctx),
308 Err(VerifyError::InvalidSignature) => {
309 last_error = Some(VerifyError::InvalidSignature);
310 continue;
311 }
312 Err(e) => return Err(e),
313 }
314 }
315
316 Err(last_error.unwrap_or(VerifyError::InvalidSignature))
317 }
318}
319
320pub struct MultiKeyVerifierBuilder {
322 keys: Vec<RotationKey>,
323 issuer: String,
324 audience: String,
325 require_origin: bool,
326 cleanup_interval: Duration,
327}
328
329impl MultiKeyVerifierBuilder {
330 pub fn new(issuer: impl Into<String>, audience: impl Into<String>) -> Self {
332 Self {
333 keys: Vec::new(),
334 issuer: issuer.into(),
335 audience: audience.into(),
336 require_origin: false,
337 cleanup_interval: Duration::from_secs(3600),
338 }
339 }
340
341 pub fn with_primary_key(mut self, key: VerifyingKey, key_id: impl Into<String>) -> Self {
343 self.keys.push(RotationKey::primary(key, key_id));
344 self
345 }
346
347 pub fn with_secondary_key(
349 mut self,
350 key: VerifyingKey,
351 key_id: impl Into<String>,
352 grace_period: Duration,
353 ) -> Self {
354 self.keys
355 .push(RotationKey::secondary(key, key_id, grace_period));
356 self
357 }
358
359 pub fn with_origin_validation(mut self) -> Self {
361 self.require_origin = true;
362 self
363 }
364
365 pub fn with_cleanup_interval(mut self, interval: Duration) -> Self {
367 self.cleanup_interval = interval;
368 self
369 }
370
371 pub fn build(self) -> MultiKeyVerifier {
373 let mut verifier = MultiKeyVerifier::new(self.keys, self.issuer, self.audience);
374 if self.require_origin {
375 verifier = verifier.with_origin_validation();
376 }
377 verifier.with_cleanup_interval(self.cleanup_interval)
378 }
379}
380
381#[cfg(test)]
382mod tests {
383 use super::*;
384 use crate::claims::{KeyClass, SessionClaims};
385 use crate::keys::SigningKey;
386 use crate::token::TokenSigner;
387
388 #[tokio::test]
389 async fn test_multi_key_verifier_single_key() {
390 let signing_key = SigningKey::generate();
391 let verifying_key = signing_key.verifying_key();
392
393 let signer = TokenSigner::new(signing_key, "test-issuer");
394 let verifier = MultiKeyVerifier::from_single_key(
395 verifying_key,
396 "key-1",
397 "test-issuer",
398 "test-audience",
399 );
400
401 let claims = SessionClaims::builder("test-issuer", "test-subject", "test-audience")
402 .with_scope("read")
403 .with_metering_key("meter-123")
404 .with_key_class(KeyClass::Publishable)
405 .build();
406
407 let token = signer.sign(claims).unwrap();
408 let context = verifier.verify(&token, None, None).await.unwrap();
409
410 assert_eq!(context.subject, "test-subject");
411 assert_eq!(verifier.primary_key_id().await, Some("key-1".to_string()));
412 }
413
414 #[tokio::test]
415 async fn test_key_rotation() {
416 let old_signing_key = SigningKey::generate();
418 let old_verifying_key = old_signing_key.verifying_key();
419 let old_signer = TokenSigner::new(old_signing_key, "test-issuer");
420
421 let new_signing_key = SigningKey::generate();
423 let new_verifying_key = new_signing_key.verifying_key();
424 let new_signer = TokenSigner::new(new_signing_key, "test-issuer");
425
426 let old_key = RotationKey::primary(old_verifying_key.clone(), "key-old");
428 let verifier = MultiKeyVerifier::new(vec![old_key], "test-issuer", "test-audience");
429
430 let old_claims = SessionClaims::builder("test-issuer", "subject-1", "test-audience")
432 .with_scope("read")
433 .with_metering_key("meter-1")
434 .with_key_class(KeyClass::Publishable)
435 .build();
436 let old_token = old_signer.sign(old_claims).unwrap();
437
438 let ctx = verifier.verify(&old_token, None, None).await.unwrap();
440 assert_eq!(ctx.subject, "subject-1");
441
442 let new_key = RotationKey::primary(new_verifying_key, "key-new");
444 verifier.add_key(new_key).await;
445
446 let ctx = verifier.verify(&old_token, None, None).await.unwrap();
448 assert_eq!(ctx.subject, "subject-1");
449
450 let new_claims = SessionClaims::builder("test-issuer", "subject-2", "test-audience")
452 .with_scope("read")
453 .with_metering_key("meter-2")
454 .with_key_class(KeyClass::Publishable)
455 .build();
456 let new_token = new_signer.sign(new_claims).unwrap();
457
458 let ctx = verifier.verify(&new_token, None, None).await.unwrap();
459 assert_eq!(ctx.subject, "subject-2");
460
461 assert_eq!(verifier.primary_key_id().await, Some("key-new".to_string()));
463
464 let key_ids = verifier.key_ids().await;
466 assert!(key_ids.contains(&"key-old".to_string()));
467 assert!(key_ids.contains(&"key-new".to_string()));
468 }
469
470 #[tokio::test]
471 async fn test_verifier_builder() {
472 let signing_key = SigningKey::generate();
473 let verifying_key = signing_key.verifying_key();
474
475 let verifier = MultiKeyVerifierBuilder::new("test-issuer", "test-audience")
476 .with_primary_key(verifying_key, "key-1")
477 .with_origin_validation()
478 .build();
479
480 let signer = TokenSigner::new(signing_key, "test-issuer");
481 let claims = SessionClaims::builder("test-issuer", "test-subject", "test-audience")
482 .with_scope("read")
483 .with_origin("https://trusted.example.com")
484 .with_key_class(KeyClass::Secret)
485 .build();
486
487 let token = signer.sign(claims).unwrap();
488 let ctx = verifier
489 .verify(&token, Some("https://trusted.example.com"), None)
490 .await
491 .unwrap();
492 assert_eq!(ctx.subject, "test-subject");
493 }
494
495 #[tokio::test]
496 async fn test_invalid_signature_with_multiple_keys() {
497 let key1_signing = SigningKey::generate();
499 let key1_verifying = key1_signing.verifying_key();
500
501 let key2_signing = SigningKey::generate();
502 let _key2_verifying = key2_signing.verifying_key();
503
504 let signer = TokenSigner::new(key1_signing, "test-issuer");
505
506 let verifier = MultiKeyVerifier::from_single_key(
508 key2_signing.verifying_key(),
509 "key-2",
510 "test-issuer",
511 "test-audience",
512 );
513
514 let claims = SessionClaims::builder("test-issuer", "test-subject", "test-audience")
515 .with_scope("read")
516 .with_key_class(KeyClass::Publishable)
517 .build();
518
519 let token = signer.sign(claims).unwrap();
520
521 let result = verifier.verify(&token, None, None).await;
523 assert!(matches!(result, Err(VerifyError::InvalidSignature)));
524 }
525
526 #[tokio::test]
527 async fn test_jwks_key_rotation_grace_period() {
528 use crate::token::{Jwk, Jwks};
529 use base64::Engine;
530
531 let old_signing_key = SigningKey::generate();
533 let old_verifying_key = old_signing_key.verifying_key();
534 let old_kid = old_verifying_key.key_id();
535 let old_signer = TokenSigner::new(old_signing_key, "test-issuer");
536
537 let new_signing_key = SigningKey::generate();
539 let new_verifying_key = new_signing_key.verifying_key();
540 let new_kid = new_verifying_key.key_id();
541 let new_signer = TokenSigner::new(new_signing_key, "test-issuer");
542
543 let old_key_b64 =
545 base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(old_verifying_key.to_bytes());
546 let new_key_b64 =
547 base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(new_verifying_key.to_bytes());
548
549 let jwks = Jwks {
550 keys: vec![
551 Jwk {
552 kty: "OKP".to_string(),
553 use_: Some("sig".to_string()),
554 kid: old_kid,
555 x: old_key_b64,
556 },
557 Jwk {
558 kty: "OKP".to_string(),
559 use_: Some("sig".to_string()),
560 kid: new_kid,
561 x: new_key_b64,
562 },
563 ],
564 };
565
566 let verifier =
568 crate::verifier::AsyncVerifier::with_jwks(jwks, "test-issuer", "test-audience");
569
570 let old_claims = SessionClaims::builder("test-issuer", "subject-old", "test-audience")
572 .with_scope("read")
573 .with_key_class(KeyClass::Secret)
574 .build();
575 let old_token = old_signer.sign(old_claims).unwrap();
576
577 let ctx = verifier.verify(&old_token, None, None).await.unwrap();
579 assert_eq!(ctx.subject, "subject-old");
580
581 let new_claims = SessionClaims::builder("test-issuer", "subject-new", "test-audience")
583 .with_scope("read")
584 .with_key_class(KeyClass::Secret)
585 .build();
586 let new_token = new_signer.sign(new_claims).unwrap();
587
588 let ctx = verifier.verify(&new_token, None, None).await.unwrap();
590 assert_eq!(ctx.subject, "subject-new");
591 }
592
593 #[tokio::test]
594 async fn test_jwks_key_not_found() {
595 use crate::token::{Jwk, Jwks};
596 use base64::Engine;
597
598 let signing_key = SigningKey::generate();
600 let _verifying_key = signing_key.verifying_key();
601 let signer = TokenSigner::new(signing_key, "test-issuer");
602
603 let different_key = SigningKey::generate();
605 let different_verifying_key = different_key.verifying_key();
606 let different_key_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD
607 .encode(different_verifying_key.to_bytes());
608
609 let jwks = Jwks {
610 keys: vec![Jwk {
611 kty: "OKP".to_string(),
612 use_: Some("sig".to_string()),
613 kid: "different-key".to_string(),
614 x: different_key_b64,
615 }],
616 };
617
618 let verifier =
619 crate::verifier::AsyncVerifier::with_jwks(jwks, "test-issuer", "test-audience");
620
621 let claims = SessionClaims::builder("test-issuer", "test-subject", "test-audience")
622 .with_scope("read")
623 .with_key_class(KeyClass::Secret)
624 .build();
625 let token = signer.sign(claims).unwrap();
626
627 let result = verifier.verify(&token, None, None).await;
629 assert!(matches!(result, Err(VerifyError::KeyNotFound(_))));
630 }
631
632 #[tokio::test]
633 async fn test_jwks_with_origin_validation() {
634 use crate::token::{Jwk, Jwks};
635 use base64::Engine;
636
637 let signing_key = SigningKey::generate();
638 let verifying_key = signing_key.verifying_key();
639 let kid = verifying_key.key_id();
640 let signer = TokenSigner::new(signing_key, "test-issuer");
641
642 let key_b64 =
643 base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(verifying_key.to_bytes());
644
645 let jwks = Jwks {
646 keys: vec![Jwk {
647 kty: "OKP".to_string(),
648 use_: Some("sig".to_string()),
649 kid,
650 x: key_b64,
651 }],
652 };
653
654 let verifier =
656 crate::verifier::AsyncVerifier::with_jwks(jwks, "test-issuer", "test-audience")
657 .with_origin_validation();
658
659 let claims = SessionClaims::builder("test-issuer", "test-subject", "test-audience")
661 .with_scope("read")
662 .with_key_class(KeyClass::Secret)
663 .with_origin("https://trusted.example.com")
664 .build();
665 let token = signer.sign(claims).unwrap();
666
667 let ctx = verifier
669 .verify(&token, Some("https://trusted.example.com"), None)
670 .await
671 .unwrap();
672 assert_eq!(ctx.subject, "test-subject");
673
674 let result = verifier
676 .verify(&token, Some("https://evil.example.com"), None)
677 .await;
678 assert!(matches!(result, Err(VerifyError::OriginMismatch { .. })));
679 }
680}