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    /// 令牌用途(v4.8.0 新增,修复 Critical C-2 令牌类型混淆)
55    ///
56    /// - `Some("access")`:访问令牌
57    /// - `Some("refresh")`:刷新令牌
58    /// - `None`:旧版令牌(仅 verify_token 兼容;refresh_token 严格拒绝)
59    ///
60    /// 修复前访问/刷新令牌共用同一 claims 结构且无类型区分,黑帽审计实证:
61    /// 窃取的访问令牌可送入 refresh 端点无限续期(exp 永续重置)。
62    #[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    /// 设置用户 ID(v0.2.1 新增)
96    pub fn with_user_id(mut self, user_id: i64) -> Self {
97        self.user_id = Some(user_id);
98        self
99    }
100
101    /// 设置令牌用途(v4.8.0 新增,修复 Critical C-2):`"access"` 或 `"refresh"`
102    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    /// 是否为指定用途的令牌
108    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        // Verify signature using constant-time comparison (v0.2.2 修复 H-3)。
166        //
167        // 原实现 `signature_b64 != expected_signature_b64` 使用 `String::ne`,
168        // 该方法逐字节比较并在第一个不匹配处短路返回,导致比较时间与匹配前缀长度成正比,
169        // 攻击者可通过测量响应时间逐字节恢复有效签名(时序攻击)。
170        //
171        // 修复:使用 `subtle::ConstantTimeEq`(RustCrypto audited crate),
172        // 确保无论匹配多少字节,比较时间恒定。
173        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        // 长度不同直接拒绝(长度信息非敏感,可短路)
181        if sig_bytes.len() != expected_bytes.len() {
182            return Err(AuthError::TokenInvalid("Invalid signature".to_string()));
183        }
184        // 常量时间比较字节数组
185        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        // Decode header
191        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        // Decode claims
210        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
231// ============================================================================
232// base64 URL-safe encoding (no padding) per RFC 4648 Section 5
233//
234// v0.2.2 重构(P2-5):使用 RustCrypto audited `base64` crate 替代手写实现。
235// ============================================================================
236
237fn 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    // JWT 使用 unpadded base64url,显式拒绝 `=` padding
243    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// ============================================================================
252// SHA-256 per FIPS 180-4
253//
254// v0.2.2 重构(P2-5):使用 RustCrypto audited `sha2` crate 替代手写实现。
255// ============================================================================
256
257#[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
268// ============================================================================
269// HMAC-SHA256 per RFC 2104
270//
271// v0.2.2 重构(P2-5):使用 RustCrypto audited `hmac` crate 替代手写实现。
272// ============================================================================
273
274fn 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// ============================================================================
284// v3.8.0:JWT 密钥轮换(多 kid 并存,无停机轮换)
285// ============================================================================
286
287/// JWT 密钥集:支持多密钥并存,以 kid 标识,实现无停机密钥轮换
288#[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    /// 创建密钥集,校验所有密钥长度 ≥ 32 字节,active_kid 存在于 keys
298    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    /// 轮换密钥:新增 kid 设为 active,保留旧 kid(旧令牌用旧密钥验证直至过期)
327    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    /// 移除非 active 的 kid
345    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    /// 获取当前 active kid
361    pub fn active_kid(&self) -> String {
362        self.active_kid.read().unwrap().clone()
363    }
364
365    /// 按 kid 获取密钥
366    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/// 带 kid 的 JWT 编解码器
375#[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    /// 签发令牌:用 active_kid 密钥签发,header 携带 kid
387    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    /// 验证令牌:解析 header.kid,按 kid 查找密钥验证签名
410    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    // SHA-256 known answer tests (FIPS 180-2 examples)
469
470    #[test]
471    fn test_sha256_empty() {
472        // sha256("") = e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855
473        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        // sha256("abc") = ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad
481        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        // sha256("abcdbcdecdefdefgefghfghighijhijkijkljklmklmnlmnomnopnopq")
489        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    // HMAC-SHA256 known answer test (RFC 4231 Test Case 1)
496
497    #[test]
498    fn test_hmac_sha256_rfc4231_case1() {
499        // Key = 0x0b repeated 20 times, Data = "Hi There"
500        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        // Key = "Jefe", Data = "what do ya want for nothing?"
510        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        // Key longer than 64 bytes should be hashed first (RFC 4231 Case 6 uses 131 bytes)
518        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    // base64url tests
526
527    #[test]
528    fn test_base64_url_encode_known() {
529        // RFC 4648 Section 10 (with URL alphabet, no padding):
530        // "" -> "", "f" -> "Zg", "fo" -> "Zm8", "foo" -> "Zm9v",
531        // "foob" -> "Zm9vYg", "fooba" -> "Zm9vYmE", "foobar" -> "Zm9vYmFy"
532        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    // JWT encode/decode tests
583
584    #[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        // Header should decode to {"alg":"HS256","typ":"JWT"}
620        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        // Header.payload should be the same, but signature differs.
636        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]); // header
639        assert_eq!(parts_a[1], parts_b[1]); // payload
640        assert_ne!(parts_a[2], parts_b[2]); // signature
641    }
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); // expired 100s ago
658        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        // Tamper with the payload by replacing it with a different valid base64url string.
687        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        // Flip the first character of the signature
704        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}