Skip to main content

sz_orm_auth/
mfa.rs

1//! 多因素认证(MFA)框架
2//!
3//! 实现基于 TOTP(RFC 6238)的 MFA 验证:
4//! - 生成密钥和 otpauth URI
5//! - 生成当前时间窗口的 6 位 TOTP 码
6//! - 验证用户提交的 TOTP 码(允许 ±1 个时间窗口漂移)
7//!
8//! TOTP 算法基于 HOTP(RFC 4226),使用 HMAC-SHA1 和 30 秒时间步长。
9
10use crate::error::AuthError;
11use hmac::{Hmac, Mac};
12use sha1::Sha1;
13
14type HmacSha1 = Hmac<Sha1>;
15
16/// TOTP 时间步长(秒),RFC 6239 默认 30 秒
17const TIME_STEP: u64 = 30;
18/// TOTP 码位数
19const CODE_DIGITS: usize = 6;
20/// 允许的时间窗口漂移(前后各 1 个窗口)
21const ALLOWED_DRIFT: u64 = 1;
22
23/// MFA 密钥(Base32 编码的随机字节)
24#[derive(Debug, Clone)]
25pub struct MfaSecret {
26    /// Base32 编码的密钥
27    pub base32_secret: String,
28    /// 关联的账户名
29    pub account: String,
30    /// 发行方名称
31    pub issuer: String,
32}
33
34impl MfaSecret {
35    /// 从原始字节创建 MFA 密钥
36    pub fn new(account: impl Into<String>, issuer: impl Into<String>) -> Self {
37        let raw = generate_random_bytes(20);
38        Self {
39            base32_secret: base32_encode(&raw),
40            account: account.into(),
41            issuer: issuer.into(),
42        }
43    }
44
45    /// 从已有 Base32 密钥创建
46    pub fn from_base32(
47        base32_secret: impl Into<String>,
48        account: impl Into<String>,
49        issuer: impl Into<String>,
50    ) -> Self {
51        Self {
52            base32_secret: base32_secret.into(),
53            account: account.into(),
54            issuer: issuer.into(),
55        }
56    }
57
58    /// 生成 otpauth URI(用于二维码扫描)
59    pub fn to_uri(&self) -> String {
60        format!(
61            "otpauth://totp/{}:{}?secret={}&issuer={}",
62            self.issuer, self.account, self.base32_secret, self.issuer
63        )
64    }
65}
66
67/// TOTP 验证器
68pub struct TotpVerifier {
69    time_step: u64,
70    digits: usize,
71    drift: u64,
72}
73
74impl TotpVerifier {
75    /// 创建默认 TOTP 验证器(30 秒步长,6 位码,±1 窗口漂移)
76    pub fn new() -> Self {
77        Self {
78            time_step: TIME_STEP,
79            digits: CODE_DIGITS,
80            drift: ALLOWED_DRIFT,
81        }
82    }
83
84    /// 配置时间步长
85    pub fn with_time_step(mut self, step: u64) -> Self {
86        self.time_step = step.max(1);
87        self
88    }
89
90    /// 配置允许的时间窗口漂移
91    pub fn with_drift(mut self, drift: u64) -> Self {
92        self.drift = drift;
93        self
94    }
95
96    /// 生成指定时间戳的 TOTP 码
97    pub fn generate_at(&self, base32_secret: &str, timestamp: u64) -> String {
98        let counter = timestamp / self.time_step;
99        self.generate_hotp(base32_secret, counter)
100    }
101
102    /// 生成当前时间的 TOTP 码
103    pub fn generate_now(&self, base32_secret: &str) -> String {
104        self.generate_at(base32_secret, current_secs())
105    }
106
107    /// 验证 TOTP 码(允许 ±drift 个时间窗口漂移)
108    pub fn verify(&self, base32_secret: &str, code: &str) -> bool {
109        self.verify_at(base32_secret, code, current_secs())
110    }
111
112    /// 验证指定时间戳的 TOTP 码
113    pub fn verify_at(&self, base32_secret: &str, code: &str, timestamp: u64) -> bool {
114        // v4.8.0 修复 M-10:空/非法 base32 密钥的 HOTP 输出恒为 "000000",
115        // 且该恒定码会被 verify 接受(黑帽实证:空密钥账户一次即过 MFA)。
116        // 验证入口直接拒绝空密钥,杜绝恒定码绕过。
117        let key = base32_decode(base32_secret).unwrap_or_default();
118        if key.is_empty() {
119            return false;
120        }
121
122        let counter = timestamp / self.time_step;
123        // 检查当前窗口及前后 drift 个窗口
124        for offset in 0..=self.drift {
125            let test_counter = counter.saturating_sub(offset);
126            if constant_time_eq(
127                self.generate_hotp(base32_secret, test_counter).as_bytes(),
128                code.as_bytes(),
129            ) {
130                return true;
131            }
132            if offset > 0 {
133                let test_counter = counter.saturating_add(offset);
134                if constant_time_eq(
135                    self.generate_hotp(base32_secret, test_counter).as_bytes(),
136                    code.as_bytes(),
137                ) {
138                    return true;
139                }
140            }
141        }
142        false
143    }
144
145    /// HOTP 算法(RFC 4226)
146    fn generate_hotp(&self, base32_secret: &str, counter: u64) -> String {
147        let key = base32_decode(base32_secret).unwrap_or_default();
148        if key.is_empty() {
149            return "0".repeat(self.digits);
150        }
151
152        let mut mac = match <HmacSha1 as Mac>::new_from_slice(&key) {
153            Ok(m) => m,
154            Err(_) => return "0".repeat(self.digits),
155        };
156
157        let counter_bytes = counter.to_be_bytes();
158        mac.update(&counter_bytes);
159        let hash = mac.finalize().into_bytes();
160
161        // Dynamic truncation
162        let offset = (hash[hash.len() - 1] & 0x0F) as usize;
163        let truncated: u32 = (((hash[offset] & 0x7F) as u32) << 24)
164            | ((hash[offset + 1] as u32) << 16)
165            | ((hash[offset + 2] as u32) << 8)
166            | (hash[offset + 3] as u32);
167
168        let code = truncated % (10u32.pow(self.digits as u32));
169        format!("{:0width$}", code, width = self.digits)
170    }
171}
172
173impl Default for TotpVerifier {
174    fn default() -> Self {
175        Self::new()
176    }
177}
178
179/// MFA 管理器:管理用户的 MFA 密钥和验证状态
180pub struct MfaManager {
181    verifier: TotpVerifier,
182    secrets: parking_lot::Mutex<std::collections::HashMap<String, MfaSecret>>,
183}
184
185impl MfaManager {
186    pub fn new() -> Self {
187        Self {
188            verifier: TotpVerifier::new(),
189            secrets: parking_lot::Mutex::new(std::collections::HashMap::new()),
190        }
191    }
192
193    /// 为用户生成新的 MFA 密钥
194    pub fn generate_secret(
195        &self,
196        user_id: &str,
197        account: impl Into<String>,
198        issuer: impl Into<String>,
199    ) -> MfaSecret {
200        let secret = MfaSecret::new(account, issuer);
201        self.secrets
202            .lock()
203            .insert(user_id.to_string(), secret.clone());
204        secret
205    }
206
207    /// 为用户绑定已有密钥
208    pub fn bind_secret(&self, user_id: &str, secret: MfaSecret) {
209        self.secrets.lock().insert(user_id.to_string(), secret);
210    }
211
212    /// 验证用户的 TOTP 码
213    pub fn verify(&self, user_id: &str, code: &str) -> Result<bool, AuthError> {
214        let secrets = self.secrets.lock();
215        let secret = secrets
216            .get(user_id)
217            .ok_or_else(|| AuthError::Config(format!("No MFA secret for user: {}", user_id)))?;
218        Ok(self.verifier.verify(&secret.base32_secret, code))
219    }
220
221    /// 生成用户当前的 TOTP 码(用于测试或重置流程)
222    pub fn generate_code(&self, user_id: &str) -> Result<String, AuthError> {
223        let secrets = self.secrets.lock();
224        let secret = secrets
225            .get(user_id)
226            .ok_or_else(|| AuthError::Config(format!("No MFA secret for user: {}", user_id)))?;
227        Ok(self.verifier.generate_now(&secret.base32_secret))
228    }
229
230    /// 移除用户的 MFA 密钥
231    pub fn remove_secret(&self, user_id: &str) -> bool {
232        self.secrets.lock().remove(user_id).is_some()
233    }
234
235    /// 检查用户是否已绑定 MFA
236    pub fn has_mfa(&self, user_id: &str) -> bool {
237        self.secrets.lock().contains_key(user_id)
238    }
239
240    /// 获取用户 MFA 密钥的 otpauth URI
241    pub fn get_uri(&self, user_id: &str) -> Result<String, AuthError> {
242        let secrets = self.secrets.lock();
243        let secret = secrets
244            .get(user_id)
245            .ok_or_else(|| AuthError::Config(format!("No MFA secret for user: {}", user_id)))?;
246        Ok(secret.to_uri())
247    }
248}
249
250impl Default for MfaManager {
251    fn default() -> Self {
252        Self::new()
253    }
254}
255
256// ============================================================================
257// 辅助函数
258// ============================================================================
259
260fn current_secs() -> u64 {
261    use std::time::{SystemTime, UNIX_EPOCH};
262    SystemTime::now()
263        .duration_since(UNIX_EPOCH)
264        .unwrap_or_default()
265        .as_secs()
266}
267
268/// 生成密码学安全的随机字节序列
269///
270/// v1.2.1 修复 Critical C-1(CWE-338):使用 `OsRng`(密码学安全 RNG)替代
271/// `DefaultHasher` + 时间种子。原实现可被预测,攻击者可在 ±1 秒窗口内
272/// 暴力枚举纳秒种子复原 MFA 密钥,绕过 TOTP 二因素认证(违反 RFC 6238 §4)。
273///
274/// 与 `sz-orm-crypto::random_bytes` 实现保持一致。
275fn generate_random_bytes(len: usize) -> Vec<u8> {
276    use rand::rngs::OsRng;
277    use rand::RngCore;
278    let mut result = vec![0u8; len];
279    OsRng.fill_bytes(&mut result);
280    result
281}
282
283/// Base32 编码(RFC 4648,无填充)
284fn base32_encode(data: &[u8]) -> String {
285    const ALPHABET: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZ234567";
286    let mut result = String::new();
287    let mut buffer: u32 = 0;
288    let mut bits_left = 0;
289    for &byte in data {
290        buffer = (buffer << 8) | (byte as u32);
291        bits_left += 8;
292        while bits_left >= 5 {
293            bits_left -= 5;
294            let idx = ((buffer >> bits_left) & 0x1F) as usize;
295            result.push(ALPHABET[idx] as char);
296        }
297    }
298    if bits_left > 0 {
299        let idx = ((buffer << (5 - bits_left)) & 0x1F) as usize;
300        result.push(ALPHABET[idx] as char);
301    }
302    result
303}
304
305/// Base32 解码(RFC 4648,无填充)
306fn base32_decode(data: &str) -> Result<Vec<u8>, ()> {
307    const ALPHABET: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZ234567";
308    let mut result = Vec::new();
309    let mut buffer: u32 = 0;
310    let mut bits_left: u32 = 0;
311    for ch in data.chars() {
312        let upper = ch.to_ascii_uppercase();
313        let idx = ALPHABET.iter().position(|&c| c == upper as u8).ok_or(())?;
314        buffer = (buffer << 5) | (idx as u32);
315        bits_left += 5;
316        if bits_left >= 8 {
317            bits_left -= 8;
318            let byte = ((buffer >> bits_left) & 0xFF) as u8;
319            result.push(byte);
320        }
321    }
322    Ok(result)
323}
324
325/// 常量时间比较
326fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
327    use subtle::ConstantTimeEq;
328    a.ct_eq(b).into()
329}
330
331#[cfg(test)]
332mod tests {
333    use super::*;
334
335    #[test]
336    fn test_base32_encode_decode_roundtrip() {
337        let original = b"Hello, MFA!";
338        let encoded = base32_encode(original);
339        let decoded = base32_decode(&encoded).unwrap();
340        assert_eq!(decoded, original);
341    }
342
343    #[test]
344    fn test_base32_encode_known() {
345        // RFC 4648 测试向量(去除填充)
346        assert_eq!(base32_encode(b""), "");
347        assert_eq!(base32_encode(b"f"), "MY");
348        assert_eq!(base32_encode(b"fo"), "MZXQ");
349        assert_eq!(base32_encode(b"foo"), "MZXW6");
350        assert_eq!(base32_encode(b"foob"), "MZXW6YQ");
351        assert_eq!(base32_encode(b"fooba"), "MZXW6YTB");
352        assert_eq!(base32_encode(b"foobar"), "MZXW6YTBOI");
353    }
354
355    #[test]
356    fn test_base32_decode_known() {
357        assert_eq!(base32_decode("").unwrap(), b"");
358        assert_eq!(base32_decode("MY").unwrap(), b"f");
359        assert_eq!(base32_decode("MZXQ").unwrap(), b"fo");
360        assert_eq!(base32_decode("MZXW6").unwrap(), b"foo");
361    }
362
363    #[test]
364    fn test_base32_decode_lowercase() {
365        assert_eq!(base32_decode("mzxw6").unwrap(), b"foo");
366    }
367
368    #[test]
369    fn test_base32_decode_invalid_char() {
370        assert!(base32_decode("INVALID!").is_err());
371        assert!(base32_decode("1").is_err());
372    }
373
374    #[test]
375    fn test_mfa_secret_new() {
376        let secret = MfaSecret::new("user@test.com", "TestApp");
377        assert!(!secret.base32_secret.is_empty());
378        assert_eq!(secret.account, "user@test.com");
379        assert_eq!(secret.issuer, "TestApp");
380    }
381
382    #[test]
383    fn test_mfa_secret_from_base32() {
384        let secret = MfaSecret::from_base32("JBSWY3DPEHPK3PXP", "alice", "MyApp");
385        assert_eq!(secret.base32_secret, "JBSWY3DPEHPK3PXP");
386        assert_eq!(secret.account, "alice");
387    }
388
389    #[test]
390    fn test_mfa_secret_to_uri() {
391        let secret = MfaSecret::from_base32("JBSWY3DPEHPK3PXP", "alice", "MyApp");
392        let uri = secret.to_uri();
393        assert!(uri.starts_with("otpauth://totp/MyApp:alice?"));
394        assert!(uri.contains("secret=JBSWY3DPEHPK3PXP"));
395        assert!(uri.contains("issuer=MyApp"));
396    }
397
398    #[test]
399    fn test_totp_verifier_generate_format() {
400        let verifier = TotpVerifier::new();
401        let code = verifier.generate_now("JBSWY3DPEHPK3PXP");
402        assert_eq!(code.len(), 6);
403        assert!(code.chars().all(|c| c.is_ascii_digit()));
404    }
405
406    #[test]
407    fn test_totp_verifier_generate_at_deterministic() {
408        let verifier = TotpVerifier::new();
409        let code1 = verifier.generate_at("JBSWY3DPEHPK3PXP", 1000000);
410        let code2 = verifier.generate_at("JBSWY3DPEHPK3PXP", 1000000);
411        assert_eq!(code1, code2);
412    }
413
414    #[test]
415    fn test_totp_verifier_generate_different_timestamps() {
416        let verifier = TotpVerifier::new();
417        // 1000000 / 30 = 33333, 1000030 / 30 = 33334(不同计数器)
418        let code1 = verifier.generate_at("JBSWY3DPEHPK3PXP", 1000000);
419        let code2 = verifier.generate_at("JBSWY3DPEHPK3PXP", 1000030);
420        assert_ne!(code1, code2);
421    }
422
423    #[test]
424    fn test_totp_verifier_verify_correct() {
425        let verifier = TotpVerifier::new();
426        let timestamp = 1000000u64;
427        let code = verifier.generate_at("JBSWY3DPEHPK3PXP", timestamp);
428        assert!(verifier.verify_at("JBSWY3DPEHPK3PXP", &code, timestamp));
429    }
430
431    #[test]
432    fn test_totp_verifier_verify_wrong_code() {
433        let verifier = TotpVerifier::new();
434        assert!(!verifier.verify_at("JBSWY3DPEHPK3PXP", "000000", 1000000));
435    }
436
437    #[test]
438    fn test_totp_verifier_verify_within_drift() {
439        let verifier = TotpVerifier::new();
440        let timestamp = 1000000u64;
441        let code = verifier.generate_at("JBSWY3DPEHPK3PXP", timestamp);
442        // 前后 1 个窗口(30 秒)内应验证通过
443        assert!(verifier.verify_at("JBSWY3DPEHPK3PXP", &code, timestamp + 30));
444        assert!(verifier.verify_at("JBSWY3DPEHPK3PXP", &code, timestamp - 30));
445    }
446
447    #[test]
448    fn test_totp_verifier_verify_outside_drift() {
449        let verifier = TotpVerifier::new();
450        let timestamp = 1000000u64;
451        let code = verifier.generate_at("JBSWY3DPEHPK3PXP", timestamp);
452        // 超出漂移范围(2 个窗口 = 60 秒)
453        assert!(!verifier.verify_at("JBSWY3DPEHPK3PXP", &code, timestamp + 60));
454        assert!(!verifier.verify_at("JBSWY3DPEHPK3PXP", &code, timestamp - 60));
455    }
456
457    #[test]
458    fn test_totp_verifier_different_secrets_different_codes() {
459        let verifier = TotpVerifier::new();
460        let code1 = verifier.generate_at("JBSWY3DPEHPK3PXP", 1000000);
461        let code2 = verifier.generate_at("JBSWY3DPEHPK3PXP", 1000000);
462        let code3 = verifier.generate_at("GEZDGNBVGY3TQOJQ", 1000000);
463        assert_eq!(code1, code2);
464        assert_ne!(code1, code3);
465    }
466
467    #[test]
468    fn test_totp_verifier_with_drift_zero() {
469        let verifier = TotpVerifier::new().with_drift(0);
470        let timestamp = 1000000u64;
471        let code = verifier.generate_at("JBSWY3DPEHPK3PXP", timestamp);
472        assert!(verifier.verify_at("JBSWY3DPEHPK3PXP", &code, timestamp));
473        // 0 漂移:前后 30 秒应失败
474        assert!(!verifier.verify_at("JBSWY3DPEHPK3PXP", &code, timestamp + 30));
475    }
476
477    #[test]
478    fn test_totp_verifier_with_custom_time_step() {
479        let verifier = TotpVerifier::new().with_time_step(60);
480        let timestamp = 1000000u64;
481        let code = verifier.generate_at("JBSWY3DPEHPK3PXP", timestamp);
482        assert_eq!(code.len(), 6);
483        assert!(verifier.verify_at("JBSWY3DPEHPK3PXP", &code, timestamp));
484        // 60 秒步长:30 秒后仍在同一窗口
485        assert!(verifier.verify_at("JBSWY3DPEHPK3PXP", &code, timestamp + 30));
486    }
487
488    #[test]
489    fn test_totp_verifier_empty_secret() {
490        let verifier = TotpVerifier::new();
491        let code = verifier.generate_at("", 1000000);
492        assert_eq!(code, "000000");
493    }
494
495    #[test]
496    fn test_mfa_manager_generate_secret() {
497        let mgr = MfaManager::new();
498        let secret = mgr.generate_secret("user1", "user1@test.com", "TestApp");
499        assert!(!secret.base32_secret.is_empty());
500        assert!(mgr.has_mfa("user1"));
501    }
502
503    #[test]
504    fn test_mfa_manager_verify_correct() {
505        let mgr = MfaManager::new();
506        mgr.generate_secret("user1", "user1@test.com", "TestApp");
507        let code = mgr.generate_code("user1").unwrap();
508        assert!(mgr.verify("user1", &code).unwrap());
509    }
510
511    #[test]
512    fn test_mfa_manager_verify_wrong_code() {
513        let mgr = MfaManager::new();
514        mgr.generate_secret("user1", "user1@test.com", "TestApp");
515        assert!(!mgr.verify("user1", "000000").unwrap());
516    }
517
518    #[test]
519    fn test_mfa_manager_no_secret_errors() {
520        let mgr = MfaManager::new();
521        assert!(mgr.verify("unknown", "123456").is_err());
522        assert!(mgr.generate_code("unknown").is_err());
523        assert!(mgr.get_uri("unknown").is_err());
524    }
525
526    #[test]
527    fn test_mfa_manager_remove_secret() {
528        let mgr = MfaManager::new();
529        mgr.generate_secret("user1", "user1@test.com", "TestApp");
530        assert!(mgr.has_mfa("user1"));
531        assert!(mgr.remove_secret("user1"));
532        assert!(!mgr.has_mfa("user1"));
533    }
534
535    #[test]
536    fn test_mfa_manager_remove_nonexistent() {
537        let mgr = MfaManager::new();
538        assert!(!mgr.remove_secret("unknown"));
539    }
540
541    #[test]
542    fn test_mfa_manager_bind_secret() {
543        let mgr = MfaManager::new();
544        let secret = MfaSecret::from_base32("JBSWY3DPEHPK3PXP", "bob", "App");
545        mgr.bind_secret("bob", secret);
546        assert!(mgr.has_mfa("bob"));
547        let code = mgr.generate_code("bob").unwrap();
548        assert!(mgr.verify("bob", &code).unwrap());
549    }
550
551    #[test]
552    fn test_mfa_manager_get_uri() {
553        let mgr = MfaManager::new();
554        mgr.generate_secret("user1", "user1@test.com", "TestApp");
555        let uri = mgr.get_uri("user1").unwrap();
556        assert!(uri.starts_with("otpauth://totp/"));
557        assert!(uri.contains("user1@test.com"));
558    }
559
560    #[test]
561    fn test_mfa_manager_default() {
562        let mgr = MfaManager::default();
563        assert!(!mgr.has_mfa("anyone"));
564    }
565}