1826 lines
60 KiB
Rust
1826 lines
60 KiB
Rust
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");
|
||
}
|
||
}
|