use async_trait::async_trait;
use sha2::{Digest, Sha256};
use sqlx::{PgPool, Row};
use subtle::ConstantTimeEq as _;
use super::{MAX_VERIFY_ATTEMPTS, OTP_RATE_MAX, OTP_RATE_WINDOW_SECS, OTP_TTL_SECS, OtpStore};
use crate::{
error::{AuthError, Result},
session::unix_now,
};
pub const PG_OTP_SCHEMA_SQL: &str = r"
CREATE SCHEMA IF NOT EXISTS core;
CREATE TABLE IF NOT EXISTS core.tb_otp_code (
email TEXT PRIMARY KEY,
code_hash BYTEA NOT NULL,
expires_at BIGINT NOT NULL,
attempts INTEGER NOT NULL DEFAULT 0
);
CREATE TABLE IF NOT EXISTS core.tb_otp_send_budget (
email TEXT PRIMARY KEY,
count INTEGER NOT NULL DEFAULT 0,
window_start BIGINT NOT NULL
);
REVOKE ALL ON core.tb_otp_code FROM PUBLIC;
REVOKE ALL ON core.tb_otp_send_budget FROM PUBLIC;
";
fn code_hash(code: &str) -> Vec<u8> {
Sha256::digest(code.as_bytes()).to_vec()
}
fn db_err(e: sqlx::Error) -> AuthError {
AuthError::DatabaseError {
message: e.to_string(),
}
}
#[derive(Debug, Clone)]
pub struct PgOtpStore {
db: PgPool,
}
impl PgOtpStore {
#[must_use]
pub const fn new(db: PgPool) -> Self {
Self { db }
}
pub async fn init(&self) -> Result<()> {
for stmt in PG_OTP_SCHEMA_SQL.split(';').map(str::trim).filter(|s| !s.is_empty()) {
sqlx::query(stmt).execute(&self.db).await.map_err(db_err)?;
}
Ok(())
}
}
#[async_trait]
impl OtpStore for PgOtpStore {
#[allow(clippy::cast_possible_wrap)] async fn sweep_expired(&self) -> Result<u64> {
let now = unix_now()? as i64;
let codes = sqlx::query("DELETE FROM core.tb_otp_code WHERE expires_at <= $1")
.bind(now)
.execute(&self.db)
.await
.map_err(db_err)?
.rows_affected();
let window = i64::try_from(crate::otp::OTP_RATE_WINDOW_SECS).unwrap_or(i64::MAX);
let budgets =
sqlx::query("DELETE FROM core.tb_otp_send_budget WHERE window_start + $1 <= $2")
.bind(window)
.bind(now)
.execute(&self.db)
.await
.map_err(db_err)?
.rows_affected();
Ok(codes + budgets)
}
#[allow(clippy::cast_possible_wrap, clippy::cast_sign_loss)] async fn create_otp(&self, email: &str) -> Result<String> {
let now = unix_now()?;
let reserved = sqlx::query(
"INSERT INTO core.tb_otp_send_budget (email, count, window_start)
VALUES ($1, 1, $2)
ON CONFLICT (email) DO UPDATE SET
count = CASE
WHEN $2 - core.tb_otp_send_budget.window_start >= $3 THEN 1
ELSE core.tb_otp_send_budget.count + 1 END,
window_start = CASE
WHEN $2 - core.tb_otp_send_budget.window_start >= $3 THEN $2
ELSE core.tb_otp_send_budget.window_start END
RETURNING count, window_start",
)
.bind(email)
.bind(now as i64)
.bind(OTP_RATE_WINDOW_SECS as i64)
.fetch_one(&self.db)
.await
.map_err(db_err)?;
let count: i32 = reserved.get("count");
let window_start: i64 = reserved.get("window_start");
if count as u32 > OTP_RATE_MAX {
return Err(AuthError::RateLimited {
retry_after_secs: (window_start as u64 + OTP_RATE_WINDOW_SECS).saturating_sub(now),
});
}
let code = format!("{:06}", rand::Rng::random_range(&mut rand::rng(), 0u32..1_000_000));
sqlx::query(
"INSERT INTO core.tb_otp_code (email, code_hash, expires_at, attempts)
VALUES ($1, $2, $3, 0)
ON CONFLICT (email) DO UPDATE SET
code_hash = EXCLUDED.code_hash,
expires_at = EXCLUDED.expires_at,
attempts = 0",
)
.bind(email)
.bind(code_hash(&code))
.bind((now + OTP_TTL_SECS) as i64)
.execute(&self.db)
.await
.map_err(db_err)?;
Ok(code)
}
#[allow(clippy::cast_sign_loss, clippy::cast_possible_wrap)] async fn verify_otp(&self, email: &str, code: &str) -> Result<()> {
let now = unix_now()?;
let row = sqlx::query(
"UPDATE core.tb_otp_code SET attempts = attempts + 1 WHERE email = $1
RETURNING code_hash, expires_at, attempts",
)
.bind(email)
.fetch_optional(&self.db)
.await
.map_err(db_err)?
.ok_or_else(|| AuthError::InvalidToken {
reason: "no pending OTP for email".into(),
})?;
let stored_hash: Vec<u8> = row.get("code_hash");
let expires_at: i64 = row.get("expires_at");
let attempts: i32 = row.get("attempts");
let consume = || async {
sqlx::query("DELETE FROM core.tb_otp_code WHERE email = $1")
.bind(email)
.execute(&self.db)
.await
.map_err(db_err)
.map(|_| ())
};
if now >= expires_at as u64 {
consume().await?;
return Err(AuthError::InvalidToken {
reason: "OTP has expired".into(),
});
}
if attempts > MAX_VERIFY_ATTEMPTS as i32 {
consume().await?;
return Err(AuthError::RateLimited {
retry_after_secs: OTP_RATE_WINDOW_SECS,
});
}
let presented = code_hash(code);
if !bool::from(presented.ct_eq(&stored_hash)) {
return Err(AuthError::InvalidToken {
reason: "invalid OTP code".into(),
});
}
consume().await?;
Ok(())
}
}