1use std::sync::Arc;
23
24use async_trait::async_trait;
25use axum::{
26 Json,
27 extract::State,
28 http::StatusCode,
29 response::{IntoResponse, Response},
30};
31use dashmap::DashMap;
32use rand::RngCore as _;
33use serde::{Deserialize, Serialize};
34use totp_rs::{Algorithm, Secret, TOTP};
35
36use crate::{
37 audit::logger::{AuditEventType, SecretType, get_audit_logger},
38 error::{AuthError, Result},
39 session::{SessionStore, unix_now},
40};
41
42const RECOVERY_CODE_COUNT: usize = 8;
46
47const RECOVERY_CODE_HEX_LEN: usize = 16;
49
50const CHALLENGE_TTL_SECS: u64 = 300;
52
53const TOTP_STEP_TOLERANCE: u8 = 1;
55
56#[cfg(not(test))]
61const BCRYPT_COST: u32 = 12;
62#[cfg(test)]
63const BCRYPT_COST: u32 = 4;
64
65#[derive(Debug, Clone)]
69pub struct TotpEnrollment {
70 pub secret_base32: String,
72 pub recovery_code_hashes: Vec<String>,
74 pub confirmed: bool,
76}
77
78#[derive(Debug, Clone)]
80struct ChallengeRecord {
81 user_id: String,
83 expires: u64,
85}
86
87#[async_trait]
93pub trait MfaStore: Send + Sync {
94 async fn begin_enrollment(
104 &self,
105 user_id: &str,
106 issuer: &str,
107 account_name: &str,
108 ) -> Result<EnrollmentResponse>;
109
110 async fn confirm_enrollment(&self, user_id: &str, totp_code: &str) -> Result<()>;
117
118 async fn create_challenge(&self, user_id: &str) -> Result<String>;
124
125 async fn verify_challenge(&self, challenge_token: &str, code: &str) -> Result<String>;
133
134 async fn unenroll(&self, user_id: &str, code: &str) -> Result<()>;
140
141 async fn is_enrolled(&self, user_id: &str) -> bool;
143}
144
145#[derive(Debug)]
147pub struct EnrollmentResponse {
148 pub secret_base32: String,
150 pub otpauth_uri: String,
152 pub recovery_codes: Vec<String>,
154}
155
156pub struct InMemoryMfaStore {
160 enrollments: DashMap<String, TotpEnrollment>,
162 challenges: DashMap<String, ChallengeRecord>,
164}
165
166impl InMemoryMfaStore {
167 #[must_use]
169 pub fn new() -> Self {
170 Self {
171 enrollments: DashMap::new(),
172 challenges: DashMap::new(),
173 }
174 }
175
176 #[must_use]
178 pub fn has_pending_enrollment(&self, user_id: &str) -> bool {
179 self.enrollments.get(user_id).is_some_and(|e| !e.confirmed)
180 }
181}
182
183impl Default for InMemoryMfaStore {
184 fn default() -> Self {
185 Self::new()
186 }
187}
188
189fn build_totp(secret_base32: &str, issuer: Option<&str>, account_name: &str) -> Result<TOTP> {
197 let secret_bytes =
198 Secret::Encoded(secret_base32.to_string())
199 .to_bytes()
200 .map_err(|e| AuthError::Internal {
201 message: format!("bad TOTP secret: {e}"),
202 })?;
203 TOTP::new(
204 Algorithm::SHA1,
205 6, TOTP_STEP_TOLERANCE,
207 30, secret_bytes,
209 issuer.map(str::to_string),
210 account_name.to_string(),
211 )
212 .map_err(|e| AuthError::Internal {
213 message: format!("TOTP init error: {e}"),
214 })
215}
216
217fn verify_totp_code(secret_base32: &str, code: &str) -> Result<bool> {
219 let totp = build_totp(secret_base32, None, "")?;
220 Ok(totp.check_current(code).unwrap_or(false))
221}
222
223fn generate_recovery_code() -> String {
225 let byte_count = RECOVERY_CODE_HEX_LEN / 2;
228 let mut bytes = vec![0u8; byte_count];
229 rand::rng().fill_bytes(&mut bytes);
230 bytes.iter().fold(String::new(), |mut s, b| {
231 use std::fmt::Write as _;
232 let _ = write!(s, "{b:02x}");
233 s
234 })
235}
236
237fn generate_challenge_token() -> String {
239 use base64::Engine as _;
240 let mut bytes = [0u8; 32];
242 rand::rng().fill_bytes(&mut bytes);
243 base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(bytes)
244}
245
246fn check_recovery_code(candidate: &str, hashes: &[String]) -> Option<usize> {
250 for (i, hash) in hashes.iter().enumerate() {
251 if bcrypt::verify(candidate, hash).unwrap_or(false) {
252 return Some(i);
253 }
254 }
255 None
256}
257
258#[async_trait]
260impl MfaStore for InMemoryMfaStore {
261 async fn begin_enrollment(
262 &self,
263 user_id: &str,
264 issuer: &str,
265 account_name: &str,
266 ) -> Result<EnrollmentResponse> {
267 let secret = Secret::generate_secret();
269 let secret_base32 = secret.to_encoded().to_string();
270
271 let totp = build_totp(&secret_base32, Some(issuer), account_name)?;
273 let otpauth_uri = totp.get_url();
274
275 let mut recovery_codes_plain = Vec::with_capacity(RECOVERY_CODE_COUNT);
277 let mut recovery_code_hashes = Vec::with_capacity(RECOVERY_CODE_COUNT);
278 for _ in 0..RECOVERY_CODE_COUNT {
279 let code = generate_recovery_code();
280 let hash = bcrypt::hash(&code, BCRYPT_COST).map_err(|e| AuthError::Internal {
281 message: format!("bcrypt error: {e}"),
282 })?;
283 recovery_codes_plain.push(code);
284 recovery_code_hashes.push(hash);
285 }
286
287 self.enrollments.insert(
288 user_id.to_string(),
289 TotpEnrollment {
290 secret_base32: secret_base32.clone(),
291 recovery_code_hashes,
292 confirmed: false,
293 },
294 );
295
296 Ok(EnrollmentResponse {
297 secret_base32,
298 otpauth_uri,
299 recovery_codes: recovery_codes_plain,
300 })
301 }
302
303 async fn confirm_enrollment(&self, user_id: &str, totp_code: &str) -> Result<()> {
304 let mut record =
305 self.enrollments.get_mut(user_id).ok_or_else(|| AuthError::InvalidToken {
306 reason: "no pending MFA enrollment for user".into(),
307 })?;
308
309 if !verify_totp_code(&record.secret_base32, totp_code)? {
310 return Err(AuthError::InvalidToken {
311 reason: "invalid TOTP code".into(),
312 });
313 }
314 record.confirmed = true;
315 Ok(())
316 }
317
318 async fn create_challenge(&self, user_id: &str) -> Result<String> {
319 let expires = unix_now()? + CHALLENGE_TTL_SECS;
320 let token = generate_challenge_token();
321 self.challenges.insert(
322 token.clone(),
323 ChallengeRecord {
324 user_id: user_id.to_string(),
325 expires,
326 },
327 );
328 Ok(token)
329 }
330
331 async fn verify_challenge(&self, challenge_token: &str, code: &str) -> Result<String> {
332 let now = unix_now()?;
333
334 let record =
335 self.challenges.get(challenge_token).ok_or_else(|| AuthError::InvalidToken {
336 reason: "unknown challenge token".into(),
337 })?;
338
339 if now >= record.expires {
340 drop(record);
341 self.challenges.remove(challenge_token);
342 return Err(AuthError::InvalidToken {
343 reason: "challenge token expired".into(),
344 });
345 }
346
347 let user_id = record.user_id.clone();
348 drop(record);
349
350 let mut enrollment =
352 self.enrollments.get_mut(&user_id).ok_or_else(|| AuthError::InvalidToken {
353 reason: "user has no MFA enrollment".into(),
354 })?;
355
356 if !enrollment.confirmed {
357 return Err(AuthError::InvalidToken {
358 reason: "MFA enrollment not confirmed".into(),
359 });
360 }
361
362 if verify_totp_code(&enrollment.secret_base32, code)? {
364 drop(enrollment);
365 self.challenges.remove(challenge_token);
366 return Ok(user_id);
367 }
368
369 let idx = check_recovery_code(code, &enrollment.recovery_code_hashes);
371 if let Some(i) = idx {
372 enrollment.recovery_code_hashes.remove(i);
374 drop(enrollment);
375 self.challenges.remove(challenge_token);
376 return Ok(user_id);
377 }
378
379 Err(AuthError::InvalidToken {
380 reason: "invalid TOTP or recovery code".into(),
381 })
382 }
383
384 async fn unenroll(&self, user_id: &str, code: &str) -> Result<()> {
385 let enrollment = self.enrollments.get(user_id).ok_or_else(|| AuthError::InvalidToken {
386 reason: "user has no MFA enrollment".into(),
387 })?;
388
389 if !enrollment.confirmed {
390 return Err(AuthError::InvalidToken {
391 reason: "MFA enrollment not confirmed".into(),
392 });
393 }
394
395 let totp_ok = verify_totp_code(&enrollment.secret_base32, code)?;
397 let recovery_ok =
398 !totp_ok && check_recovery_code(code, &enrollment.recovery_code_hashes).is_some();
399
400 if !totp_ok && !recovery_ok {
401 return Err(AuthError::InvalidToken {
402 reason: "re-authentication failed — invalid TOTP or recovery code".into(),
403 });
404 }
405
406 drop(enrollment);
407 self.enrollments.remove(user_id);
408 Ok(())
409 }
410
411 async fn is_enrolled(&self, user_id: &str) -> bool {
412 self.enrollments.get(user_id).map_or(false, |e| e.confirmed)
413 }
414}
415
416#[derive(Clone)]
420pub struct MfaRouteState {
421 pub mfa_store: Arc<dyn MfaStore>,
423 pub session_store: Arc<dyn SessionStore>,
425 pub issuer: String,
427}
428
429#[derive(Debug, Deserialize)]
433pub struct MfaEnrollRequest {
434 pub user_id: String,
436 pub account_name: String,
438}
439
440#[derive(Debug, Serialize)]
442pub struct MfaEnrollResponse {
443 pub otpauth_uri: String,
445 pub recovery_codes: Vec<String>,
447}
448
449#[derive(Debug, Deserialize)]
451pub struct MfaChallengeRequest {
452 pub user_id: String,
454}
455
456#[derive(Debug, Serialize)]
458pub struct MfaChallengeResponse {
459 pub challenge_token: String,
461}
462
463#[derive(Debug, Deserialize)]
465pub struct MfaVerifyRequest {
466 pub challenge_token: String,
468 pub code: String,
470}
471
472#[derive(Debug, Deserialize)]
474pub struct MfaUnenrollRequest {
475 pub user_id: String,
477 pub code: String,
479}
480
481pub async fn mfa_enroll(
491 State(state): State<Arc<MfaRouteState>>,
492 Json(req): Json<MfaEnrollRequest>,
493) -> Response {
494 match state
495 .mfa_store
496 .begin_enrollment(&req.user_id, &state.issuer, &req.account_name)
497 .await
498 {
499 Ok(resp) => (
500 StatusCode::OK,
501 Json(MfaEnrollResponse {
502 otpauth_uri: resp.otpauth_uri,
503 recovery_codes: resp.recovery_codes,
504 }),
505 )
506 .into_response(),
507 Err(e) => {
508 tracing::error!(error = %e, "MFA enroll error");
509 (StatusCode::INTERNAL_SERVER_ERROR, "enrollment failed").into_response()
510 },
511 }
512}
513
514pub async fn mfa_challenge(
522 State(state): State<Arc<MfaRouteState>>,
523 Json(req): Json<MfaChallengeRequest>,
524) -> Response {
525 if !state.mfa_store.is_enrolled(&req.user_id).await {
526 return (
527 StatusCode::NOT_FOUND,
528 Json(serde_json::json!({
529 "error": "not_enrolled",
530 "message": "user has no active MFA enrollment"
531 })),
532 )
533 .into_response();
534 }
535
536 match state.mfa_store.create_challenge(&req.user_id).await {
537 Ok(token) => (
538 StatusCode::OK,
539 Json(MfaChallengeResponse {
540 challenge_token: token,
541 }),
542 )
543 .into_response(),
544 Err(e) => {
545 tracing::error!(error = %e, "MFA challenge error");
546 (StatusCode::INTERNAL_SERVER_ERROR, "challenge creation failed").into_response()
547 },
548 }
549}
550
551pub async fn mfa_verify(
559 State(state): State<Arc<MfaRouteState>>,
560 Json(req): Json<MfaVerifyRequest>,
561) -> Response {
562 let logger = get_audit_logger();
563
564 let user_id = match state.mfa_store.verify_challenge(&req.challenge_token, &req.code).await {
565 Ok(uid) => uid,
566 Err(e) => {
567 logger.log_failure(
568 AuditEventType::AuthFailure,
569 SecretType::SessionToken,
570 None,
571 "mfa_verify",
572 &e.to_string(),
573 );
574 return (
575 StatusCode::UNPROCESSABLE_ENTITY,
576 Json(serde_json::json!({
577 "error": "invalid_mfa",
578 "message": "invalid or expired MFA code"
579 })),
580 )
581 .into_response();
582 },
583 };
584
585 let expires_at = match unix_now() {
586 Ok(now) => now + 3_600, Err(_) => return (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response(),
588 };
589
590 let tokens = match state.session_store.create_session(&user_id, expires_at).await {
591 Ok(t) => t,
592 Err(e) => {
593 tracing::error!(error = %e, "Session creation failed after MFA verify");
594 return (StatusCode::INTERNAL_SERVER_ERROR, "session creation failed").into_response();
595 },
596 };
597
598 logger.log_success(
599 AuditEventType::AuthSuccess,
600 SecretType::SessionToken,
601 Some(user_id),
602 "mfa_verify",
603 );
604
605 (StatusCode::OK, Json(tokens)).into_response()
606}
607
608pub async fn mfa_unenroll(
616 State(state): State<Arc<MfaRouteState>>,
617 Json(req): Json<MfaUnenrollRequest>,
618) -> Response {
619 let logger = get_audit_logger();
620
621 match state.mfa_store.unenroll(&req.user_id, &req.code).await {
622 Ok(()) => {
623 logger.log_success(
624 AuditEventType::SessionTokenRevoked,
625 SecretType::SessionToken,
626 Some(req.user_id),
627 "mfa_unenroll",
628 );
629 StatusCode::OK.into_response()
630 },
631 Err(e) => {
632 logger.log_failure(
633 AuditEventType::AuthFailure,
634 SecretType::SessionToken,
635 Some(req.user_id.clone()),
636 "mfa_unenroll",
637 &e.to_string(),
638 );
639 (
640 StatusCode::UNPROCESSABLE_ENTITY,
641 Json(serde_json::json!({
642 "error": "invalid_code",
643 "message": "re-authentication failed"
644 })),
645 )
646 .into_response()
647 },
648 }
649}
650
651#[allow(clippy::unwrap_used)] #[cfg(test)]
655mod tests;