Skip to main content

arete_auth/
multi_key.rs

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/// A key with its metadata for rotation
11#[derive(Clone)]
12pub struct RotationKey {
13    /// The verifying key
14    pub key: VerifyingKey,
15    /// Key ID for JWKS compatibility
16    pub key_id: String,
17    /// When this key was added
18    pub added_at: Instant,
19    /// Optional: when this key should be removed (for grace period rotation)
20    pub expires_at: Option<Instant>,
21    /// Whether this is the primary (current) key
22    pub is_primary: bool,
23}
24
25impl RotationKey {
26    /// Create a new primary key
27    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    /// Create a secondary (rotating out) key with expiration
38    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    /// Check if this key has expired
49    pub fn is_expired(&self) -> bool {
50        self.expires_at
51            .map(|exp| Instant::now() > exp)
52            .unwrap_or(false)
53    }
54}
55
56/// Multi-key verifier supporting graceful key rotation
57///
58/// This verifier maintains multiple keys and attempts verification with each
59/// until one succeeds. This allows zero-downtime key rotation:
60///
61/// 1. Generate new key pair
62/// 2. Add new key as primary, mark old key as secondary with grace period
63/// 3. Update JWKS to include both keys
64/// 4. After grace period, remove old key
65///
66/// # Example
67/// ```rust
68/// use arete_auth::{MultiKeyVerifier, RotationKey, SigningKey};
69/// use std::time::Duration;
70///
71/// // Generate key pairs
72/// let old_signing_key = SigningKey::generate();
73/// let old_verifying_key = old_signing_key.verifying_key();
74/// let new_signing_key = SigningKey::generate();
75/// let new_verifying_key = new_signing_key.verifying_key();
76///
77/// // Create rotation keys
78/// let old_key = RotationKey::secondary(old_verifying_key, "key-1", Duration::from_secs(86400));
79/// let new_key = RotationKey::primary(new_verifying_key, "key-2");
80///
81/// let verifier = MultiKeyVerifier::new(vec![old_key, new_key], "issuer", "audience")
82///     .with_cleanup_interval(Duration::from_secs(3600));
83/// ```
84/// Build the per-key verifier used for one rotation candidate, carrying the
85/// accepted-audience set and origin policy through unchanged.
86fn 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    /// Create a new multi-key verifier
115    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), // 1 hour default
129            last_cleanup: Arc::new(RwLock::new(Instant::now())),
130        }
131    }
132
133    /// Create from a single key (for backward compatibility)
134    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    /// Require origin validation
144    pub fn with_origin_validation(mut self) -> Self {
145        self.require_origin = true;
146        self
147    }
148
149    /// Set cleanup interval for expired keys
150    pub fn with_cleanup_interval(mut self, interval: Duration) -> Self {
151        self.cleanup_interval = interval;
152        self
153    }
154
155    /// Add a new key to the verifier
156    pub async fn add_key(&self, key: RotationKey) {
157        let mut keys = self.keys.write().await;
158
159        // If adding a primary key, demote existing primary to secondary
160        if key.is_primary {
161            for existing in keys.values_mut() {
162                if existing.is_primary {
163                    existing.is_primary = false;
164                    // Set grace period for old primary
165                    existing.expires_at = Some(Instant::now() + Duration::from_secs(86400));
166                    // 24 hours
167                }
168            }
169        }
170
171        keys.insert(key.key_id.clone(), key);
172    }
173
174    /// Remove a key by ID
175    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    /// Get all key IDs
181    pub async fn key_ids(&self) -> Vec<String> {
182        let keys = self.keys.read().await;
183        keys.keys().cloned().collect()
184    }
185
186    /// Get primary key ID
187    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    /// Clean up expired keys
195    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        // Update last cleanup time
217        let mut last = self.last_cleanup.write().await;
218        *last = Instant::now();
219    }
220
221    /// Verify a token against all keys
222    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        // Clean up expired keys periodically
229        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        // Try primary key first, then secondary keys
240        let mut key_order: Vec<&RotationKey> = keys.values().collect();
241        key_order.sort_by_key(|k| !k.is_primary); // Primary first
242
243        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                    // Wrong key, try next
261                    last_error = Some(VerifyError::InvalidSignature);
262                    continue;
263                }
264                Err(e) => {
265                    // Other errors (expired, invalid format, etc.) - don't try other keys
266                    return Err(e);
267                }
268            }
269        }
270
271        // All keys failed
272        Err(last_error.unwrap_or(VerifyError::InvalidSignature))
273    }
274
275    /// Verify without cleaning up (for high-throughput scenarios)
276    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        // Try primary key first, then secondary keys
291        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
320/// Builder for constructing a MultiKeyVerifier with rotation support
321pub 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    /// Create a new builder
331    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    /// Add a primary key
342    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    /// Add a secondary key with grace period
348    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    /// Require origin validation
360    pub fn with_origin_validation(mut self) -> Self {
361        self.require_origin = true;
362        self
363    }
364
365    /// Set cleanup interval
366    pub fn with_cleanup_interval(mut self, interval: Duration) -> Self {
367        self.cleanup_interval = interval;
368        self
369    }
370
371    /// Build the verifier
372    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        // Create old key pair
417        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        // Create new key pair
422        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        // Start with old key as primary
427        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        // Sign token with old key
431        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        // Verify old token works
439        let ctx = verifier.verify(&old_token, None, None).await.unwrap();
440        assert_eq!(ctx.subject, "subject-1");
441
442        // Rotate: add new key as primary (old key becomes secondary)
443        let new_key = RotationKey::primary(new_verifying_key, "key-new");
444        verifier.add_key(new_key).await;
445
446        // Verify old token still works (grace period)
447        let ctx = verifier.verify(&old_token, None, None).await.unwrap();
448        assert_eq!(ctx.subject, "subject-1");
449
450        // Sign and verify new token
451        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        // Check that new key is now primary
462        assert_eq!(verifier.primary_key_id().await, Some("key-new".to_string()));
463
464        // Both keys should be present
465        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        // Create two different key pairs
498        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        // Create verifier with only key2
507        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        // Should fail because token was signed with key1, verifier only has key2
522        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        // Create old key pair with specific key ID
532        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        // Create new key pair with specific key ID
538        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        // Create JWKS with both keys using their actual key IDs
544        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        // Create verifier from JWKS
567        let verifier =
568            crate::verifier::AsyncVerifier::with_jwks(jwks, "test-issuer", "test-audience");
569
570        // Sign and verify token with old key
571        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        // Old token should still verify during rotation
578        let ctx = verifier.verify(&old_token, None, None).await.unwrap();
579        assert_eq!(ctx.subject, "subject-old");
580
581        // Sign and verify token with new key
582        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        // New token should also verify
589        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        // Create a key pair
599        let signing_key = SigningKey::generate();
600        let _verifying_key = signing_key.verifying_key();
601        let signer = TokenSigner::new(signing_key, "test-issuer");
602
603        // Create JWKS with a different key (not the one used for signing)
604        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        // Should fail with key not found
628        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        // Create verifier with origin validation
655        let verifier =
656            crate::verifier::AsyncVerifier::with_jwks(jwks, "test-issuer", "test-audience")
657                .with_origin_validation();
658
659        // Token with matching origin
660        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        // Should succeed with matching origin
668        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        // Should fail with wrong origin
675        let result = verifier
676            .verify(&token, Some("https://evil.example.com"), None)
677            .await;
678        assert!(matches!(result, Err(VerifyError::OriginMismatch { .. })));
679    }
680}