use std::time::{Duration, SystemTime};
use tracing::{debug, info};
use crate::db::Database;
use acme_proxy_core::random::random_token;
#[derive(Debug)]
pub struct Nonce {
pub value: String,
pub created_at: i64,
}
impl Default for Nonce {
fn default() -> Self {
Self::new()
}
}
#[must_use]
pub fn fingerprint(value: &str) -> &str {
value.get(..8).unwrap_or(value)
}
pub fn now_secs() -> i64 {
SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.unwrap_or_default()
.as_secs() as i64
}
impl Nonce {
#[must_use]
pub fn new() -> Self {
Nonce {
value: random_token(),
created_at: now_secs(),
}
}
pub async fn save(&self, database: &Database) -> Result<(), sqlx::Error> {
crate::sql::query("INSERT INTO nonces VALUES (?, ?);")
.bind(self.value.clone())
.bind(self.created_at)
.execute(database)
.await?;
debug!(event = "db_nonce_saved",
outcome = "success",
nonce_fp = %fingerprint(&self.value),
created_at = ?self.created_at);
Ok(())
}
pub async fn verify(
nonce: &str,
database: &Database,
ttl: Duration,
) -> Result<bool, sqlx::Error> {
let cutoff = now_secs().saturating_sub(ttl.as_secs() as i64);
let result = crate::sql::query("DELETE FROM nonces WHERE value = ? AND created_at > ?;")
.bind(nonce)
.bind(cutoff)
.execute(database)
.await?;
let is_valid = result.rows_affected() == 1;
if is_valid {
debug!(
event = "db_nonce_verified_valid",
outcome = "success",
nonce_fp = %fingerprint(nonce),
cutoff = cutoff,
ttl_seconds = ttl.as_secs(),
);
} else {
debug!(
event = "db_nonce_verified_invalid",
outcome = "failure",
nonce_fp = %fingerprint(nonce),
cutoff = cutoff,
ttl_seconds = ttl.as_secs(),
);
}
Ok(is_valid)
}
pub async fn count(database: &Database) -> Result<i64, sqlx::Error> {
crate::sql::query("SELECT COUNT(*) FROM nonces;")
.fetch_one(database)
.await?
.try_get(0usize)
}
pub async fn cleanup(database: &Database, ttl: Duration) -> Result<u64, sqlx::Error> {
let cutoff = now_secs().saturating_sub(ttl.as_secs() as i64);
let result = crate::sql::query("DELETE FROM nonces WHERE created_at <= ?;")
.bind(cutoff)
.execute(database)
.await?;
info!(event = "db_nonce_cleanup_completed",
outcome = "success",
rows_removed = ?result.rows_affected(),
cutoff = ?cutoff,
ttl_seconds = ?ttl.as_secs());
Ok(result.rows_affected())
}
}
#[cfg(test)]
mod tests {
use super::*;
use base64::prelude::*;
use std::sync::Arc;
const TTL: Duration = Duration::from_secs(300);
async fn nonce_count(database: &Arc<Database>) -> i64 {
crate::sql::query("SELECT COUNT(*) FROM nonces;")
.fetch_one(database)
.await
.unwrap()
.try_get(0usize)
.unwrap()
}
#[tokio::test]
async fn verify_accepts_fresh_nonce_exactly_once() {
let database = Arc::new(Database::connect_for_test().await.unwrap());
let nonce = Nonce::new();
let value = nonce.value.clone();
nonce.save(&database).await.unwrap();
assert!(Nonce::verify(&value, &database, TTL).await.unwrap());
assert!(!Nonce::verify(&value, &database, TTL).await.unwrap());
}
#[tokio::test]
async fn verify_rejects_unknown_nonce() {
let database = Arc::new(Database::connect_for_test().await.unwrap());
assert!(!Nonce::verify("never-issued", &database, TTL).await.unwrap());
}
#[tokio::test]
async fn verify_rejects_expired_nonce() {
let database = Arc::new(Database::connect_for_test().await.unwrap());
Nonce {
value: "stale".to_string(),
created_at: now_secs() - 600,
}
.save(&database)
.await
.unwrap();
assert!(!Nonce::verify("stale", &database, TTL).await.unwrap());
}
#[tokio::test]
async fn verify_accepts_nonce_near_edge_of_window() {
let database = Arc::new(Database::connect_for_test().await.unwrap());
#[allow(clippy::cast_possible_wrap)]
let created_at = now_secs() - (TTL.as_secs() as i64 - 2);
Nonce {
value: "edge".to_string(),
created_at,
}
.save(&database)
.await
.unwrap();
assert!(Nonce::verify("edge", &database, TTL).await.unwrap());
}
#[tokio::test]
async fn verify_rejects_nonce_at_exact_cutoff_boundary() {
let database = Arc::new(Database::connect_for_test().await.unwrap());
#[allow(clippy::cast_possible_wrap)]
let created_at = now_secs() - TTL.as_secs() as i64;
Nonce {
value: "boundary".to_string(),
created_at,
}
.save(&database)
.await
.unwrap();
assert!(!Nonce::verify("boundary", &database, TTL).await.unwrap());
}
#[tokio::test]
async fn cleanup_removes_only_stale_nonces() {
let database = Arc::new(Database::connect_for_test().await.unwrap());
Nonce {
value: "stale".to_string(),
created_at: now_secs() - 600,
}
.save(&database)
.await
.unwrap();
let fresh = Nonce::new();
let fresh_value = fresh.value.clone();
fresh.save(&database).await.unwrap();
assert_eq!(nonce_count(&database).await, 2);
let removed = Nonce::cleanup(&database, TTL).await.unwrap();
assert_eq!(removed, 1);
assert_eq!(nonce_count(&database).await, 1);
assert!(Nonce::verify(&fresh_value, &database, TTL).await.unwrap());
}
#[test]
fn a_minted_nonce_is_base64url_over_32_bytes() {
let value = Nonce::new().value;
assert!(
value
.bytes()
.all(|b| b.is_ascii_alphanumeric() || b == b'-' || b == b'_'),
"outside the base64url alphabet: {value}"
);
assert_eq!(
BASE64_URL_SAFE_NO_PAD.decode(&value).unwrap().len(),
32,
"256 bits, like every other non-guessable value here"
);
}
#[test]
fn the_default_nonce_is_a_fresh_one() {
let nonce = Nonce::default();
assert!(!nonce.value.is_empty());
assert_ne!(nonce.value, Nonce::default().value, "each must be unique");
}
#[tokio::test]
async fn count_reports_the_table_size_and_follows_cleanup() {
let db = Arc::new(Database::connect_for_test().await.unwrap());
assert_eq!(Nonce::count(&db).await.unwrap(), 0);
for _ in 0..3 {
Nonce::new().save(&db).await.unwrap();
}
assert_eq!(Nonce::count(&db).await.unwrap(), 3);
let stale = Nonce {
value: "stale".to_string(),
created_at: now_secs() - 10_000,
};
stale.save(&db).await.unwrap();
assert_eq!(Nonce::count(&db).await.unwrap(), 4);
Nonce::cleanup(&db, TTL).await.unwrap();
assert_eq!(Nonce::count(&db).await.unwrap(), 3);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn concurrent_verify_of_one_nonce_succeeds_exactly_once() {
let file =
std::env::temp_dir().join(format!("acme-proxy-test-{}.db", uuid::Uuid::now_v7()));
let url = format!("sqlite://{}", file.display());
let database = Arc::new(Database::connect_and_migrate(&url).await.unwrap());
let nonce = Nonce::new();
nonce.save(&database).await.unwrap();
const RACERS: usize = 8;
let barrier = Arc::new(tokio::sync::Barrier::new(RACERS));
let mut tasks = Vec::with_capacity(RACERS);
for _ in 0..RACERS {
let database = database.clone();
let barrier = barrier.clone();
let value = nonce.value.clone();
tasks.push(tokio::spawn(async move {
barrier.wait().await;
Nonce::verify(&value, &database, TTL).await
}));
}
let mut accepted = 0;
for task in tasks {
if task.await.unwrap().unwrap() {
accepted += 1;
}
}
assert_eq!(accepted, 1, "a nonce may be spent exactly once");
assert_eq!(nonce_count(&database).await, 0, "the row is consumed");
database.close().await;
for suffix in ["", "-wal", "-shm"] {
let _ = std::fs::remove_file(format!("{}{suffix}", file.display()));
}
}
}