1use std::sync::Arc;
15use uuid::Uuid;
16
17use tracing::{info, warn};
18
19use crate::admin::{password, recovery, totp};
20use acme_proxy_store::admin_recovery_code::AdminRecoveryCode;
21use acme_proxy_store::admin_session::AdminSession;
22use acme_proxy_store::admin_user::AdminUser;
23use acme_proxy_store::db::Database;
24use acme_proxy_store::nonce::now_secs;
25
26#[derive(Debug, Clone, Copy, PartialEq, Eq)]
28pub enum MfaMethod {
29 Totp,
30 RecoveryCode,
31}
32
33impl MfaMethod {
34 #[must_use]
37 pub fn as_str(self) -> &'static str {
38 match self {
39 MfaMethod::Totp => "totp",
40 MfaMethod::RecoveryCode => "recovery_code",
41 }
42 }
43}
44
45#[derive(Debug, Clone, PartialEq, Eq)]
51pub enum MfaOutcome {
52 Accepted {
53 via: MfaMethod,
54 recovery_codes_left: i64,
57 },
58 Rejected,
60 Replayed,
64 NotEnrolled,
67}
68
69impl MfaOutcome {
70 #[must_use]
72 pub fn reason(&self) -> &'static str {
73 match self {
74 MfaOutcome::Accepted { .. } => "",
75 MfaOutcome::Rejected => "wrong_code",
76 MfaOutcome::Replayed => "replayed",
77 MfaOutcome::NotEnrolled => "no_factor",
78 }
79 }
80}
81
82pub async fn verify_second_factor(
95 user: &mut AdminUser,
96 submitted: &str,
97 database: Arc<Database>,
98) -> Result<MfaOutcome, sqlx::Error> {
99 let Some(secret) = user.totp_secret.clone() else {
100 return Ok(MfaOutcome::NotEnrolled);
101 };
102
103 let trimmed = submitted.trim();
104 if let Some(step) = totp::verify(&secret, trimmed, now_secs()) {
105 if !user.claim_totp_step(step, &database).await? {
106 return Ok(MfaOutcome::Replayed);
107 }
108 let left = AdminRecoveryCode::count_unused(user.id, &database).await?;
109 return Ok(MfaOutcome::Accepted {
110 via: MfaMethod::Totp,
111 recovery_codes_left: left,
112 });
113 }
114
115 let candidate = recovery::normalize(trimmed);
116 if !recovery::is_well_formed(&candidate) {
117 return Ok(MfaOutcome::Rejected);
118 }
119
120 for code in AdminRecoveryCode::list_unused(user.id, &database).await? {
121 match password::verify_password_off_runtime(&code.code_hash, &candidate).await {
122 Ok(true) => {
123 if !AdminRecoveryCode::consume(code.id, &database).await? {
124 return Ok(MfaOutcome::Rejected);
127 }
128 let left = AdminRecoveryCode::count_unused(user.id, &database).await?;
129 warn!(event = "admin_mfa_recovery_code_used",
130 outcome = "success",
131 user_id = %user.id,
132 username = %user.username,
133 remaining = left);
134 return Ok(MfaOutcome::Accepted {
135 via: MfaMethod::RecoveryCode,
136 recovery_codes_left: left,
137 });
138 }
139 Ok(false) => {}
140 Err(error) => {
141 warn!(event = "admin_recovery_code_hash_unreadable",
145 outcome = "failure",
146 user_id = %user.id,
147 code_id = %code.id,
148 error = %error);
149 }
150 }
151 }
152
153 Ok(MfaOutcome::Rejected)
154}
155
156pub async fn begin_totp_enrolment(
163 user: &mut AdminUser,
164 base_url: &str,
165 database: Arc<Database>,
166) -> Result<totp::Enrolment, sqlx::Error> {
167 let account = totp::account_label(&user.username, base_url);
168 let enrolment = totp::begin_enrolment(totp::ISSUER, &account);
169 user.set_totp_pending(&enrolment.secret, &database).await?;
170 Ok(enrolment)
171}
172
173pub async fn resume_or_begin_totp_enrolment(
180 user: &mut AdminUser,
181 base_url: &str,
182 database: Arc<Database>,
183) -> Result<totp::Enrolment, sqlx::Error> {
184 let Some(secret) = user.totp_pending_secret.clone() else {
185 return begin_totp_enrolment(user, base_url, database).await;
186 };
187
188 let account = totp::account_label(&user.username, base_url);
189 let secret_base32 = totp::base32_encode(&secret);
190 let uri = totp::provisioning_uri(&secret_base32, totp::ISSUER, &account);
191 Ok(totp::Enrolment {
192 secret,
193 secret_base32,
194 uri,
195 })
196}
197
198pub async fn confirm_totp_enrolment(
208 user: &mut AdminUser,
209 code: &str,
210 keep_session: Option<&str>,
211 database: Arc<Database>,
212) -> Result<Option<Vec<String>>, sqlx::Error> {
213 let Some(pending) = user.totp_pending_secret.clone() else {
214 return Ok(None);
215 };
216
217 let Some(step) = totp::verify(&pending, code.trim(), now_secs()) else {
218 return Ok(None);
219 };
220
221 user.confirm_totp(&database).await?;
222 user.claim_totp_step(step, &database).await?;
225
226 let codes = issue_recovery_codes(user, database.clone()).await?;
227 revoke_other_sessions(user, keep_session, database).await?;
228
229 info!(event = "admin_mfa_enabled",
230 outcome = "success",
231 user_id = %user.id,
232 username = %user.username,
233 recovery_codes = codes.len());
234 Ok(Some(codes))
235}
236
237pub async fn disable_totp(
245 user: &mut AdminUser,
246 keep_session: Option<&str>,
247 database: Arc<Database>,
248) -> Result<(), sqlx::Error> {
249 user.clear_totp(&database).await?;
250 AdminRecoveryCode::delete_for_user(user.id, &database).await?;
251 revoke_other_sessions(user, keep_session, database).await?;
252
253 info!(event = "admin_mfa_disabled", outcome = "success", user_id = %user.id, username = %user.username);
254 Ok(())
255}
256
257pub async fn regenerate_recovery_codes(
260 user: &AdminUser,
261 database: Arc<Database>,
262) -> Result<Vec<String>, sqlx::Error> {
263 let codes = issue_recovery_codes(user, database).await?;
264 info!(event = "admin_mfa_recovery_codes_regenerated",
265 outcome = "success",
266 user_id = %user.id,
267 username = %user.username,
268 minted = codes.len());
269 Ok(codes)
270}
271
272pub async fn recovery_codes_remaining(
274 user_id: Uuid,
275 database: Arc<Database>,
276) -> Result<i64, sqlx::Error> {
277 AdminRecoveryCode::count_unused(user_id, &database).await
278}
279
280pub async fn operators_without_a_factor(database: Arc<Database>) -> Result<usize, sqlx::Error> {
283 Ok(AdminUser::list_all(&database)
284 .await?
285 .iter()
286 .filter(|user| !user.has_totp())
287 .count())
288}
289
290async fn issue_recovery_codes(
291 user: &AdminUser,
292 database: Arc<Database>,
293) -> Result<Vec<String>, sqlx::Error> {
294 let codes = recovery::generate_codes();
295 let hashes: Vec<String> = codes
298 .iter()
299 .map(|code| password::hash_generated_secret(&recovery::normalize(code)))
300 .collect();
301 AdminRecoveryCode::replace_all(user.id, &hashes, &database).await?;
302 Ok(codes)
303}
304
305async fn revoke_other_sessions(
306 user: &AdminUser,
307 keep_session: Option<&str>,
308 database: Arc<Database>,
309) -> Result<u64, sqlx::Error> {
310 match keep_session {
311 Some(token_hash) => {
312 AdminSession::delete_for_user_except(user.id, token_hash, &database).await
313 }
314 None => AdminSession::delete_for_user(user.id, &database).await,
315 }
316}
317
318#[cfg(test)]
319mod tests {
320 use super::*;
321 use crate::admin::totp::{DIGITS, step_at, totp_at};
322 use acme_proxy_store::admin_session::NewSession;
323
324 async fn db() -> Arc<Database> {
325 Arc::new(Database::connect_in_memory().await.unwrap())
326 }
327
328 async fn operator(database: Arc<Database>) -> AdminUser {
329 AdminUser::create("alice", "hash", None, &database)
330 .await
331 .unwrap()
332 }
333
334 async fn enrolled(database: Arc<Database>) -> (AdminUser, Vec<u8>) {
337 let mut user = operator(database.clone()).await;
338 let enrolment = begin_totp_enrolment(&mut user, "http://localhost:3001", database.clone())
339 .await
340 .unwrap();
341 let code = totp_at(&enrolment.secret, step_at(now_secs()), DIGITS);
342 let codes = confirm_totp_enrolment(&mut user, &code, None, database)
343 .await
344 .unwrap()
345 .expect("a freshly generated code must confirm its own enrolment");
346 assert_eq!(codes.len(), recovery::CODE_COUNT);
347 (user, enrolment.secret)
348 }
349
350 #[tokio::test]
351 async fn an_operator_with_no_factor_is_not_enrolled() {
352 let db = db().await;
353 let mut user = operator(db.clone()).await;
354
355 assert_eq!(
356 verify_second_factor(&mut user, "123456", db).await.unwrap(),
357 MfaOutcome::NotEnrolled
358 );
359 }
360
361 #[tokio::test]
362 async fn enrolment_is_two_steps_and_a_wrong_code_finishes_neither() {
363 let db = db().await;
364 let mut user = operator(db.clone()).await;
365
366 let enrolment = begin_totp_enrolment(&mut user, "http://localhost:3001", db.clone())
367 .await
368 .unwrap();
369 assert!(
370 !user.has_totp(),
371 "a pending enrolment is not a second factor"
372 );
373 assert!(user.has_pending_totp());
374
375 assert!(
377 confirm_totp_enrolment(&mut user, "000000", None, db.clone())
378 .await
379 .unwrap()
380 .is_none()
381 );
382 assert!(!user.has_totp());
383 assert!(user.has_pending_totp());
384 assert_eq!(
385 recovery_codes_remaining(user.id, db.clone()).await.unwrap(),
386 0
387 );
388
389 let code = totp_at(&enrolment.secret, step_at(now_secs()), DIGITS);
390 let codes = confirm_totp_enrolment(&mut user, &code, None, db.clone())
391 .await
392 .unwrap()
393 .unwrap();
394
395 assert!(user.has_totp());
396 assert!(
397 !user.has_pending_totp(),
398 "confirming must clear the pending column, not leave two secrets live"
399 );
400 assert_eq!(codes.len(), recovery::CODE_COUNT);
401 assert_eq!(
402 recovery_codes_remaining(user.id, db).await.unwrap(),
403 recovery::CODE_COUNT as i64
404 );
405 }
406
407 #[tokio::test]
410 async fn a_correct_code_is_accepted_once_and_replayed_thereafter() {
411 let db = db().await;
412 let (mut user, secret) = enrolled(db.clone()).await;
413
414 let claimed = user.totp_last_step.expect("enrolment claims its own step");
419 let code = totp_at(&secret, claimed, DIGITS);
420 assert_eq!(
421 verify_second_factor(&mut user, &code, db.clone())
422 .await
423 .unwrap(),
424 MfaOutcome::Replayed
425 );
426
427 let next = totp_at(&secret, step_at(now_secs()) + 1, DIGITS);
429 assert_eq!(
430 verify_second_factor(&mut user, &next, db.clone())
431 .await
432 .unwrap(),
433 MfaOutcome::Accepted {
434 via: MfaMethod::Totp,
435 recovery_codes_left: recovery::CODE_COUNT as i64,
436 }
437 );
438 assert_eq!(
439 verify_second_factor(&mut user, &next, db).await.unwrap(),
440 MfaOutcome::Replayed
441 );
442 }
443
444 #[tokio::test]
445 async fn a_wrong_code_is_rejected_without_touching_the_replay_guard() {
446 let db = db().await;
447 let (mut user, secret) = enrolled(db.clone()).await;
448 let claimed = user.totp_last_step;
449
450 for wrong in ["000000", "12345", "abcdef", ""] {
451 assert_eq!(
452 verify_second_factor(&mut user, wrong, db.clone())
453 .await
454 .unwrap(),
455 MfaOutcome::Rejected,
456 "submission {wrong:?}"
457 );
458 }
459 assert_eq!(
460 user.totp_last_step, claimed,
461 "a wrong code must not advance the guard, or it would lock out the right one"
462 );
463
464 let next = totp_at(&secret, step_at(now_secs()) + 1, DIGITS);
466 assert!(matches!(
467 verify_second_factor(&mut user, &next, db).await.unwrap(),
468 MfaOutcome::Accepted { .. }
469 ));
470 }
471
472 #[tokio::test]
473 async fn a_recovery_code_is_accepted_once_and_decrements_the_count() {
474 let db = db().await;
475 let mut user = operator(db.clone()).await;
476 let enrolment = begin_totp_enrolment(&mut user, "http://localhost:3001", db.clone())
477 .await
478 .unwrap();
479 let code = totp_at(&enrolment.secret, step_at(now_secs()), DIGITS);
480 let codes = confirm_totp_enrolment(&mut user, &code, None, db.clone())
481 .await
482 .unwrap()
483 .unwrap();
484
485 assert_eq!(
487 verify_second_factor(&mut user, &codes[0], db.clone())
488 .await
489 .unwrap(),
490 MfaOutcome::Accepted {
491 via: MfaMethod::RecoveryCode,
492 recovery_codes_left: recovery::CODE_COUNT as i64 - 1,
493 }
494 );
495 assert_eq!(
496 verify_second_factor(&mut user, &codes[0], db.clone())
497 .await
498 .unwrap(),
499 MfaOutcome::Rejected,
500 "single-use: a spent code is worth nothing"
501 );
502
503 assert_eq!(
505 verify_second_factor(&mut user, &codes[1].to_lowercase(), db.clone())
506 .await
507 .unwrap(),
508 MfaOutcome::Accepted {
509 via: MfaMethod::RecoveryCode,
510 recovery_codes_left: recovery::CODE_COUNT as i64 - 2,
511 }
512 );
513 assert_eq!(
514 verify_second_factor(&mut user, &codes[2].replace('-', " "), db.clone())
515 .await
516 .unwrap(),
517 MfaOutcome::Accepted {
518 via: MfaMethod::RecoveryCode,
519 recovery_codes_left: recovery::CODE_COUNT as i64 - 3,
520 }
521 );
522
523 assert_eq!(
524 recovery_codes_remaining(user.id, db).await.unwrap(),
525 recovery::CODE_COUNT as i64 - 3
526 );
527 }
528
529 #[tokio::test]
530 async fn regenerating_supersedes_the_previous_set() {
531 let db = db().await;
532 let (mut user, _) = enrolled(db.clone()).await;
533
534 let first = AdminRecoveryCode::list_unused(user.id, &db).await.unwrap();
535 let second = regenerate_recovery_codes(&user, db.clone()).await.unwrap();
536 assert_eq!(second.len(), recovery::CODE_COUNT);
537
538 for code in AdminRecoveryCode::list_unused(user.id, &db).await.unwrap() {
542 assert!(first.iter().all(|old| old.id != code.id));
543 }
544 assert_eq!(
545 verify_second_factor(&mut user, &second[0], db.clone())
546 .await
547 .unwrap(),
548 MfaOutcome::Accepted {
549 via: MfaMethod::RecoveryCode,
550 recovery_codes_left: recovery::CODE_COUNT as i64 - 1,
551 }
552 );
553 }
554
555 #[tokio::test]
556 async fn disabling_clears_every_column_and_every_code() {
557 let db = db().await;
558 let (mut user, _) = enrolled(db.clone()).await;
559 assert!(user.has_totp());
560
561 disable_totp(&mut user, None, db.clone()).await.unwrap();
562
563 assert!(!user.has_totp());
564 assert!(!user.has_pending_totp());
565 assert_eq!(user.totp_last_step, None);
566 assert_eq!(
567 recovery_codes_remaining(user.id, db.clone()).await.unwrap(),
568 0,
569 "a recovery code for a factor that no longer exists is a second password"
570 );
571
572 let reloaded = AdminUser::find_by_id(user.id, &db).await.unwrap().unwrap();
575 assert!(!reloaded.has_totp());
576 assert!(!reloaded.has_pending_totp());
577 assert_eq!(reloaded.totp_last_step, None);
578 }
579
580 #[tokio::test]
581 async fn a_factor_change_revokes_every_other_session() {
582 let db = db().await;
583 let mut user = operator(db.clone()).await;
584
585 let kept = AdminSession::create(
586 NewSession {
587 user_id: user.id,
588 token_hash: "kept-hash",
589 csrf_token: "csrf",
590 created_ip: None,
591 user_agent: None,
592 },
593 std::time::Duration::from_secs(3600),
594 &db,
595 )
596 .await
597 .unwrap();
598 AdminSession::create(
599 NewSession {
600 user_id: user.id,
601 token_hash: "other-hash",
602 csrf_token: "csrf",
603 created_ip: None,
604 user_agent: None,
605 },
606 std::time::Duration::from_secs(3600),
607 &db,
608 )
609 .await
610 .unwrap();
611
612 let enrolment = begin_totp_enrolment(&mut user, "http://localhost:3001", db.clone())
613 .await
614 .unwrap();
615 let code = totp_at(&enrolment.secret, step_at(now_secs()), DIGITS);
616 confirm_totp_enrolment(&mut user, &code, Some(&kept.token_hash), db.clone())
617 .await
618 .unwrap()
619 .unwrap();
620
621 let live = AdminSession::list_all(Some(user.id), &db).await.unwrap();
622 assert_eq!(live.len(), 1, "every other browser must be signed out");
623 assert_eq!(live[0].token_hash, "kept-hash");
624
625 disable_totp(&mut user, None, db.clone()).await.unwrap();
627 assert!(
628 AdminSession::list_all(Some(user.id), &db)
629 .await
630 .unwrap()
631 .is_empty()
632 );
633 }
634
635 #[tokio::test]
636 async fn every_outcome_names_itself_for_the_log() {
637 assert_eq!(
638 MfaOutcome::Accepted {
639 via: MfaMethod::Totp,
640 recovery_codes_left: 10
641 }
642 .reason(),
643 ""
644 );
645 assert_eq!(MfaOutcome::Rejected.reason(), "wrong_code");
646 assert_eq!(MfaOutcome::Replayed.reason(), "replayed");
647 assert_eq!(MfaOutcome::NotEnrolled.reason(), "no_factor");
648 assert_eq!(MfaMethod::Totp.as_str(), "totp");
649 assert_eq!(MfaMethod::RecoveryCode.as_str(), "recovery_code");
650 }
651
652 #[tokio::test]
653 async fn the_startup_count_sees_only_confirmed_factors() {
654 let db = db().await;
655 let mut alice = operator(db.clone()).await;
656 AdminUser::create("bob", "hash", None, &db).await.unwrap();
657 assert_eq!(operators_without_a_factor(db.clone()).await.unwrap(), 2);
658
659 let enrolment = begin_totp_enrolment(&mut alice, "http://localhost:3001", db.clone())
661 .await
662 .unwrap();
663 assert_eq!(operators_without_a_factor(db.clone()).await.unwrap(), 2);
664
665 let code = totp_at(&enrolment.secret, step_at(now_secs()), DIGITS);
666 confirm_totp_enrolment(&mut alice, &code, None, db.clone())
667 .await
668 .unwrap()
669 .unwrap();
670 assert_eq!(operators_without_a_factor(db).await.unwrap(), 1);
671 }
672}