use sqlx::Row;
use sqlx::sqlite::SqliteRow;
use tracing::{info, warn};
use uuid::Uuid;
use crate::sqlite::db::Database;
use crate::sqlite::nonce::now_secs;
use crate::sqlite::order::rfc3339;
#[derive(Debug, Clone)]
pub struct AdminRecoveryCode {
pub id: Uuid,
pub user_id: Uuid,
pub code_hash: String,
pub created_at: i64,
pub used_at: Option<i64>,
}
impl AdminRecoveryCode {
fn from_row(row: SqliteRow) -> Result<Self, sqlx::Error> {
Ok(AdminRecoveryCode {
id: row.try_get("id")?,
user_id: row.try_get("user_id")?,
code_hash: row.try_get("code_hash")?,
created_at: row.try_get("created_at")?,
used_at: row.try_get("used_at")?,
})
}
pub async fn replace_all(
user_id: Uuid,
hashes: &[String],
database: &Database,
) -> Result<(), sqlx::Error> {
let now = now_secs();
let mut tx = database.pool.begin().await?;
sqlx::query("DELETE FROM admin_recovery_codes WHERE user_id = ?;")
.bind(user_id)
.execute(&mut *tx)
.await?;
for hash in hashes {
sqlx::query(
"INSERT INTO admin_recovery_codes (id, user_id, code_hash, created_at, used_at) \
VALUES (?, ?, ?, ?, NULL);",
)
.bind(crate::sqlite::id::mint())
.bind(user_id)
.bind(hash)
.bind(now)
.execute(&mut *tx)
.await?;
}
tx.commit().await?;
info!(
event = "db_admin_recovery_codes_replaced",
outcome = "success",
user_id = %user_id,
minted = hashes.len()
);
Ok(())
}
pub async fn list_unused(
user_id: Uuid,
database: &Database,
) -> Result<Vec<AdminRecoveryCode>, sqlx::Error> {
let rows = sqlx::query(
"SELECT id, user_id, code_hash, created_at, used_at \
FROM admin_recovery_codes WHERE user_id = ? AND used_at IS NULL \
ORDER BY created_at ASC, id ASC;",
)
.bind(user_id)
.fetch_all(&database.pool)
.await?;
rows.into_iter().map(AdminRecoveryCode::from_row).collect()
}
pub async fn count_unused(user_id: Uuid, database: &Database) -> Result<i64, sqlx::Error> {
let row = sqlx::query(
"SELECT COUNT(*) AS total FROM admin_recovery_codes \
WHERE user_id = ? AND used_at IS NULL;",
)
.bind(user_id)
.fetch_one(&database.pool)
.await?;
row.try_get("total")
}
pub async fn consume(id: Uuid, database: &Database) -> Result<bool, sqlx::Error> {
let result = sqlx::query(
"UPDATE admin_recovery_codes SET used_at = ? WHERE id = ? AND used_at IS NULL;",
)
.bind(now_secs())
.bind(id)
.execute(&database.pool)
.await?;
let consumed = result.rows_affected() == 1;
if !consumed {
warn!(event = "db_admin_recovery_code_already_used", outcome = "failure", code_id = %id);
}
Ok(consumed)
}
pub async fn delete_for_user(user_id: Uuid, database: &Database) -> Result<u64, sqlx::Error> {
let result = sqlx::query("DELETE FROM admin_recovery_codes WHERE user_id = ?;")
.bind(user_id)
.execute(&database.pool)
.await?;
Ok(result.rows_affected())
}
#[must_use]
pub fn to_json(&self) -> serde_json::Value {
serde_json::json!({
"id": self.id,
"createdAt": rfc3339(self.created_at),
"usedAt": self.used_at.map(rfc3339),
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::sqlite::admin_user::AdminUser;
use std::sync::Arc;
async fn db_with_user() -> (Arc<Database>, AdminUser) {
let db = Arc::new(Database::connect_in_memory().await.unwrap());
let user = AdminUser::create("alice", "hash", None, &db).await.unwrap();
(db, user)
}
fn hashes(count: usize) -> Vec<String> {
(0..count).map(|index| format!("hash-{index}")).collect()
}
#[tokio::test]
async fn replace_all_mints_a_set_and_supersedes_the_previous_one() {
let (db, user) = db_with_user().await;
AdminRecoveryCode::replace_all(user.id, &hashes(10), &db)
.await
.unwrap();
assert_eq!(
AdminRecoveryCode::count_unused(user.id, &db).await.unwrap(),
10
);
let first = AdminRecoveryCode::list_unused(user.id, &db).await.unwrap()[0].id;
assert!(AdminRecoveryCode::consume(first, &db).await.unwrap());
assert_eq!(
AdminRecoveryCode::count_unused(user.id, &db).await.unwrap(),
9
);
AdminRecoveryCode::replace_all(user.id, &hashes(10), &db)
.await
.unwrap();
let after = AdminRecoveryCode::list_unused(user.id, &db).await.unwrap();
assert_eq!(after.len(), 10);
assert!(
after.iter().all(|code| code.id != first),
"no row of the superseded set may survive"
);
}
#[tokio::test]
async fn a_code_can_be_consumed_exactly_once() {
let (db, user) = db_with_user().await;
AdminRecoveryCode::replace_all(user.id, &hashes(3), &db)
.await
.unwrap();
let code = AdminRecoveryCode::list_unused(user.id, &db).await.unwrap()[0].clone();
assert!(AdminRecoveryCode::consume(code.id, &db).await.unwrap());
assert!(
!AdminRecoveryCode::consume(code.id, &db).await.unwrap(),
"a second consumption of one code must fail, whatever raced it"
);
assert!(
!AdminRecoveryCode::consume(crate::sqlite::id::mint(), &db)
.await
.unwrap()
);
let remaining = AdminRecoveryCode::list_unused(user.id, &db).await.unwrap();
assert_eq!(remaining.len(), 2);
assert!(remaining.iter().all(|other| other.id != code.id));
}
#[tokio::test]
async fn deleting_the_operator_cascades_to_their_codes() {
let (db, user) = db_with_user().await;
AdminRecoveryCode::replace_all(user.id, &hashes(10), &db)
.await
.unwrap();
assert!(AdminUser::delete(user.id, &db).await.unwrap());
assert_eq!(
AdminRecoveryCode::count_unused(user.id, &db).await.unwrap(),
0
);
}
#[tokio::test]
async fn delete_for_user_removes_the_whole_set() {
let (db, user) = db_with_user().await;
AdminRecoveryCode::replace_all(user.id, &hashes(10), &db)
.await
.unwrap();
assert_eq!(
AdminRecoveryCode::delete_for_user(user.id, &db)
.await
.unwrap(),
10
);
assert_eq!(
AdminRecoveryCode::count_unused(user.id, &db).await.unwrap(),
0
);
}
#[tokio::test]
async fn to_json_leaks_no_hash() {
let (db, user) = db_with_user().await;
AdminRecoveryCode::replace_all(user.id, &["a-secret-hash".to_string()], &db)
.await
.unwrap();
let code = AdminRecoveryCode::list_unused(user.id, &db).await.unwrap()[0].clone();
let rendered = code.to_json().to_string();
assert!(!rendered.contains("a-secret-hash"));
assert!(!rendered.contains("codeHash"));
assert!(rendered.contains(&code.id.to_string()));
assert!(rendered.contains("\"usedAt\":null"));
}
}