use std::sync::Arc;
use async_trait::async_trait;
use axum::{
Json,
extract::State,
http::StatusCode,
response::{IntoResponse, Response},
};
use serde::{Deserialize, Serialize};
use tokio::sync::RwLock;
use crate::{account_linking::AccountStore, otp::OtpStore, session::SessionStore};
pub(crate) const OTP_TTL_SECS: u64 = 600;
const OTP_LENGTH: usize = 6;
const MAX_PHONE_LEN: usize = 16;
const MIN_PHONE_LEN: usize = 8;
#[async_trait]
pub trait SmsSender: Send + Sync {
async fn send_sms_otp(&self, to: &str, code: &str) -> crate::error::Result<()>;
}
#[derive(Debug)]
pub struct InMemorySmsSender {
pub messages: RwLock<Vec<(String, String)>>,
}
impl InMemorySmsSender {
#[must_use]
pub fn new() -> Self {
Self {
messages: RwLock::new(Vec::new()),
}
}
pub async fn sms_count(&self) -> usize {
self.messages.read().await.len()
}
pub async fn last_otp_for(&self, phone: &str) -> Option<String> {
let messages = self.messages.read().await;
messages.iter().rev().find(|(to, _)| to == phone).map(|(_, code)| code.clone())
}
}
impl Default for InMemorySmsSender {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl SmsSender for InMemorySmsSender {
async fn send_sms_otp(&self, to: &str, code: &str) -> crate::error::Result<()> {
let mut messages = self.messages.write().await;
messages.push((to.to_string(), code.to_string()));
Ok(())
}
}
#[must_use]
pub fn normalize_e164(phone: &str) -> Option<String> {
let cleaned: String = phone.chars().filter(|c| c.is_ascii_digit() || *c == '+').collect();
if cleaned.is_empty() {
return None;
}
let normalized = if cleaned.starts_with('+') {
cleaned
} else {
format!("+{cleaned}")
};
let digits = &normalized[1..];
if digits.len() < 7 || digits.len() > 15 {
return None;
}
if !digits.chars().all(|c| c.is_ascii_digit()) {
return None;
}
Some(normalized)
}
fn phone_otp_key(e164: &str) -> String {
format!("sms:{e164}")
}
pub(crate) fn unix_now() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
}
#[derive(Debug, Deserialize)]
pub struct SmsOtpRequest {
pub phone: String,
}
#[derive(Debug, Serialize)]
pub struct SmsOtpResponse {
pub status: String,
pub expires_in: u64,
}
#[derive(Debug, Deserialize)]
pub struct SmsVerifyRequest {
pub phone: String,
pub code: String,
}
#[derive(Debug, Serialize)]
pub struct SmsVerifyResponse {
pub access_token: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub refresh_token: Option<String>,
pub token_type: String,
pub expires_in: u64,
}
#[derive(Clone)]
pub struct SmsOtpAuthState {
pub otp_store: Arc<dyn OtpStore>,
pub sms_sender: Arc<dyn SmsSender>,
pub session_store: Arc<dyn SessionStore>,
pub user_store: Option<Arc<dyn AccountStore>>,
}
fn json_error(status: StatusCode, message: &str) -> Response {
(status, Json(serde_json::json!({ "error": message }))).into_response()
}
pub async fn send_sms_otp(
State(state): State<Arc<SmsOtpAuthState>>,
Json(req): Json<SmsOtpRequest>,
) -> Response {
if req.phone.is_empty() {
return json_error(StatusCode::BAD_REQUEST, "phone is required");
}
if req.phone.len() > MAX_PHONE_LEN * 2 {
return json_error(StatusCode::BAD_REQUEST, "phone number too long");
}
let Some(e164) = normalize_e164(&req.phone) else {
return json_error(StatusCode::BAD_REQUEST, "invalid phone number format");
};
if e164.len() < MIN_PHONE_LEN {
return json_error(StatusCode::BAD_REQUEST, "phone number too short");
}
let key = phone_otp_key(&e164);
let code = match state.otp_store.create_otp(&key).await {
Ok(c) => c,
Err(crate::error::AuthError::RateLimited { .. }) => {
return Json(SmsOtpResponse {
status: "otp_sent".to_string(),
expires_in: OTP_TTL_SECS,
})
.into_response();
},
Err(e) => {
tracing::error!(error = %e, "OTP store failed for SMS");
return Json(SmsOtpResponse {
status: "otp_sent".to_string(),
expires_in: OTP_TTL_SECS,
})
.into_response();
},
};
if let Err(e) = state.sms_sender.send_sms_otp(&e164, &code).await {
tracing::error!(error = %e, "SMS delivery failed");
}
Json(SmsOtpResponse {
status: "otp_sent".to_string(),
expires_in: OTP_TTL_SECS,
})
.into_response()
}
pub async fn verify_sms_otp(
State(state): State<Arc<SmsOtpAuthState>>,
Json(req): Json<SmsVerifyRequest>,
) -> Response {
if req.phone.is_empty() {
return json_error(StatusCode::BAD_REQUEST, "phone is required");
}
let Some(e164) = normalize_e164(&req.phone) else {
return json_error(StatusCode::BAD_REQUEST, "invalid phone number format");
};
if req.code.len() != OTP_LENGTH {
return json_error(StatusCode::BAD_REQUEST, "invalid OTP code format");
}
let key = phone_otp_key(&e164);
match state.otp_store.verify_otp(&key, &req.code).await {
Ok(()) => {},
Err(crate::error::AuthError::RateLimited { retry_after_secs }) => {
return (
StatusCode::TOO_MANY_REQUESTS,
Json(serde_json::json!({
"error": "too many verification attempts",
"retry_after_secs": retry_after_secs,
})),
)
.into_response();
},
Err(crate::error::AuthError::InvalidToken { .. }) => {
return json_error(StatusCode::BAD_REQUEST, "invalid or expired OTP code");
},
Err(e) => {
tracing::error!(error = %e, "SMS OTP verification error");
return json_error(StatusCode::INTERNAL_SERVER_ERROR, "verification failed");
},
}
let user_id = if let Some(account_store) = &state.user_store {
let email = format!("{e164}@phone.local");
match account_store.link_or_create_user(&email, "phone", &e164).await {
Ok(result) => result.user_id,
Err(e) => {
tracing::error!(error = %e, "account store lookup failed");
return json_error(StatusCode::INTERNAL_SERVER_ERROR, "user resolution failed");
},
}
} else {
e164
};
let session_expiry = unix_now() + (7 * 24 * 60 * 60);
match state.session_store.create_session(&user_id, session_expiry).await {
Ok(tokens) => Json(SmsVerifyResponse {
access_token: tokens.access_token,
refresh_token: Some(tokens.refresh_token),
token_type: "Bearer".to_string(),
expires_in: tokens.expires_in,
})
.into_response(),
Err(e) => {
tracing::error!(error = %e, "session creation failed");
json_error(StatusCode::INTERNAL_SERVER_ERROR, "session could not be created")
},
}
}