1use crate::error::AuthError;
11use hmac::{Hmac, Mac};
12use sha1::Sha1;
13
14type HmacSha1 = Hmac<Sha1>;
15
16const TIME_STEP: u64 = 30;
18const CODE_DIGITS: usize = 6;
20const ALLOWED_DRIFT: u64 = 1;
22
23#[derive(Debug, Clone)]
25pub struct MfaSecret {
26 pub base32_secret: String,
28 pub account: String,
30 pub issuer: String,
32}
33
34impl MfaSecret {
35 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 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 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
67pub struct TotpVerifier {
69 time_step: u64,
70 digits: usize,
71 drift: u64,
72}
73
74impl TotpVerifier {
75 pub fn new() -> Self {
77 Self {
78 time_step: TIME_STEP,
79 digits: CODE_DIGITS,
80 drift: ALLOWED_DRIFT,
81 }
82 }
83
84 pub fn with_time_step(mut self, step: u64) -> Self {
86 self.time_step = step.max(1);
87 self
88 }
89
90 pub fn with_drift(mut self, drift: u64) -> Self {
92 self.drift = drift;
93 self
94 }
95
96 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 pub fn generate_now(&self, base32_secret: &str) -> String {
104 self.generate_at(base32_secret, current_secs())
105 }
106
107 pub fn verify(&self, base32_secret: &str, code: &str) -> bool {
109 self.verify_at(base32_secret, code, current_secs())
110 }
111
112 pub fn verify_at(&self, base32_secret: &str, code: &str, timestamp: u64) -> bool {
114 let counter = timestamp / self.time_step;
115 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 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 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
171pub 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 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 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 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 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 pub fn remove_secret(&self, user_id: &str) -> bool {
228 self.secrets.lock().unwrap().remove(user_id).is_some()
229 }
230
231 pub fn has_mfa(&self, user_id: &str) -> bool {
233 self.secrets.lock().unwrap().contains_key(user_id)
234 }
235
236 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
252fn 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
282fn 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
304fn 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
324fn 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 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 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 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 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 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 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}