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