Skip to main content

sz_orm_auth/
jwt.rs

1//! Real JWT (HS256) implementation using RustCrypto audited crates.
2//!
3//! v0.2.2 重构(P2-5):将手写 SHA-256 / HMAC-SHA256 / base64url 替换为
4//! RustCrypto audited crate(`sha2`、`hmac`、`base64`),降低密码学实现风险。
5//!
6//! Token format: `base64url(header).base64url(payload).base64url(signature)`
7//! where signature = HMAC-SHA256(secret, `header.payload`).
8
9use 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
16/// HMAC-SHA256 类型别名(来自 RustCrypto `hmac` crate)
17type 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    /// 用户 ID(v0.2.1 新增,修复 Critical S-2)
49    ///
50    /// - `Some(id)`:携带用户 ID,verify_token 可恢复正确 user.id
51    /// - `None`:兼容旧 token;verify_token 会回退为 0 并通过 tracing 警告
52    #[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    /// 设置用户 ID(v0.2.1 新增)
85    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        // Verify signature using constant-time comparison (v0.2.2 修复 H-3)。
144        //
145        // 原实现 `signature_b64 != expected_signature_b64` 使用 `String::ne`,
146        // 该方法逐字节比较并在第一个不匹配处短路返回,导致比较时间与匹配前缀长度成正比,
147        // 攻击者可通过测量响应时间逐字节恢复有效签名(时序攻击)。
148        //
149        // 修复:使用 `subtle::ConstantTimeEq`(RustCrypto audited crate),
150        // 确保无论匹配多少字节,比较时间恒定。
151        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        // 长度不同直接拒绝(长度信息非敏感,可短路)
159        if sig_bytes.len() != expected_bytes.len() {
160            return Err(AuthError::TokenInvalid("Invalid signature".to_string()));
161        }
162        // 常量时间比较字节数组
163        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        // Decode header
169        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        // Decode claims
188        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
209// ============================================================================
210// base64 URL-safe encoding (no padding) per RFC 4648 Section 5
211//
212// v0.2.2 重构(P2-5):使用 RustCrypto audited `base64` crate 替代手写实现。
213// ============================================================================
214
215fn 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    // JWT 使用 unpadded base64url,显式拒绝 `=` padding
221    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// ============================================================================
230// SHA-256 per FIPS 180-4
231//
232// v0.2.2 重构(P2-5):使用 RustCrypto audited `sha2` crate 替代手写实现。
233// ============================================================================
234
235#[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
246// ============================================================================
247// HMAC-SHA256 per RFC 2104
248//
249// v0.2.2 重构(P2-5):使用 RustCrypto audited `hmac` crate 替代手写实现。
250// ============================================================================
251
252fn 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// ============================================================================
262// v3.8.0:JWT 密钥轮换(多 kid 并存,无停机轮换)
263// ============================================================================
264
265/// JWT 密钥集:支持多密钥并存,以 kid 标识,实现无停机密钥轮换
266#[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    /// 创建密钥集,校验所有密钥长度 ≥ 32 字节,active_kid 存在于 keys
276    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    /// 轮换密钥:新增 kid 设为 active,保留旧 kid(旧令牌用旧密钥验证直至过期)
305    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    /// 移除非 active 的 kid
323    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    /// 获取当前 active kid
339    pub fn active_kid(&self) -> String {
340        self.active_kid.read().unwrap().clone()
341    }
342
343    /// 按 kid 获取密钥
344    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/// 带 kid 的 JWT 编解码器
353#[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    /// 签发令牌:用 active_kid 密钥签发,header 携带 kid
365    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    /// 验证令牌:解析 header.kid,按 kid 查找密钥验证签名
388    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    // SHA-256 known answer tests (FIPS 180-2 examples)
447
448    #[test]
449    fn test_sha256_empty() {
450        // sha256("") = e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855
451        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        // sha256("abc") = ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad
459        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        // sha256("abcdbcdecdefdefgefghfghighijhijkijkljklmklmnlmnomnopnopq")
467        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    // HMAC-SHA256 known answer test (RFC 4231 Test Case 1)
474
475    #[test]
476    fn test_hmac_sha256_rfc4231_case1() {
477        // Key = 0x0b repeated 20 times, Data = "Hi There"
478        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        // Key = "Jefe", Data = "what do ya want for nothing?"
488        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        // Key longer than 64 bytes should be hashed first (RFC 4231 Case 6 uses 131 bytes)
496        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    // base64url tests
504
505    #[test]
506    fn test_base64_url_encode_known() {
507        // RFC 4648 Section 10 (with URL alphabet, no padding):
508        // "" -> "", "f" -> "Zg", "fo" -> "Zm8", "foo" -> "Zm9v",
509        // "foob" -> "Zm9vYg", "fooba" -> "Zm9vYmE", "foobar" -> "Zm9vYmFy"
510        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    // JWT encode/decode tests
561
562    #[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        // Header should decode to {"alg":"HS256","typ":"JWT"}
598        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        // Header.payload should be the same, but signature differs.
614        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]); // header
617        assert_eq!(parts_a[1], parts_b[1]); // payload
618        assert_ne!(parts_a[2], parts_b[2]); // signature
619    }
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); // expired 100s ago
636        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        // Tamper with the payload by replacing it with a different valid base64url string.
665        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        // Flip the first character of the signature
682        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}