Files
ystp/src/api/user.rs
2026-07-26 05:47:19 +08:00

1826 lines
60 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
use crate::api::context;
use crate::api::envelope::Envelope;
use crate::auth;
use crate::error::{AppError, ErrorCode};
use crate::services::billing;
use crate::services::{credentials, mail, settings};
use crate::state::AppState;
use axum::extract::{ConnectInfo, Path, Query, State};
use axum::http::HeaderMap;
use axum::routing::{delete, get, post, put};
use axum::{Json, Router};
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
use chrono::{DateTime, Duration, Utc};
use rand::RngCore;
use serde::{Deserialize, Serialize};
use sqlx::FromRow;
use std::collections::HashMap;
use std::net::SocketAddr;
use uuid::Uuid;
pub fn router() -> Router<AppState> {
Router::new()
.route("/user/profile", get(get_profile))
.route("/user/profile", put(update_profile))
.route("/user/password", put(update_password))
.route("/user/history", get(list_history))
.route("/user/api-keys", get(list_api_keys))
.route("/user/api-keys", post(create_api_key))
.route("/user/api-keys/{key_id}/rotate", post(rotate_api_key))
.route("/user/api-keys/{key_id}", delete(disable_api_key))
}
#[derive(Debug, FromRow, Serialize)]
struct ApiKeyView {
id: Uuid,
name: String,
key_prefix: String,
permissions: serde_json::Value,
rate_limit: i32,
is_active: bool,
last_used_at: Option<chrono::DateTime<chrono::Utc>>,
last_used_ip: Option<String>,
created_at: chrono::DateTime<chrono::Utc>,
}
#[derive(Debug, Serialize)]
struct ApiKeyListResponse {
api_keys: Vec<ApiKeyView>,
}
#[derive(Debug, Serialize)]
struct UserView {
id: Uuid,
email: String,
username: String,
role: String,
email_verified: bool,
pending_email: Option<String>,
}
#[derive(Debug, Serialize)]
struct MessageResponse {
message: String,
}
async fn get_profile(
State(state): State<AppState>,
jar: axum_extra::extract::cookie::CookieJar,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
headers: HeaderMap,
) -> Result<Json<Envelope<UserView>>, AppError> {
let ip = context::client_ip(&headers, addr.ip());
let (_jar, principal) = context::authenticate(&state, jar, &headers, ip).await?;
let user_id = match principal {
context::Principal::User { user_id, .. } => user_id,
_ => return Err(AppError::new(ErrorCode::Unauthorized, "未登录")),
};
#[derive(Debug, FromRow)]
struct UserRow {
id: Uuid,
email: String,
username: String,
role: String,
email_verified_at: Option<DateTime<Utc>>,
pending_email: Option<String>,
}
let user = sqlx::query_as::<_, UserRow>(
r#"
SELECT u.id, u.email, u.username, u.role::text AS role, u.email_verified_at,
(
SELECT r.new_email
FROM email_change_requests r
WHERE r.user_id = u.id
AND r.confirmed_at IS NULL
AND r.canceled_at IS NULL
AND r.expires_at > NOW()
ORDER BY r.created_at DESC
LIMIT 1
) AS pending_email
FROM users u
WHERE u.id = $1
"#,
)
.bind(user_id)
.fetch_one(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "查询用户失败").with_source(err))?;
let verification_required = settings::email_verification_required(&state).await?;
Ok(Json(Envelope {
success: true,
data: UserView {
id: user.id,
email: user.email,
username: user.username,
role: user.role,
email_verified: user.email_verified_at.is_some() || !verification_required,
pending_email: user.pending_email,
},
}))
}
#[derive(Debug, Deserialize)]
struct UpdateProfileRequest {
email: Option<String>,
username: Option<String>,
current_password: Option<String>,
}
#[derive(Debug, Serialize)]
struct UpdateProfileResponse {
user: UserView,
message: String,
#[serde(skip_serializing_if = "Option::is_none")]
token: Option<String>,
}
async fn update_profile(
State(state): State<AppState>,
jar: axum_extra::extract::cookie::CookieJar,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
headers: HeaderMap,
Json(req): Json<UpdateProfileRequest>,
) -> Result<Json<Envelope<UpdateProfileResponse>>, AppError> {
let ip = context::client_ip(&headers, addr.ip());
let (_jar, principal) = context::authenticate(&state, jar, &headers, ip).await?;
let user_id = match principal {
context::Principal::User { user_id, .. } => user_id,
_ => return Err(AppError::new(ErrorCode::Unauthorized, "未登录")),
};
if req.email.is_none() && req.username.is_none() {
return Err(AppError::new(ErrorCode::InvalidRequest, "未提供可更新字段"));
}
let verification_required = settings::email_verification_required(&state).await?;
let mut tx = state
.db
.begin()
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "开启事务失败").with_source(err))?;
#[derive(Debug, FromRow)]
struct UserRow {
id: Uuid,
email: String,
username: String,
password_hash: String,
role: String,
email_verified_at: Option<DateTime<Utc>>,
token_version: i32,
pending_email: Option<String>,
}
let user = sqlx::query_as::<_, UserRow>(
r#"
SELECT u.id, u.email, u.username, u.password_hash,
u.role::text AS role, u.email_verified_at, u.token_version,
(
SELECT r.new_email
FROM email_change_requests r
WHERE r.user_id = u.id
AND r.confirmed_at IS NULL
AND r.canceled_at IS NULL
AND r.expires_at > NOW()
ORDER BY r.created_at DESC
LIMIT 1
) AS pending_email
FROM users u
WHERE u.id = $1
FOR UPDATE OF u
"#,
)
.bind(user_id)
.fetch_one(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "查询用户失败").with_source(err))?;
let next_email = match req.email.as_ref() {
Some(email) => {
let email = email.trim().to_lowercase();
credentials::validate_email(&email)?;
email
}
None => user.email.clone(),
};
let next_username = match req.username.as_ref() {
Some(username) => {
let username = username.trim().to_string();
credentials::validate_username(&username)?;
username
}
None => user.username.clone(),
};
let email_changed = next_email != user.email;
let username_changed = next_username != user.username;
if !email_changed && !username_changed {
tx.rollback().await.ok();
return Ok(Json(Envelope {
success: true,
data: UpdateProfileResponse {
user: UserView {
id: user.id,
email: user.email,
username: user.username,
role: user.role,
email_verified: user.email_verified_at.is_some() || !verification_required,
pending_email: user.pending_email,
},
message: "暂无更新".to_string(),
token: None,
},
}));
}
if email_changed {
let current_password = req
.current_password
.as_deref()
.filter(|value| !value.is_empty())
.ok_or_else(|| AppError::new(ErrorCode::Unauthorized, "修改邮箱需要当前密码"))?;
if !credentials::verify_password(current_password, &user.password_hash).await? {
return Err(AppError::new(ErrorCode::Unauthorized, "当前密码不正确"));
}
let email_in_use: bool =
sqlx::query_scalar("SELECT EXISTS(SELECT 1 FROM users WHERE email = $1 AND id <> $2)")
.bind(&next_email)
.bind(user_id)
.fetch_one(&mut *tx)
.await
.map_err(|err| {
AppError::new(ErrorCode::Internal, "检查邮箱失败").with_source(err)
})?;
if email_in_use {
return Err(AppError::new(ErrorCode::InvalidRequest, "邮箱已存在"));
}
}
let mut verification_link: Option<String> = None;
let mut pending_email = user.pending_email.clone();
let updated: UserRow;
if email_changed && verification_required {
let token = credentials::generate_token();
let token_hash = credentials::sha256_hex(&token);
let expires_at = Utc::now() + Duration::hours(24);
sqlx::query(
r#"
UPDATE email_change_requests
SET canceled_at = NOW()
WHERE confirmed_at IS NULL
AND canceled_at IS NULL
AND (user_id = $1 OR (new_email = $2 AND expires_at <= NOW()))
"#,
)
.bind(user_id)
.bind(&next_email)
.execute(&mut *tx)
.await
.map_err(|err| {
AppError::new(ErrorCode::Internal, "撤销旧邮箱变更请求失败").with_source(err)
})?;
sqlx::query(
r#"
INSERT INTO email_change_requests (user_id, new_email, token_hash, expires_at)
VALUES ($1, $2, $3, $4)
"#,
)
.bind(user_id)
.bind(&next_email)
.bind(token_hash)
.bind(expires_at)
.execute(&mut *tx)
.await
.map_err(map_unique_violation)?;
updated = sqlx::query_as::<_, UserRow>(
r#"
UPDATE users
SET username = $2, updated_at = NOW()
WHERE id = $1
RETURNING id, email, username, password_hash, role::text AS role,
email_verified_at, token_version, $3::text AS pending_email
"#,
)
.bind(user_id)
.bind(&next_username)
.bind(&next_email)
.fetch_one(&mut *tx)
.await
.map_err(map_unique_violation)?;
pending_email = Some(next_email.clone());
verification_link = Some(format!(
"{}/verify-email?token={}",
state.config.public_base_url, token
));
} else if email_changed {
updated = sqlx::query_as::<_, UserRow>(
r#"
UPDATE users
SET email = $2,
username = $3,
email_verified_at = NOW(),
token_version = token_version + 1,
updated_at = NOW()
WHERE id = $1
RETURNING id, email, username, password_hash, role::text AS role,
email_verified_at, token_version, NULL::text AS pending_email
"#,
)
.bind(user_id)
.bind(&next_email)
.bind(&next_username)
.fetch_one(&mut *tx)
.await
.map_err(map_unique_violation)?;
sqlx::query(
"UPDATE email_change_requests SET canceled_at = NOW() WHERE user_id = $1 AND confirmed_at IS NULL AND canceled_at IS NULL",
)
.bind(user_id)
.execute(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "撤销邮箱变更请求失败").with_source(err))?;
sqlx::query(
"UPDATE password_resets SET used_at = NOW() WHERE user_id = $1 AND used_at IS NULL",
)
.bind(user_id)
.execute(&mut *tx)
.await
.map_err(|err| {
AppError::new(ErrorCode::Internal, "撤销密码重置请求失败").with_source(err)
})?;
pending_email = None;
} else {
updated = sqlx::query_as::<_, UserRow>(
r#"
UPDATE users
SET username = $2, updated_at = NOW()
WHERE id = $1
RETURNING id, email, username, password_hash, role::text AS role,
email_verified_at, token_version, $3::text AS pending_email
"#,
)
.bind(user_id)
.bind(&next_username)
.bind(pending_email.as_deref())
.fetch_one(&mut *tx)
.await
.map_err(map_unique_violation)?;
}
tx.commit()
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "提交事务失败").with_source(err))?;
let verification_mail_sent = if let Some(link) = verification_link.as_deref() {
match mail::send_verification_email(&state, &next_email, &updated.username, link).await {
Ok(()) => true,
Err(err) => {
tracing::error!(user_id = %user_id, error = ?err, "email change verification delivery failed");
false
}
}
} else {
true
};
if email_changed && !verification_required {
if let Err(err) = mail::send_email_change_notice(&state, &user.email).await {
tracing::warn!(user_id = %user_id, error = ?err, "email change notice delivery failed");
}
}
let message = if email_changed && verification_required {
if verification_mail_sent {
"资料已更新,请验证新邮箱;确认前仍使用原邮箱登录和找回密码".to_string()
} else {
"新邮箱已进入待确认状态,但验证邮件发送失败,请重新提交邮箱变更".to_string()
}
} else if email_changed {
"资料已更新,其他登录状态已失效".to_string()
} else {
"资料已更新".to_string()
};
let token = if email_changed && !verification_required {
Some(
auth::issue_jwt(
&state.config.jwt_secret,
state.config.jwt_expiry_hours,
updated.id,
&updated.role,
updated.token_version,
)?
.0,
)
} else {
None
};
Ok(Json(Envelope {
success: true,
data: UpdateProfileResponse {
user: UserView {
id: updated.id,
email: updated.email,
username: updated.username,
role: updated.role,
email_verified: updated.email_verified_at.is_some() || !verification_required,
pending_email,
},
message,
token,
},
}))
}
#[derive(Debug, Deserialize)]
struct UpdatePasswordRequest {
current_password: String,
new_password: String,
}
async fn update_password(
State(state): State<AppState>,
jar: axum_extra::extract::cookie::CookieJar,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
headers: HeaderMap,
Json(req): Json<UpdatePasswordRequest>,
) -> Result<Json<Envelope<MessageResponse>>, AppError> {
let ip = context::client_ip(&headers, addr.ip());
let (_jar, principal) = context::authenticate(&state, jar, &headers, ip).await?;
let user_id = match principal {
context::Principal::User { user_id, .. } => user_id,
_ => return Err(AppError::new(ErrorCode::Unauthorized, "未登录")),
};
credentials::validate_password(&req.new_password)?;
#[derive(Debug, FromRow)]
struct PasswordRow {
password_hash: String,
}
let new_hash = credentials::hash_password(&req.new_password).await?;
let now = Utc::now();
let mut tx = state
.db
.begin()
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "开启事务失败").with_source(err))?;
let row = sqlx::query_as::<_, PasswordRow>(
"SELECT password_hash FROM users WHERE id = $1 FOR UPDATE",
)
.bind(user_id)
.fetch_one(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "查询用户失败").with_source(err))?;
if !credentials::verify_password(&req.current_password, &row.password_hash).await? {
return Err(AppError::new(ErrorCode::Unauthorized, "密码错误"));
}
sqlx::query(
"UPDATE users SET password_hash = $2, token_version = token_version + 1, updated_at = NOW() WHERE id = $1",
)
.bind(user_id)
.bind(new_hash)
.execute(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "更新密码失败").with_source(err))?;
credentials::invalidate_account_recovery(&mut tx, user_id, now).await?;
tx.commit()
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "提交密码更新失败").with_source(err))?;
Ok(Json(Envelope {
success: true,
data: MessageResponse {
message: "密码已更新,请重新登录以确保安全".to_string(),
},
}))
}
#[derive(Debug, Deserialize)]
struct HistoryQuery {
page: Option<u32>,
limit: Option<u32>,
status: Option<String>,
}
#[derive(Debug, Serialize)]
struct HistoryFileView {
file_id: Uuid,
original_name: String,
original_size: i64,
compressed_size: Option<i64>,
saved_percent: Option<f64>,
status: String,
output_format: String,
error_message: Option<String>,
download_url: Option<String>,
}
#[derive(Debug, Serialize)]
struct HistoryTaskView {
task_id: Uuid,
status: String,
source: String,
progress: i32,
total_files: i32,
completed_files: i32,
failed_files: i32,
created_at: DateTime<Utc>,
completed_at: Option<DateTime<Utc>>,
expires_at: DateTime<Utc>,
download_all_url: Option<String>,
files: Vec<HistoryFileView>,
}
#[derive(Debug, Serialize)]
struct HistoryResponse {
tasks: Vec<HistoryTaskView>,
page: u32,
limit: u32,
total: i64,
}
async fn list_history(
State(state): State<AppState>,
jar: axum_extra::extract::cookie::CookieJar,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
headers: HeaderMap,
Query(query): Query<HistoryQuery>,
) -> Result<Json<Envelope<HistoryResponse>>, AppError> {
let ip = context::client_ip(&headers, addr.ip());
let (_jar, principal) = context::authenticate(&state, jar, &headers, ip).await?;
let user_id = match principal {
context::Principal::User { user_id, .. } => user_id,
_ => return Err(AppError::new(ErrorCode::Unauthorized, "未登录")),
};
let limit = query.limit.unwrap_or(20).clamp(1, 100);
let page = query.page.unwrap_or(1).max(1);
let offset = (page - 1) * limit;
let status = super::normalize_task_status(query.status.as_deref())?;
let total: i64 = if let Some(status) = status {
sqlx::query_scalar(
"SELECT COUNT(*) FROM tasks WHERE user_id = $1 AND status = $2::task_status",
)
.bind(user_id)
.bind(status)
.fetch_one(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "查询历史失败").with_source(err))?
} else {
sqlx::query_scalar("SELECT COUNT(*) FROM tasks WHERE user_id = $1")
.bind(user_id)
.fetch_one(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "查询历史失败").with_source(err))?
};
#[derive(Debug, FromRow)]
struct TaskRow {
id: Uuid,
status: String,
source: String,
total_files: i32,
completed_files: i32,
failed_files: i32,
created_at: DateTime<Utc>,
completed_at: Option<DateTime<Utc>>,
expires_at: DateTime<Utc>,
}
let tasks: Vec<TaskRow> = if let Some(status) = status {
sqlx::query_as::<_, TaskRow>(
r#"
SELECT
id,
status::text AS status,
source::text AS source,
total_files,
completed_files,
failed_files,
created_at,
completed_at,
expires_at
FROM tasks
WHERE user_id = $1 AND status = $2::task_status
ORDER BY created_at DESC
LIMIT $3 OFFSET $4
"#,
)
.bind(user_id)
.bind(status)
.bind(limit as i64)
.bind(offset as i64)
.fetch_all(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "查询历史失败").with_source(err))?
} else {
sqlx::query_as::<_, TaskRow>(
r#"
SELECT
id,
status::text AS status,
source::text AS source,
total_files,
completed_files,
failed_files,
created_at,
completed_at,
expires_at
FROM tasks
WHERE user_id = $1
ORDER BY created_at DESC
LIMIT $2 OFFSET $3
"#,
)
.bind(user_id)
.bind(limit as i64)
.bind(offset as i64)
.fetch_all(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "查询历史失败").with_source(err))?
};
#[derive(Debug, FromRow)]
struct FileRow {
task_id: Uuid,
id: Uuid,
original_name: String,
original_size: i64,
compressed_size: Option<i64>,
saved_percent: Option<f64>,
status: String,
output_format: String,
error_message: Option<String>,
has_storage: bool,
}
let now = Utc::now();
let task_ids = tasks.iter().map(|task| task.id).collect::<Vec<_>>();
let files = if task_ids.is_empty() {
Vec::new()
} else {
sqlx::query_as::<_, FileRow>(
r#"
SELECT
task_id,
id,
original_name,
original_size,
compressed_size,
saved_percent::float8 AS saved_percent,
status::text AS status,
output_format,
error_message,
COALESCE(storage_key, storage_path) IS NOT NULL AS has_storage
FROM task_files
WHERE task_id = ANY($1)
ORDER BY task_id, created_at ASC
"#,
)
.bind(&task_ids)
.fetch_all(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "查询任务文件失败").with_source(err))?
};
let mut files_by_task = HashMap::<Uuid, Vec<FileRow>>::new();
for file in files {
files_by_task.entry(file.task_id).or_default().push(file);
}
let mut result_tasks = Vec::with_capacity(tasks.len());
for task in tasks {
let file_views = files_by_task
.remove(&task.id)
.unwrap_or_default()
.into_iter()
.map(|file| HistoryFileView {
file_id: file.id,
original_name: file.original_name,
original_size: file.original_size,
compressed_size: file.compressed_size,
saved_percent: file.saved_percent,
status: file.status.clone(),
output_format: file.output_format,
error_message: file.error_message,
download_url: if file.status == "completed"
&& file.has_storage
&& task.expires_at > now
{
Some(format!("/downloads/{}", file.id))
} else {
None
},
})
.collect::<Vec<_>>();
let progress = if task.total_files > 0 {
((task.completed_files + task.failed_files) * 100 / task.total_files).clamp(0, 100)
} else {
0
};
let download_all_url = if task.status == "completed" && task.expires_at > now {
Some(format!("/downloads/tasks/{}", task.id))
} else {
None
};
result_tasks.push(HistoryTaskView {
task_id: task.id,
status: task.status,
source: task.source,
progress,
total_files: task.total_files,
completed_files: task.completed_files,
failed_files: task.failed_files,
created_at: task.created_at,
completed_at: task.completed_at,
expires_at: task.expires_at,
download_all_url,
files: file_views,
});
}
Ok(Json(Envelope {
success: true,
data: HistoryResponse {
tasks: result_tasks,
page,
limit,
total,
},
}))
}
async fn list_api_keys(
State(state): State<AppState>,
jar: axum_extra::extract::cookie::CookieJar,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
headers: HeaderMap,
) -> Result<Json<Envelope<ApiKeyListResponse>>, AppError> {
let ip = context::client_ip(&headers, addr.ip());
let (_jar, principal) = context::authenticate(&state, jar, &headers, ip).await?;
let (user_id, _email_verified) = match principal {
context::Principal::User {
user_id,
email_verified,
..
} => (user_id, email_verified),
_ => return Err(AppError::new(ErrorCode::Unauthorized, "未登录")),
};
let rows = sqlx::query_as::<_, ApiKeyView>(
r#"
SELECT
id,
name,
key_prefix,
permissions,
rate_limit,
is_active,
last_used_at,
last_used_ip::text AS last_used_ip,
created_at
FROM api_keys
WHERE user_id = $1
ORDER BY created_at DESC
"#,
)
.bind(user_id)
.fetch_all(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "查询 API Key 失败").with_source(err))?;
Ok(Json(Envelope {
success: true,
data: ApiKeyListResponse { api_keys: rows },
}))
}
#[derive(Debug, Deserialize)]
struct CreateApiKeyRequest {
name: String,
permissions: Option<Vec<String>>,
}
#[derive(Debug, Serialize)]
struct CreateApiKeyResponse {
id: Uuid,
name: String,
key_prefix: String,
key: String,
message: String,
}
async fn create_api_key(
State(state): State<AppState>,
jar: axum_extra::extract::cookie::CookieJar,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
headers: HeaderMap,
Json(req): Json<CreateApiKeyRequest>,
) -> Result<Json<Envelope<CreateApiKeyResponse>>, AppError> {
if req.name.trim().is_empty() || req.name.len() > 100 {
return Err(AppError::new(ErrorCode::InvalidRequest, "name 不合法"));
}
let ip = context::client_ip(&headers, addr.ip());
let (_jar, principal) = context::authenticate(&state, jar, &headers, ip).await?;
let (user_id, email_verified) = match principal {
context::Principal::User {
user_id,
email_verified,
..
} => (user_id, email_verified),
_ => return Err(AppError::new(ErrorCode::Unauthorized, "未登录")),
};
if !email_verified {
return Err(AppError::new(ErrorCode::EmailNotVerified, "请先验证邮箱"));
}
if !settings::runtime_policy(&state)
.await?
.features
.api_key_enabled
{
return Err(AppError::new(
ErrorCode::Forbidden,
"API Key 功能当前已关闭",
));
}
let billing = billing::get_user_billing(&state, user_id).await?;
if !billing.plan.feature_api_enabled {
return Err(AppError::new(
ErrorCode::Forbidden,
"当前套餐未开通 API Key",
));
}
let permissions = normalize_permissions(req.permissions)?;
let (full_key, key_prefix) = generate_api_key();
let key_hash = context::api_key_hash(&full_key, &state.config.api_key_pepper)?;
let row_id: Uuid = sqlx::query_scalar(
r#"
INSERT INTO api_keys (user_id, name, key_prefix, key_hash, permissions, rate_limit)
VALUES ($1, $2, $3, $4, $5, 100)
RETURNING id
"#,
)
.bind(user_id)
.bind(req.name.trim())
.bind(&key_prefix)
.bind(key_hash)
.bind(&permissions)
.fetch_one(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "创建 API Key 失败").with_source(err))?;
Ok(Json(Envelope {
success: true,
data: CreateApiKeyResponse {
id: row_id,
name: req.name.trim().to_string(),
key_prefix,
key: full_key,
message: "请保存此 Key它只会显示一次".to_string(),
},
}))
}
async fn disable_api_key(
State(state): State<AppState>,
jar: axum_extra::extract::cookie::CookieJar,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
headers: HeaderMap,
Path(key_id): Path<Uuid>,
) -> Result<Json<Envelope<serde_json::Value>>, AppError> {
let ip = context::client_ip(&headers, addr.ip());
let (_jar, principal) = context::authenticate(&state, jar, &headers, ip).await?;
let user_id = match principal {
context::Principal::User { user_id, .. } => user_id,
_ => return Err(AppError::new(ErrorCode::Unauthorized, "未登录")),
};
let result =
sqlx::query("UPDATE api_keys SET is_active = false WHERE id = $1 AND user_id = $2")
.bind(key_id)
.bind(user_id)
.execute(&state.db)
.await
.map_err(|err| {
AppError::new(ErrorCode::Internal, "更新 API Key 失败").with_source(err)
})?;
if result.rows_affected() == 0 {
return Err(AppError::new(ErrorCode::NotFound, "API Key 不存在"));
}
Ok(Json(Envelope {
success: true,
data: serde_json::json!({ "message": "已禁用" }),
}))
}
async fn rotate_api_key(
State(state): State<AppState>,
jar: axum_extra::extract::cookie::CookieJar,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
headers: HeaderMap,
Path(key_id): Path<Uuid>,
) -> Result<Json<Envelope<CreateApiKeyResponse>>, AppError> {
let ip = context::client_ip(&headers, addr.ip());
let (_jar, principal) = context::authenticate(&state, jar, &headers, ip).await?;
let (user_id, email_verified) = match principal {
context::Principal::User {
user_id,
email_verified,
..
} => (user_id, email_verified),
_ => return Err(AppError::new(ErrorCode::Unauthorized, "未登录")),
};
if !email_verified {
return Err(AppError::new(ErrorCode::EmailNotVerified, "请先验证邮箱"));
}
if !settings::runtime_policy(&state)
.await?
.features
.api_key_enabled
{
return Err(AppError::new(
ErrorCode::Forbidden,
"API Key 功能当前已关闭",
));
}
let (full_key, key_prefix) = generate_api_key();
let key_hash = context::api_key_hash(&full_key, &state.config.api_key_pepper)?;
#[derive(Debug, FromRow)]
struct RotateRow {
id: Uuid,
name: String,
}
let row = sqlx::query_as::<_, RotateRow>(
r#"
UPDATE api_keys
SET key_prefix = $1,
key_hash = $2,
is_active = true
WHERE id = $3 AND user_id = $4
RETURNING id, name
"#,
)
.bind(&key_prefix)
.bind(key_hash)
.bind(key_id)
.bind(user_id)
.fetch_optional(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "更新 API Key 失败").with_source(err))?
.ok_or_else(|| AppError::new(ErrorCode::NotFound, "API Key 不存在"))?;
Ok(Json(Envelope {
success: true,
data: CreateApiKeyResponse {
id: row.id,
name: row.name,
key_prefix,
key: full_key,
message: "请保存此 Key它只会显示一次".to_string(),
},
}))
}
fn generate_api_key() -> (String, String) {
let mut prefix_bytes = [0u8; 4];
rand::rngs::OsRng.fill_bytes(&mut prefix_bytes);
let prefix = hex::encode(prefix_bytes);
let key_prefix = format!("if_live_{prefix}");
let mut secret_bytes = [0u8; 32];
rand::rngs::OsRng.fill_bytes(&mut secret_bytes);
let secret = URL_SAFE_NO_PAD.encode(secret_bytes);
let full = format!("{key_prefix}_{secret}");
(full, key_prefix)
}
fn normalize_permissions(input: Option<Vec<String>>) -> Result<serde_json::Value, AppError> {
let allowed = ["compress", "batch_compress"];
let mut perms = Vec::<String>::new();
if let Some(values) = input {
for value in values {
let v = value.trim().to_ascii_lowercase();
if v.is_empty() {
continue;
}
if !allowed.contains(&v.as_str()) {
return Err(AppError::new(
ErrorCode::InvalidRequest,
format!("不支持的权限: {v}"),
));
}
if !perms.contains(&v) {
perms.push(v);
}
}
}
if perms.is_empty() {
perms.push("compress".to_string());
}
Ok(serde_json::json!(perms))
}
fn map_unique_violation(err: sqlx::Error) -> AppError {
if let sqlx::Error::Database(db_err) = &err {
if let Some(code) = db_err.code() {
if code == "23505" {
return AppError::new(ErrorCode::InvalidRequest, "邮箱或用户名已存在");
}
}
}
AppError::new(ErrorCode::Internal, "数据库操作失败").with_source(err)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::Config;
use crate::services::mail::Mailer;
use axum::body::{to_bytes, Body};
use axum::http::{Method, Request, StatusCode};
use sqlx::postgres::PgPoolOptions;
use std::sync::Arc;
use tokio::sync::Semaphore;
use tower::ServiceExt;
async fn json_request(
app: &axum::Router,
method: Method,
uri: &str,
token: Option<&str>,
payload: serde_json::Value,
) -> (StatusCode, serde_json::Value) {
let mut builder = Request::builder()
.method(method)
.uri(uri)
.header(axum::http::header::CONTENT_TYPE, "application/json");
if let Some(token) = token {
builder = builder.header(axum::http::header::AUTHORIZATION, format!("Bearer {token}"));
}
let mut request = builder
.body(Body::from(payload.to_string()))
.expect("build test request");
request.extensions_mut().insert(ConnectInfo(
"127.0.0.1:41000"
.parse::<SocketAddr>()
.expect("parse test address"),
));
let response = app
.clone()
.oneshot(request)
.await
.expect("execute test request");
let status = response.status();
let body = to_bytes(response.into_body(), usize::MAX)
.await
.expect("read test response");
let json = serde_json::from_slice(&body)
.unwrap_or_else(|_| panic!("response is not JSON: {}", String::from_utf8_lossy(&body)));
(status, json)
}
async fn replace_pending_token(pool: &sqlx::PgPool, user_id: Uuid, token: &str) {
let updated = sqlx::query(
r#"
UPDATE email_change_requests
SET token_hash = $2
WHERE user_id = $1
AND confirmed_at IS NULL
AND canceled_at IS NULL
"#,
)
.bind(user_id)
.bind(credentials::sha256_hex(token))
.execute(pool)
.await
.expect("replace pending email token");
assert_eq!(updated.rows_affected(), 1);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
#[ignore = "requires isolated IMAGEFORGE_TEST_DATABASE_URL and IMAGEFORGE_TEST_REDIS_URL"]
async fn pending_email_cannot_take_over_user_or_admin_recovery() {
let database_url = std::env::var("IMAGEFORGE_TEST_DATABASE_URL")
.expect("IMAGEFORGE_TEST_DATABASE_URL must be set");
assert!(
database_url.to_ascii_lowercase().contains("test"),
"refusing to run destructive integration test outside a test database"
);
let redis_url = std::env::var("IMAGEFORGE_TEST_REDIS_URL")
.expect("IMAGEFORGE_TEST_REDIS_URL must be set");
let pool = PgPoolOptions::new()
.max_connections(16)
.connect(&database_url)
.await
.expect("connect test database");
sqlx::migrate!().run(&pool).await.expect("run migrations");
sqlx::query(
r#"
INSERT INTO system_config (key, value, description)
VALUES (
'auth_config',
'{"email_verification_required":true}'::jsonb,
'security integration test'
)
ON CONFLICT (key) DO UPDATE
SET value = EXCLUDED.value, updated_at = NOW()
"#,
)
.execute(&pool)
.await
.expect("enable email verification");
let mut config = Config::from_env().expect("load test config");
config.database_url = database_url;
config.redis_url = redis_url;
config.mail_enabled = false;
config.mail_log_links_when_disabled = false;
let redis = redis::Client::open(config.redis_url.clone())
.expect("create test redis client")
.get_connection_manager()
.await
.expect("connect test redis");
let state = AppState {
mailer: Arc::new(Mailer::new(&config).expect("create disabled test mailer")),
image_processing_semaphore: Arc::new(Semaphore::new(2)),
zip_build_semaphore: Arc::new(Semaphore::new(2)),
runtime_policy_cache: crate::services::settings::RuntimePolicyCache::new(),
storage_cache: crate::services::storage::StorageCache::new(),
config,
db: pool.clone(),
redis,
};
let app = axum::Router::new()
.nest("/auth", crate::api::auth::router())
.merge(router())
.with_state(state.clone());
let marker = Uuid::new_v4().simple().to_string();
let password = "Original9!";
let password_hash = credentials::hash_password(password)
.await
.expect("hash test password");
let user_id = Uuid::new_v4();
let old_email = format!("old-{marker}@example.test");
let pending_email = format!("pending-{marker}@example.test");
sqlx::query(
r#"
INSERT INTO users (
id, email, username, password_hash, email_verified_at
) VALUES ($1, $2, $3, $4, NOW())
"#,
)
.bind(user_id)
.bind(&old_email)
.bind(format!("user_{marker}"))
.bind(&password_hash)
.execute(&pool)
.await
.expect("insert test user");
let (old_token, _) = auth::issue_jwt(
&state.config.jwt_secret,
state.config.jwt_expiry_hours,
user_id,
"user",
0,
)
.expect("issue test jwt");
let (status, _) = json_request(
&app,
Method::PUT,
"/user/profile",
Some(&old_token),
serde_json::json!({ "email": pending_email }),
)
.await;
assert_eq!(status, StatusCode::UNAUTHORIZED, "JWT alone changed email");
let (status, _) = json_request(
&app,
Method::PUT,
"/user/profile",
Some(&old_token),
serde_json::json!({
"email": pending_email,
"current_password": "WrongPassword9!"
}),
)
.await;
assert_eq!(status, StatusCode::UNAUTHORIZED);
let (status, response) = json_request(
&app,
Method::PUT,
"/user/profile",
Some(&old_token),
serde_json::json!({
"email": pending_email,
"current_password": password
}),
)
.await;
assert_eq!(status, StatusCode::OK, "{response}");
assert_eq!(response["data"]["user"]["email"], old_email);
assert_eq!(response["data"]["user"]["pending_email"], pending_email);
let persisted_email: String = sqlx::query_scalar("SELECT email FROM users WHERE id = $1")
.bind(user_id)
.fetch_one(&pool)
.await
.expect("query persisted email");
assert_eq!(persisted_email, old_email);
let (status, _) = json_request(
&app,
Method::POST,
"/auth/forgot-password",
None,
serde_json::json!({ "email": pending_email }),
)
.await;
assert_eq!(status, StatusCode::OK);
let active_resets: i64 = sqlx::query_scalar(
"SELECT COUNT(*) FROM password_resets WHERE user_id = $1 AND used_at IS NULL",
)
.bind(user_id)
.fetch_one(&pool)
.await
.expect("count pending-email resets");
assert_eq!(active_resets, 0, "pending email became a recovery address");
let (status, _) = json_request(
&app,
Method::POST,
"/auth/forgot-password",
None,
serde_json::json!({ "email": old_email }),
)
.await;
assert_eq!(status, StatusCode::OK);
let active_resets: i64 = sqlx::query_scalar(
"SELECT COUNT(*) FROM password_resets WHERE user_id = $1 AND used_at IS NULL",
)
.bind(user_id)
.fetch_one(&pool)
.await
.expect("count old-email resets");
assert_eq!(
active_resets, 1,
"old email stopped being the recovery address early"
);
let confirm_token = format!("confirm-{marker}");
replace_pending_token(&pool, user_id, &confirm_token).await;
let (status, response) = json_request(
&app,
Method::POST,
"/auth/verify-email",
None,
serde_json::json!({ "token": confirm_token }),
)
.await;
assert_eq!(status, StatusCode::OK, "{response}");
assert_eq!(response["data"]["session_invalidated"], true);
let active_resets: i64 = sqlx::query_scalar(
"SELECT COUNT(*) FROM password_resets WHERE user_id = $1 AND used_at IS NULL",
)
.bind(user_id)
.fetch_one(&pool)
.await
.expect("count invalidated resets");
assert_eq!(active_resets, 0);
let (status, _) = json_request(
&app,
Method::GET,
"/user/profile",
Some(&old_token),
serde_json::json!({}),
)
.await;
assert_eq!(status, StatusCode::UNAUTHORIZED, "old JWT remained valid");
let (status, response) = json_request(
&app,
Method::POST,
"/auth/login",
None,
serde_json::json!({ "email": pending_email, "password": password }),
)
.await;
assert_eq!(status, StatusCode::OK, "{response}");
let admin_id = Uuid::new_v4();
let admin_old_email = format!("admin-old-{marker}@example.test");
let admin_pending_email = format!("admin-pending-{marker}@example.test");
sqlx::query(
r#"
INSERT INTO users (
id, email, username, password_hash, role, email_verified_at
) VALUES ($1, $2, $3, $4, 'admin', NOW())
"#,
)
.bind(admin_id)
.bind(&admin_old_email)
.bind(format!("admin_{marker}"))
.bind(&password_hash)
.execute(&pool)
.await
.expect("insert test admin");
let (admin_token, _) = auth::issue_jwt(
&state.config.jwt_secret,
state.config.jwt_expiry_hours,
admin_id,
"admin",
0,
)
.expect("issue admin jwt");
let (status, _) = json_request(
&app,
Method::PUT,
"/user/profile",
Some(&admin_token),
serde_json::json!({ "email": admin_pending_email }),
)
.await;
assert_eq!(status, StatusCode::UNAUTHORIZED);
let (status, response) = json_request(
&app,
Method::PUT,
"/user/profile",
Some(&admin_token),
serde_json::json!({
"email": admin_pending_email,
"current_password": password
}),
)
.await;
assert_eq!(status, StatusCode::OK, "{response}");
assert_eq!(response["data"]["user"]["email"], admin_old_email);
let admin_confirm_token = format!("admin-confirm-{marker}");
replace_pending_token(&pool, admin_id, &admin_confirm_token).await;
let (status, response) = json_request(
&app,
Method::POST,
"/auth/verify-email",
None,
serde_json::json!({ "token": admin_confirm_token }),
)
.await;
assert_eq!(status, StatusCode::OK, "{response}");
let (status, _) = json_request(
&app,
Method::GET,
"/user/profile",
Some(&admin_token),
serde_json::json!({}),
)
.await;
assert_eq!(status, StatusCode::UNAUTHORIZED);
let admin: (String, String) =
sqlx::query_as("SELECT email, role::text FROM users WHERE id = $1")
.bind(admin_id)
.fetch_one(&pool)
.await
.expect("query test admin");
assert_eq!(admin, (admin_pending_email, "admin".to_string()));
let reset_user_id = Uuid::new_v4();
let reset_old_email = format!("reset-old-{marker}@example.test");
let reset_new_email = format!("reset-new-{marker}@example.test");
let reset_token_a = format!("reset-a-{marker}");
let reset_token_b = format!("reset-b-{marker}");
let reset_email_token = format!("reset-email-{marker}");
sqlx::query(
r#"
INSERT INTO users (id, email, username, password_hash, email_verified_at)
VALUES ($1, $2, $3, $4, NOW())
"#,
)
.bind(reset_user_id)
.bind(&reset_old_email)
.bind(format!("reset_{marker}"))
.bind(&password_hash)
.execute(&pool)
.await
.expect("insert multi-reset user");
sqlx::query(
r#"
INSERT INTO password_resets (user_id, token_hash, expires_at)
VALUES
($1, $2, NOW() + INTERVAL '1 hour'),
($1, $3, NOW() + INTERVAL '1 hour')
"#,
)
.bind(reset_user_id)
.bind(credentials::sha256_hex(&reset_token_a))
.bind(credentials::sha256_hex(&reset_token_b))
.execute(&pool)
.await
.expect("insert two password resets");
sqlx::query(
r#"
INSERT INTO email_change_requests (user_id, new_email, token_hash, expires_at)
VALUES ($1, $2, $3, NOW() + INTERVAL '1 hour')
"#,
)
.bind(reset_user_id)
.bind(&reset_new_email)
.bind(credentials::sha256_hex(&reset_email_token))
.execute(&pool)
.await
.expect("insert pending email change before reset");
let (status, response) = json_request(
&app,
Method::POST,
"/auth/reset-password",
None,
serde_json::json!({
"token": reset_token_a,
"new_password": "Replacement9!"
}),
)
.await;
assert_eq!(status, StatusCode::OK, "{response}");
let (status, response) = json_request(
&app,
Method::POST,
"/auth/reset-password",
None,
serde_json::json!({
"token": reset_token_b,
"new_password": "SecondReplacement9!"
}),
)
.await;
assert_eq!(status, StatusCode::BAD_REQUEST, "{response}");
assert_eq!(response["error"]["code"], "INVALID_TOKEN");
let (status, response) = json_request(
&app,
Method::POST,
"/auth/verify-email",
None,
serde_json::json!({ "token": reset_email_token }),
)
.await;
assert_eq!(status, StatusCode::BAD_REQUEST, "{response}");
assert_eq!(response["error"]["code"], "INVALID_TOKEN");
let recovery_state: (i64, i64) = sqlx::query_as(
r#"
SELECT
(SELECT COUNT(*) FROM password_resets WHERE user_id = $1 AND used_at IS NULL),
(SELECT COUNT(*) FROM email_change_requests
WHERE user_id = $1 AND confirmed_at IS NULL AND canceled_at IS NULL)
"#,
)
.bind(reset_user_id)
.fetch_one(&pool)
.await
.expect("query recovery invalidation state");
assert_eq!(recovery_state, (0, 0));
let password_user_id = Uuid::new_v4();
let password_email = format!("password-{marker}@example.test");
let password_reset_token = format!("password-reset-{marker}");
let password_email_token = format!("password-email-{marker}");
sqlx::query(
r#"
INSERT INTO users (id, email, username, password_hash, email_verified_at)
VALUES ($1, $2, $3, $4, NOW())
"#,
)
.bind(password_user_id)
.bind(&password_email)
.bind(format!("password_{marker}"))
.bind(&password_hash)
.execute(&pool)
.await
.expect("insert password-update user");
sqlx::query(
r#"
INSERT INTO password_resets (user_id, token_hash, expires_at)
VALUES ($1, $2, NOW() + INTERVAL '1 hour')
"#,
)
.bind(password_user_id)
.bind(credentials::sha256_hex(&password_reset_token))
.execute(&pool)
.await
.expect("insert reset before authenticated password update");
sqlx::query(
r#"
INSERT INTO email_change_requests (user_id, new_email, token_hash, expires_at)
VALUES ($1, $2, $3, NOW() + INTERVAL '1 hour')
"#,
)
.bind(password_user_id)
.bind(format!("password-new-{marker}@example.test"))
.bind(credentials::sha256_hex(&password_email_token))
.execute(&pool)
.await
.expect("insert email change before authenticated password update");
let (password_token, _) = auth::issue_jwt(
&state.config.jwt_secret,
state.config.jwt_expiry_hours,
password_user_id,
"user",
0,
)
.expect("issue password-update JWT");
let (status, response) = json_request(
&app,
Method::PUT,
"/user/password",
Some(&password_token),
serde_json::json!({
"current_password": password,
"new_password": "AuthenticatedReplacement9!"
}),
)
.await;
assert_eq!(status, StatusCode::OK, "{response}");
let (status, response) = json_request(
&app,
Method::POST,
"/auth/reset-password",
None,
serde_json::json!({
"token": password_reset_token,
"new_password": "StaleReset9!"
}),
)
.await;
assert_eq!(status, StatusCode::BAD_REQUEST, "{response}");
assert_eq!(response["error"]["code"], "INVALID_TOKEN");
let (status, response) = json_request(
&app,
Method::POST,
"/auth/verify-email",
None,
serde_json::json!({ "token": password_email_token }),
)
.await;
assert_eq!(status, StatusCode::BAD_REQUEST, "{response}");
assert_eq!(response["error"]["code"], "INVALID_TOKEN");
let reset_race_user_id = Uuid::new_v4();
let reset_race_old_email = format!("reset-race-old-{marker}@example.test");
let reset_race_new_email = format!("reset-race-new-{marker}@example.test");
let reset_race_token = format!("reset-race-{marker}");
let reset_race_email_token = format!("reset-race-email-{marker}");
sqlx::query(
r#"
INSERT INTO users (id, email, username, password_hash, email_verified_at)
VALUES ($1, $2, $3, $4, NOW())
"#,
)
.bind(reset_race_user_id)
.bind(&reset_race_old_email)
.bind(format!("reset_race_{marker}"))
.bind(&password_hash)
.execute(&pool)
.await
.expect("insert reset-email race user");
sqlx::query(
r#"
INSERT INTO password_resets (user_id, token_hash, expires_at)
VALUES ($1, $2, NOW() + INTERVAL '1 hour')
"#,
)
.bind(reset_race_user_id)
.bind(credentials::sha256_hex(&reset_race_token))
.execute(&pool)
.await
.expect("insert racing reset");
sqlx::query(
r#"
INSERT INTO email_change_requests (user_id, new_email, token_hash, expires_at)
VALUES ($1, $2, $3, NOW() + INTERVAL '1 hour')
"#,
)
.bind(reset_race_user_id)
.bind(&reset_race_new_email)
.bind(credentials::sha256_hex(&reset_race_email_token))
.execute(&pool)
.await
.expect("insert racing email change");
let mut blocker = pool.begin().await.expect("begin reset-email race blocker");
let _: Uuid = sqlx::query_scalar("SELECT id FROM users WHERE id = $1 FOR UPDATE")
.bind(reset_race_user_id)
.fetch_one(&mut *blocker)
.await
.expect("lock reset-email race user");
let reset_join = {
let app = app.clone();
let token = reset_race_token.clone();
tokio::spawn(async move {
json_request(
&app,
Method::POST,
"/auth/reset-password",
None,
serde_json::json!({
"token": token,
"new_password": "RaceReplacement9!"
}),
)
.await
})
};
let confirm_join = {
let app = app.clone();
let token = reset_race_email_token.clone();
tokio::spawn(async move {
json_request(
&app,
Method::POST,
"/auth/verify-email",
None,
serde_json::json!({ "token": token }),
)
.await
})
};
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
blocker
.commit()
.await
.expect("release reset-email race user");
let reset_result = reset_join.await.expect("join racing reset");
let confirm_result = confirm_join.await.expect("join racing confirmation");
let successes = [reset_result.0, confirm_result.0]
.into_iter()
.filter(|status| *status == StatusCode::OK)
.count();
assert_eq!(successes, 1, "reset and email confirmation both committed");
for (status, response) in [&reset_result, &confirm_result] {
if *status != StatusCode::OK {
assert_eq!(*status, StatusCode::BAD_REQUEST, "{response}");
assert_eq!(response["error"]["code"], "INVALID_TOKEN");
}
}
let (race_email, race_hash): (String, String) =
sqlx::query_as("SELECT email, password_hash FROM users WHERE id = $1")
.bind(reset_race_user_id)
.fetch_one(&pool)
.await
.expect("query reset-email race result");
if reset_result.0 == StatusCode::OK {
assert_eq!(race_email, reset_race_old_email);
assert!(
credentials::verify_password("RaceReplacement9!", &race_hash)
.await
.expect("verify racing reset password")
);
} else {
assert_eq!(race_email, reset_race_new_email);
assert!(credentials::verify_password(password, &race_hash)
.await
.expect("verify original password after email confirmation"));
}
let race_user_id = Uuid::new_v4();
let race_old_email = format!("race-old-{marker}@example.test");
let race_new_email = format!("race-new-{marker}@example.test");
let race_confirm_token = format!("race-confirm-{marker}");
sqlx::query(
r#"
INSERT INTO users (
id, email, username, password_hash, email_verified_at
) VALUES ($1, $2, $3, $4, NOW())
"#,
)
.bind(race_user_id)
.bind(&race_old_email)
.bind(format!("race_{marker}"))
.bind(&password_hash)
.execute(&pool)
.await
.expect("insert recovery race user");
sqlx::query(
r#"
INSERT INTO email_change_requests (
user_id, new_email, token_hash, expires_at
) VALUES ($1, $2, $3, NOW() + INTERVAL '1 hour')
"#,
)
.bind(race_user_id)
.bind(&race_new_email)
.bind(credentials::sha256_hex(&race_confirm_token))
.execute(&pool)
.await
.expect("insert recovery race email change");
let mut blocker = pool.begin().await.expect("begin recovery race blocker");
let _: Uuid = sqlx::query_scalar("SELECT id FROM users WHERE id = $1 FOR UPDATE")
.bind(race_user_id)
.fetch_one(&mut *blocker)
.await
.expect("lock recovery race user");
let verify_join = {
let app = app.clone();
let token = race_confirm_token.clone();
tokio::spawn(async move {
json_request(
&app,
Method::POST,
"/auth/verify-email",
None,
serde_json::json!({ "token": token }),
)
.await
})
};
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
let forgot_join = {
let app = app.clone();
let email = race_old_email.clone();
tokio::spawn(async move {
json_request(
&app,
Method::POST,
"/auth/forgot-password",
None,
serde_json::json!({ "email": email }),
)
.await
})
};
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
blocker.commit().await.expect("release recovery race user");
let (verify_status, verify_response) =
verify_join.await.expect("join racing email confirmation");
let (forgot_status, forgot_response) =
forgot_join.await.expect("join racing password recovery");
assert_eq!(verify_status, StatusCode::OK, "{verify_response}");
assert_eq!(forgot_status, StatusCode::OK, "{forgot_response}");
let race_result: (String, i64) = sqlx::query_as(
r#"
SELECT u.email,
COUNT(r.id) FILTER (WHERE r.used_at IS NULL)
FROM users u
LEFT JOIN password_resets r ON r.user_id = u.id
WHERE u.id = $1
GROUP BY u.email
"#,
)
.bind(race_user_id)
.fetch_one(&pool)
.await
.expect("query recovery race result");
assert_eq!(race_result, (race_new_email, 0));
sqlx::query("DELETE FROM users WHERE id IN ($1, $2, $3, $4, $5, $6)")
.bind(user_id)
.bind(admin_id)
.bind(race_user_id)
.bind(reset_user_id)
.bind(password_user_id)
.bind(reset_race_user_id)
.execute(&pool)
.await
.expect("delete account recovery test users");
}
}