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: 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 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 pub fn bind_secret(&self, user_id: &str, secret: MfaSecret) {
201 self.secrets.lock().insert(user_id.to_string(), secret);
202 }
203
204 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 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 pub fn remove_secret(&self, user_id: &str) -> bool {
224 self.secrets.lock().remove(user_id).is_some()
225 }
226
227 pub fn has_mfa(&self, user_id: &str) -> bool {
229 self.secrets.lock().contains_key(user_id)
230 }
231
232 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
248fn 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
260fn 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
275fn 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
297fn 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
317fn 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 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 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 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 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 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 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}