use std::time::Duration;
use chrono::{TimeDelta, Utc};
use sha2::{Digest, Sha256};
use sqlx::Row;
use subtle::ConstantTimeEq;
use zeroize::Zeroize;
use super::dialect::{ResetPool, sql, stored_time};
use super::error::PasswordResetError;
use super::migrate;
use super::token::{
ID_BYTES, IssuedPasswordReset, PlaintextReset, SECRET_BYTES, format_plaintext, parse_plaintext,
};
const ISSUE_ATTEMPTS: u32 = 8;
#[derive(Clone, Debug)]
#[non_exhaustive]
pub struct PasswordResets {
pool: ResetPool,
}
impl PasswordResets {
#[must_use]
pub fn new(pool: ResetPool) -> Self {
Self { pool }
}
#[must_use]
pub fn pool(&self) -> &ResetPool {
&self.pool
}
pub async fn migrate(&self) -> Result<(), PasswordResetError> {
migrate::apply(&self.pool).await
}
pub async fn issue(
&self,
subject: &str,
ttl: Duration,
) -> Result<IssuedPasswordReset, PasswordResetError> {
let expires_at = deadline(ttl)?;
let created_at = Utc::now();
let mut secret = [0u8; SECRET_BYTES];
fill_random(&mut secret)?;
let digest = digest_of(&secret);
sqlx::query(sql::DELETE_FOR)
.bind(subject)
.execute(&self.pool)
.await?;
for _ in 0..ISSUE_ATTEMPTS {
let mut id_bytes = [0u8; ID_BYTES];
fill_random(&mut id_bytes)?;
let written = sqlx::query(sql::INSERT_NEW)
.bind(id_bytes.to_vec())
.bind(digest.to_vec())
.bind(subject)
.bind(stored_time(expires_at))
.bind(stored_time(created_at))
.execute(&self.pool)
.await?
.rows_affected();
if written == 0 {
continue;
}
let plaintext = PlaintextReset::new(format_plaintext(&id_bytes, &secret));
secret.zeroize();
return Ok(IssuedPasswordReset::new(
subject.to_owned(),
expires_at,
plaintext,
));
}
secret.zeroize();
Err(PasswordResetError::IdCollision {
attempts: ISSUE_ATTEMPTS,
})
}
pub async fn consume(&self, presented: &str) -> Result<Option<String>, PasswordResetError> {
let Some((id, mut secret)) = parse_plaintext(presented) else {
return Ok(None);
};
let mut presented_digest = digest_of(&secret);
secret.zeroize();
let found = sqlx::query(sql::FIND_LIVE)
.bind(id.as_bytes().to_vec())
.fetch_optional(&self.pool)
.await?;
let Some(row) = found else {
presented_digest.zeroize();
return Ok(None);
};
let stored: Vec<u8> = row.try_get(0)?;
let matches: bool = presented_digest.ct_eq(stored.as_slice()).into();
presented_digest.zeroize();
if !matches {
return Ok(None);
}
let subject: String = row.try_get(1)?;
let spent = sqlx::query(sql::DELETE_FOR)
.bind(&subject)
.execute(&self.pool)
.await?
.rows_affected();
if spent == 0 {
return Ok(None);
}
Ok(Some(subject))
}
pub async fn revoke_all_for(&self, subject: &str) -> Result<u64, PasswordResetError> {
let result = sqlx::query(sql::DELETE_FOR)
.bind(subject)
.execute(&self.pool)
.await?;
Ok(result.rows_affected())
}
pub async fn sweep_expired(&self) -> Result<u64, PasswordResetError> {
let result = sqlx::query(sql::DELETE_EXPIRED).execute(&self.pool).await?;
Ok(result.rows_affected())
}
}
fn deadline(ttl: Duration) -> Result<chrono::DateTime<Utc>, PasswordResetError> {
let describe = || format!("now + {ttl:?}");
let delta = TimeDelta::from_std(ttl).map_err(|_| PasswordResetError::Expiry(describe()))?;
Utc::now()
.checked_add_signed(delta)
.ok_or_else(|| PasswordResetError::Expiry(describe()))
}
fn digest_of(secret: &[u8]) -> [u8; 32] {
Sha256::digest(secret).into()
}
fn fill_random(buffer: &mut [u8]) -> Result<(), PasswordResetError> {
getrandom::fill(buffer).map_err(|_| PasswordResetError::Entropy)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_digest_is_not_the_secret() {
let secret = [0xa5u8; SECRET_BYTES];
let digest = digest_of(&secret);
assert_eq!(digest.len(), 32);
assert_ne!(&digest[..], &secret[..]);
}
#[test]
fn the_digest_is_the_documented_sha256() {
let digest = crate::crypt::base64url::encode(&digest_of(&[0u8; 32]));
assert_eq!(digest, "Zmh6rfhivXdsj8GLjp-OIAiXFIVu4jOzkCpZHQ1fKSU");
}
#[test]
fn two_secrets_that_differ_in_one_bit_give_different_digests() {
let a = [0u8; SECRET_BYTES];
let mut b = [0u8; SECRET_BYTES];
b[SECRET_BYTES - 1] = 1;
assert_ne!(digest_of(&a), digest_of(&b));
assert_eq!(digest_of(&a), digest_of(&[0u8; SECRET_BYTES]));
}
#[test]
fn the_random_source_fills_the_whole_buffer() {
let mut first = [0u8; SECRET_BYTES];
let mut second = [0u8; SECRET_BYTES];
fill_random(&mut first).expect("the OS randomness source is available");
fill_random(&mut second).expect("the OS randomness source is available");
assert_ne!(first, [0u8; SECRET_BYTES]);
assert_ne!(first, second);
}
#[test]
fn an_ordinary_ttl_lands_in_the_future() {
let before = Utc::now();
let at = deadline(Duration::from_secs(3600)).expect("an hour is representable");
assert!(at > before);
assert!(at <= before + TimeDelta::seconds(3601));
}
#[test]
fn a_ttl_too_wide_to_represent_is_refused_rather_than_wrapped() {
let error = deadline(Duration::from_secs(u64::MAX)).expect_err("not representable");
assert!(matches!(error, PasswordResetError::Expiry(_)));
}
#[test]
fn a_zero_ttl_is_representable_and_already_expired() {
let at = deadline(Duration::ZERO).expect("zero is representable");
assert!(at <= Utc::now());
}
}