use std::time::Duration;
use chrono::{DateTime, TimeDelta, Utc};
use sha2::{Digest, Sha256};
use sqlx::Row;
use subtle::ConstantTimeEq;
use zeroize::Zeroize;
use super::dialect::{RememberPool, sql, stored_time};
use super::error::RememberTokenError;
use super::migrate;
use super::token::{
IssuedRememberToken, PlaintextRememberToken, SECRET_BYTES, SERIES_BYTES, format_plaintext,
parse_plaintext,
};
const ISSUE_ATTEMPTS: u32 = 8;
const DEFAULT_GRACE: Duration = Duration::from_secs(60);
#[derive(Debug)]
#[non_exhaustive]
pub enum RememberOutcome {
Unrecognised,
#[non_exhaustive]
Accepted {
subject: String,
replacement: Option<PlaintextRememberToken>,
},
#[non_exhaustive]
Theft {
subject: String,
},
}
impl RememberOutcome {
#[must_use]
pub fn subject(&self) -> Option<&str> {
match self {
Self::Unrecognised => None,
Self::Accepted { subject, .. } | Self::Theft { subject } => Some(subject),
}
}
}
#[derive(Clone, Debug)]
#[non_exhaustive]
pub struct RememberTokens {
pool: RememberPool,
grace: Duration,
}
impl RememberTokens {
#[must_use]
pub fn new(pool: RememberPool) -> Self {
Self {
pool,
grace: DEFAULT_GRACE,
}
}
#[must_use]
pub fn grace(mut self, grace: Duration) -> Self {
self.grace = grace;
self
}
#[must_use]
pub fn pool(&self) -> &RememberPool {
&self.pool
}
pub async fn migrate(&self) -> Result<(), RememberTokenError> {
migrate::apply(&self.pool).await
}
pub async fn issue(
&self,
subject: &str,
ttl: Duration,
) -> Result<IssuedRememberToken, RememberTokenError> {
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);
for _ in 0..ISSUE_ATTEMPTS {
let mut series = [0u8; SERIES_BYTES];
fill_random(&mut series)?;
let written = sqlx::query(sql::INSERT_NEW)
.bind(series.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 = PlaintextRememberToken::new(format_plaintext(&series, &secret));
secret.zeroize();
return Ok(IssuedRememberToken::new(
subject.to_owned(),
expires_at,
plaintext,
));
}
secret.zeroize();
Err(RememberTokenError::SeriesCollision {
attempts: ISSUE_ATTEMPTS,
})
}
pub async fn present(&self, presented: &str) -> Result<RememberOutcome, RememberTokenError> {
let Some((series, mut secret)) = parse_plaintext(presented) else {
return Ok(RememberOutcome::Unrecognised);
};
let mut presented_digest = digest_of(&secret);
secret.zeroize();
let found = sqlx::query(sql::FIND_LIVE)
.bind(stored_time(self.grace_cutoff()))
.bind(series.as_bytes().to_vec())
.fetch_optional(&self.pool)
.await?;
let Some(row) = found else {
presented_digest.zeroize();
return Ok(RememberOutcome::Unrecognised);
};
let current: Vec<u8> = row.try_get(0)?;
let previous: Option<Vec<u8>> = row.try_get(1)?;
let rotation_is_recent: bool = row.try_get::<i64, _>(2)? != 0;
let subject: String = row.try_get(3)?;
let retired = previous.unwrap_or_else(|| vec![0u8; 32]);
let matches_current: bool = presented_digest.ct_eq(current.as_slice()).into();
let matches_retired: bool = presented_digest.ct_eq(retired.as_slice()).into();
presented_digest.zeroize();
if matches_current {
return self.rotate(&series, current, subject).await;
}
if matches_retired && rotation_is_recent {
return Ok(RememberOutcome::Accepted {
subject,
replacement: None,
});
}
sqlx::query(sql::DELETE_FOR)
.bind(&subject)
.execute(&self.pool)
.await?;
Ok(RememberOutcome::Theft { subject })
}
async fn rotate(
&self,
series: &super::token::SeriesId,
expected: Vec<u8>,
subject: String,
) -> Result<RememberOutcome, RememberTokenError> {
let mut next = [0u8; SECRET_BYTES];
fill_random(&mut next)?;
let next_digest = digest_of(&next);
let rotated = sqlx::query(sql::ROTATE)
.bind(next_digest.to_vec())
.bind(stored_time(Utc::now()))
.bind(series.as_bytes().to_vec())
.bind(expected)
.execute(&self.pool)
.await?
.rows_affected();
if rotated == 0 {
next.zeroize();
return Ok(RememberOutcome::Accepted {
subject,
replacement: None,
});
}
let plaintext = PlaintextRememberToken::new(format_plaintext(series.as_bytes(), &next));
next.zeroize();
Ok(RememberOutcome::Accepted {
subject,
replacement: Some(plaintext),
})
}
pub async fn revoke_all_for(&self, subject: &str) -> Result<u64, RememberTokenError> {
let result = sqlx::query(sql::DELETE_FOR)
.bind(subject)
.execute(&self.pool)
.await?;
Ok(result.rows_affected())
}
pub async fn revoke(&self, presented: &str) -> Result<bool, RememberTokenError> {
let Some((series, mut secret)) = parse_plaintext(presented) else {
return Ok(false);
};
secret.zeroize();
let deleted = sqlx::query(sql::DELETE_SERIES)
.bind(series.as_bytes().to_vec())
.execute(&self.pool)
.await?
.rows_affected();
Ok(deleted > 0)
}
pub async fn sweep_expired(&self) -> Result<u64, RememberTokenError> {
let result = sqlx::query(sql::DELETE_EXPIRED).execute(&self.pool).await?;
Ok(result.rows_affected())
}
fn grace_cutoff(&self) -> DateTime<Utc> {
let delta = TimeDelta::from_std(self.grace).unwrap_or(TimeDelta::MAX);
Utc::now()
.checked_sub_signed(delta)
.unwrap_or(DateTime::<Utc>::MIN_UTC)
}
}
fn deadline(ttl: Duration) -> Result<DateTime<Utc>, RememberTokenError> {
let describe = || format!("now + {ttl:?}");
let delta = TimeDelta::from_std(ttl).map_err(|_| RememberTokenError::Expiry(describe()))?;
Utc::now()
.checked_add_signed(delta)
.ok_or_else(|| RememberTokenError::Expiry(describe()))
}
pub(super) fn digest_of(secret: &[u8]) -> [u8; 32] {
Sha256::digest(secret).into()
}
fn fill_random(buffer: &mut [u8]) -> Result<(), RememberTokenError> {
getrandom::fill(buffer).map_err(|_| RememberTokenError::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 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 a_rotation_never_writes_the_digest_it_replaces() {
let mut a = [0u8; SECRET_BYTES];
let mut b = [0u8; SECRET_BYTES];
fill_random(&mut a).expect("the OS randomness source is available");
fill_random(&mut b).expect("the OS randomness source is available");
assert_ne!(digest_of(&a), digest_of(&b));
}
#[test]
fn an_ordinary_ttl_lands_in_the_future() {
let before = Utc::now();
let at = deadline(Duration::from_secs(60 * 60 * 24 * 30)).expect("a month is fine");
assert!(at > before);
}
#[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, RememberTokenError::Expiry(_)));
}
#[test]
fn the_default_grace_window_is_in_the_past_and_close_to_it() {
let store_grace = DEFAULT_GRACE;
let cutoff = Utc::now() - TimeDelta::from_std(store_grace).expect("a minute is fine");
assert!(cutoff < Utc::now());
assert!(cutoff > Utc::now() - TimeDelta::seconds(120));
}
#[test]
fn a_zero_grace_window_puts_the_cutoff_at_now() {
let cutoff = Utc::now() - TimeDelta::from_std(Duration::ZERO).expect("zero is fine");
assert!(cutoff <= Utc::now());
}
#[test]
fn an_outcome_names_its_subject_except_when_there_is_none() {
assert_eq!(RememberOutcome::Unrecognised.subject(), None);
assert_eq!(
RememberOutcome::Accepted {
subject: "user@example.test".to_owned(),
replacement: None,
}
.subject(),
Some("user@example.test")
);
assert_eq!(
RememberOutcome::Theft {
subject: "user@example.test".to_owned(),
}
.subject(),
Some("user@example.test")
);
}
#[test]
fn a_debug_of_an_accepted_outcome_does_not_leak_the_replacement() {
let outcome = RememberOutcome::Accepted {
subject: "user@example.test".to_owned(),
replacement: Some(PlaintextRememberToken::new("arcrmb_dead.beef".to_owned())),
};
let rendered = format!("{outcome:?}");
assert!(rendered.contains("user@example.test"));
assert!(!rendered.contains("beef"), "{rendered}");
}
}