use std::sync::Arc;
use async_trait::async_trait;
use axum::{
Json,
extract::State,
http::StatusCode,
response::{IntoResponse, Response},
};
use dashmap::DashMap;
use rand::RngCore as _;
use serde::{Deserialize, Serialize};
use totp_rs::{Algorithm, Secret, TOTP};
use crate::{
audit::logger::{AuditEventType, SecretType, get_audit_logger},
error::{AuthError, Result},
session::{SessionStore, unix_now},
};
const RECOVERY_CODE_COUNT: usize = 8;
const RECOVERY_CODE_HEX_LEN: usize = 16;
const CHALLENGE_TTL_SECS: u64 = 300;
const TOTP_STEP_TOLERANCE: u8 = 1;
#[cfg(not(test))]
const BCRYPT_COST: u32 = 12;
#[cfg(test)]
const BCRYPT_COST: u32 = 4;
#[derive(Debug, Clone)]
pub struct TotpEnrollment {
pub secret_base32: String,
pub recovery_code_hashes: Vec<String>,
pub confirmed: bool,
}
#[derive(Debug, Clone)]
struct ChallengeRecord {
user_id: String,
expires: u64,
}
#[async_trait]
pub trait MfaStore: Send + Sync {
async fn begin_enrollment(
&self,
user_id: &str,
issuer: &str,
account_name: &str,
) -> Result<EnrollmentResponse>;
async fn confirm_enrollment(&self, user_id: &str, totp_code: &str) -> Result<()>;
async fn create_challenge(&self, user_id: &str) -> Result<String>;
async fn verify_challenge(&self, challenge_token: &str, code: &str) -> Result<String>;
async fn unenroll(&self, user_id: &str, code: &str) -> Result<()>;
async fn is_enrolled(&self, user_id: &str) -> bool;
}
#[derive(Debug)]
pub struct EnrollmentResponse {
pub secret_base32: String,
pub otpauth_uri: String,
pub recovery_codes: Vec<String>,
}
pub struct InMemoryMfaStore {
enrollments: DashMap<String, TotpEnrollment>,
challenges: DashMap<String, ChallengeRecord>,
}
impl InMemoryMfaStore {
#[must_use]
pub fn new() -> Self {
Self {
enrollments: DashMap::new(),
challenges: DashMap::new(),
}
}
#[must_use]
pub fn has_pending_enrollment(&self, user_id: &str) -> bool {
self.enrollments.get(user_id).is_some_and(|e| !e.confirmed)
}
}
impl Default for InMemoryMfaStore {
fn default() -> Self {
Self::new()
}
}
fn build_totp(secret_base32: &str, issuer: Option<&str>, account_name: &str) -> Result<TOTP> {
let secret_bytes =
Secret::Encoded(secret_base32.to_string())
.to_bytes()
.map_err(|e| AuthError::Internal {
message: format!("bad TOTP secret: {e}"),
})?;
TOTP::new(
Algorithm::SHA1,
6, TOTP_STEP_TOLERANCE,
30, secret_bytes,
issuer.map(str::to_string),
account_name.to_string(),
)
.map_err(|e| AuthError::Internal {
message: format!("TOTP init error: {e}"),
})
}
fn verify_totp_code(secret_base32: &str, code: &str) -> Result<bool> {
let totp = build_totp(secret_base32, None, "")?;
Ok(totp.check_current(code).unwrap_or(false))
}
fn generate_recovery_code() -> String {
let byte_count = RECOVERY_CODE_HEX_LEN / 2;
let mut bytes = vec![0u8; byte_count];
rand::rng().fill_bytes(&mut bytes);
bytes.iter().fold(String::new(), |mut s, b| {
use std::fmt::Write as _;
let _ = write!(s, "{b:02x}");
s
})
}
fn generate_challenge_token() -> String {
use base64::Engine as _;
let mut bytes = [0u8; 32];
rand::rng().fill_bytes(&mut bytes);
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(bytes)
}
fn check_recovery_code(candidate: &str, hashes: &[String]) -> Option<usize> {
for (i, hash) in hashes.iter().enumerate() {
if bcrypt::verify(candidate, hash).unwrap_or(false) {
return Some(i);
}
}
None
}
#[async_trait]
impl MfaStore for InMemoryMfaStore {
async fn begin_enrollment(
&self,
user_id: &str,
issuer: &str,
account_name: &str,
) -> Result<EnrollmentResponse> {
let secret = Secret::generate_secret();
let secret_base32 = secret.to_encoded().to_string();
let totp = build_totp(&secret_base32, Some(issuer), account_name)?;
let otpauth_uri = totp.get_url();
let mut recovery_codes_plain = Vec::with_capacity(RECOVERY_CODE_COUNT);
let mut recovery_code_hashes = Vec::with_capacity(RECOVERY_CODE_COUNT);
for _ in 0..RECOVERY_CODE_COUNT {
let code = generate_recovery_code();
let hash = bcrypt::hash(&code, BCRYPT_COST).map_err(|e| AuthError::Internal {
message: format!("bcrypt error: {e}"),
})?;
recovery_codes_plain.push(code);
recovery_code_hashes.push(hash);
}
self.enrollments.insert(
user_id.to_string(),
TotpEnrollment {
secret_base32: secret_base32.clone(),
recovery_code_hashes,
confirmed: false,
},
);
Ok(EnrollmentResponse {
secret_base32,
otpauth_uri,
recovery_codes: recovery_codes_plain,
})
}
async fn confirm_enrollment(&self, user_id: &str, totp_code: &str) -> Result<()> {
let mut record =
self.enrollments.get_mut(user_id).ok_or_else(|| AuthError::InvalidToken {
reason: "no pending MFA enrollment for user".into(),
})?;
if !verify_totp_code(&record.secret_base32, totp_code)? {
return Err(AuthError::InvalidToken {
reason: "invalid TOTP code".into(),
});
}
record.confirmed = true;
Ok(())
}
async fn create_challenge(&self, user_id: &str) -> Result<String> {
let expires = unix_now()? + CHALLENGE_TTL_SECS;
let token = generate_challenge_token();
self.challenges.insert(
token.clone(),
ChallengeRecord {
user_id: user_id.to_string(),
expires,
},
);
Ok(token)
}
async fn verify_challenge(&self, challenge_token: &str, code: &str) -> Result<String> {
let now = unix_now()?;
let record =
self.challenges.get(challenge_token).ok_or_else(|| AuthError::InvalidToken {
reason: "unknown challenge token".into(),
})?;
if now >= record.expires {
drop(record);
self.challenges.remove(challenge_token);
return Err(AuthError::InvalidToken {
reason: "challenge token expired".into(),
});
}
let user_id = record.user_id.clone();
drop(record);
let mut enrollment =
self.enrollments.get_mut(&user_id).ok_or_else(|| AuthError::InvalidToken {
reason: "user has no MFA enrollment".into(),
})?;
if !enrollment.confirmed {
return Err(AuthError::InvalidToken {
reason: "MFA enrollment not confirmed".into(),
});
}
if verify_totp_code(&enrollment.secret_base32, code)? {
drop(enrollment);
self.challenges.remove(challenge_token);
return Ok(user_id);
}
let idx = check_recovery_code(code, &enrollment.recovery_code_hashes);
if let Some(i) = idx {
enrollment.recovery_code_hashes.remove(i);
drop(enrollment);
self.challenges.remove(challenge_token);
return Ok(user_id);
}
Err(AuthError::InvalidToken {
reason: "invalid TOTP or recovery code".into(),
})
}
async fn unenroll(&self, user_id: &str, code: &str) -> Result<()> {
let enrollment = self.enrollments.get(user_id).ok_or_else(|| AuthError::InvalidToken {
reason: "user has no MFA enrollment".into(),
})?;
if !enrollment.confirmed {
return Err(AuthError::InvalidToken {
reason: "MFA enrollment not confirmed".into(),
});
}
let totp_ok = verify_totp_code(&enrollment.secret_base32, code)?;
let recovery_ok =
!totp_ok && check_recovery_code(code, &enrollment.recovery_code_hashes).is_some();
if !totp_ok && !recovery_ok {
return Err(AuthError::InvalidToken {
reason: "re-authentication failed — invalid TOTP or recovery code".into(),
});
}
drop(enrollment);
self.enrollments.remove(user_id);
Ok(())
}
async fn is_enrolled(&self, user_id: &str) -> bool {
self.enrollments.get(user_id).map_or(false, |e| e.confirmed)
}
}
#[derive(Clone)]
pub struct MfaRouteState {
pub mfa_store: Arc<dyn MfaStore>,
pub session_store: Arc<dyn SessionStore>,
pub issuer: String,
}
#[derive(Debug, Deserialize)]
pub struct MfaEnrollRequest {
pub user_id: String,
pub account_name: String,
}
#[derive(Debug, Serialize)]
pub struct MfaEnrollResponse {
pub otpauth_uri: String,
pub recovery_codes: Vec<String>,
}
#[derive(Debug, Deserialize)]
pub struct MfaChallengeRequest {
pub user_id: String,
}
#[derive(Debug, Serialize)]
pub struct MfaChallengeResponse {
pub challenge_token: String,
}
#[derive(Debug, Deserialize)]
pub struct MfaVerifyRequest {
pub challenge_token: String,
pub code: String,
}
#[derive(Debug, Deserialize)]
pub struct MfaUnenrollRequest {
pub user_id: String,
pub code: String,
}
pub async fn mfa_enroll(
State(state): State<Arc<MfaRouteState>>,
Json(req): Json<MfaEnrollRequest>,
) -> Response {
match state
.mfa_store
.begin_enrollment(&req.user_id, &state.issuer, &req.account_name)
.await
{
Ok(resp) => (
StatusCode::OK,
Json(MfaEnrollResponse {
otpauth_uri: resp.otpauth_uri,
recovery_codes: resp.recovery_codes,
}),
)
.into_response(),
Err(e) => {
tracing::error!(error = %e, "MFA enroll error");
(StatusCode::INTERNAL_SERVER_ERROR, "enrollment failed").into_response()
},
}
}
pub async fn mfa_challenge(
State(state): State<Arc<MfaRouteState>>,
Json(req): Json<MfaChallengeRequest>,
) -> Response {
if !state.mfa_store.is_enrolled(&req.user_id).await {
return (
StatusCode::NOT_FOUND,
Json(serde_json::json!({
"error": "not_enrolled",
"message": "user has no active MFA enrollment"
})),
)
.into_response();
}
match state.mfa_store.create_challenge(&req.user_id).await {
Ok(token) => (
StatusCode::OK,
Json(MfaChallengeResponse {
challenge_token: token,
}),
)
.into_response(),
Err(e) => {
tracing::error!(error = %e, "MFA challenge error");
(StatusCode::INTERNAL_SERVER_ERROR, "challenge creation failed").into_response()
},
}
}
pub async fn mfa_verify(
State(state): State<Arc<MfaRouteState>>,
Json(req): Json<MfaVerifyRequest>,
) -> Response {
let logger = get_audit_logger();
let user_id = match state.mfa_store.verify_challenge(&req.challenge_token, &req.code).await {
Ok(uid) => uid,
Err(e) => {
logger.log_failure(
AuditEventType::AuthFailure,
SecretType::SessionToken,
None,
"mfa_verify",
&e.to_string(),
);
return (
StatusCode::UNPROCESSABLE_ENTITY,
Json(serde_json::json!({
"error": "invalid_mfa",
"message": "invalid or expired MFA code"
})),
)
.into_response();
},
};
let expires_at = match unix_now() {
Ok(now) => now + 3_600, Err(_) => return (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response(),
};
let tokens = match state.session_store.create_session(&user_id, expires_at).await {
Ok(t) => t,
Err(e) => {
tracing::error!(error = %e, "Session creation failed after MFA verify");
return (StatusCode::INTERNAL_SERVER_ERROR, "session creation failed").into_response();
},
};
logger.log_success(
AuditEventType::AuthSuccess,
SecretType::SessionToken,
Some(user_id),
"mfa_verify",
);
(StatusCode::OK, Json(tokens)).into_response()
}
pub async fn mfa_unenroll(
State(state): State<Arc<MfaRouteState>>,
Json(req): Json<MfaUnenrollRequest>,
) -> Response {
let logger = get_audit_logger();
match state.mfa_store.unenroll(&req.user_id, &req.code).await {
Ok(()) => {
logger.log_success(
AuditEventType::SessionTokenRevoked,
SecretType::SessionToken,
Some(req.user_id),
"mfa_unenroll",
);
StatusCode::OK.into_response()
},
Err(e) => {
logger.log_failure(
AuditEventType::AuthFailure,
SecretType::SessionToken,
Some(req.user_id.clone()),
"mfa_unenroll",
&e.to_string(),
);
(
StatusCode::UNPROCESSABLE_ENTITY,
Json(serde_json::json!({
"error": "invalid_code",
"message": "re-authentication failed"
})),
)
.into_response()
},
}
}
#[allow(clippy::unwrap_used)] #[cfg(test)]
mod tests;