Skip to main content

yggdrasil/api/
auth.rs

1//! 认证相关的 Dioxus server function 与辅助函数。
2//!
3//! 提供注册、登录、登出、获取当前用户等接口,
4//! 通过 HttpOnly Cookie 维护会话,首个注册用户自动成为 admin。
5//! 所有 server function 均在 `#[server(Name, "/api")]` 下注册,供客户端与服务端调用。
6//! 仅在 `feature = "server"` 启用的服务端构建中执行数据库操作与 Cookie 写入。
7
8#![allow(clippy::unused_unit, deprecated)]
9
10use dioxus::prelude::*;
11#[cfg(feature = "server")]
12use http::header::{HeaderValue, SET_COOKIE};
13
14#[cfg(feature = "server")]
15use crate::api::error::AppError;
16#[cfg(feature = "server")]
17use crate::auth::session::get_session_from_ctx;
18#[cfg(feature = "server")]
19use crate::auth::{password, session};
20#[cfg(feature = "server")]
21use crate::db::pool::get_conn;
22use crate::models::user::PublicUser;
23#[cfg(feature = "server")]
24use crate::models::user::{SessionUser, UserRole};
25
26#[cfg(feature = "server")]
27fn validate_username(username: &str) -> Result<(), String> {
28    if username.len() < 3 || username.len() > 50 {
29        return Err("用户名长度必须在 3-50 字符之间".to_string());
30    }
31    if !username.chars().all(|c| c.is_alphanumeric() || c == '_') {
32        return Err("用户名只能包含字母、数字和下划线".to_string());
33    }
34    Ok(())
35}
36
37#[cfg(feature = "server")]
38fn validate_email(email: &str) -> Result<(), String> {
39    if !crate::utils::server::EMAIL_REGEX.is_match(email) {
40        return Err("邮箱格式不正确".to_string());
41    }
42    Ok(())
43}
44
45#[cfg(feature = "server")]
46fn validate_password(password: &str) -> Result<(), String> {
47    if password.len() < 8 {
48        return Err("密码长度至少 8 位".to_string());
49    }
50    Ok(())
51}
52
53#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
54/// 认证接口统一响应结构。
55pub struct AuthResponse {
56    /// 操作是否成功。
57    pub success: bool,
58    /// 提示信息。
59    pub message: String,
60    /// 登录成功后的会话 token(已废弃,实际通过 Cookie 传递)。
61    pub token: Option<String>,
62}
63
64/// 用户注册。
65///
66/// 校验用户名、邮箱、密码,首个注册用户自动设为 admin;
67/// 已有 admin 时返回 "Registration is closed"。
68/// Dioxus server function,注册在 `/api` 路径下。
69#[server(Register, "/api")]
70pub async fn register(
71    username: String,
72    email: String,
73    password: String,
74) -> Result<AuthResponse, ServerFnError> {
75    // 服务端构建时先进行严格限流检查。
76    #[cfg(feature = "server")]
77    {
78        if let Some(ctx) = dioxus::fullstack::FullstackContext::current() {
79            let parts = ctx.parts_mut();
80            let ip = crate::api::rate_limit::get_client_ip(&parts.headers);
81            if let Err(msg) = crate::api::rate_limit::check_strict_limit(&ip) {
82                return Ok(AuthResponse {
83                    success: false,
84                    message: msg,
85                    token: None,
86                });
87            }
88        }
89    }
90
91    if let Err(e) = validate_username(&username) {
92        return Ok(AuthResponse {
93            success: false,
94            message: e,
95            token: None,
96        });
97    }
98    if let Err(e) = validate_email(&email) {
99        return Ok(AuthResponse {
100            success: false,
101            message: e,
102            token: None,
103        });
104    }
105    if let Err(e) = validate_password(&password) {
106        return Ok(AuthResponse {
107            success: false,
108            message: e,
109            token: None,
110        });
111    }
112
113    let client = get_conn().await.map_err(AppError::db_conn)?;
114
115    // Argon2 是 memory-hard 计算,必须在 spawn_blocking 中执行,避免阻塞 Tokio worker。
116    let pw_for_hash = password.clone();
117    let password_hash = tokio::task::spawn_blocking(move || password::hash_password(&pw_for_hash))
118        .await
119        .map_err(|_| AppError::Internal("密码处理任务失败"))?
120        .map_err(|_| AppError::Internal("密码处理失败"))?;
121
122    // 使用 INSERT ON CONFLICT 原子性地完成“首个用户成为 admin”的竞争。
123    // 若已有 admin 或用户名/邮箱冲突,RETURNING 将返回空。
124    let result = client
125        .query_opt(
126            "INSERT INTO users (username, email, password_hash, role)
127             VALUES ($1, $2, $3, 'admin')
128             ON CONFLICT DO NOTHING
129             RETURNING id",
130            &[&username, &email, &password_hash],
131        )
132        .await
133        .map_err(AppError::query)?;
134
135    if result.is_some() {
136        return Ok(AuthResponse {
137            success: true,
138            message: "注册成功".to_string(),
139            token: None,
140        });
141    }
142
143    // 插入失败:区分是已有 admin 还是用户名/邮箱冲突。
144    let admin_exists: bool = client
145        .query_one(
146            "SELECT EXISTS (SELECT 1 FROM users WHERE role = 'admin')",
147            &[],
148        )
149        .await
150        .map_err(AppError::query)?
151        .get(0);
152
153    let message = if admin_exists {
154        "Registration is closed".to_string()
155    } else {
156        "用户名或邮箱已存在".to_string()
157    };
158
159    Ok(AuthResponse {
160        success: false,
161        message,
162        token: None,
163    })
164}
165
166/// 用户登录。
167///
168/// 验证用户名/邮箱与密码,生成会话并写入 HttpOnly Cookie;
169/// 同一用户活跃会话数超过 `MAX_SESSIONS_PER_USER` 时删除最早会话。
170/// Dioxus server function,注册在 `/api` 路径下。
171#[server(Login, "/api")]
172pub async fn login(username: String, password: String) -> Result<AuthResponse, ServerFnError> {
173    // 服务端构建时先进行严格限流检查。
174    #[cfg(feature = "server")]
175    {
176        if let Some(ctx) = dioxus::fullstack::FullstackContext::current() {
177            let parts = ctx.parts_mut();
178            let ip = crate::api::rate_limit::get_client_ip(&parts.headers);
179            if let Err(msg) = crate::api::rate_limit::check_strict_limit(&ip) {
180                return Ok(AuthResponse {
181                    success: false,
182                    message: msg,
183                    token: None,
184                });
185            }
186        }
187    }
188
189    let mut client = get_conn().await.map_err(AppError::db_conn)?;
190
191    let row = match client
192        .query_opt(
193            "SELECT id, username, email, password_hash, role, created_at FROM users WHERE username = $1 OR email = $1",
194            &[&username],
195        )
196        .await
197    {
198        Ok(Some(row)) => row,
199        Ok(None) => {
200            // 用户不存在时也执行一次 Argon2 verify,抹平「用户不存在」与
201            // 「密码错误」的响应时序差,防止通过响应时间枚举账号(L2)。
202            // 用固定合法哈希做 verify(必然失败),耗时与真实校验一致。
203            const DUMMY_HASH: &str =
204                "$argon2id$v=19$m=19456,t=2,p=1$j3rNaAXzdExYaL94WBWtfg$n1S75LUQKaYJwaRl5bkFF/f/N1tLfRYR/7TuQxKP94c";
205            let dummy_pw = password.clone();
206            let _ = tokio::task::spawn_blocking(move || {
207                crate::auth::password::verify_password(&dummy_pw, DUMMY_HASH)
208            })
209            .await;
210            return Ok(AuthResponse {
211                success: false,
212                message: "Invalid credentials".to_string(),
213                token: None,
214            });
215        }
216        Err(e) => {
217            return Err(AppError::query(e).into());
218        }
219    };
220
221    let password_hash: String = row.get("password_hash");
222    // Argon2 校验同样在 spawn_blocking 中执行。
223    let pw_for_verify = password.clone();
224    let hash_for_verify = password_hash.clone();
225    let valid = tokio::task::spawn_blocking(move || {
226        password::verify_password(&pw_for_verify, &hash_for_verify)
227    })
228    .await
229    .map_err(|_| AppError::Internal("密码处理任务失败"))?
230    .map_err(|_| AppError::Internal("密码处理失败"))?;
231
232    if !valid {
233        return Ok(AuthResponse {
234            success: false,
235            message: "Invalid credentials".to_string(),
236            token: None,
237        });
238    }
239
240    let user_id: i32 = row.get("id");
241    let token = session::generate_token();
242    let token_hash = session::hash_token(&token);
243    let expires_at = session::default_expiry();
244
245    let max_sessions = std::env::var("MAX_SESSIONS_PER_USER")
246        .ok()
247        .and_then(|s| s.parse::<i64>().ok())
248        .unwrap_or(5)
249        .max(1);
250
251    // 用事务 + 对 users 行加 FOR UPDATE 锁,串行化同一用户的并发登录,
252    // 避免 COUNT→DELETE→INSERT 之间的竞态导致超出上限(M1)。
253    let tx = client.transaction().await.map_err(AppError::query)?;
254    // 锁住该用户行,并发登录在此排队。
255    tx.execute("SELECT 1 FROM users WHERE id = $1 FOR UPDATE", &[&user_id])
256        .await
257        .map_err(AppError::query)?;
258
259    let session_count: i64 = tx
260        .query_one(
261            "SELECT COUNT(*) FROM sessions WHERE user_id = $1 AND expires_at > NOW()",
262            &[&user_id],
263        )
264        .await
265        .map_err(AppError::query)?
266        .get(0);
267
268    if session_count >= max_sessions {
269        tx.execute(
270            "DELETE FROM sessions WHERE id IN (
271                SELECT id FROM sessions
272                WHERE user_id = $1 AND expires_at > NOW()
273                ORDER BY created_at ASC
274                LIMIT 1
275            )",
276            &[&user_id],
277        )
278        .await
279        .map_err(AppError::query)?;
280    }
281
282    tx.execute(
283        "INSERT INTO sessions (user_id, token_hash, user_agent, expires_at) VALUES ($1, $2, $3, $4)",
284        &[&user_id, &token_hash, &None::<String>, &expires_at],
285    )
286    .await
287    .map_err(AppError::query)?;
288
289    tx.commit().await.map_err(AppError::query)?;
290
291    let cookie = session::session_cookie(&token, 30 * 24 * 60 * 60, session::cookie_secure());
292    // 通过 Dioxus FullstackContext 设置 HttpOnly Cookie 响应头。
293    if let Some(ctx) = dioxus::fullstack::FullstackContext::current() {
294        if let Ok(value) = HeaderValue::try_from(cookie.as_str()) {
295            ctx.add_response_header(SET_COOKIE, value);
296        }
297    }
298
299    Ok(AuthResponse {
300        success: true,
301        message: "登录成功".to_string(),
302        token: None,
303    })
304}
305
306/// 用户登出。
307///
308/// 清空客户端 session Cookie,并删除数据库中对应会话记录。
309/// Dioxus server function,注册在 `/api` 路径下。
310#[server(Logout, "/api")]
311pub async fn logout() -> Result<AuthResponse, ServerFnError> {
312    let token = get_session_from_ctx();
313
314    let client = get_conn().await.map_err(AppError::db_conn)?;
315
316    // 设置过期时间为 0 的 Cookie,通知浏览器清除会话。
317    let cookie = session::session_cookie("", 0, session::cookie_secure());
318    if let Some(ctx) = dioxus::fullstack::FullstackContext::current() {
319        if let Ok(value) = HeaderValue::try_from(cookie.as_str()) {
320            ctx.add_response_header(SET_COOKIE, value);
321        }
322    }
323
324    if let Some(t) = token {
325        let token_hash = session::hash_token(&t);
326        crate::cache::invalidate_session_user(&token_hash).await;
327        client
328            .execute("DELETE FROM sessions WHERE token_hash = $1", &[&token_hash])
329            .await
330            .map_err(AppError::query)?;
331    }
332
333    Ok(AuthResponse {
334        success: true,
335        message: "登出成功".to_string(),
336        token: None,
337    })
338}
339
340#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
341/// 当前用户查询响应。
342pub struct CurrentUserResponse {
343    /// 当前已登录用户的公开信息;未登录时为 `None`。
344    pub user: Option<PublicUser>,
345}
346
347#[cfg(feature = "server")]
348/// 根据会话 token 查询对应用户(不含密码哈希,供会话缓存使用)。
349///
350/// 优先命中内存缓存,避免每次请求都执行 DB JOIN;未命中时回查数据库并回填缓存。
351/// 缓存命中后仍回查 `users.session_generation`:若用户已被降级/封禁(generation 被
352/// bump),缓存的旧 SessionUser.generation 不再匹配,此时逐出缓存并视为未登录,
353/// 消除权限残留窗口(见 H2)。仅服务端内部使用,不会暴露给前端。
354pub async fn get_user_by_token(token: &str) -> Result<Option<SessionUser>, ServerFnError> {
355    let token_hash = session::hash_token(token);
356
357    if let Some(cached) = crate::cache::get_session_user(&token_hash).await {
358        // 缓存命中后校验世代号:bump 后该用户所有 session 应失效。
359        // 查询走主键,亚毫秒级,代价可接受。
360        let current_gen: Option<i32> = get_conn()
361            .await
362            .map_err(AppError::db_conn)?
363            .query_opt(
364                "SELECT session_generation FROM users WHERE id = $1",
365                &[&cached.id],
366            )
367            .await
368            .map_err(AppError::query)?
369            .map(|r| r.get::<_, i32>(0));
370        match current_gen {
371            Some(gen) if gen == cached.session_generation => return Ok(Some(cached)),
372            _ => {
373                // 世代不匹配或用户已删:逐出缓存,落入下方重新查询。
374                crate::cache::invalidate_session_user(&token_hash).await;
375            }
376        }
377    }
378
379    let client = get_conn().await.map_err(AppError::db_conn)?;
380
381    let row = client
382        .query_opt(
383            "SELECT u.id, u.username, u.email, u.role, u.created_at, u.session_generation
384             FROM sessions s
385             JOIN users u ON s.user_id = u.id
386             WHERE s.token_hash = $1 AND s.expires_at > NOW()",
387            &[&token_hash],
388        )
389        .await
390        .map_err(AppError::query)?;
391
392    let user = match row {
393        Some(row) => {
394            let role_str: String = row.get("role");
395            let role = UserRole::from_str(&role_str).unwrap_or(UserRole::Blocked);
396            Some(SessionUser {
397                id: row.get("id"),
398                username: row.get("username"),
399                email: row.get("email"),
400                role,
401                created_at: row.get("created_at"),
402                session_generation: row.get("session_generation"),
403            })
404        }
405        None => None,
406    };
407
408    if let Some(ref u) = user {
409        crate::cache::set_session_user(&token_hash, u.clone()).await;
410    }
411
412    Ok(user)
413}
414
415/// 获取当前登录用户的公开信息。
416///
417/// Dioxus server function,注册在 `/api` 路径下。
418#[server(GetCurrentUser, "/api")]
419pub async fn get_current_user() -> Result<CurrentUserResponse, ServerFnError> {
420    let token = match get_session_from_ctx() {
421        Some(t) => t,
422        None => return Ok(CurrentUserResponse { user: None }),
423    };
424
425    let user = get_user_by_token(&token).await?.map(PublicUser::from);
426
427    Ok(CurrentUserResponse { user })
428}
429
430#[cfg(feature = "server")]
431/// 获取当前登录用户并要求其为 admin,否则返回 401/403。
432///
433/// 供其它服务端接口内部调用。
434pub async fn get_current_admin_user() -> Result<SessionUser, AppError> {
435    let token = get_session_from_ctx().ok_or(AppError::Unauthorized("未登录"))?;
436
437    let session_user = get_user_by_token(&token)
438        .await
439        .map_err(AppError::query)?
440        .ok_or(AppError::Unauthorized("会话已过期"))?;
441
442    if session_user.role != UserRole::Admin {
443        return Err(AppError::Forbidden("权限不足"));
444    }
445
446    Ok(session_user)
447}
448
449#[cfg(all(test, feature = "server"))]
450mod tests {
451    use super::*;
452
453    #[test]
454    fn validate_username_valid() {
455        assert!(validate_username("admin").is_ok());
456        assert!(validate_username("user_123").is_ok());
457        assert!(validate_username("abc").is_ok());
458    }
459
460    #[test]
461    fn validate_username_too_short() {
462        assert!(validate_username("ab").is_err());
463    }
464
465    #[test]
466    fn validate_username_too_long() {
467        assert!(validate_username(&"a".repeat(51)).is_err());
468    }
469
470    #[test]
471    fn validate_username_max_length() {
472        assert!(validate_username(&"a".repeat(50)).is_ok());
473    }
474
475    #[test]
476    fn validate_username_special_chars() {
477        assert!(validate_username("user name").is_err());
478        assert!(validate_username("user@name").is_err());
479        assert!(validate_username("user-name").is_err());
480    }
481
482    #[test]
483    fn validate_username_unicode() {
484        assert!(validate_username("用户名").is_ok());
485    }
486
487    #[test]
488    fn validate_email_valid() {
489        assert!(validate_email("user@example.com").is_ok());
490        assert!(validate_email("a.b+c@domain.co").is_ok());
491    }
492
493    #[test]
494    fn validate_email_invalid() {
495        assert!(validate_email("notanemail").is_err());
496        assert!(validate_email("@domain.com").is_err());
497        assert!(validate_email("user@").is_err());
498        assert!(validate_email("user@.com").is_err());
499        assert!(validate_email("").is_err());
500    }
501
502    #[test]
503    fn validate_password_valid() {
504        assert!(validate_password("12345678").is_ok());
505        assert!(validate_password("a very long password with spaces").is_ok());
506    }
507
508    #[test]
509    fn validate_password_too_short() {
510        assert!(validate_password("1234567").is_err());
511    }
512
513    #[test]
514    fn validate_password_exactly_8() {
515        assert!(validate_password("12345678").is_ok());
516    }
517
518    #[test]
519    fn validate_password_empty() {
520        assert!(validate_password("").is_err());
521    }
522}