use std::sync::Arc;
use async_trait::async_trait;
use axum::{
Json,
extract::State,
http::StatusCode,
response::{IntoResponse, Response},
};
use dashmap::DashMap;
use rand::Rng;
use serde::{Deserialize, Serialize};
use crate::{
audit::logger::{AuditEventType, SecretType, get_audit_logger},
error::{AuthError, Result},
session::{SessionStore, unix_now},
};
const OTP_TTL_SECS: u64 = 600;
const MAX_VERIFY_ATTEMPTS: u32 = 3;
const OTP_RATE_WINDOW_SECS: u64 = 900;
const OTP_RATE_MAX: u32 = 3;
#[derive(Debug, Clone)]
struct OtpRecord {
code: String,
expires: u64,
attempts: u32,
}
impl OtpRecord {
fn is_expired(&self) -> bool {
Self::is_expired_at(self.expires, unix_now().ok())
}
fn is_expired_at(expires: u64, now: Option<u64>) -> bool {
now.is_none_or(|now| now >= expires)
}
}
#[derive(Debug, Clone)]
struct RateRecord {
count: u32,
window_start: u64,
}
#[async_trait]
pub trait OtpStore: Send + Sync {
async fn create_otp(&self, email: &str) -> Result<String>;
async fn verify_otp(&self, email: &str, code: &str) -> Result<()>;
}
pub struct InMemoryOtpStore {
codes: DashMap<String, OtpRecord>,
rate_limits: DashMap<String, RateRecord>,
}
impl InMemoryOtpStore {
#[must_use]
pub fn new() -> Self {
Self {
codes: DashMap::new(),
rate_limits: DashMap::new(),
}
}
#[must_use]
pub fn pending_count(&self) -> usize {
self.codes.len()
}
}
impl Default for InMemoryOtpStore {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl OtpStore for InMemoryOtpStore {
async fn create_otp(&self, email: &str) -> Result<String> {
let now = unix_now()?;
{
let mut entry = self.rate_limits.entry(email.to_string()).or_insert(RateRecord {
count: 0,
window_start: now,
});
if now >= entry.window_start + OTP_RATE_WINDOW_SECS {
entry.count = 0;
entry.window_start = now;
}
if entry.count >= OTP_RATE_MAX {
return Err(AuthError::RateLimited {
retry_after_secs: (entry.window_start + OTP_RATE_WINDOW_SECS)
.saturating_sub(now),
});
}
entry.count += 1;
}
let code = format!("{:06}", rand::rng().random_range(0u32..1_000_000));
let expires = now + OTP_TTL_SECS;
self.codes.insert(
email.to_string(),
OtpRecord {
code: code.clone(),
expires,
attempts: 0,
},
);
Ok(code)
}
async fn verify_otp(&self, email: &str, code: &str) -> Result<()> {
let mut entry = self.codes.get_mut(email).ok_or_else(|| AuthError::InvalidToken {
reason: "no pending OTP for email".into(),
})?;
if entry.is_expired() {
drop(entry);
self.codes.remove(email);
return Err(AuthError::InvalidToken {
reason: "OTP has expired".into(),
});
}
entry.attempts += 1;
if entry.attempts > MAX_VERIFY_ATTEMPTS {
drop(entry);
self.codes.remove(email);
return Err(AuthError::RateLimited {
retry_after_secs: OTP_RATE_WINDOW_SECS,
});
}
if entry.code != code {
return Err(AuthError::InvalidToken {
reason: "invalid OTP code".into(),
});
}
drop(entry);
self.codes.remove(email);
Ok(())
}
}
#[async_trait]
pub trait EmailDelivery: Send + Sync {
async fn send_otp(&self, email: &str, code: &str) -> Result<String>;
}
pub struct NoopEmailDelivery;
#[async_trait]
impl EmailDelivery for NoopEmailDelivery {
async fn send_otp(&self, email: &str, code: &str) -> Result<String> {
tracing::info!(email, code, "NoopEmailDelivery: OTP code (NOT sent via real email)");
Ok(format!("noop-{email}-{code}"))
}
}
#[derive(Clone)]
pub struct OtpRouteState {
pub otp_store: Arc<dyn OtpStore>,
pub email_delivery: Arc<dyn EmailDelivery>,
pub session_store: Arc<dyn SessionStore>,
}
#[derive(Debug, Deserialize)]
pub struct OtpRequest {
pub email: String,
}
#[derive(Debug, Serialize)]
pub struct OtpResponse {
pub message_id: String,
}
#[derive(Debug, Deserialize)]
pub struct VerifyRequest {
pub email: String,
pub code: String,
}
pub async fn otp_send(
State(state): State<Arc<OtpRouteState>>,
Json(req): Json<OtpRequest>,
) -> Response {
let email = req.email.trim().to_lowercase();
if email.is_empty() {
return (
StatusCode::UNPROCESSABLE_ENTITY,
Json(serde_json::json!({
"error": "invalid_email",
"message": "email must not be blank"
})),
)
.into_response();
}
let code = match state.otp_store.create_otp(&email).await {
Ok(c) => c,
Err(AuthError::RateLimited { retry_after_secs }) => {
return (
StatusCode::TOO_MANY_REQUESTS,
Json(serde_json::json!({
"error": "rate_limited",
"retry_after_secs": retry_after_secs
})),
)
.into_response();
},
Err(e) => {
tracing::error!(error = %e, "OTP store error");
return (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response();
},
};
let message_id = match state.email_delivery.send_otp(&email, &code).await {
Ok(id) => id,
Err(e) => {
tracing::error!(error = %e, "Email delivery failed");
return (StatusCode::INTERNAL_SERVER_ERROR, "delivery failed").into_response();
},
};
let logger = get_audit_logger();
logger.log_success(
AuditEventType::OauthStart,
SecretType::CsrfToken,
None,
&format!("otp_send:{email}"),
);
(StatusCode::OK, Json(OtpResponse { message_id })).into_response()
}
pub async fn otp_verify(
State(state): State<Arc<OtpRouteState>>,
Json(req): Json<VerifyRequest>,
) -> Response {
let email = req.email.trim().to_lowercase();
let logger = get_audit_logger();
match state.otp_store.verify_otp(&email, &req.code).await {
Ok(()) => {},
Err(AuthError::RateLimited { retry_after_secs }) => {
return (
StatusCode::TOO_MANY_REQUESTS,
Json(serde_json::json!({
"error": "rate_limited",
"retry_after_secs": retry_after_secs
})),
)
.into_response();
},
Err(e) => {
logger.log_failure(
AuditEventType::AuthFailure,
SecretType::CsrfToken,
None,
"otp_verify",
&e.to_string(),
);
return (
StatusCode::UNPROCESSABLE_ENTITY,
Json(serde_json::json!({
"error": "invalid_otp",
"message": "invalid or expired OTP code"
})),
)
.into_response();
},
}
let user_id = format!("otp:{email}");
let expires_at = match unix_now() {
Ok(now) => now + 3_600, Err(_) => return (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response(),
};
let tokens = match state.session_store.create_session(&user_id, expires_at).await {
Ok(t) => t,
Err(e) => {
tracing::error!(error = %e, "Session creation failed after OTP verify");
return (StatusCode::INTERNAL_SERVER_ERROR, "session creation failed").into_response();
},
};
logger.log_success(
AuditEventType::AuthSuccess,
SecretType::SessionToken,
Some(user_id),
"otp_verify",
);
(StatusCode::OK, Json(tokens)).into_response()
}
#[cfg(test)]
mod tests;