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: std::sync::Mutex<std::collections::HashMap<String, MfaSecret>>,
175}
176
177impl MfaManager {
178    pub fn new() -> Self {
179        Self {
180            verifier: TotpVerifier::new(),
181            secrets: std::sync::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            .unwrap()
196            .insert(user_id.to_string(), secret.clone());
197        secret
198    }
199
200    /// 为用户绑定已有密钥
201    pub fn bind_secret(&self, user_id: &str, secret: MfaSecret) {
202        self.secrets
203            .lock()
204            .unwrap()
205            .insert(user_id.to_string(), secret);
206    }
207
208    /// 验证用户的 TOTP 码
209    pub fn verify(&self, user_id: &str, code: &str) -> Result<bool, AuthError> {
210        let secrets = self.secrets.lock().unwrap();
211        let secret = secrets
212            .get(user_id)
213            .ok_or_else(|| AuthError::Config(format!("No MFA secret for user: {}", user_id)))?;
214        Ok(self.verifier.verify(&secret.base32_secret, code))
215    }
216
217    /// 生成用户当前的 TOTP 码(用于测试或重置流程)
218    pub fn generate_code(&self, user_id: &str) -> Result<String, AuthError> {
219        let secrets = self.secrets.lock().unwrap();
220        let secret = secrets
221            .get(user_id)
222            .ok_or_else(|| AuthError::Config(format!("No MFA secret for user: {}", user_id)))?;
223        Ok(self.verifier.generate_now(&secret.base32_secret))
224    }
225
226    /// 移除用户的 MFA 密钥
227    pub fn remove_secret(&self, user_id: &str) -> bool {
228        self.secrets.lock().unwrap().remove(user_id).is_some()
229    }
230
231    /// 检查用户是否已绑定 MFA
232    pub fn has_mfa(&self, user_id: &str) -> bool {
233        self.secrets.lock().unwrap().contains_key(user_id)
234    }
235
236    /// 获取用户 MFA 密钥的 otpauth URI
237    pub fn get_uri(&self, user_id: &str) -> Result<String, AuthError> {
238        let secrets = self.secrets.lock().unwrap();
239        let secret = secrets
240            .get(user_id)
241            .ok_or_else(|| AuthError::Config(format!("No MFA secret for user: {}", user_id)))?;
242        Ok(secret.to_uri())
243    }
244}
245
246impl Default for MfaManager {
247    fn default() -> Self {
248        Self::new()
249    }
250}
251
252// ============================================================================
253// 辅助函数
254// ============================================================================
255
256fn current_secs() -> u64 {
257    use std::time::{SystemTime, UNIX_EPOCH};
258    SystemTime::now()
259        .duration_since(UNIX_EPOCH)
260        .unwrap_or_default()
261        .as_secs()
262}
263
264fn generate_random_bytes(len: usize) -> Vec<u8> {
265    use std::collections::hash_map::DefaultHasher;
266    use std::hash::{Hash, Hasher};
267    let mut result = Vec::with_capacity(len);
268    let mut seed = std::time::SystemTime::now()
269        .duration_since(std::time::UNIX_EPOCH)
270        .unwrap_or_default()
271        .as_nanos();
272    for i in 0..len {
273        let mut hasher = DefaultHasher::new();
274        seed.wrapping_add(i as u128).hash(&mut hasher);
275        let h = hasher.finish();
276        result.push((h & 0xFF) as u8);
277        seed = seed.wrapping_add(h as u128);
278    }
279    result
280}
281
282/// Base32 编码(RFC 4648,无填充)
283fn base32_encode(data: &[u8]) -> String {
284    const ALPHABET: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZ234567";
285    let mut result = String::new();
286    let mut buffer: u32 = 0;
287    let mut bits_left = 0;
288    for &byte in data {
289        buffer = (buffer << 8) | (byte as u32);
290        bits_left += 8;
291        while bits_left >= 5 {
292            bits_left -= 5;
293            let idx = ((buffer >> bits_left) & 0x1F) as usize;
294            result.push(ALPHABET[idx] as char);
295        }
296    }
297    if bits_left > 0 {
298        let idx = ((buffer << (5 - bits_left)) & 0x1F) as usize;
299        result.push(ALPHABET[idx] as char);
300    }
301    result
302}
303
304/// Base32 解码(RFC 4648,无填充)
305fn base32_decode(data: &str) -> Result<Vec<u8>, ()> {
306    const ALPHABET: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZ234567";
307    let mut result = Vec::new();
308    let mut buffer: u32 = 0;
309    let mut bits_left: u32 = 0;
310    for ch in data.chars() {
311        let upper = ch.to_ascii_uppercase();
312        let idx = ALPHABET.iter().position(|&c| c == upper as u8).ok_or(())?;
313        buffer = (buffer << 5) | (idx as u32);
314        bits_left += 5;
315        if bits_left >= 8 {
316            bits_left -= 8;
317            let byte = ((buffer >> bits_left) & 0xFF) as u8;
318            result.push(byte);
319        }
320    }
321    Ok(result)
322}
323
324/// 常量时间比较
325fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
326    use subtle::ConstantTimeEq;
327    a.ct_eq(b).into()
328}
329
330#[cfg(test)]
331mod tests {
332    use super::*;
333
334    #[test]
335    fn test_base32_encode_decode_roundtrip() {
336        let original = b"Hello, MFA!";
337        let encoded = base32_encode(original);
338        let decoded = base32_decode(&encoded).unwrap();
339        assert_eq!(decoded, original);
340    }
341
342    #[test]
343    fn test_base32_encode_known() {
344        // RFC 4648 测试向量(去除填充)
345        assert_eq!(base32_encode(b""), "");
346        assert_eq!(base32_encode(b"f"), "MY");
347        assert_eq!(base32_encode(b"fo"), "MZXQ");
348        assert_eq!(base32_encode(b"foo"), "MZXW6");
349        assert_eq!(base32_encode(b"foob"), "MZXW6YQ");
350        assert_eq!(base32_encode(b"fooba"), "MZXW6YTB");
351        assert_eq!(base32_encode(b"foobar"), "MZXW6YTBOI");
352    }
353
354    #[test]
355    fn test_base32_decode_known() {
356        assert_eq!(base32_decode("").unwrap(), b"");
357        assert_eq!(base32_decode("MY").unwrap(), b"f");
358        assert_eq!(base32_decode("MZXQ").unwrap(), b"fo");
359        assert_eq!(base32_decode("MZXW6").unwrap(), b"foo");
360    }
361
362    #[test]
363    fn test_base32_decode_lowercase() {
364        assert_eq!(base32_decode("mzxw6").unwrap(), b"foo");
365    }
366
367    #[test]
368    fn test_base32_decode_invalid_char() {
369        assert!(base32_decode("INVALID!").is_err());
370        assert!(base32_decode("1").is_err());
371    }
372
373    #[test]
374    fn test_mfa_secret_new() {
375        let secret = MfaSecret::new("user@test.com", "TestApp");
376        assert!(!secret.base32_secret.is_empty());
377        assert_eq!(secret.account, "user@test.com");
378        assert_eq!(secret.issuer, "TestApp");
379    }
380
381    #[test]
382    fn test_mfa_secret_from_base32() {
383        let secret = MfaSecret::from_base32("JBSWY3DPEHPK3PXP", "alice", "MyApp");
384        assert_eq!(secret.base32_secret, "JBSWY3DPEHPK3PXP");
385        assert_eq!(secret.account, "alice");
386    }
387
388    #[test]
389    fn test_mfa_secret_to_uri() {
390        let secret = MfaSecret::from_base32("JBSWY3DPEHPK3PXP", "alice", "MyApp");
391        let uri = secret.to_uri();
392        assert!(uri.starts_with("otpauth://totp/MyApp:alice?"));
393        assert!(uri.contains("secret=JBSWY3DPEHPK3PXP"));
394        assert!(uri.contains("issuer=MyApp"));
395    }
396
397    #[test]
398    fn test_totp_verifier_generate_format() {
399        let verifier = TotpVerifier::new();
400        let code = verifier.generate_now("JBSWY3DPEHPK3PXP");
401        assert_eq!(code.len(), 6);
402        assert!(code.chars().all(|c| c.is_ascii_digit()));
403    }
404
405    #[test]
406    fn test_totp_verifier_generate_at_deterministic() {
407        let verifier = TotpVerifier::new();
408        let code1 = verifier.generate_at("JBSWY3DPEHPK3PXP", 1000000);
409        let code2 = verifier.generate_at("JBSWY3DPEHPK3PXP", 1000000);
410        assert_eq!(code1, code2);
411    }
412
413    #[test]
414    fn test_totp_verifier_generate_different_timestamps() {
415        let verifier = TotpVerifier::new();
416        // 1000000 / 30 = 33333, 1000030 / 30 = 33334(不同计数器)
417        let code1 = verifier.generate_at("JBSWY3DPEHPK3PXP", 1000000);
418        let code2 = verifier.generate_at("JBSWY3DPEHPK3PXP", 1000030);
419        assert_ne!(code1, code2);
420    }
421
422    #[test]
423    fn test_totp_verifier_verify_correct() {
424        let verifier = TotpVerifier::new();
425        let timestamp = 1000000u64;
426        let code = verifier.generate_at("JBSWY3DPEHPK3PXP", timestamp);
427        assert!(verifier.verify_at("JBSWY3DPEHPK3PXP", &code, timestamp));
428    }
429
430    #[test]
431    fn test_totp_verifier_verify_wrong_code() {
432        let verifier = TotpVerifier::new();
433        assert!(!verifier.verify_at("JBSWY3DPEHPK3PXP", "000000", 1000000));
434    }
435
436    #[test]
437    fn test_totp_verifier_verify_within_drift() {
438        let verifier = TotpVerifier::new();
439        let timestamp = 1000000u64;
440        let code = verifier.generate_at("JBSWY3DPEHPK3PXP", timestamp);
441        // 前后 1 个窗口(30 秒)内应验证通过
442        assert!(verifier.verify_at("JBSWY3DPEHPK3PXP", &code, timestamp + 30));
443        assert!(verifier.verify_at("JBSWY3DPEHPK3PXP", &code, timestamp - 30));
444    }
445
446    #[test]
447    fn test_totp_verifier_verify_outside_drift() {
448        let verifier = TotpVerifier::new();
449        let timestamp = 1000000u64;
450        let code = verifier.generate_at("JBSWY3DPEHPK3PXP", timestamp);
451        // 超出漂移范围(2 个窗口 = 60 秒)
452        assert!(!verifier.verify_at("JBSWY3DPEHPK3PXP", &code, timestamp + 60));
453        assert!(!verifier.verify_at("JBSWY3DPEHPK3PXP", &code, timestamp - 60));
454    }
455
456    #[test]
457    fn test_totp_verifier_different_secrets_different_codes() {
458        let verifier = TotpVerifier::new();
459        let code1 = verifier.generate_at("JBSWY3DPEHPK3PXP", 1000000);
460        let code2 = verifier.generate_at("JBSWY3DPEHPK3PXP", 1000000);
461        let code3 = verifier.generate_at("GEZDGNBVGY3TQOJQ", 1000000);
462        assert_eq!(code1, code2);
463        assert_ne!(code1, code3);
464    }
465
466    #[test]
467    fn test_totp_verifier_with_drift_zero() {
468        let verifier = TotpVerifier::new().with_drift(0);
469        let timestamp = 1000000u64;
470        let code = verifier.generate_at("JBSWY3DPEHPK3PXP", timestamp);
471        assert!(verifier.verify_at("JBSWY3DPEHPK3PXP", &code, timestamp));
472        // 0 漂移:前后 30 秒应失败
473        assert!(!verifier.verify_at("JBSWY3DPEHPK3PXP", &code, timestamp + 30));
474    }
475
476    #[test]
477    fn test_totp_verifier_with_custom_time_step() {
478        let verifier = TotpVerifier::new().with_time_step(60);
479        let timestamp = 1000000u64;
480        let code = verifier.generate_at("JBSWY3DPEHPK3PXP", timestamp);
481        assert_eq!(code.len(), 6);
482        assert!(verifier.verify_at("JBSWY3DPEHPK3PXP", &code, timestamp));
483        // 60 秒步长:30 秒后仍在同一窗口
484        assert!(verifier.verify_at("JBSWY3DPEHPK3PXP", &code, timestamp + 30));
485    }
486
487    #[test]
488    fn test_totp_verifier_empty_secret() {
489        let verifier = TotpVerifier::new();
490        let code = verifier.generate_at("", 1000000);
491        assert_eq!(code, "000000");
492    }
493
494    #[test]
495    fn test_mfa_manager_generate_secret() {
496        let mgr = MfaManager::new();
497        let secret = mgr.generate_secret("user1", "user1@test.com", "TestApp");
498        assert!(!secret.base32_secret.is_empty());
499        assert!(mgr.has_mfa("user1"));
500    }
501
502    #[test]
503    fn test_mfa_manager_verify_correct() {
504        let mgr = MfaManager::new();
505        mgr.generate_secret("user1", "user1@test.com", "TestApp");
506        let code = mgr.generate_code("user1").unwrap();
507        assert!(mgr.verify("user1", &code).unwrap());
508    }
509
510    #[test]
511    fn test_mfa_manager_verify_wrong_code() {
512        let mgr = MfaManager::new();
513        mgr.generate_secret("user1", "user1@test.com", "TestApp");
514        assert!(!mgr.verify("user1", "000000").unwrap());
515    }
516
517    #[test]
518    fn test_mfa_manager_no_secret_errors() {
519        let mgr = MfaManager::new();
520        assert!(mgr.verify("unknown", "123456").is_err());
521        assert!(mgr.generate_code("unknown").is_err());
522        assert!(mgr.get_uri("unknown").is_err());
523    }
524
525    #[test]
526    fn test_mfa_manager_remove_secret() {
527        let mgr = MfaManager::new();
528        mgr.generate_secret("user1", "user1@test.com", "TestApp");
529        assert!(mgr.has_mfa("user1"));
530        assert!(mgr.remove_secret("user1"));
531        assert!(!mgr.has_mfa("user1"));
532    }
533
534    #[test]
535    fn test_mfa_manager_remove_nonexistent() {
536        let mgr = MfaManager::new();
537        assert!(!mgr.remove_secret("unknown"));
538    }
539
540    #[test]
541    fn test_mfa_manager_bind_secret() {
542        let mgr = MfaManager::new();
543        let secret = MfaSecret::from_base32("JBSWY3DPEHPK3PXP", "bob", "App");
544        mgr.bind_secret("bob", secret);
545        assert!(mgr.has_mfa("bob"));
546        let code = mgr.generate_code("bob").unwrap();
547        assert!(mgr.verify("bob", &code).unwrap());
548    }
549
550    #[test]
551    fn test_mfa_manager_get_uri() {
552        let mgr = MfaManager::new();
553        mgr.generate_secret("user1", "user1@test.com", "TestApp");
554        let uri = mgr.get_uri("user1").unwrap();
555        assert!(uri.starts_with("otpauth://totp/"));
556        assert!(uri.contains("user1@test.com"));
557    }
558
559    #[test]
560    fn test_mfa_manager_default() {
561        let mgr = MfaManager::default();
562        assert!(!mgr.has_mfa("anyone"));
563    }
564}