use crate::utils::config::TraceResult;
use axum::{
Router,
extract::{Path, State},
response::{IntoResponse, Redirect, Response},
};
use futures_util::future::BoxFuture;
use serde::Serialize;
use std::{sync::Arc, time::Duration};
use async_trait::async_trait;
use crate::auth::session::logout;
use crate::auth::user::BuiltinUserEntity;
use crate::auth::user_trait::RuniqueUser;
use crate::context::template::Request;
use crate::forms::{
Forms,
field::RuniqueForm,
fields::{hidden::HiddenField, text::TextField},
};
use crate::utils::{
aliases::{AppResult, StrMap},
trad::{current_lang, t, tf},
};
use crate::{context_update, impl_form_access};
#[derive(Serialize, Debug, Clone)]
#[serde(transparent)]
pub struct ForgotPasswordForm {
pub form: Forms,
}
impl RuniqueForm for ForgotPasswordForm {
fn register_fields(form: &mut Forms) {
form.field(
&TextField::text("email")
.label(&t("reset.email_label"))
.required(),
);
}
impl_form_access!();
}
#[derive(Serialize, Debug, Clone)]
#[serde(transparent)]
pub struct PasswordResetForm {
pub form: Forms,
}
#[async_trait]
impl RuniqueForm for PasswordResetForm {
fn register_fields(form: &mut Forms) {
form.field(&HiddenField::new("token"));
form.field(&HiddenField::new("encrypted_email"));
form.field(
&TextField::text("email")
.label(&t("reset.email_label"))
.required(),
);
form.field(
&TextField::password("password")
.label(&t("reset.new_password_label"))
.required(),
);
form.field(
&TextField::password("confirm")
.label(&t("reset.confirm_label"))
.required(),
);
}
async fn clean(&mut self) -> Result<(), StrMap> {
let token = self.cleaned_string("token").unwrap_or_default();
let encrypted = self.cleaned_string("encrypted_email").unwrap_or_default();
let email = self.cleaned_string("email").unwrap_or_default();
let password = self.cleaned_string("password").unwrap_or_default();
let confirm = self.cleaned_string("confirm").unwrap_or_default();
let mut errors = StrMap::new();
match crate::utils::reset_token::decrypt_email(&token, &encrypted) {
Some(ref expected) if expected.to_lowercase() == email.trim().to_lowercase() => {}
Some(_) => {
errors.insert("email".to_string(), t("reset.email_mismatch").to_string());
}
None => {
errors.insert("token".to_string(), t("reset.invalid_link").to_string());
}
}
const SPECIAL: &str = "!@#$%^&*()_+-=[]{}|;':\",./<>?";
if password.len() < 10 {
errors.insert(
"password".to_string(),
tf("reset.password_min_length", &["10"]).clone(),
);
} else if !password.chars().any(|c| c.is_uppercase())
|| !password.chars().any(|c| c.is_lowercase())
|| !password.chars().any(|c| c.is_ascii_digit())
|| !password.chars().any(|c| SPECIAL.contains(c))
{
errors.insert("password".to_string(), t("reset.password_weak").to_string());
}
if password != confirm {
errors.insert(
"confirm".to_string(),
t("reset.password_mismatch").to_string(),
);
}
if errors.is_empty() {
Ok(())
} else {
Err(errors)
}
}
impl_form_access!();
}
pub type ExtraContextFn = Arc<dyn for<'a> Fn(&'a mut Request) -> BoxFuture<'a, ()> + Send + Sync>;
#[derive(Clone)]
pub struct PasswordResetConfig {
pub forgot_route: String,
pub reset_route: String,
pub forgot_template: String,
pub reset_template: String,
pub email_template: Option<String>,
pub success_redirect: String,
pub max_requests: u64,
pub retry_after: u64,
pub token_ttl: Duration,
pub extra_context: Option<ExtraContextFn>,
}
impl Default for PasswordResetConfig {
fn default() -> Self {
Self {
forgot_route: "/forgot-password".to_string(),
reset_route: "/reset-password".to_string(),
forgot_template: "auth/forgot_password.html".to_string(),
reset_template: "auth/reset_password.html".to_string(),
email_template: None,
success_redirect: "/".to_string(),
max_requests: 5,
retry_after: 300,
token_ttl: Duration::from_secs(3600),
extra_context: None,
}
}
}
impl PasswordResetConfig {
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn forgot_route(mut self, route: &str) -> Self {
self.forgot_route = route.to_string();
self
}
#[must_use]
pub fn reset_route(mut self, route: &str) -> Self {
self.reset_route = route.to_string();
self
}
#[must_use]
pub fn forgot_template(mut self, template: &str) -> Self {
self.forgot_template = template.to_string();
self
}
#[must_use]
pub fn reset_template(mut self, template: &str) -> Self {
self.reset_template = template.to_string();
self
}
#[must_use]
pub fn success_redirect(mut self, redirect: &str) -> Self {
self.success_redirect = redirect.to_string();
self
}
#[must_use]
pub fn email_template(mut self, template: &str) -> Self {
self.email_template = Some(template.to_string());
self
}
#[must_use]
pub fn token_ttl(mut self, ttl: Duration) -> Self {
self.token_ttl = ttl;
self
}
#[must_use]
pub fn extra_context(mut self, hook: ExtraContextFn) -> Self {
self.extra_context = Some(hook);
self
}
}
pub(crate) fn public_link_base(
configured: Option<&str>,
headers: &axum::http::HeaderMap,
debug: bool,
) -> Option<String> {
if let Some(base) = configured {
return Some(base.trim_end_matches('/').to_string());
}
if !debug {
return None;
}
let host = headers
.get(axum::http::header::HOST)
.and_then(|v| v.to_str().ok())?;
tracing::warn!(
"reset link built from the request's Host header: set .with_public_url() before production"
);
Some(format!("http://{host}"))
}
async fn apply_extra_context(request: &mut Request, hook: &Option<ExtraContextFn>) {
if let Some(hook) = hook {
hook(request).await;
}
}
pub async fn handle_forgot_password(
request: &mut Request,
form: ForgotPasswordForm,
config: &PasswordResetConfig,
) -> AppResult<Response> {
let template = config.forgot_template.as_str();
let forgot_route = config.forgot_route.as_str();
let reset_path = config.reset_route.as_str();
let email_template = config.email_template.as_deref();
let token_ttl = config.token_ttl;
request.context.insert("lang", ¤t_lang().code());
let form = match crate::forms::ValidationForm::try_new(form, request).await {
Ok(validated) => validated.into_form(),
Err(form) => {
apply_extra_context(request, &config.extra_context).await;
context_update!(request => {
"title" => t("reset.forgot_title").as_ref(),
"forgot_form" => &form,
});
return request.render(template);
}
};
let email = form.cleaned_string("email").unwrap_or_default();
let email = email.trim().to_lowercase();
let db = request.engine.db.clone();
let user = BuiltinUserEntity::find_by_email(&db, &email).await;
let blocked = user
.as_ref()
.is_some_and(|u| u.activated_at.is_some() && !u.is_active);
if blocked {
if let Some(user) = &user
&& crate::utils::mailer_configured()
{
let mail = crate::utils::Email::new()
.to(email.clone())
.subject(t("reset.blocked_subject").to_string())
.html(tf("reset.blocked_body", &[user.username()]).to_string());
let log_level = crate::utils::runique_log::get_log()
.auth
.as_ref()
.and_then(|a| a.reset);
tokio::spawn(async move {
mail.send().await.trace_or(
log_level,
tracing::Level::WARN,
"blocked account email send",
);
});
}
} else if let Some(user) = user
&& let Some(token) = crate::utils::reset_token::generate(&db, user.user_id(), token_ttl)
.await
.trace_or(
crate::utils::runique_log::get_log()
.auth
.as_ref()
.and_then(|a| a.reset),
tracing::Level::ERROR,
"reset token generation failed",
)
{
let encrypted_email = crate::utils::reset_token::encrypt_email(&token, &email);
if let Some(level) = crate::utils::runique_log::get_log()
.auth
.as_ref()
.and_then(|a| a.reset)
{
crate::runique_log!(level, %email, "reset token generated");
}
let Some(host) = request.public_url() else {
request
.notices
.success(t("reset.check_inbox").to_string())
.await;
return Ok(Redirect::to(forgot_route).into_response());
};
let reset_url = format!(
"{}/{}/{}/{}",
host,
reset_path.trim_matches('/'),
token,
encrypted_email
);
if crate::utils::mailer_configured() {
let username = user.username().to_string();
let subject = t("reset.email_subject").to_string();
let mail = crate::utils::Email::new()
.to(email.clone())
.subject(&subject);
if let Some(tpl) = email_template {
use tera::Context as TeraCtx;
let mut ctx = TeraCtx::new();
ctx.insert("username", &username);
ctx.insert("reset_url", &reset_url);
if let Ok(msg) = mail.template(&request.engine.tera, tpl, ctx) {
let log_level = crate::utils::runique_log::get_log()
.auth
.as_ref()
.and_then(|a| a.reset);
tokio::spawn(async move {
if msg
.send()
.await
.trace_or(log_level, tracing::Level::WARN, "reset email send")
.is_some()
&& let Some(level) = log_level
{
crate::runique_log!(level, "reset email sent");
}
});
} else {
tracing::warn!(
template = %tpl,
"reset email template render failed — email not sent"
);
}
} else {
let body = tf("reset.email_body", &[&username, &reset_url, &reset_url]).clone();
let log_level = crate::utils::runique_log::get_log()
.auth
.as_ref()
.and_then(|a| a.reset);
tokio::spawn(async move {
if mail
.html(body)
.send()
.await
.trace_or(log_level, tracing::Level::WARN, "reset email send")
.is_some()
&& let Some(level) = log_level
{
crate::runique_log!(level, "reset email sent");
}
});
}
}
}
request
.notices
.success(t("reset.check_inbox").to_string())
.await;
Ok(Redirect::to(forgot_route).into_response())
}
pub async fn handle_password_reset(
request: &mut Request,
form: PasswordResetForm,
token: String,
encrypted_email: String,
config: &PasswordResetConfig,
) -> AppResult<Response> {
let template = config.reset_template.as_str();
request.context.insert("lang", ¤t_lang().code());
logout(&request.session, None).await.trace(
crate::utils::runique_log::get_log()
.session
.as_ref()
.and_then(|s| s.store),
"logout before password reset",
);
let Some(email) = crate::utils::reset_token::decrypt_email(&token, &encrypted_email) else {
request
.notices
.error(t("reset.invalid_or_expired").to_string())
.await;
return Ok(Redirect::to("/").into_response());
};
let db = request.engine.db.clone();
if !crate::utils::reset_token::peek(&db, &token).await {
if let Some(level) = crate::utils::runique_log::get_log()
.auth
.as_ref()
.and_then(|a| a.reset)
{
crate::runique_log!(level, %email, "reset token invalid or expired");
}
request
.notices
.error(t("reset.invalid_or_expired").to_string())
.await;
return Ok(Redirect::to("/").into_response());
}
let mut form = match crate::forms::ValidationForm::try_new(form, request).await {
Ok(validated) => validated.into_form(),
Err(mut form) => {
if request.method.is_safe() {
form.get_form_mut().add_value("token", &token);
form.get_form_mut()
.add_value("encrypted_email", &encrypted_email);
}
apply_extra_context(request, &config.extra_context).await;
context_update!(request => {
"title" => t("reset.reset_title").as_ref(),
"reset_form" => &form,
"token" => &token,
"encrypted_email" => &encrypted_email,
});
return request.render(template);
}
};
let Some(user_id) = crate::utils::reset_token::consume(&db, &token).await else {
if let Some(level) = crate::utils::runique_log::get_log()
.auth
.as_ref()
.and_then(|a| a.reset)
{
crate::runique_log!(level, %email, "reset token consume failed");
}
request
.notices
.error(t("reset.invalid_or_expired").to_string())
.await;
return Ok(Redirect::to("/").into_response());
};
let Some(user) = BuiltinUserEntity::find_by_id(&db, user_id).await else {
request
.notices
.error(t("reset.invalid_or_expired").to_string())
.await;
return Ok(Redirect::to("/").into_response());
};
if user.email().to_lowercase() != email.to_lowercase() {
request
.notices
.error(t("reset.invalid_or_expired").to_string())
.await;
return Ok(Redirect::to("/").into_response());
}
let email_clean = form.cleaned_string("email").unwrap_or_default();
let new_hash = form.cleaned_string("password").unwrap_or_default();
match BuiltinUserEntity::set_password_and_activate(&db, user_id, &new_hash).await {
Ok(()) => {
request.engine.close_user_sessions(user_id).await;
if let Some(level) = crate::utils::runique_log::get_log()
.auth
.as_ref()
.and_then(|a| a.reset)
{
crate::runique_log!(level, email = %email_clean, "password reset ok");
}
form.clear();
apply_extra_context(request, &config.extra_context).await;
context_update!(request => {
"title" => t("reset.success_title").as_ref(),
"reset_form" => &form,
"reset_done" => &true,
"token" => &token,
"encrypted_email" => &encrypted_email,
});
return request.render(template);
}
Err(e) => {
if let Some(level) = crate::utils::runique_log::get_log()
.auth
.as_ref()
.and_then(|a| a.reset)
{
crate::runique_log!(level, email = %email_clean, error = %e, "password reset db error");
}
form.get_form_mut().database_error(&e);
}
}
apply_extra_context(request, &config.extra_context).await;
context_update!(request => {
"title" => t("reset.reset_title").as_ref(),
"reset_form" => &form,
"token" => &token,
"encrypted_email" => &encrypted_email,
});
request.render(template)
}
#[derive(Clone)]
struct ForgotState {
config: Arc<PasswordResetConfig>,
}
#[derive(Clone)]
struct ResetState {
config: Arc<PasswordResetConfig>,
}
async fn forgot_view(
State(state): State<ForgotState>,
mut request: Request,
) -> AppResult<Response> {
let form: ForgotPasswordForm = request.form();
handle_forgot_password(&mut request, form, &state.config).await
}
async fn reset_view(
State(state): State<ResetState>,
Path((token, encrypted_email)): Path<(String, String)>,
mut request: Request,
) -> AppResult<Response> {
let form: PasswordResetForm = request.form();
handle_password_reset(&mut request, form, token, encrypted_email, &state.config).await
}
pub fn build_router(config: Arc<PasswordResetConfig>) -> Router {
use crate::middleware::security::rate_limit::{RateLimiter, rate_limit_middleware};
use axum::middleware;
use axum::routing::any;
let limiter = Arc::new(
RateLimiter::new()
.max_requests(u32::try_from(config.max_requests).unwrap_or(u32::MAX))
.retry_after(config.retry_after),
);
let forgot_state = ForgotState {
config: config.clone(),
};
let reset_state = ResetState { config };
let forgot_route = Router::new()
.route(&forgot_state.config.forgot_route, any(forgot_view))
.with_state(forgot_state)
.route_layer(middleware::from_fn_with_state(
limiter.clone(),
rate_limit_middleware,
));
let reset_path = format!(
"{}/{{token}}/{{encrypted_email}}",
reset_state.config.reset_route.trim_end_matches('/')
);
let reset_route = Router::new()
.route(&reset_path, any(reset_view))
.with_state(reset_state)
.route_layer(middleware::from_fn_with_state(
limiter,
rate_limit_middleware,
));
forgot_route.merge(reset_route)
}
pub struct PasswordResetStaging {
pub config: PasswordResetConfig,
}
#[cfg(test)]
mod public_link_base_tests {
use super::public_link_base;
use axum::http::{HeaderMap, HeaderValue, header::HOST};
fn with_host(host: &str) -> HeaderMap {
let mut headers = HeaderMap::new();
headers.insert(HOST, HeaderValue::from_str(host).unwrap());
headers
}
#[test]
fn the_configured_base_wins_over_the_host() {
let headers = with_host("evil.com");
for debug in [true, false] {
assert_eq!(
public_link_base(Some("https://mysite.com/"), &headers, debug).as_deref(),
Some("https://mysite.com")
);
}
}
#[test]
fn production_never_uses_the_host() {
assert_eq!(public_link_base(None, &with_host("evil.com"), false), None);
}
#[test]
fn debug_falls_back_on_the_host() {
assert_eq!(
public_link_base(None, &with_host("localhost:3000"), true).as_deref(),
Some("http://localhost:3000")
);
assert_eq!(public_link_base(None, &HeaderMap::new(), true), None);
}
}