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 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 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 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 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
179pub 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 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 pub fn bind_secret(&self, user_id: &str, secret: MfaSecret) {
209 self.secrets.lock().insert(user_id.to_string(), secret);
210 }
211
212 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 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 pub fn remove_secret(&self, user_id: &str) -> bool {
232 self.secrets.lock().remove(user_id).is_some()
233 }
234
235 pub fn has_mfa(&self, user_id: &str) -> bool {
237 self.secrets.lock().contains_key(user_id)
238 }
239
240 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
256fn 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
268fn 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
283fn 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
305fn 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
325fn 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 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 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 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 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 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 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}