use std::sync::Arc;
use async_trait::async_trait;
use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD};
use chrono::{DateTime, Duration, Utc};
use sha2::{Digest, Sha256};
use sqlx::Row;
use super::{LOCAL_PROVIDER, LocalPasswordAuthenticator, db_error, validate_password};
use crate::{
account_linking::normalize_email,
audit::logger::{AuditEventType, SecretType, get_audit_logger},
constant_time::ConstantTimeOps,
error::{AuthError, Result},
session::SessionStore,
};
pub const RESET_TOKEN_TTL_SECS: i64 = 3600;
const SELECTOR_LEN: usize = 16;
const VERIFIER_LEN: usize = 32;
pub const PASSWORD_RESET_SCHEMA_SQL: &str = r"
CREATE SCHEMA IF NOT EXISTS core;
CREATE TABLE IF NOT EXISTS core.tb_password_reset_token (
pk_password_reset_token BIGINT GENERATED ALWAYS AS IDENTITY PRIMARY KEY,
id UUID NOT NULL DEFAULT gen_random_uuid(),
fk_user BIGINT NOT NULL REFERENCES core.tb_user (pk_user) ON DELETE CASCADE,
user_id TEXT NOT NULL,
selector TEXT NOT NULL,
verifier_hash BYTEA NOT NULL,
expires_at TIMESTAMPTZ NOT NULL,
used_at TIMESTAMPTZ,
tenant_id UUID,
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
UNIQUE (selector)
);
CREATE INDEX IF NOT EXISTS idx_password_reset_token_fk_user
ON core.tb_password_reset_token (fk_user);
-- RLS deny-by-default (mirrors core.tb_user / core.tb_auth_identity / tb_password_credential).
ALTER TABLE core.tb_password_reset_token ENABLE ROW LEVEL SECURITY;
DROP POLICY IF EXISTS p_password_reset_token_tenant_read ON core.tb_password_reset_token;
CREATE POLICY p_password_reset_token_tenant_read ON core.tb_password_reset_token
FOR SELECT USING (tenant_id = NULLIF(current_setting('fraiseql.tenant_id', true), '')::uuid);
DROP POLICY IF EXISTS p_password_reset_token_insert ON core.tb_password_reset_token;
CREATE POLICY p_password_reset_token_insert ON core.tb_password_reset_token
FOR INSERT WITH CHECK (true);
-- Least-privilege baseline: never world-readable. RLS is defence-in-depth on top.
REVOKE ALL ON core.tb_password_reset_token FROM PUBLIC;
";
#[async_trait]
pub trait ResetEmailSender: Send + Sync {
async fn send_reset_link(&self, to: &str, token: &str) -> Result<()>;
}
struct ResetToken {
selector: [u8; SELECTOR_LEN],
verifier: [u8; VERIFIER_LEN],
}
struct ParsedToken {
selector: String,
verifier_hash: Vec<u8>,
}
impl ResetToken {
fn generate() -> Self {
use rand::RngCore as _;
let mut selector = [0u8; SELECTOR_LEN];
let mut verifier = [0u8; VERIFIER_LEN];
rand::rng().fill_bytes(&mut selector);
rand::rng().fill_bytes(&mut verifier);
Self { selector, verifier }
}
fn selector_b64(&self) -> String {
URL_SAFE_NO_PAD.encode(self.selector)
}
fn verifier_hash(&self) -> Vec<u8> {
Sha256::digest(self.verifier).to_vec()
}
fn to_token_string(&self) -> String {
format!(
"{}.{}",
URL_SAFE_NO_PAD.encode(self.selector),
URL_SAFE_NO_PAD.encode(self.verifier)
)
}
fn parse(token: &str) -> Result<ParsedToken> {
let (selector_b64, verifier_b64) =
token.split_once('.').ok_or_else(|| AuthError::InvalidToken {
reason: "reset token is not in selector.verifier form".to_string(),
})?;
let selector =
URL_SAFE_NO_PAD.decode(selector_b64).map_err(|_| AuthError::InvalidToken {
reason: "reset token selector is not valid base64url".to_string(),
})?;
let verifier =
URL_SAFE_NO_PAD.decode(verifier_b64).map_err(|_| AuthError::InvalidToken {
reason: "reset token verifier is not valid base64url".to_string(),
})?;
if selector.len() != SELECTOR_LEN || verifier.len() != VERIFIER_LEN {
return Err(AuthError::InvalidToken {
reason: "reset token has an unexpected length".to_string(),
});
}
Ok(ParsedToken {
selector: selector_b64.to_string(),
verifier_hash: Sha256::digest(&verifier).to_vec(),
})
}
}
fn invalid_reset_token() -> AuthError {
AuthError::InvalidToken {
reason: "invalid, expired, or already-used password reset token".to_string(),
}
}
impl LocalPasswordAuthenticator {
#[must_use]
pub fn with_email_sender(mut self, sender: Arc<dyn ResetEmailSender>) -> Self {
self.email_sender = Some(sender);
self
}
#[must_use]
pub fn with_session_store(mut self, store: Arc<dyn SessionStore>) -> Self {
self.session_store = Some(store);
self
}
pub async fn start_password_reset(&self, email: &str) -> Result<()> {
let normalized = normalize_email(email);
let logger = get_audit_logger();
let row = sqlx::query(
"SELECT c.fk_user, c.user_id \
FROM core.tb_password_credential c \
JOIN core.tb_auth_identity i ON i.fk_user = c.fk_user \
WHERE i.provider = $1 AND i.provider_id = $2",
)
.bind(LOCAL_PROVIDER)
.bind(&normalized)
.fetch_optional(&self.db)
.await
.map_err(|e| db_error("lookup credential for reset", &e))?;
let Some(row) = row else {
logger.log_failure(
AuditEventType::AuthFailure,
SecretType::SessionToken,
None,
"password_reset_start",
"no_local_credential",
);
return Ok(());
};
let fk_user: i64 = row.get("fk_user");
let user_id: String = row.get("user_id");
let token = ResetToken::generate();
let expires_at = Utc::now() + Duration::seconds(RESET_TOKEN_TTL_SECS);
sqlx::query(
"INSERT INTO core.tb_password_reset_token \
(fk_user, user_id, selector, verifier_hash, expires_at) \
VALUES ($1, $2, $3, $4, $5)",
)
.bind(fk_user)
.bind(&user_id)
.bind(token.selector_b64())
.bind(token.verifier_hash())
.bind(expires_at)
.execute(&self.db)
.await
.map_err(|e| db_error("insert reset token", &e))?;
if let Some(sender) = self.email_sender.clone() {
let to = normalized;
let token_str = token.to_token_string();
tokio::spawn(async move {
if let Err(e) = sender.send_reset_link(&to, &token_str).await {
tracing::warn!("password_reset_start: reset email dispatch failed: {e}");
}
});
} else {
tracing::warn!(
"password_reset_start: token issued but no ResetEmailSender is configured; \
the reset link was not delivered"
);
}
logger.log_success(
AuditEventType::AuthSuccess,
SecretType::SessionToken,
Some(user_id),
"password_reset_start",
);
Ok(())
}
pub async fn confirm_password_reset(&self, token: &str, new_password: &str) -> Result<()> {
validate_password(new_password)?;
let logger = get_audit_logger();
let Ok(parsed) = ResetToken::parse(token) else {
logger.log_failure(
AuditEventType::AuthFailure,
SecretType::SessionToken,
None,
"password_reset_confirm",
"malformed_token",
);
return Err(invalid_reset_token());
};
let row = sqlx::query(
"SELECT pk_password_reset_token AS pk, fk_user, user_id, verifier_hash, \
expires_at, used_at \
FROM core.tb_password_reset_token WHERE selector = $1",
)
.bind(&parsed.selector)
.fetch_optional(&self.db)
.await
.map_err(|e| db_error("lookup reset token", &e))?;
let Some(row) = row else {
logger.log_failure(
AuditEventType::AuthFailure,
SecretType::SessionToken,
None,
"password_reset_confirm",
"unknown_selector",
);
return Err(invalid_reset_token());
};
let pk: i64 = row.get("pk");
let fk_user: i64 = row.get("fk_user");
let user_id: String = row.get("user_id");
let stored_hash: Vec<u8> = row.get("verifier_hash");
let expires_at: DateTime<Utc> = row.get("expires_at");
let used_at: Option<DateTime<Utc>> = row.get("used_at");
if !ConstantTimeOps::compare(&stored_hash, &parsed.verifier_hash) {
logger.log_failure(
AuditEventType::AuthFailure,
SecretType::SessionToken,
Some(user_id),
"password_reset_confirm",
"bad_verifier",
);
return Err(invalid_reset_token());
}
if used_at.is_some() {
logger.log_failure(
AuditEventType::AuthFailure,
SecretType::SessionToken,
Some(user_id),
"password_reset_confirm",
"used",
);
return Err(invalid_reset_token());
}
if expires_at <= Utc::now() {
logger.log_failure(
AuditEventType::AuthFailure,
SecretType::SessionToken,
Some(user_id),
"password_reset_confirm",
"expired",
);
return Err(invalid_reset_token());
}
let new_hash = self.hash_password(new_password)?;
let mut tx = self.db.begin().await.map_err(|e| db_error("begin reset transaction", &e))?;
let consumed = sqlx::query(
"UPDATE core.tb_password_reset_token SET used_at = now() \
WHERE pk_password_reset_token = $1 AND used_at IS NULL AND expires_at > now()",
)
.bind(pk)
.execute(&mut *tx)
.await
.map_err(|e| db_error("consume reset token", &e))?;
if consumed.rows_affected() == 0 {
tx.rollback().await.map_err(|e| db_error("rollback reset transaction", &e))?;
logger.log_failure(
AuditEventType::AuthFailure,
SecretType::SessionToken,
Some(user_id),
"password_reset_confirm",
"race",
);
return Err(invalid_reset_token());
}
let updated = sqlx::query(
"UPDATE core.tb_password_credential SET password_hash = $1, updated_at = now() \
WHERE fk_user = $2",
)
.bind(&new_hash)
.bind(fk_user)
.execute(&mut *tx)
.await
.map_err(|e| db_error("update credential on reset", &e))?;
if updated.rows_affected() == 0 {
tx.rollback().await.map_err(|e| db_error("rollback reset transaction", &e))?;
return Err(AuthError::Internal {
message: "reset token resolved to a user with no local credential".to_string(),
});
}
sqlx::query(
"UPDATE core.tb_password_reset_token SET used_at = now() \
WHERE fk_user = $1 AND used_at IS NULL",
)
.bind(fk_user)
.execute(&mut *tx)
.await
.map_err(|e| db_error("invalidate sibling reset tokens", &e))?;
tx.commit().await.map_err(|e| db_error("commit reset transaction", &e))?;
if let Some(store) = self.session_store.as_ref() {
if let Err(e) = store.revoke_all_sessions(&user_id).await {
tracing::warn!(
"password_reset_confirm: session revocation failed for {user_id}: {e}"
);
}
} else {
tracing::warn!(
"password_reset_confirm: no session store configured; outstanding sessions for \
{user_id} were not revoked"
);
}
logger.log_success(
AuditEventType::AuthSuccess,
SecretType::SessionToken,
Some(user_id),
"password_reset_confirm",
);
Ok(())
}
}
#[allow(clippy::unwrap_used)] #[cfg(test)]
mod tests;