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}
24
25impl Default for JwtHeader {
26    fn default() -> Self {
27        Self {
28            alg: "HS256".to_string(),
29            typ: "JWT".to_string(),
30        }
31    }
32}
33
34#[derive(Debug, Clone, Serialize, Deserialize)]
35pub struct JwtClaims {
36    pub sub: String,
37    pub exp: i64,
38    pub iat: i64,
39    #[serde(skip_serializing_if = "Option::is_none")]
40    pub iss: Option<String>,
41    /// 接收人(audience)— P1-SEC-10 新增
42    ///
43    /// 用于防止跨服务令牌重用攻击:其他服务签发的 JWT 因 `aud` 不匹配而被拒绝。
44    #[serde(skip_serializing_if = "Option::is_none")]
45    pub aud: Option<String>,
46    #[serde(default)]
47    pub roles: Vec<String>,
48    #[serde(default)]
49    pub permissions: Vec<String>,
50    /// 用户 ID(v0.2.1 新增,修复 Critical S-2)
51    ///
52    /// - `Some(id)`:携带用户 ID,verify_token 可恢复正确 user.id
53    /// - `None`:兼容旧 token;verify_token 会回退为 0 并通过 tracing 警告
54    #[serde(default, skip_serializing_if = "Option::is_none")]
55    pub user_id: Option<i64>,
56}
57
58impl JwtClaims {
59    pub fn new(sub: impl Into<String>, exp: i64) -> Self {
60        Self {
61            sub: sub.into(),
62            exp,
63            iat: current_timestamp(),
64            iss: None,
65            aud: None,
66            roles: Vec::new(),
67            permissions: Vec::new(),
68            user_id: None,
69        }
70    }
71
72    pub fn with_issuer(mut self, iss: impl Into<String>) -> Self {
73        self.iss = Some(iss.into());
74        self
75    }
76
77    /// 设置接收人(audience)— P1-SEC-10
78    pub fn with_audience(mut self, aud: impl Into<String>) -> Self {
79        self.aud = Some(aud.into());
80        self
81    }
82
83    pub fn with_roles(mut self, roles: Vec<String>) -> Self {
84        self.roles = roles;
85        self
86    }
87
88    pub fn with_permissions(mut self, permissions: Vec<String>) -> Self {
89        self.permissions = permissions;
90        self
91    }
92
93    /// 设置用户 ID(v0.2.1 新增)
94    pub fn with_user_id(mut self, user_id: i64) -> Self {
95        self.user_id = Some(user_id);
96        self
97    }
98
99    pub fn is_expired(&self) -> bool {
100        current_timestamp() > self.exp
101    }
102}
103
104pub struct JwtEncoder {
105    secret: String,
106}
107
108impl JwtEncoder {
109    pub fn new(secret: impl Into<String>) -> Self {
110        Self {
111            secret: secret.into(),
112        }
113    }
114
115    pub fn secret(&self) -> &str {
116        &self.secret
117    }
118
119    pub fn encode(&self, claims: &JwtClaims) -> Result<String, AuthError> {
120        let header = JwtHeader::default();
121        let header_json = serde_json::to_string(&header)
122            .map_err(|e| AuthError::TokenInvalid(format!("Header serialization failed: {}", e)))?;
123        let claims_json = serde_json::to_string(claims)
124            .map_err(|e| AuthError::TokenInvalid(format!("Claims serialization failed: {}", e)))?;
125
126        let header_b64 = base64_url_encode(header_json.as_bytes());
127        let claims_b64 = base64_url_encode(claims_json.as_bytes());
128
129        let signing_input = format!("{}.{}", header_b64, claims_b64);
130        let signature = hmac_sha256(self.secret.as_bytes(), signing_input.as_bytes());
131        let signature_b64 = base64_url_encode(&signature);
132
133        Ok(format!("{}.{}.{}", header_b64, claims_b64, signature_b64))
134    }
135
136    pub fn decode(&self, token: &str) -> Result<JwtClaims, AuthError> {
137        if token.is_empty() {
138            return Err(AuthError::TokenInvalid("Token is empty".to_string()));
139        }
140
141        let parts: Vec<&str> = token.split('.').collect();
142        if parts.len() != 3 {
143            return Err(AuthError::TokenInvalid(
144                "Invalid JWT format: expected 3 parts".to_string(),
145            ));
146        }
147
148        let header_b64 = parts[0];
149        let claims_b64 = parts[1];
150        let signature_b64 = parts[2];
151
152        // Verify signature using constant-time comparison (v0.2.2 修复 H-3)。
153        //
154        // 原实现 `signature_b64 != expected_signature_b64` 使用 `String::ne`,
155        // 该方法逐字节比较并在第一个不匹配处短路返回,导致比较时间与匹配前缀长度成正比,
156        // 攻击者可通过测量响应时间逐字节恢复有效签名(时序攻击)。
157        //
158        // 修复:使用 `subtle::ConstantTimeEq`(RustCrypto audited crate),
159        // 确保无论匹配多少字节,比较时间恒定。
160        use subtle::ConstantTimeEq;
161        let signing_input = format!("{}.{}", header_b64, claims_b64);
162        let expected_signature = hmac_sha256(self.secret.as_bytes(), signing_input.as_bytes());
163        let expected_signature_b64 = base64_url_encode(&expected_signature);
164
165        let sig_bytes = signature_b64.as_bytes();
166        let expected_bytes = expected_signature_b64.as_bytes();
167        // 长度不同直接拒绝(长度信息非敏感,可短路)
168        if sig_bytes.len() != expected_bytes.len() {
169            return Err(AuthError::TokenInvalid("Invalid signature".to_string()));
170        }
171        // 常量时间比较字节数组
172        let sig_match: bool = sig_bytes.ct_eq(expected_bytes).into();
173        if !sig_match {
174            return Err(AuthError::TokenInvalid("Invalid signature".to_string()));
175        }
176
177        // Decode header
178        let header_bytes = base64_url_decode(header_b64)
179            .map_err(|e| AuthError::TokenInvalid(format!("Header decode failed: {}", e)))?;
180        let header: JwtHeader = serde_json::from_slice(&header_bytes)
181            .map_err(|e| AuthError::TokenInvalid(format!("Header parse failed: {}", e)))?;
182
183        if header.alg != "HS256" {
184            return Err(AuthError::TokenInvalid(format!(
185                "Unsupported algorithm: {}",
186                header.alg
187            )));
188        }
189        if header.typ != "JWT" {
190            return Err(AuthError::TokenInvalid(format!(
191                "Unsupported token type: {}",
192                header.typ
193            )));
194        }
195
196        // Decode claims
197        let claims_bytes = base64_url_decode(claims_b64)
198            .map_err(|e| AuthError::TokenInvalid(format!("Claims decode failed: {}", e)))?;
199        let claims: JwtClaims = serde_json::from_slice(&claims_bytes)
200            .map_err(|e| AuthError::TokenInvalid(format!("Claims parse failed: {}", e)))?;
201
202        if claims.is_expired() {
203            return Err(AuthError::TokenExpired("Token has expired".to_string()));
204        }
205
206        Ok(claims)
207    }
208}
209
210fn current_timestamp() -> i64 {
211    use std::time::{SystemTime, UNIX_EPOCH};
212    SystemTime::now()
213        .duration_since(UNIX_EPOCH)
214        .unwrap_or_default()
215        .as_secs() as i64
216}
217
218// ============================================================================
219// base64 URL-safe encoding (no padding) per RFC 4648 Section 5
220//
221// v0.2.2 重构(P2-5):使用 RustCrypto audited `base64` crate 替代手写实现。
222// ============================================================================
223
224fn base64_url_encode(input: &[u8]) -> String {
225    URL_SAFE_NO_PAD.encode(input)
226}
227
228fn base64_url_decode(input: &str) -> Result<Vec<u8>, String> {
229    // JWT 使用 unpadded base64url,显式拒绝 `=` padding
230    if input.contains('=') {
231        return Err("base64url must not contain padding '='".to_string());
232    }
233    URL_SAFE_NO_PAD
234        .decode(input)
235        .map_err(|e| format!("Invalid base64url: {}", e))
236}
237
238// ============================================================================
239// SHA-256 per FIPS 180-4
240//
241// v0.2.2 重构(P2-5):使用 RustCrypto audited `sha2` crate 替代手写实现。
242// ============================================================================
243
244#[cfg(test)]
245fn sha256(data: &[u8]) -> [u8; 32] {
246    use sha2::Digest;
247    let mut hasher = Sha256::new();
248    hasher.update(data);
249    let result = hasher.finalize();
250    let mut out = [0u8; 32];
251    out.copy_from_slice(&result);
252    out
253}
254
255// ============================================================================
256// HMAC-SHA256 per RFC 2104
257//
258// v0.2.2 重构(P2-5):使用 RustCrypto audited `hmac` crate 替代手写实现。
259// ============================================================================
260
261fn hmac_sha256(key: &[u8], message: &[u8]) -> [u8; 32] {
262    let mut mac = HmacSha256::new_from_slice(key).expect("HMAC accepts any key length");
263    mac.update(message);
264    let result = mac.finalize().into_bytes();
265    let mut out = [0u8; 32];
266    out.copy_from_slice(&result);
267    out
268}
269
270#[cfg(test)]
271mod tests {
272    use super::*;
273
274    fn now() -> i64 {
275        current_timestamp()
276    }
277
278    // SHA-256 known answer tests (FIPS 180-2 examples)
279
280    #[test]
281    fn test_sha256_empty() {
282        // sha256("") = e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855
283        let hash = sha256(b"");
284        let expected_hex = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855";
285        assert_eq!(hex_str(&hash), expected_hex);
286    }
287
288    #[test]
289    fn test_sha256_abc() {
290        // sha256("abc") = ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad
291        let hash = sha256(b"abc");
292        let expected_hex = "ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad";
293        assert_eq!(hex_str(&hash), expected_hex);
294    }
295
296    #[test]
297    fn test_sha256_longer_message() {
298        // sha256("abcdbcdecdefdefgefghfghighijhijkijkljklmklmnlmnomnopnopq")
299        let input = b"abcdbcdecdefdefgefghfghighijhijkijkljklmklmnlmnomnopnopq";
300        let hash = sha256(input);
301        let expected_hex = "248d6a61d20638b8e5c026930c3e6039a33ce45964ff2167f6ecedd419db06c1";
302        assert_eq!(hex_str(&hash), expected_hex);
303    }
304
305    // HMAC-SHA256 known answer test (RFC 4231 Test Case 1)
306
307    #[test]
308    fn test_hmac_sha256_rfc4231_case1() {
309        // Key = 0x0b repeated 20 times, Data = "Hi There"
310        let key = [0x0bu8; 20];
311        let message = b"Hi There";
312        let mac = hmac_sha256(&key, message);
313        let expected_hex = "b0344c61d8db38535ca8afceaf0bf12b881dc200c9833da726e9376c2e32cff7";
314        assert_eq!(hex_str(&mac), expected_hex);
315    }
316
317    #[test]
318    fn test_hmac_sha256_rfc4231_case2() {
319        // Key = "Jefe", Data = "what do ya want for nothing?"
320        let mac = hmac_sha256(b"Jefe", b"what do ya want for nothing?");
321        let expected_hex = "5bdcc146bf60754e6a042426089575c75a003f089d2739839dec58b964ec3843";
322        assert_eq!(hex_str(&mac), expected_hex);
323    }
324
325    #[test]
326    fn test_hmac_sha256_long_key() {
327        // Key longer than 64 bytes should be hashed first (RFC 4231 Case 6 uses 131 bytes)
328        let key = [0xaau8; 131];
329        let data = b"Test Using Larger Than Block-Size Key - Hash Key First";
330        let mac = hmac_sha256(&key, data);
331        let expected_hex = "60e431591ee0b67f0d8a26aacbf5b77f8e0bc6213728c5140546040f0ee37f54";
332        assert_eq!(hex_str(&mac), expected_hex);
333    }
334
335    // base64url tests
336
337    #[test]
338    fn test_base64_url_encode_known() {
339        // RFC 4648 Section 10 (with URL alphabet, no padding):
340        // "" -> "", "f" -> "Zg", "fo" -> "Zm8", "foo" -> "Zm9v",
341        // "foob" -> "Zm9vYg", "fooba" -> "Zm9vYmE", "foobar" -> "Zm9vYmFy"
342        assert_eq!(base64_url_encode(b""), "");
343        assert_eq!(base64_url_encode(b"f"), "Zg");
344        assert_eq!(base64_url_encode(b"fo"), "Zm8");
345        assert_eq!(base64_url_encode(b"foo"), "Zm9v");
346        assert_eq!(base64_url_encode(b"foob"), "Zm9vYg");
347        assert_eq!(base64_url_encode(b"fooba"), "Zm9vYmE");
348        assert_eq!(base64_url_encode(b"foobar"), "Zm9vYmFy");
349    }
350
351    #[test]
352    fn test_base64_url_decode_known() {
353        assert_eq!(base64_url_decode("").unwrap(), b"");
354        assert_eq!(base64_url_decode("Zg").unwrap(), b"f");
355        assert_eq!(base64_url_decode("Zm8").unwrap(), b"fo");
356        assert_eq!(base64_url_decode("Zm9v").unwrap(), b"foo");
357        assert_eq!(base64_url_decode("Zm9vYg").unwrap(), b"foob");
358        assert_eq!(base64_url_decode("Zm9vYmE").unwrap(), b"fooba");
359        assert_eq!(base64_url_decode("Zm9vYmFy").unwrap(), b"foobar");
360    }
361
362    #[test]
363    fn test_base64_url_roundtrip() {
364        let cases: &[&[u8]] = &[
365            b"",
366            b"a",
367            b"ab",
368            b"abc",
369            b"abcd",
370            b"hello world",
371            &[0xffu8; 64],
372            &[0x00u8; 64],
373            &(0u8..=255).collect::<Vec<u8>>(),
374        ];
375        for c in cases {
376            let encoded = base64_url_encode(c);
377            let decoded = base64_url_decode(&encoded).unwrap();
378            assert_eq!(decoded.as_slice(), *c, "roundtrip failed for {:?}", c);
379        }
380    }
381
382    #[test]
383    fn test_base64_url_rejects_padding() {
384        assert!(base64_url_decode("Zg==").is_err());
385    }
386
387    #[test]
388    fn test_base64_url_rejects_invalid_char() {
389        assert!(base64_url_decode("Zm9v*").is_err());
390    }
391
392    // JWT encode/decode tests
393
394    #[test]
395    fn test_jwt_encode_decode_roundtrip() {
396        let encoder = JwtEncoder::new("my-secret");
397        let claims = JwtClaims::new("user123", now() + 3600)
398            .with_issuer("test-issuer")
399            .with_roles(vec!["user".to_string(), "editor".to_string()])
400            .with_permissions(vec!["read:posts".to_string(), "write:posts".to_string()]);
401
402        let token = encoder.encode(&claims).expect("encode");
403        assert!(!token.is_empty());
404
405        let parts: Vec<&str> = token.split('.').collect();
406        assert_eq!(parts.len(), 3);
407
408        let decoded = encoder.decode(&token).expect("decode");
409        assert_eq!(decoded.sub, "user123");
410        assert_eq!(decoded.iss, Some("test-issuer".to_string()));
411        assert_eq!(
412            decoded.roles,
413            vec!["user".to_string(), "editor".to_string()]
414        );
415        assert_eq!(
416            decoded.permissions,
417            vec!["read:posts".to_string(), "write:posts".to_string()]
418        );
419    }
420
421    #[test]
422    fn test_jwt_format_is_header_payload_signature() {
423        let encoder = JwtEncoder::new("secret");
424        let claims = JwtClaims::new("alice", now() + 60);
425        let token = encoder.encode(&claims).unwrap();
426        let parts: Vec<&str> = token.split('.').collect();
427        assert_eq!(parts.len(), 3);
428
429        // Header should decode to {"alg":"HS256","typ":"JWT"}
430        let header_bytes = base64_url_decode(parts[0]).unwrap();
431        let header: JwtHeader = serde_json::from_slice(&header_bytes).unwrap();
432        assert_eq!(header.alg, "HS256");
433        assert_eq!(header.typ, "JWT");
434    }
435
436    #[test]
437    fn test_jwt_signature_changes_with_secret() {
438        let encoder_a = JwtEncoder::new("secret-a");
439        let encoder_b = JwtEncoder::new("secret-b");
440        let claims = JwtClaims::new("user", now() + 3600);
441
442        let token_a = encoder_a.encode(&claims).unwrap();
443        let token_b = encoder_b.encode(&claims).unwrap();
444
445        // Header.payload should be the same, but signature differs.
446        let parts_a: Vec<&str> = token_a.split('.').collect();
447        let parts_b: Vec<&str> = token_b.split('.').collect();
448        assert_eq!(parts_a[0], parts_b[0]); // header
449        assert_eq!(parts_a[1], parts_b[1]); // payload
450        assert_ne!(parts_a[2], parts_b[2]); // signature
451    }
452
453    #[test]
454    fn test_jwt_verify_with_wrong_secret_fails() {
455        let encoder_a = JwtEncoder::new("secret-a");
456        let encoder_b = JwtEncoder::new("secret-b");
457        let claims = JwtClaims::new("user", now() + 3600);
458
459        let token = encoder_a.encode(&claims).unwrap();
460        let result = encoder_b.decode(&token);
461        assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
462    }
463
464    #[test]
465    fn test_jwt_expired_token_rejected() {
466        let encoder = JwtEncoder::new("secret");
467        let claims = JwtClaims::new("user", now() - 100); // expired 100s ago
468        let token = encoder.encode(&claims).unwrap();
469        let result = encoder.decode(&token);
470        assert!(matches!(result, Err(AuthError::TokenExpired(_))));
471    }
472
473    #[test]
474    fn test_jwt_decode_invalid_format() {
475        let encoder = JwtEncoder::new("secret");
476        assert!(matches!(
477            encoder.decode(""),
478            Err(AuthError::TokenInvalid(_))
479        ));
480        assert!(matches!(
481            encoder.decode("not.a.jwt.token"),
482            Err(AuthError::TokenInvalid(_))
483        ));
484        assert!(matches!(
485            encoder.decode("only.two"),
486            Err(AuthError::TokenInvalid(_))
487        ));
488    }
489
490    #[test]
491    fn test_jwt_tampered_payload_rejected() {
492        let encoder = JwtEncoder::new("secret");
493        let claims = JwtClaims::new("alice", now() + 3600);
494        let token = encoder.encode(&claims).unwrap();
495
496        // Tamper with the payload by replacing it with a different valid base64url string.
497        let parts: Vec<&str> = token.split('.').collect();
498        let tampered_payload = base64_url_encode(
499            br#"{"sub":"mallory","exp":9999999999,"iat":0,"roles":[],"permissions":[]}"#,
500        );
501        let tampered = format!("{}.{}.{}", parts[0], tampered_payload, parts[2]);
502        let result = encoder.decode(&tampered);
503        assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
504    }
505
506    #[test]
507    fn test_jwt_tampered_signature_rejected() {
508        let encoder = JwtEncoder::new("secret");
509        let claims = JwtClaims::new("alice", now() + 3600);
510        let token = encoder.encode(&claims).unwrap();
511
512        let parts: Vec<&str> = token.split('.').collect();
513        // Flip the first character of the signature
514        let mut sig = parts[2].to_string();
515        let first = sig.chars().next().unwrap();
516        let replacement = if first == 'A' { 'B' } else { 'A' };
517        sig.replace_range(0..first.len_utf8(), &replacement.to_string());
518        let tampered = format!("{}.{}.{}", parts[0], parts[1], sig);
519        let result = encoder.decode(&tampered);
520        assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
521    }
522
523    fn hex_str(bytes: &[u8]) -> String {
524        let mut s = String::with_capacity(bytes.len() * 2);
525        for b in bytes {
526            s.push_str(&format!("{:02x}", b));
527        }
528        s
529    }
530}