1use crate::error::AuthError;
10use base64::engine::general_purpose::URL_SAFE_NO_PAD;
11use base64::Engine;
12use hmac::{Hmac, Mac};
13use serde::{Deserialize, Serialize};
14use sha2::Sha256;
15
16type HmacSha256 = Hmac<Sha256>;
18
19#[derive(Debug, Clone, Serialize, Deserialize)]
20pub struct JwtHeader {
21 pub alg: String,
22 pub typ: String,
23 #[serde(skip_serializing_if = "Option::is_none", default)]
24 pub kid: Option<String>,
25}
26
27impl Default for JwtHeader {
28 fn default() -> Self {
29 Self {
30 alg: "HS256".to_string(),
31 typ: "JWT".to_string(),
32 kid: None,
33 }
34 }
35}
36
37#[derive(Debug, Clone, Serialize, Deserialize)]
38pub struct JwtClaims {
39 pub sub: String,
40 pub exp: i64,
41 pub iat: i64,
42 #[serde(skip_serializing_if = "Option::is_none")]
43 pub iss: Option<String>,
44 #[serde(default)]
45 pub roles: Vec<String>,
46 #[serde(default)]
47 pub permissions: Vec<String>,
48 #[serde(default, skip_serializing_if = "Option::is_none")]
53 pub user_id: Option<i64>,
54}
55
56impl JwtClaims {
57 pub fn new(sub: impl Into<String>, exp: i64) -> Self {
58 Self {
59 sub: sub.into(),
60 exp,
61 iat: current_timestamp(),
62 iss: None,
63 roles: Vec::new(),
64 permissions: Vec::new(),
65 user_id: None,
66 }
67 }
68
69 pub fn with_issuer(mut self, iss: impl Into<String>) -> Self {
70 self.iss = Some(iss.into());
71 self
72 }
73
74 pub fn with_roles(mut self, roles: Vec<String>) -> Self {
75 self.roles = roles;
76 self
77 }
78
79 pub fn with_permissions(mut self, permissions: Vec<String>) -> Self {
80 self.permissions = permissions;
81 self
82 }
83
84 pub fn with_user_id(mut self, user_id: i64) -> Self {
86 self.user_id = Some(user_id);
87 self
88 }
89
90 pub fn is_expired(&self) -> bool {
91 current_timestamp() > self.exp
92 }
93}
94
95pub struct JwtEncoder {
96 secret: String,
97}
98
99impl JwtEncoder {
100 pub fn new(secret: impl Into<String>) -> Self {
101 Self {
102 secret: secret.into(),
103 }
104 }
105
106 pub fn secret(&self) -> &str {
107 &self.secret
108 }
109
110 pub fn encode(&self, claims: &JwtClaims) -> Result<String, AuthError> {
111 let header = JwtHeader::default();
112 let header_json = serde_json::to_string(&header)
113 .map_err(|e| AuthError::TokenInvalid(format!("Header serialization failed: {}", e)))?;
114 let claims_json = serde_json::to_string(claims)
115 .map_err(|e| AuthError::TokenInvalid(format!("Claims serialization failed: {}", e)))?;
116
117 let header_b64 = base64_url_encode(header_json.as_bytes());
118 let claims_b64 = base64_url_encode(claims_json.as_bytes());
119
120 let signing_input = format!("{}.{}", header_b64, claims_b64);
121 let signature = hmac_sha256(self.secret.as_bytes(), signing_input.as_bytes());
122 let signature_b64 = base64_url_encode(&signature);
123
124 Ok(format!("{}.{}.{}", header_b64, claims_b64, signature_b64))
125 }
126
127 pub fn decode(&self, token: &str) -> Result<JwtClaims, AuthError> {
128 if token.is_empty() {
129 return Err(AuthError::TokenInvalid("Token is empty".to_string()));
130 }
131
132 let parts: Vec<&str> = token.split('.').collect();
133 if parts.len() != 3 {
134 return Err(AuthError::TokenInvalid(
135 "Invalid JWT format: expected 3 parts".to_string(),
136 ));
137 }
138
139 let header_b64 = parts[0];
140 let claims_b64 = parts[1];
141 let signature_b64 = parts[2];
142
143 use subtle::ConstantTimeEq;
152 let signing_input = format!("{}.{}", header_b64, claims_b64);
153 let expected_signature = hmac_sha256(self.secret.as_bytes(), signing_input.as_bytes());
154 let expected_signature_b64 = base64_url_encode(&expected_signature);
155
156 let sig_bytes = signature_b64.as_bytes();
157 let expected_bytes = expected_signature_b64.as_bytes();
158 if sig_bytes.len() != expected_bytes.len() {
160 return Err(AuthError::TokenInvalid("Invalid signature".to_string()));
161 }
162 let sig_match: bool = sig_bytes.ct_eq(expected_bytes).into();
164 if !sig_match {
165 return Err(AuthError::TokenInvalid("Invalid signature".to_string()));
166 }
167
168 let header_bytes = base64_url_decode(header_b64)
170 .map_err(|e| AuthError::TokenInvalid(format!("Header decode failed: {}", e)))?;
171 let header: JwtHeader = serde_json::from_slice(&header_bytes)
172 .map_err(|e| AuthError::TokenInvalid(format!("Header parse failed: {}", e)))?;
173
174 if header.alg != "HS256" {
175 return Err(AuthError::TokenInvalid(format!(
176 "Unsupported algorithm: {}",
177 header.alg
178 )));
179 }
180 if header.typ != "JWT" {
181 return Err(AuthError::TokenInvalid(format!(
182 "Unsupported token type: {}",
183 header.typ
184 )));
185 }
186
187 let claims_bytes = base64_url_decode(claims_b64)
189 .map_err(|e| AuthError::TokenInvalid(format!("Claims decode failed: {}", e)))?;
190 let claims: JwtClaims = serde_json::from_slice(&claims_bytes)
191 .map_err(|e| AuthError::TokenInvalid(format!("Claims parse failed: {}", e)))?;
192
193 if claims.is_expired() {
194 return Err(AuthError::TokenExpired("Token has expired".to_string()));
195 }
196
197 Ok(claims)
198 }
199}
200
201fn current_timestamp() -> i64 {
202 use std::time::{SystemTime, UNIX_EPOCH};
203 SystemTime::now()
204 .duration_since(UNIX_EPOCH)
205 .unwrap_or_default()
206 .as_secs() as i64
207}
208
209fn base64_url_encode(input: &[u8]) -> String {
216 URL_SAFE_NO_PAD.encode(input)
217}
218
219fn base64_url_decode(input: &str) -> Result<Vec<u8>, String> {
220 if input.contains('=') {
222 return Err("base64url must not contain padding '='".to_string());
223 }
224 URL_SAFE_NO_PAD
225 .decode(input)
226 .map_err(|e| format!("Invalid base64url: {}", e))
227}
228
229#[cfg(test)]
236fn sha256(data: &[u8]) -> [u8; 32] {
237 use sha2::Digest;
238 let mut hasher = Sha256::new();
239 hasher.update(data);
240 let result = hasher.finalize();
241 let mut out = [0u8; 32];
242 out.copy_from_slice(&result);
243 out
244}
245
246fn hmac_sha256(key: &[u8], message: &[u8]) -> [u8; 32] {
253 let mut mac = HmacSha256::new_from_slice(key).expect("HMAC accepts any key length");
254 mac.update(message);
255 let result = mac.finalize().into_bytes();
256 let mut out = [0u8; 32];
257 out.copy_from_slice(&result);
258 out
259}
260
261#[cfg(feature = "prod-jwt-key-rotation")]
267pub struct JwtKeySet {
268 keys: std::sync::RwLock<std::collections::HashMap<String, String>>,
269 active_kid: std::sync::RwLock<String>,
270 min_secret_length: usize,
271}
272
273#[cfg(feature = "prod-jwt-key-rotation")]
274impl JwtKeySet {
275 pub fn new(
277 keys: std::collections::HashMap<String, String>,
278 active_kid: String,
279 ) -> Result<Self, AuthError> {
280 const MIN_LEN: usize = 32;
281 for (kid, secret) in &keys {
282 if secret.len() < MIN_LEN {
283 return Err(AuthError::SecretTooShort(format!(
284 "key '{}' has {} bytes, minimum {} required",
285 kid,
286 secret.len(),
287 MIN_LEN
288 )));
289 }
290 }
291 if !keys.contains_key(&active_kid) {
292 return Err(AuthError::TokenInvalid(format!(
293 "active_kid '{}' not found in keys",
294 active_kid
295 )));
296 }
297 Ok(Self {
298 keys: std::sync::RwLock::new(keys),
299 active_kid: std::sync::RwLock::new(active_kid),
300 min_secret_length: MIN_LEN,
301 })
302 }
303
304 pub fn rotate(&self, new_kid: String, new_secret: String) -> Result<(), AuthError> {
306 if new_secret.len() < self.min_secret_length {
307 return Err(AuthError::SecretTooShort(format!(
308 "new key has {} bytes, minimum {} required",
309 new_secret.len(),
310 self.min_secret_length
311 )));
312 }
313 {
314 let mut keys = self.keys.write().unwrap();
315 keys.insert(new_kid.clone(), new_secret);
316 }
317 let mut active = self.active_kid.write().unwrap();
318 *active = new_kid;
319 Ok(())
320 }
321
322 pub fn remove_kid(&self, kid: &str) -> Result<(), AuthError> {
324 let active = self.active_kid.read().unwrap();
325 if *active == kid {
326 return Err(AuthError::TokenInvalid(format!(
327 "cannot remove active kid '{}'",
328 kid
329 )));
330 }
331 let mut keys = self.keys.write().unwrap();
332 if keys.remove(kid).is_none() {
333 return Err(AuthError::TokenInvalid(format!("kid '{}' not found", kid)));
334 }
335 Ok(())
336 }
337
338 pub fn active_kid(&self) -> String {
340 self.active_kid.read().unwrap().clone()
341 }
342
343 pub fn get_secret(&self, kid: &str) -> Result<String, AuthError> {
345 let keys = self.keys.read().unwrap();
346 keys.get(kid)
347 .cloned()
348 .ok_or_else(|| AuthError::TokenInvalid(format!("kid '{}' not found", kid)))
349 }
350}
351
352#[cfg(feature = "prod-jwt-key-rotation")]
354pub struct JwtEncoderWithKid {
355 key_set: std::sync::Arc<JwtKeySet>,
356}
357
358#[cfg(feature = "prod-jwt-key-rotation")]
359impl JwtEncoderWithKid {
360 pub fn new(key_set: std::sync::Arc<JwtKeySet>) -> Self {
361 Self { key_set }
362 }
363
364 pub fn encode(&self, claims: &JwtClaims) -> Result<String, AuthError> {
366 let kid = self.key_set.active_kid();
367 let secret = self.key_set.get_secret(&kid)?;
368
369 let header = JwtHeader {
370 kid: Some(kid),
371 ..Default::default()
372 };
373 let header_json = serde_json::to_string(&header)
374 .map_err(|e| AuthError::TokenInvalid(format!("Header serialization failed: {}", e)))?;
375 let claims_json = serde_json::to_string(claims)
376 .map_err(|e| AuthError::TokenInvalid(format!("Claims serialization failed: {}", e)))?;
377
378 let header_b64 = base64_url_encode(header_json.as_bytes());
379 let claims_b64 = base64_url_encode(claims_json.as_bytes());
380 let signing_input = format!("{}.{}", header_b64, claims_b64);
381 let signature = hmac_sha256(secret.as_bytes(), signing_input.as_bytes());
382 let signature_b64 = base64_url_encode(&signature);
383
384 Ok(format!("{}.{}.{}", header_b64, claims_b64, signature_b64))
385 }
386
387 pub fn decode(&self, token: &str) -> Result<JwtClaims, AuthError> {
389 if token.is_empty() {
390 return Err(AuthError::TokenInvalid("Token is empty".to_string()));
391 }
392 let parts: Vec<&str> = token.split('.').collect();
393 if parts.len() != 3 {
394 return Err(AuthError::TokenInvalid(
395 "Invalid JWT format: expected 3 parts".to_string(),
396 ));
397 }
398 let header_b64 = parts[0];
399 let claims_b64 = parts[1];
400 let signature_b64 = parts[2];
401
402 let header_bytes = base64_url_decode(header_b64)
403 .map_err(|e| AuthError::TokenInvalid(format!("Header decode failed: {}", e)))?;
404 let header: JwtHeader = serde_json::from_slice(&header_bytes)
405 .map_err(|e| AuthError::TokenInvalid(format!("Header parse failed: {}", e)))?;
406
407 let kid = header
408 .kid
409 .ok_or_else(|| AuthError::TokenInvalid("missing kid in token".to_string()))?;
410 let secret = self.key_set.get_secret(&kid)?;
411
412 let signing_input = format!("{}.{}", header_b64, claims_b64);
413 let expected_signature = hmac_sha256(secret.as_bytes(), signing_input.as_bytes());
414 let expected_b64 = base64_url_encode(&expected_signature);
415
416 use subtle::ConstantTimeEq;
417 if signature_b64
418 .as_bytes()
419 .ct_eq(expected_b64.as_bytes())
420 .into()
421 {
422 let claims_bytes = base64_url_decode(claims_b64)
423 .map_err(|e| AuthError::TokenInvalid(format!("Claims decode failed: {}", e)))?;
424 let claims: JwtClaims = serde_json::from_slice(&claims_bytes)
425 .map_err(|e| AuthError::TokenInvalid(format!("Claims parse failed: {}", e)))?;
426 if claims.is_expired() {
427 return Err(AuthError::TokenExpired("token expired".to_string()));
428 }
429 Ok(claims)
430 } else {
431 Err(AuthError::TokenInvalid(
432 "Signature verification failed".to_string(),
433 ))
434 }
435 }
436}
437
438#[cfg(test)]
439mod tests {
440 use super::*;
441
442 fn now() -> i64 {
443 current_timestamp()
444 }
445
446 #[test]
449 fn test_sha256_empty() {
450 let hash = sha256(b"");
452 let expected_hex = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855";
453 assert_eq!(hex_str(&hash), expected_hex);
454 }
455
456 #[test]
457 fn test_sha256_abc() {
458 let hash = sha256(b"abc");
460 let expected_hex = "ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad";
461 assert_eq!(hex_str(&hash), expected_hex);
462 }
463
464 #[test]
465 fn test_sha256_longer_message() {
466 let input = b"abcdbcdecdefdefgefghfghighijhijkijkljklmklmnlmnomnopnopq";
468 let hash = sha256(input);
469 let expected_hex = "248d6a61d20638b8e5c026930c3e6039a33ce45964ff2167f6ecedd419db06c1";
470 assert_eq!(hex_str(&hash), expected_hex);
471 }
472
473 #[test]
476 fn test_hmac_sha256_rfc4231_case1() {
477 let key = [0x0bu8; 20];
479 let message = b"Hi There";
480 let mac = hmac_sha256(&key, message);
481 let expected_hex = "b0344c61d8db38535ca8afceaf0bf12b881dc200c9833da726e9376c2e32cff7";
482 assert_eq!(hex_str(&mac), expected_hex);
483 }
484
485 #[test]
486 fn test_hmac_sha256_rfc4231_case2() {
487 let mac = hmac_sha256(b"Jefe", b"what do ya want for nothing?");
489 let expected_hex = "5bdcc146bf60754e6a042426089575c75a003f089d2739839dec58b964ec3843";
490 assert_eq!(hex_str(&mac), expected_hex);
491 }
492
493 #[test]
494 fn test_hmac_sha256_long_key() {
495 let key = [0xaau8; 131];
497 let data = b"Test Using Larger Than Block-Size Key - Hash Key First";
498 let mac = hmac_sha256(&key, data);
499 let expected_hex = "60e431591ee0b67f0d8a26aacbf5b77f8e0bc6213728c5140546040f0ee37f54";
500 assert_eq!(hex_str(&mac), expected_hex);
501 }
502
503 #[test]
506 fn test_base64_url_encode_known() {
507 assert_eq!(base64_url_encode(b""), "");
511 assert_eq!(base64_url_encode(b"f"), "Zg");
512 assert_eq!(base64_url_encode(b"fo"), "Zm8");
513 assert_eq!(base64_url_encode(b"foo"), "Zm9v");
514 assert_eq!(base64_url_encode(b"foob"), "Zm9vYg");
515 assert_eq!(base64_url_encode(b"fooba"), "Zm9vYmE");
516 assert_eq!(base64_url_encode(b"foobar"), "Zm9vYmFy");
517 }
518
519 #[test]
520 fn test_base64_url_decode_known() {
521 assert_eq!(base64_url_decode("").unwrap(), b"");
522 assert_eq!(base64_url_decode("Zg").unwrap(), b"f");
523 assert_eq!(base64_url_decode("Zm8").unwrap(), b"fo");
524 assert_eq!(base64_url_decode("Zm9v").unwrap(), b"foo");
525 assert_eq!(base64_url_decode("Zm9vYg").unwrap(), b"foob");
526 assert_eq!(base64_url_decode("Zm9vYmE").unwrap(), b"fooba");
527 assert_eq!(base64_url_decode("Zm9vYmFy").unwrap(), b"foobar");
528 }
529
530 #[test]
531 fn test_base64_url_roundtrip() {
532 let cases: &[&[u8]] = &[
533 b"",
534 b"a",
535 b"ab",
536 b"abc",
537 b"abcd",
538 b"hello world",
539 &[0xffu8; 64],
540 &[0x00u8; 64],
541 &(0u8..=255).collect::<Vec<u8>>(),
542 ];
543 for c in cases {
544 let encoded = base64_url_encode(c);
545 let decoded = base64_url_decode(&encoded).unwrap();
546 assert_eq!(decoded.as_slice(), *c, "roundtrip failed for {:?}", c);
547 }
548 }
549
550 #[test]
551 fn test_base64_url_rejects_padding() {
552 assert!(base64_url_decode("Zg==").is_err());
553 }
554
555 #[test]
556 fn test_base64_url_rejects_invalid_char() {
557 assert!(base64_url_decode("Zm9v*").is_err());
558 }
559
560 #[test]
563 fn test_jwt_encode_decode_roundtrip() {
564 let encoder = JwtEncoder::new("my-secret");
565 let claims = JwtClaims::new("user123", now() + 3600)
566 .with_issuer("test-issuer")
567 .with_roles(vec!["user".to_string(), "editor".to_string()])
568 .with_permissions(vec!["read:posts".to_string(), "write:posts".to_string()]);
569
570 let token = encoder.encode(&claims).expect("encode");
571 assert!(!token.is_empty());
572
573 let parts: Vec<&str> = token.split('.').collect();
574 assert_eq!(parts.len(), 3);
575
576 let decoded = encoder.decode(&token).expect("decode");
577 assert_eq!(decoded.sub, "user123");
578 assert_eq!(decoded.iss, Some("test-issuer".to_string()));
579 assert_eq!(
580 decoded.roles,
581 vec!["user".to_string(), "editor".to_string()]
582 );
583 assert_eq!(
584 decoded.permissions,
585 vec!["read:posts".to_string(), "write:posts".to_string()]
586 );
587 }
588
589 #[test]
590 fn test_jwt_format_is_header_payload_signature() {
591 let encoder = JwtEncoder::new("secret");
592 let claims = JwtClaims::new("alice", now() + 60);
593 let token = encoder.encode(&claims).unwrap();
594 let parts: Vec<&str> = token.split('.').collect();
595 assert_eq!(parts.len(), 3);
596
597 let header_bytes = base64_url_decode(parts[0]).unwrap();
599 let header: JwtHeader = serde_json::from_slice(&header_bytes).unwrap();
600 assert_eq!(header.alg, "HS256");
601 assert_eq!(header.typ, "JWT");
602 }
603
604 #[test]
605 fn test_jwt_signature_changes_with_secret() {
606 let encoder_a = JwtEncoder::new("secret-a");
607 let encoder_b = JwtEncoder::new("secret-b");
608 let claims = JwtClaims::new("user", now() + 3600);
609
610 let token_a = encoder_a.encode(&claims).unwrap();
611 let token_b = encoder_b.encode(&claims).unwrap();
612
613 let parts_a: Vec<&str> = token_a.split('.').collect();
615 let parts_b: Vec<&str> = token_b.split('.').collect();
616 assert_eq!(parts_a[0], parts_b[0]); assert_eq!(parts_a[1], parts_b[1]); assert_ne!(parts_a[2], parts_b[2]); }
620
621 #[test]
622 fn test_jwt_verify_with_wrong_secret_fails() {
623 let encoder_a = JwtEncoder::new("secret-a");
624 let encoder_b = JwtEncoder::new("secret-b");
625 let claims = JwtClaims::new("user", now() + 3600);
626
627 let token = encoder_a.encode(&claims).unwrap();
628 let result = encoder_b.decode(&token);
629 assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
630 }
631
632 #[test]
633 fn test_jwt_expired_token_rejected() {
634 let encoder = JwtEncoder::new("secret");
635 let claims = JwtClaims::new("user", now() - 100); let token = encoder.encode(&claims).unwrap();
637 let result = encoder.decode(&token);
638 assert!(matches!(result, Err(AuthError::TokenExpired(_))));
639 }
640
641 #[test]
642 fn test_jwt_decode_invalid_format() {
643 let encoder = JwtEncoder::new("secret");
644 assert!(matches!(
645 encoder.decode(""),
646 Err(AuthError::TokenInvalid(_))
647 ));
648 assert!(matches!(
649 encoder.decode("not.a.jwt.token"),
650 Err(AuthError::TokenInvalid(_))
651 ));
652 assert!(matches!(
653 encoder.decode("only.two"),
654 Err(AuthError::TokenInvalid(_))
655 ));
656 }
657
658 #[test]
659 fn test_jwt_tampered_payload_rejected() {
660 let encoder = JwtEncoder::new("secret");
661 let claims = JwtClaims::new("alice", now() + 3600);
662 let token = encoder.encode(&claims).unwrap();
663
664 let parts: Vec<&str> = token.split('.').collect();
666 let tampered_payload = base64_url_encode(
667 br#"{"sub":"mallory","exp":9999999999,"iat":0,"roles":[],"permissions":[]}"#,
668 );
669 let tampered = format!("{}.{}.{}", parts[0], tampered_payload, parts[2]);
670 let result = encoder.decode(&tampered);
671 assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
672 }
673
674 #[test]
675 fn test_jwt_tampered_signature_rejected() {
676 let encoder = JwtEncoder::new("secret");
677 let claims = JwtClaims::new("alice", now() + 3600);
678 let token = encoder.encode(&claims).unwrap();
679
680 let parts: Vec<&str> = token.split('.').collect();
681 let mut sig = parts[2].to_string();
683 let first = sig.chars().next().unwrap();
684 let replacement = if first == 'A' { 'B' } else { 'A' };
685 sig.replace_range(0..first.len_utf8(), &replacement.to_string());
686 let tampered = format!("{}.{}.{}", parts[0], parts[1], sig);
687 let result = encoder.decode(&tampered);
688 assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
689 }
690
691 fn hex_str(bytes: &[u8]) -> String {
692 let mut s = String::with_capacity(bytes.len() * 2);
693 for b in bytes {
694 s.push_str(&format!("{:02x}", b));
695 }
696 s
697 }
698
699 #[cfg(feature = "prod-jwt-key-rotation")]
700 mod prod_jwt_key_rotation_tests {
701 use super::*;
702 use std::collections::HashMap;
703
704 fn make_secret(n: usize) -> String {
705 "a".repeat(n)
706 }
707
708 #[test]
709 fn test_key_set_new_validates_min_length() {
710 let mut keys = HashMap::new();
711 keys.insert("kid1".to_string(), make_secret(32));
712 let result = JwtKeySet::new(keys, "kid1".to_string());
713 assert!(result.is_ok());
714 }
715
716 #[test]
717 fn test_key_set_new_rejects_short_secret() {
718 let mut keys = HashMap::new();
719 keys.insert("kid1".to_string(), make_secret(31));
720 let result = JwtKeySet::new(keys, "kid1".to_string());
721 assert!(result.is_err());
722 }
723
724 #[test]
725 fn test_key_set_new_rejects_missing_active_kid() {
726 let mut keys = HashMap::new();
727 keys.insert("kid1".to_string(), make_secret(32));
728 let result = JwtKeySet::new(keys, "kid2".to_string());
729 assert!(result.is_err());
730 }
731
732 #[test]
733 fn test_encode_decode_with_kid() {
734 let mut keys = HashMap::new();
735 keys.insert("kid1".to_string(), make_secret(32));
736 keys.insert("kid2".to_string(), make_secret(32));
737 let key_set = JwtKeySet::new(keys, "kid2".to_string()).unwrap();
738 let encoder = JwtEncoderWithKid::new(std::sync::Arc::new(key_set));
739
740 let claims = JwtClaims::new("alice", current_timestamp() + 3600);
741 let token = encoder.encode(&claims).unwrap();
742 let decoded = encoder.decode(&token).unwrap();
743 assert_eq!(decoded.sub, "alice");
744 }
745
746 #[test]
747 fn test_rotate_old_token_still_valid() {
748 let mut keys = HashMap::new();
749 keys.insert("kid1".to_string(), make_secret(32));
750 let key_set = std::sync::Arc::new(JwtKeySet::new(keys, "kid1".to_string()).unwrap());
751 let encoder = JwtEncoderWithKid::new(key_set.clone());
752
753 let claims = JwtClaims::new("bob", current_timestamp() + 3600);
754 let old_token = encoder.encode(&claims).unwrap();
755
756 key_set.rotate("kid2".to_string(), make_secret(32)).unwrap();
757 let new_token = encoder.encode(&claims).unwrap();
758
759 assert!(encoder.decode(&old_token).is_ok());
760 assert!(encoder.decode(&new_token).is_ok());
761 }
762
763 #[test]
764 fn test_remove_kid_rejects_active() {
765 let mut keys = HashMap::new();
766 keys.insert("kid1".to_string(), make_secret(32));
767 keys.insert("kid2".to_string(), make_secret(32));
768 let key_set = JwtKeySet::new(keys, "kid1".to_string()).unwrap();
769 assert!(key_set.remove_kid("kid1").is_err());
770 assert!(key_set.remove_kid("kid2").is_ok());
771 }
772
773 #[test]
774 fn test_decode_missing_kid_rejected() {
775 let mut keys = HashMap::new();
776 keys.insert("kid1".to_string(), make_secret(32));
777 let key_set = JwtKeySet::new(keys, "kid1".to_string()).unwrap();
778 let encoder = JwtEncoderWithKid::new(std::sync::Arc::new(key_set));
779
780 let plain_encoder = JwtEncoder::new(make_secret(32));
781 let claims = JwtClaims::new("eve", current_timestamp() + 3600);
782 let token = plain_encoder.encode(&claims).unwrap();
783
784 assert!(encoder.decode(&token).is_err());
785 }
786 }
787}