Skip to main content

arete_auth/
token.rs

1use crate::claims::{AuthContext, SessionClaims};
2use crate::error::VerifyError;
3use crate::keys::{SigningKey, VerifyingKey};
4use base64::Engine;
5use serde::{Deserialize, Serialize};
6use serde_json;
7
8/// JWT Header for EdDSA (Ed25519) tokens
9#[derive(Debug, Clone, Serialize, Deserialize)]
10struct JwtHeader {
11    alg: String,
12    typ: String,
13    #[serde(skip_serializing_if = "Option::is_none")]
14    kid: Option<String>,
15}
16
17impl Default for JwtHeader {
18    fn default() -> Self {
19        Self {
20            alg: "EdDSA".to_string(),
21            typ: "JWT".to_string(),
22            kid: None,
23        }
24    }
25}
26
27/// Token signer for issuing session tokens using Ed25519 (EdDSA)
28pub struct TokenSigner {
29    signing_key: SigningKey,
30    issuer: String,
31}
32
33impl TokenSigner {
34    /// Create a new token signer with an Ed25519 signing key
35    ///
36    /// Uses EdDSA (Ed25519) for asymmetric signing. This is the recommended
37    /// algorithm for production use as it provides better security than HMAC.
38    pub fn new(signing_key: SigningKey, issuer: impl Into<String>) -> Self {
39        Self {
40            signing_key,
41            issuer: issuer.into(),
42        }
43    }
44
45    /// Sign a session token using Ed25519
46    pub fn sign(&self, claims: SessionClaims) -> Result<String, TokenError> {
47        // Create header with key ID
48        let header = JwtHeader {
49            kid: Some(self.signing_key.key_id()),
50            ..Default::default()
51        };
52
53        // Encode header
54        let header_json = serde_json::to_string(&header)?;
55        let header_b64 =
56            base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(header_json.as_bytes());
57
58        // Encode claims
59        let claims_json = serde_json::to_string(&claims)?;
60        let claims_b64 =
61            base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(claims_json.as_bytes());
62
63        // Create message to sign
64        let message = format!("{}.{}", header_b64, claims_b64);
65
66        // Sign with Ed25519
67        let signature = self.signing_key.sign(message.as_bytes());
68        let signature_b64 =
69            base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(signature.to_bytes());
70
71        // Combine into JWT
72        Ok(format!("{}.{}.{}", header_b64, claims_b64, signature_b64))
73    }
74
75    /// Get the issuer
76    pub fn issuer(&self) -> &str {
77        &self.issuer
78    }
79}
80
81/// Token error type
82#[derive(Debug)]
83pub enum TokenError {
84    Serialization(serde_json::Error),
85    Base64(base64::DecodeError),
86    InvalidFormat(String),
87}
88
89impl std::fmt::Display for TokenError {
90    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
91        match self {
92            TokenError::Serialization(e) => write!(f, "Serialization error: {}", e),
93            TokenError::Base64(e) => write!(f, "Base64 error: {}", e),
94            TokenError::InvalidFormat(s) => write!(f, "Invalid format: {}", s),
95        }
96    }
97}
98
99impl std::error::Error for TokenError {}
100
101impl From<serde_json::Error> for TokenError {
102    fn from(e: serde_json::Error) -> Self {
103        TokenError::Serialization(e)
104    }
105}
106
107impl From<base64::DecodeError> for TokenError {
108    fn from(e: base64::DecodeError) -> Self {
109        TokenError::Base64(e)
110    }
111}
112
113/// Token verifier for validating session tokens using Ed25519 (EdDSA)
114pub struct TokenVerifier {
115    verifying_key: VerifyingKey,
116    issuer: String,
117    audience: String,
118    require_origin: bool,
119    require_client_ip: bool,
120}
121
122impl TokenVerifier {
123    /// Create a new token verifier with an Ed25519 verifying key
124    ///
125    /// Uses EdDSA (Ed25519) for asymmetric signature verification.
126    /// This is the recommended algorithm for production use.
127    pub fn new(
128        verifying_key: VerifyingKey,
129        issuer: impl Into<String>,
130        audience: impl Into<String>,
131    ) -> Self {
132        Self {
133            verifying_key,
134            issuer: issuer.into(),
135            audience: audience.into(),
136            require_origin: false,
137            require_client_ip: false,
138        }
139    }
140
141    /// Require origin validation
142    pub fn with_origin_validation(mut self) -> Self {
143        self.require_origin = true;
144        self
145    }
146
147    /// Require client IP validation
148    pub fn with_client_ip_validation(mut self) -> Self {
149        self.require_client_ip = true;
150        self
151    }
152
153    /// Verify a token and return the auth context
154    ///
155    /// # Arguments
156    /// * `token` - The JWT token to verify
157    /// * `expected_origin` - Optional expected origin for origin validation
158    /// * `expected_client_ip` - Optional expected client IP for IP binding validation
159    pub fn verify(
160        &self,
161        token: &str,
162        expected_origin: Option<&str>,
163        expected_client_ip: Option<&str>,
164    ) -> Result<AuthContext, VerifyError> {
165        // Split token into parts
166        let parts: Vec<&str> = token.split('.').collect();
167        if parts.len() != 3 {
168            return Err(VerifyError::InvalidFormat("Invalid JWT format".to_string()));
169        }
170
171        let header_b64 = parts[0];
172        let claims_b64 = parts[1];
173        let signature_b64 = parts[2];
174
175        // Decode and verify header
176        let header_json = base64::engine::general_purpose::URL_SAFE_NO_PAD
177            .decode(header_b64)
178            .map_err(|e| VerifyError::InvalidFormat(format!("Invalid header base64: {}", e)))?;
179        let header: JwtHeader = serde_json::from_slice(&header_json)
180            .map_err(|e| VerifyError::InvalidFormat(format!("Invalid header JSON: {}", e)))?;
181
182        if header.alg != "EdDSA" {
183            return Err(VerifyError::InvalidFormat(format!(
184                "Unsupported algorithm: {}",
185                header.alg
186            )));
187        }
188
189        // Decode claims
190        let claims_json = base64::engine::general_purpose::URL_SAFE_NO_PAD
191            .decode(claims_b64)
192            .map_err(|e| VerifyError::InvalidFormat(format!("Invalid claims base64: {}", e)))?;
193        let claims: SessionClaims = serde_json::from_slice(&claims_json)
194            .map_err(|e| VerifyError::InvalidFormat(format!("Invalid claims JSON: {}", e)))?;
195
196        // Decode signature
197        let signature_bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD
198            .decode(signature_b64)
199            .map_err(|e| VerifyError::InvalidFormat(format!("Invalid signature base64: {}", e)))?;
200        if signature_bytes.len() != 64 {
201            return Err(VerifyError::InvalidFormat(
202                "Invalid signature length".to_string(),
203            ));
204        }
205        let signature = ed25519_dalek::Signature::from_bytes(&signature_bytes.try_into().unwrap());
206
207        // Verify signature
208        let message = format!("{}.{}", header_b64, claims_b64);
209        self.verifying_key
210            .verify(message.as_bytes(), &signature)
211            .map_err(|_| VerifyError::InvalidSignature)?;
212
213        // Check issuer
214        if claims.iss != self.issuer {
215            return Err(VerifyError::InvalidIssuer);
216        }
217
218        // Check audience
219        if claims.aud != self.audience {
220            return Err(VerifyError::InvalidAudience);
221        }
222
223        // Check expiration
224        use std::time::{SystemTime, UNIX_EPOCH};
225        let now = SystemTime::now()
226            .duration_since(UNIX_EPOCH)
227            .expect("time should not be before epoch")
228            .as_secs();
229
230        if claims.exp <= now {
231            return Err(VerifyError::Expired);
232        }
233
234        if claims.nbf > now || claims.iat > now {
235            return Err(VerifyError::NotYetValid);
236        }
237
238        // Validate origin if required or if token has origin binding
239        let token_has_origin = claims.origin.is_some();
240        let origin_provided = expected_origin.is_some();
241
242        if token_has_origin && origin_provided {
243            // Token is origin-bound and origin was provided - validate they match
244            let expected = expected_origin.unwrap();
245            let actual = claims.origin.as_ref().unwrap();
246
247            if actual != expected {
248                return Err(VerifyError::OriginMismatch {
249                    expected: expected.to_string(),
250                    actual: actual.clone(),
251                });
252            }
253        } else if token_has_origin && self.require_origin {
254            // Token has origin but none was provided, and origin is required
255            return Err(VerifyError::OriginRequired {
256                token_origin: claims.origin.as_ref().unwrap().clone(),
257            });
258        } else if !token_has_origin && self.require_origin {
259            // Verifier requires origin but token doesn't have one bound
260            return Err(VerifyError::MissingClaim("origin".to_string()));
261        }
262        // If token has origin but none provided, and origin is NOT required,
263        // we allow the connection (defense-in-depth is optional)
264
265        // Validate client IP if required
266        if self.require_client_ip {
267            if let Some(expected) = expected_client_ip {
268                match &claims.client_ip {
269                    Some(actual) if actual == expected => {}
270                    Some(actual) => {
271                        return Err(VerifyError::OriginMismatch {
272                            expected: expected.to_string(),
273                            actual: actual.clone(),
274                        });
275                    }
276                    None => {
277                        return Err(VerifyError::MissingClaim("client_ip".to_string()));
278                    }
279                }
280            } else if claims.client_ip.is_none() {
281                return Err(VerifyError::MissingClaim("client_ip".to_string()));
282            }
283        }
284
285        // Validate the v2 policy identity set; old tokens with none of the
286        // new fields pass unchanged.
287        claims
288            .validate_policy_claims()
289            .map_err(|error| VerifyError::InvalidPolicyClaims(error.to_string()))?;
290
291        Ok(AuthContext::from_claims(claims))
292    }
293
294    /// Get the expected issuer
295    pub fn issuer(&self) -> &str {
296        &self.issuer
297    }
298
299    /// Get the expected audience
300    pub fn audience(&self) -> &str {
301        &self.audience
302    }
303}
304
305/// JWKS structure for key rotation
306#[derive(Debug, Clone, Deserialize)]
307pub struct Jwks {
308    pub keys: Vec<Jwk>,
309}
310
311#[derive(Debug, Clone, Deserialize)]
312pub struct Jwk {
313    pub kty: String,
314    #[serde(rename = "use")]
315    pub use_: Option<String>,
316    pub kid: String,
317    pub x: String, // Base64-encoded public key
318}
319
320/// Token verifier with JWKS support for key rotation
321#[derive(Clone)]
322pub struct JwksVerifier {
323    jwks: Jwks,
324    issuer: String,
325    audience: String,
326    require_origin: bool,
327}
328
329impl JwksVerifier {
330    /// Create a new JWKS verifier
331    pub fn new(jwks: Jwks, issuer: impl Into<String>, audience: impl Into<String>) -> Self {
332        Self {
333            jwks,
334            issuer: issuer.into(),
335            audience: audience.into(),
336            require_origin: false,
337        }
338    }
339
340    /// Require origin validation
341    pub fn with_origin_validation(mut self) -> Self {
342        self.require_origin = true;
343        self
344    }
345
346    /// Verify a token using the appropriate key from JWKS
347    pub fn verify(
348        &self,
349        token: &str,
350        expected_origin: Option<&str>,
351        expected_client_ip: Option<&str>,
352    ) -> Result<AuthContext, VerifyError> {
353        // Decode header to get kid
354        let parts: Vec<&str> = token.split('.').collect();
355        if parts.len() != 3 {
356            return Err(VerifyError::InvalidFormat("Invalid JWT format".to_string()));
357        }
358
359        let header_json = base64::engine::general_purpose::URL_SAFE_NO_PAD
360            .decode(parts[0])
361            .map_err(|e| VerifyError::InvalidFormat(format!("Invalid header: {}", e)))?;
362        let header: JwtHeader = serde_json::from_slice(&header_json)
363            .map_err(|e| VerifyError::InvalidFormat(format!("Invalid header JSON: {}", e)))?;
364
365        let kid = header
366            .kid
367            .ok_or_else(|| VerifyError::MissingClaim("kid".to_string()))?;
368
369        // Find the key
370        let jwk = self
371            .jwks
372            .keys
373            .iter()
374            .find(|k| k.kid == kid)
375            .ok_or(VerifyError::KeyNotFound(kid))?;
376
377        // Decode the public key from hex (first 16 chars of hex = 8 bytes of key id)
378        // Actually, we need to decode the full public key from the JWKS
379        // The JWKS should contain the full base64-encoded public key
380        let public_key_bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD
381            .decode(&jwk.x)
382            .map_err(|_| VerifyError::InvalidFormat("Invalid public key base64".to_string()))?;
383
384        let public_key: [u8; 32] = public_key_bytes
385            .try_into()
386            .map_err(|_| VerifyError::InvalidFormat("Invalid key length".to_string()))?;
387
388        // Create verifier for this key
389        let verifying_key = VerifyingKey::from_bytes(&public_key)
390            .map_err(|e| VerifyError::InvalidFormat(e.to_string()))?;
391
392        let verifier = if self.require_origin {
393            TokenVerifier::new(verifying_key, &self.issuer, &self.audience).with_origin_validation()
394        } else {
395            TokenVerifier::new(verifying_key, &self.issuer, &self.audience)
396        };
397
398        verifier.verify(token, expected_origin, expected_client_ip)
399    }
400
401    /// Fetch JWKS from a URL
402    #[cfg(feature = "jwks")]
403    pub async fn fetch_jwks(url: &str) -> Result<Jwks, reqwest::Error> {
404        let response = reqwest::get(url).await?;
405        let jwks: Jwks = response.json().await?;
406        Ok(jwks)
407    }
408}
409
410#[cfg(test)]
411/// HMAC-based verifier for tests only
412pub struct HmacVerifier {
413    _secret: Vec<u8>,
414    _issuer: String,
415    _audience: String,
416}
417
418#[cfg(test)]
419impl HmacVerifier {
420    /// Create a new HMAC verifier (dev only)
421    pub fn new(
422        secret: impl Into<Vec<u8>>,
423        issuer: impl Into<String>,
424        audience: impl Into<String>,
425    ) -> Self {
426        Self {
427            _secret: secret.into(),
428            _issuer: issuer.into(),
429            _audience: audience.into(),
430        }
431    }
432
433    /// Verify a token using HMAC
434    pub fn verify(
435        &self,
436        token: &str,
437        _expected_origin: Option<&str>,
438    ) -> Result<AuthContext, VerifyError> {
439        // Split token
440        let parts: Vec<&str> = token.split('.').collect();
441        if parts.len() != 3 {
442            return Err(VerifyError::InvalidFormat("Invalid JWT format".to_string()));
443        }
444
445        // For HMAC, we'd need to verify the HMAC signature
446        // This is a simplified implementation - in practice you'd use hmac-sha256
447        // For now, just decode the claims without verification (dev only!)
448        let claims_json = base64::engine::general_purpose::URL_SAFE_NO_PAD
449            .decode(parts[1])
450            .map_err(|e| VerifyError::InvalidFormat(format!("Invalid claims: {}", e)))?;
451        let claims: SessionClaims = serde_json::from_slice(&claims_json)
452            .map_err(|e| VerifyError::InvalidFormat(format!("Invalid claims JSON: {}", e)))?;
453
454        claims
455            .validate_policy_claims()
456            .map_err(|error| VerifyError::InvalidPolicyClaims(error.to_string()))?;
457
458        Ok(AuthContext::from_claims(claims))
459    }
460}
461
462#[cfg(test)]
463mod tests {
464    use super::*;
465    use crate::claims::{KeyClass, Limits};
466
467    fn create_test_claims() -> SessionClaims {
468        SessionClaims::builder("test-issuer", "test-subject", "test-audience")
469            .with_ttl(300)
470            .with_scope("read")
471            .with_metering_key("meter-123")
472            .with_key_class(KeyClass::Publishable)
473            .with_limits(Limits {
474                max_connections: Some(10),
475                max_subscriptions: Some(100),
476                max_snapshot_rows: Some(1000),
477                max_messages_per_minute: Some(1000),
478                max_bytes_per_minute: Some(10_000_000),
479                max_http_requests_per_minute: Some(300),
480                max_http_batch_addresses: Some(100),
481                ..Limits::default()
482            })
483            .build()
484    }
485
486    #[test]
487    fn test_sign_and_verify() {
488        // Generate keys
489        let signing_key = crate::keys::SigningKey::generate();
490        let verifying_key = signing_key.verifying_key();
491
492        // Create signer and verifier
493        let signer = TokenSigner::new(signing_key, "test-issuer");
494        let verifier = TokenVerifier::new(verifying_key, "test-issuer", "test-audience");
495
496        // Sign token
497        let claims = create_test_claims();
498        let token = signer.sign(claims.clone()).unwrap();
499
500        // Verify token
501        let context = verifier.verify(&token, None, None).unwrap();
502
503        assert_eq!(context.subject, "test-subject");
504        assert_eq!(context.issuer, "test-issuer");
505        assert_eq!(context.metering_key, "meter-123");
506    }
507
508    #[test]
509    fn legacy_claims_without_typed_target_deserialize_and_verify() {
510        use std::time::{SystemTime, UNIX_EPOCH};
511
512        let now = SystemTime::now()
513            .duration_since(UNIX_EPOCH)
514            .unwrap()
515            .as_secs();
516        let claims: SessionClaims = serde_json::from_value(serde_json::json!({
517            "iss": "test-issuer",
518            "sub": "legacy-subject",
519            "aud": "deployment-1",
520            "iat": now,
521            "nbf": now,
522            "exp": now + 300,
523            "jti": "legacy-jti",
524            "scope": "read",
525            "metering_key": "api_key:1",
526            "deployment_id": "deployment-1",
527            "key_class": "publishable"
528        }))
529        .unwrap();
530        assert_eq!(claims.target_kind, None);
531        assert_eq!(claims.target_id, None);
532        assert_eq!(claims.program_id, None);
533        assert_eq!(claims.program_release_hash, None);
534
535        let signing_key = crate::keys::SigningKey::generate();
536        let verifying_key = signing_key.verifying_key();
537        let token = TokenSigner::new(signing_key, "test-issuer")
538            .sign(claims)
539            .unwrap();
540        let context = TokenVerifier::new(verifying_key, "test-issuer", "deployment-1")
541            .verify(&token, None, None)
542            .unwrap();
543
544        assert_eq!(context.subject, "legacy-subject");
545        assert_eq!(context.deployment_id.as_deref(), Some("deployment-1"));
546        assert_eq!(context.target_kind, None);
547    }
548
549    #[test]
550    fn verifier_rejects_partial_v2_policy_claims_and_accepts_full_sets() {
551        let signing_key = crate::keys::SigningKey::generate();
552        let verifying_key = signing_key.verifying_key();
553        let signer = TokenSigner::new(signing_key, "test-issuer");
554        let verifier = TokenVerifier::new(verifying_key, "test-issuer", "test-audience");
555
556        // A signed token with only a subset of the v2 identity fields fails
557        // verification at the authorization boundary.
558        let partial = SessionClaims::builder("test-issuer", "test-subject", "test-audience")
559            .with_metering_key("account:42")
560            .with_account_key("account:42")
561            .build();
562        let token = signer.sign(partial).unwrap();
563        assert!(matches!(
564            verifier.verify(&token, None, None),
565            Err(VerifyError::InvalidPolicyClaims(_))
566        ));
567
568        // The complete authenticated tuple verifies and resolves.
569        let full = SessionClaims::builder("test-issuer", "test-subject", "test-audience")
570            .with_metering_key("account:42")
571            .with_plan("pro")
572            .with_actor_key("user:1")
573            .with_account_key("account:42")
574            .with_consumer_key("consumer:abc123")
575            .with_policy_version(2)
576            .with_account_limits(Limits::default())
577            .build();
578        let token = signer.sign(full).unwrap();
579        let context = verifier.verify(&token, None, None).unwrap();
580        assert!(!context.is_legacy_policy());
581        assert_eq!(context.account_key(), "account:42");
582        assert_eq!(context.consumer_key(), "consumer:abc123");
583        assert_eq!(context.policy_version, Some(2));
584    }
585
586    #[test]
587    fn test_expired_token() {
588        let signing_key = crate::keys::SigningKey::generate();
589        let verifying_key = signing_key.verifying_key();
590
591        let signer = TokenSigner::new(signing_key, "test-issuer");
592        let verifier = TokenVerifier::new(verifying_key, "test-issuer", "test-audience");
593
594        // Create expired claims
595        let claims = SessionClaims::builder("test-issuer", "test-subject", "test-audience")
596            .with_ttl(0) // Already expired
597            .with_scope("read")
598            .with_metering_key("meter-123")
599            .with_key_class(KeyClass::Publishable)
600            .build();
601
602        let token = signer.sign(claims).unwrap();
603
604        // Should fail with expired error
605        let result = verifier.verify(&token, None, None);
606        assert!(matches!(result, Err(VerifyError::Expired)));
607    }
608
609    #[test]
610    fn test_future_issued_token_is_not_yet_valid() {
611        let signing_key = crate::keys::SigningKey::generate();
612        let verifying_key = signing_key.verifying_key();
613        let signer = TokenSigner::new(signing_key, "test-issuer");
614        let verifier = TokenVerifier::new(verifying_key, "test-issuer", "test-audience");
615        let mut claims = create_test_claims();
616        claims.iat += 300;
617
618        let token = signer.sign(claims).unwrap();
619
620        assert!(matches!(
621            verifier.verify(&token, None, None),
622            Err(VerifyError::NotYetValid)
623        ));
624    }
625
626    #[test]
627    fn test_invalid_signature() {
628        let signing_key = crate::keys::SigningKey::generate();
629        let wrong_signing_key = crate::keys::SigningKey::generate();
630        let wrong_verifying_key = wrong_signing_key.verifying_key();
631
632        let signer = TokenSigner::new(signing_key, "test-issuer");
633        let verifier = TokenVerifier::new(wrong_verifying_key, "test-issuer", "test-audience");
634
635        let claims = create_test_claims();
636        let token = signer.sign(claims).unwrap();
637
638        // Should fail with invalid signature
639        let result = verifier.verify(&token, None, None);
640        assert!(matches!(result, Err(VerifyError::InvalidSignature)));
641    }
642
643    #[test]
644    fn test_wrong_issuer() {
645        let signing_key = crate::keys::SigningKey::generate();
646        let verifying_key = signing_key.verifying_key();
647
648        let signer = TokenSigner::new(signing_key, "wrong-issuer");
649        let verifier = TokenVerifier::new(verifying_key, "test-issuer", "test-audience");
650
651        // Create claims with the wrong issuer
652        let claims = SessionClaims::builder("wrong-issuer", "test-subject", "test-audience")
653            .with_ttl(300)
654            .with_scope("read")
655            .with_metering_key("meter-123")
656            .with_key_class(KeyClass::Publishable)
657            .build();
658        let token = signer.sign(claims).unwrap();
659
660        // Should fail with invalid issuer
661        let result = verifier.verify(&token, None, None);
662        assert!(matches!(result, Err(VerifyError::InvalidIssuer)));
663    }
664
665    #[test]
666    fn test_wrong_audience() {
667        let signing_key = crate::keys::SigningKey::generate();
668        let verifying_key = signing_key.verifying_key();
669
670        let signer = TokenSigner::new(signing_key, "test-issuer");
671        let verifier = TokenVerifier::new(verifying_key, "test-issuer", "expected-audience");
672
673        let claims = SessionClaims::builder("test-issuer", "test-subject", "wrong-audience")
674            .with_ttl(300)
675            .with_scope("read")
676            .with_metering_key("meter-123")
677            .with_key_class(KeyClass::Publishable)
678            .build();
679        let token = signer.sign(claims).unwrap();
680
681        let result = verifier.verify(&token, None, None);
682        assert!(matches!(result, Err(VerifyError::InvalidAudience)));
683    }
684
685    #[test]
686    fn test_origin_mismatch() {
687        let signing_key = crate::keys::SigningKey::generate();
688        let verifying_key = signing_key.verifying_key();
689
690        let signer = TokenSigner::new(signing_key, "test-issuer");
691        let verifier = TokenVerifier::new(verifying_key, "test-issuer", "test-audience")
692            .with_origin_validation();
693
694        let claims = SessionClaims::builder("test-issuer", "test-subject", "test-audience")
695            .with_ttl(300)
696            .with_scope("read")
697            .with_metering_key("meter-123")
698            .with_origin("https://allowed.example")
699            .with_key_class(KeyClass::Publishable)
700            .build();
701        let token = signer.sign(claims).unwrap();
702
703        let result = verifier.verify(&token, Some("https://other.example"), None);
704        assert!(matches!(result, Err(VerifyError::OriginMismatch { .. })));
705    }
706
707    #[test]
708    fn test_origin_validation_success() {
709        let signing_key = crate::keys::SigningKey::generate();
710        let verifying_key = signing_key.verifying_key();
711
712        let signer = TokenSigner::new(signing_key, "test-issuer");
713        let verifier = TokenVerifier::new(verifying_key, "test-issuer", "test-audience")
714            .with_origin_validation();
715
716        let claims = SessionClaims::builder("test-issuer", "test-subject", "test-audience")
717            .with_ttl(300)
718            .with_scope("read")
719            .with_metering_key("meter-123")
720            .with_origin("https://allowed.example")
721            .with_key_class(KeyClass::Publishable)
722            .build();
723        let token = signer.sign(claims).unwrap();
724
725        let context = verifier
726            .verify(&token, Some("https://allowed.example"), None)
727            .unwrap();
728        assert_eq!(context.origin.as_deref(), Some("https://allowed.example"));
729    }
730
731    #[test]
732    fn test_origin_validation_requires_origin_claim() {
733        let signing_key = crate::keys::SigningKey::generate();
734        let verifying_key = signing_key.verifying_key();
735
736        let signer = TokenSigner::new(signing_key, "test-issuer");
737        let verifier = TokenVerifier::new(verifying_key, "test-issuer", "test-audience")
738            .with_origin_validation();
739
740        let claims = SessionClaims::builder("test-issuer", "test-subject", "test-audience")
741            .with_ttl(300)
742            .with_scope("read")
743            .with_metering_key("meter-123")
744            .with_key_class(KeyClass::Publishable)
745            .build();
746        let token = signer.sign(claims).unwrap();
747
748        let result = verifier.verify(&token, None, None);
749        assert!(matches!(
750            result,
751            Err(VerifyError::MissingClaim(ref claim)) if claim == "origin"
752        ));
753    }
754
755    #[test]
756    fn test_client_ip_validation_requires_client_ip_claim() {
757        let signing_key = crate::keys::SigningKey::generate();
758        let verifying_key = signing_key.verifying_key();
759
760        let signer = TokenSigner::new(signing_key, "test-issuer");
761        let verifier = TokenVerifier::new(verifying_key, "test-issuer", "test-audience")
762            .with_client_ip_validation();
763
764        let claims = SessionClaims::builder("test-issuer", "test-subject", "test-audience")
765            .with_ttl(300)
766            .with_scope("read")
767            .with_metering_key("meter-123")
768            .with_key_class(KeyClass::Publishable)
769            .build();
770        let token = signer.sign(claims).unwrap();
771
772        let result = verifier.verify(&token, None, None);
773        assert!(matches!(
774            result,
775            Err(VerifyError::MissingClaim(ref claim)) if claim == "client_ip"
776        ));
777    }
778
779    #[test]
780    fn test_origin_bound_token_allows_no_origin_when_not_required() {
781        // This tests the non-browser client scenario (Rust, Python, etc.)
782        // where the client doesn't send an Origin header, but the JWT has
783        // an origin claim from when the token was minted via browser/API.
784        // When require_origin is false, the connection should still be allowed
785        // for defense-in-depth flexibility.
786        let signing_key = crate::keys::SigningKey::generate();
787        let verifying_key = signing_key.verifying_key();
788
789        let signer = TokenSigner::new(signing_key, "test-issuer");
790        // Verifier WITHOUT origin validation (the default for public stacks)
791        let verifier = TokenVerifier::new(verifying_key, "test-issuer", "test-audience");
792
793        let claims = SessionClaims::builder("test-issuer", "test-subject", "test-audience")
794            .with_ttl(300)
795            .with_scope("read")
796            .with_metering_key("meter-123")
797            .with_origin("https://example.com") // Token has origin claim
798            .with_key_class(KeyClass::Publishable)
799            .build();
800        let token = signer.sign(claims).unwrap();
801
802        // No origin provided, but require_origin is false - should succeed
803        let context = verifier.verify(&token, None, None).unwrap();
804        assert_eq!(context.origin.as_deref(), Some("https://example.com"));
805    }
806
807    #[test]
808    fn test_origin_bound_token_validates_when_origin_provided_even_when_not_required() {
809        // When origin IS provided, it should still be validated against the token
810        // even when require_origin is false (defense-in-depth)
811        let signing_key = crate::keys::SigningKey::generate();
812        let verifying_key = signing_key.verifying_key();
813
814        let signer = TokenSigner::new(signing_key, "test-issuer");
815        // Verifier WITHOUT origin validation (the default)
816        let verifier = TokenVerifier::new(verifying_key, "test-issuer", "test-audience");
817
818        let claims = SessionClaims::builder("test-issuer", "test-subject", "test-audience")
819            .with_ttl(300)
820            .with_scope("read")
821            .with_metering_key("meter-123")
822            .with_origin("https://allowed.example")
823            .with_key_class(KeyClass::Publishable)
824            .build();
825        let token = signer.sign(claims).unwrap();
826
827        // Origin provided and matches - should succeed
828        let context = verifier
829            .verify(&token, Some("https://allowed.example"), None)
830            .unwrap();
831        assert_eq!(context.origin.as_deref(), Some("https://allowed.example"));
832
833        // Origin provided but doesn't match - should fail
834        let result = verifier.verify(&token, Some("https://evil.example"), None);
835        assert!(matches!(result, Err(VerifyError::OriginMismatch { .. })));
836    }
837}