use std::sync::Arc;
use uuid::Uuid;
use tracing::{info, warn};
use crate::admin::{password, recovery, totp};
use acme_proxy_store::admin_recovery_code::AdminRecoveryCode;
use acme_proxy_store::admin_session::AdminSession;
use acme_proxy_store::admin_user::AdminUser;
use acme_proxy_store::db::Database;
use acme_proxy_store::nonce::now_secs;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum MfaMethod {
Totp,
RecoveryCode,
}
impl MfaMethod {
#[must_use]
pub fn as_str(self) -> &'static str {
match self {
MfaMethod::Totp => "totp",
MfaMethod::RecoveryCode => "recovery_code",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum MfaOutcome {
Accepted {
via: MfaMethod,
recovery_codes_left: i64,
},
Rejected,
Replayed,
NotEnrolled,
}
impl MfaOutcome {
#[must_use]
pub fn reason(&self) -> &'static str {
match self {
MfaOutcome::Accepted { .. } => "",
MfaOutcome::Rejected => "wrong_code",
MfaOutcome::Replayed => "replayed",
MfaOutcome::NotEnrolled => "no_factor",
}
}
}
pub async fn verify_second_factor(
user: &mut AdminUser,
submitted: &str,
database: Arc<Database>,
) -> Result<MfaOutcome, sqlx::Error> {
let Some(secret) = user.totp_secret.clone() else {
return Ok(MfaOutcome::NotEnrolled);
};
let trimmed = submitted.trim();
if let Some(step) = totp::verify(&secret, trimmed, now_secs()) {
if !user.claim_totp_step(step, &database).await? {
return Ok(MfaOutcome::Replayed);
}
let left = AdminRecoveryCode::count_unused(user.id, &database).await?;
return Ok(MfaOutcome::Accepted {
via: MfaMethod::Totp,
recovery_codes_left: left,
});
}
let candidate = recovery::normalize(trimmed);
if !recovery::is_well_formed(&candidate) {
return Ok(MfaOutcome::Rejected);
}
for code in AdminRecoveryCode::list_unused(user.id, &database).await? {
match password::verify_password_off_runtime(&code.code_hash, &candidate).await {
Ok(true) => {
if !AdminRecoveryCode::consume(code.id, &database).await? {
return Ok(MfaOutcome::Rejected);
}
let left = AdminRecoveryCode::count_unused(user.id, &database).await?;
warn!(event = "admin_mfa_recovery_code_used",
outcome = "success",
user_id = %user.id,
username = %user.username,
remaining = left);
return Ok(MfaOutcome::Accepted {
via: MfaMethod::RecoveryCode,
recovery_codes_left: left,
});
}
Ok(false) => {}
Err(error) => {
warn!(event = "admin_recovery_code_hash_unreadable",
outcome = "failure",
user_id = %user.id,
code_id = %code.id,
error = %error);
}
}
}
Ok(MfaOutcome::Rejected)
}
pub async fn begin_totp_enrolment(
user: &mut AdminUser,
base_url: &str,
database: Arc<Database>,
) -> Result<totp::Enrolment, sqlx::Error> {
let account = totp::account_label(&user.username, base_url);
let enrolment = totp::begin_enrolment(totp::ISSUER, &account);
user.set_totp_pending(&enrolment.secret, &database).await?;
Ok(enrolment)
}
pub async fn resume_or_begin_totp_enrolment(
user: &mut AdminUser,
base_url: &str,
database: Arc<Database>,
) -> Result<totp::Enrolment, sqlx::Error> {
let Some(secret) = user.totp_pending_secret.clone() else {
return begin_totp_enrolment(user, base_url, database).await;
};
let account = totp::account_label(&user.username, base_url);
let secret_base32 = totp::base32_encode(&secret);
let uri = totp::provisioning_uri(&secret_base32, totp::ISSUER, &account);
Ok(totp::Enrolment {
secret,
secret_base32,
uri,
})
}
pub async fn confirm_totp_enrolment(
user: &mut AdminUser,
code: &str,
keep_session: Option<&str>,
database: Arc<Database>,
) -> Result<Option<Vec<String>>, sqlx::Error> {
let Some(pending) = user.totp_pending_secret.clone() else {
return Ok(None);
};
let Some(step) = totp::verify(&pending, code.trim(), now_secs()) else {
return Ok(None);
};
user.confirm_totp(&database).await?;
user.claim_totp_step(step, &database).await?;
let codes = issue_recovery_codes(user, database.clone()).await?;
revoke_other_sessions(user, keep_session, database).await?;
info!(event = "admin_mfa_enabled",
outcome = "success",
user_id = %user.id,
username = %user.username,
recovery_codes = codes.len());
Ok(Some(codes))
}
pub async fn disable_totp(
user: &mut AdminUser,
keep_session: Option<&str>,
database: Arc<Database>,
) -> Result<(), sqlx::Error> {
user.clear_totp(&database).await?;
AdminRecoveryCode::delete_for_user(user.id, &database).await?;
revoke_other_sessions(user, keep_session, database).await?;
info!(event = "admin_mfa_disabled", outcome = "success", user_id = %user.id, username = %user.username);
Ok(())
}
pub async fn regenerate_recovery_codes(
user: &AdminUser,
database: Arc<Database>,
) -> Result<Vec<String>, sqlx::Error> {
let codes = issue_recovery_codes(user, database).await?;
info!(event = "admin_mfa_recovery_codes_regenerated",
outcome = "success",
user_id = %user.id,
username = %user.username,
minted = codes.len());
Ok(codes)
}
pub async fn recovery_codes_remaining(
user_id: Uuid,
database: Arc<Database>,
) -> Result<i64, sqlx::Error> {
AdminRecoveryCode::count_unused(user_id, &database).await
}
pub async fn operators_without_a_factor(database: Arc<Database>) -> Result<usize, sqlx::Error> {
Ok(AdminUser::list_all(&database)
.await?
.iter()
.filter(|user| !user.has_totp())
.count())
}
async fn issue_recovery_codes(
user: &AdminUser,
database: Arc<Database>,
) -> Result<Vec<String>, sqlx::Error> {
let codes = recovery::generate_codes();
let hashes: Vec<String> = codes
.iter()
.map(|code| password::hash_generated_secret(&recovery::normalize(code)))
.collect();
AdminRecoveryCode::replace_all(user.id, &hashes, &database).await?;
Ok(codes)
}
async fn revoke_other_sessions(
user: &AdminUser,
keep_session: Option<&str>,
database: Arc<Database>,
) -> Result<u64, sqlx::Error> {
match keep_session {
Some(token_hash) => {
AdminSession::delete_for_user_except(user.id, token_hash, &database).await
}
None => AdminSession::delete_for_user(user.id, &database).await,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::admin::totp::{DIGITS, step_at, totp_at};
use acme_proxy_store::admin_session::NewSession;
async fn db() -> Arc<Database> {
Arc::new(Database::connect_in_memory().await.unwrap())
}
async fn operator(database: Arc<Database>) -> AdminUser {
AdminUser::create("alice", "hash", None, &database)
.await
.unwrap()
}
async fn enrolled(database: Arc<Database>) -> (AdminUser, Vec<u8>) {
let mut user = operator(database.clone()).await;
let enrolment = begin_totp_enrolment(&mut user, "http://localhost:3001", database.clone())
.await
.unwrap();
let code = totp_at(&enrolment.secret, step_at(now_secs()), DIGITS);
let codes = confirm_totp_enrolment(&mut user, &code, None, database)
.await
.unwrap()
.expect("a freshly generated code must confirm its own enrolment");
assert_eq!(codes.len(), recovery::CODE_COUNT);
(user, enrolment.secret)
}
#[tokio::test]
async fn an_operator_with_no_factor_is_not_enrolled() {
let db = db().await;
let mut user = operator(db.clone()).await;
assert_eq!(
verify_second_factor(&mut user, "123456", db).await.unwrap(),
MfaOutcome::NotEnrolled
);
}
#[tokio::test]
async fn enrolment_is_two_steps_and_a_wrong_code_finishes_neither() {
let db = db().await;
let mut user = operator(db.clone()).await;
let enrolment = begin_totp_enrolment(&mut user, "http://localhost:3001", db.clone())
.await
.unwrap();
assert!(
!user.has_totp(),
"a pending enrolment is not a second factor"
);
assert!(user.has_pending_totp());
assert!(
confirm_totp_enrolment(&mut user, "000000", None, db.clone())
.await
.unwrap()
.is_none()
);
assert!(!user.has_totp());
assert!(user.has_pending_totp());
assert_eq!(
recovery_codes_remaining(user.id, db.clone()).await.unwrap(),
0
);
let code = totp_at(&enrolment.secret, step_at(now_secs()), DIGITS);
let codes = confirm_totp_enrolment(&mut user, &code, None, db.clone())
.await
.unwrap()
.unwrap();
assert!(user.has_totp());
assert!(
!user.has_pending_totp(),
"confirming must clear the pending column, not leave two secrets live"
);
assert_eq!(codes.len(), recovery::CODE_COUNT);
assert_eq!(
recovery_codes_remaining(user.id, db).await.unwrap(),
recovery::CODE_COUNT as i64
);
}
#[tokio::test]
async fn a_correct_code_is_accepted_once_and_replayed_thereafter() {
let db = db().await;
let (mut user, secret) = enrolled(db.clone()).await;
let claimed = user.totp_last_step.expect("enrolment claims its own step");
let code = totp_at(&secret, claimed, DIGITS);
assert_eq!(
verify_second_factor(&mut user, &code, db.clone())
.await
.unwrap(),
MfaOutcome::Replayed
);
let next = totp_at(&secret, step_at(now_secs()) + 1, DIGITS);
assert_eq!(
verify_second_factor(&mut user, &next, db.clone())
.await
.unwrap(),
MfaOutcome::Accepted {
via: MfaMethod::Totp,
recovery_codes_left: recovery::CODE_COUNT as i64,
}
);
assert_eq!(
verify_second_factor(&mut user, &next, db).await.unwrap(),
MfaOutcome::Replayed
);
}
#[tokio::test]
async fn a_wrong_code_is_rejected_without_touching_the_replay_guard() {
let db = db().await;
let (mut user, secret) = enrolled(db.clone()).await;
let claimed = user.totp_last_step;
for wrong in ["000000", "12345", "abcdef", ""] {
assert_eq!(
verify_second_factor(&mut user, wrong, db.clone())
.await
.unwrap(),
MfaOutcome::Rejected,
"submission {wrong:?}"
);
}
assert_eq!(
user.totp_last_step, claimed,
"a wrong code must not advance the guard, or it would lock out the right one"
);
let next = totp_at(&secret, step_at(now_secs()) + 1, DIGITS);
assert!(matches!(
verify_second_factor(&mut user, &next, db).await.unwrap(),
MfaOutcome::Accepted { .. }
));
}
#[tokio::test]
async fn a_recovery_code_is_accepted_once_and_decrements_the_count() {
let db = db().await;
let mut user = operator(db.clone()).await;
let enrolment = begin_totp_enrolment(&mut user, "http://localhost:3001", db.clone())
.await
.unwrap();
let code = totp_at(&enrolment.secret, step_at(now_secs()), DIGITS);
let codes = confirm_totp_enrolment(&mut user, &code, None, db.clone())
.await
.unwrap()
.unwrap();
assert_eq!(
verify_second_factor(&mut user, &codes[0], db.clone())
.await
.unwrap(),
MfaOutcome::Accepted {
via: MfaMethod::RecoveryCode,
recovery_codes_left: recovery::CODE_COUNT as i64 - 1,
}
);
assert_eq!(
verify_second_factor(&mut user, &codes[0], db.clone())
.await
.unwrap(),
MfaOutcome::Rejected,
"single-use: a spent code is worth nothing"
);
assert_eq!(
verify_second_factor(&mut user, &codes[1].to_lowercase(), db.clone())
.await
.unwrap(),
MfaOutcome::Accepted {
via: MfaMethod::RecoveryCode,
recovery_codes_left: recovery::CODE_COUNT as i64 - 2,
}
);
assert_eq!(
verify_second_factor(&mut user, &codes[2].replace('-', " "), db.clone())
.await
.unwrap(),
MfaOutcome::Accepted {
via: MfaMethod::RecoveryCode,
recovery_codes_left: recovery::CODE_COUNT as i64 - 3,
}
);
assert_eq!(
recovery_codes_remaining(user.id, db).await.unwrap(),
recovery::CODE_COUNT as i64 - 3
);
}
#[tokio::test]
async fn regenerating_supersedes_the_previous_set() {
let db = db().await;
let (mut user, _) = enrolled(db.clone()).await;
let first = AdminRecoveryCode::list_unused(user.id, &db).await.unwrap();
let second = regenerate_recovery_codes(&user, db.clone()).await.unwrap();
assert_eq!(second.len(), recovery::CODE_COUNT);
for code in AdminRecoveryCode::list_unused(user.id, &db).await.unwrap() {
assert!(first.iter().all(|old| old.id != code.id));
}
assert_eq!(
verify_second_factor(&mut user, &second[0], db.clone())
.await
.unwrap(),
MfaOutcome::Accepted {
via: MfaMethod::RecoveryCode,
recovery_codes_left: recovery::CODE_COUNT as i64 - 1,
}
);
}
#[tokio::test]
async fn disabling_clears_every_column_and_every_code() {
let db = db().await;
let (mut user, _) = enrolled(db.clone()).await;
assert!(user.has_totp());
disable_totp(&mut user, None, db.clone()).await.unwrap();
assert!(!user.has_totp());
assert!(!user.has_pending_totp());
assert_eq!(user.totp_last_step, None);
assert_eq!(
recovery_codes_remaining(user.id, db.clone()).await.unwrap(),
0,
"a recovery code for a factor that no longer exists is a second password"
);
let reloaded = AdminUser::find_by_id(user.id, &db).await.unwrap().unwrap();
assert!(!reloaded.has_totp());
assert!(!reloaded.has_pending_totp());
assert_eq!(reloaded.totp_last_step, None);
}
#[tokio::test]
async fn a_factor_change_revokes_every_other_session() {
let db = db().await;
let mut user = operator(db.clone()).await;
let kept = AdminSession::create(
NewSession {
user_id: user.id,
token_hash: "kept-hash",
csrf_token: "csrf",
created_ip: None,
user_agent: None,
},
std::time::Duration::from_secs(3600),
&db,
)
.await
.unwrap();
AdminSession::create(
NewSession {
user_id: user.id,
token_hash: "other-hash",
csrf_token: "csrf",
created_ip: None,
user_agent: None,
},
std::time::Duration::from_secs(3600),
&db,
)
.await
.unwrap();
let enrolment = begin_totp_enrolment(&mut user, "http://localhost:3001", db.clone())
.await
.unwrap();
let code = totp_at(&enrolment.secret, step_at(now_secs()), DIGITS);
confirm_totp_enrolment(&mut user, &code, Some(&kept.token_hash), db.clone())
.await
.unwrap()
.unwrap();
let live = AdminSession::list_all(Some(user.id), &db).await.unwrap();
assert_eq!(live.len(), 1, "every other browser must be signed out");
assert_eq!(live[0].token_hash, "kept-hash");
disable_totp(&mut user, None, db.clone()).await.unwrap();
assert!(
AdminSession::list_all(Some(user.id), &db)
.await
.unwrap()
.is_empty()
);
}
#[tokio::test]
async fn every_outcome_names_itself_for_the_log() {
assert_eq!(
MfaOutcome::Accepted {
via: MfaMethod::Totp,
recovery_codes_left: 10
}
.reason(),
""
);
assert_eq!(MfaOutcome::Rejected.reason(), "wrong_code");
assert_eq!(MfaOutcome::Replayed.reason(), "replayed");
assert_eq!(MfaOutcome::NotEnrolled.reason(), "no_factor");
assert_eq!(MfaMethod::Totp.as_str(), "totp");
assert_eq!(MfaMethod::RecoveryCode.as_str(), "recovery_code");
}
#[tokio::test]
async fn the_startup_count_sees_only_confirmed_factors() {
let db = db().await;
let mut alice = operator(db.clone()).await;
AdminUser::create("bob", "hash", None, &db).await.unwrap();
assert_eq!(operators_without_a_factor(db.clone()).await.unwrap(), 2);
let enrolment = begin_totp_enrolment(&mut alice, "http://localhost:3001", db.clone())
.await
.unwrap();
assert_eq!(operators_without_a_factor(db.clone()).await.unwrap(), 2);
let code = totp_at(&enrolment.secret, step_at(now_secs()), DIGITS);
confirm_totp_enrolment(&mut alice, &code, None, db.clone())
.await
.unwrap()
.unwrap();
assert_eq!(operators_without_a_factor(db).await.unwrap(), 1);
}
}