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