1#![allow(clippy::unused_unit, deprecated)]
11
12use dioxus::prelude::*;
13#[cfg(feature = "server")]
14use http::header::{HeaderValue, SET_COOKIE};
15
16#[cfg(feature = "server")]
17use crate::api::error::AppError;
18#[cfg(feature = "server")]
19use crate::auth::session::get_session_from_ctx;
20#[cfg(feature = "server")]
21use crate::auth::{password, session};
22#[cfg(feature = "server")]
23use crate::db::pool::get_conn;
24use crate::models::user::PublicUser;
25#[cfg(feature = "server")]
26use crate::models::user::{SessionUser, UserRole};
27
28#[cfg(feature = "server")]
29fn validate_username(username: &str) -> Result<(), String> {
30 if username.len() < 3 || username.len() > 50 {
31 return Err("用户名长度必须在 3-50 字符之间".to_string());
32 }
33 if !username.chars().all(|c| c.is_alphanumeric() || c == '_') {
34 return Err("用户名只能包含字母、数字和下划线".to_string());
35 }
36 Ok(())
37}
38
39#[cfg(feature = "server")]
40pub(crate) fn validate_email(email: &str) -> Result<(), String> {
41 if !crate::utils::server::EMAIL_REGEX.is_match(email) {
42 return Err("邮箱格式不正确".to_string());
43 }
44 Ok(())
45}
46
47#[cfg(feature = "server")]
48const MAX_PASSWORD_BYTES: usize = 256;
51
52#[cfg(feature = "server")]
53pub(crate) fn validate_password(password: &str) -> Result<(), String> {
54 if password.len() < 8 {
55 return Err("密码长度至少 8 位".to_string());
56 }
57 if password.len() > MAX_PASSWORD_BYTES {
58 return Err("密码长度不能超过 256 位".to_string());
59 }
60 Ok(())
61}
62
63#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
64pub struct AuthResponse {
66 pub success: bool,
68 pub message: String,
70 pub token: Option<String>,
72}
73
74#[server(Register, "/api")]
80pub async fn register(
81 username: String,
82 email: String,
83 password: String,
84) -> Result<AuthResponse, ServerFnError> {
85 #[cfg(feature = "server")]
87 {
88 if let Some(ctx) = dioxus::fullstack::FullstackContext::current() {
89 let headers = ctx.parts_mut().headers.clone();
90 let ip = crate::api::rate_limit::get_client_ip(&headers).await;
91 if let Err(msg) = crate::api::rate_limit::check_strict_limit(&ip) {
92 return Ok(AuthResponse {
93 success: false,
94 message: msg,
95 token: None,
96 });
97 }
98 }
99 }
100
101 if let Err(e) = validate_username(&username) {
102 return Ok(AuthResponse {
103 success: false,
104 message: e,
105 token: None,
106 });
107 }
108 if let Err(e) = validate_email(&email) {
109 return Ok(AuthResponse {
110 success: false,
111 message: e,
112 token: None,
113 });
114 }
115 if let Err(e) = validate_password(&password) {
116 return Ok(AuthResponse {
117 success: false,
118 message: e,
119 token: None,
120 });
121 }
122
123 let client = get_conn().await.map_err(AppError::db_conn)?;
124
125 let pw_for_hash = password.clone();
127 let password_hash = tokio::task::spawn_blocking(move || password::hash_password(&pw_for_hash))
128 .await
129 .map_err(|_| AppError::Internal("密码处理任务失败"))?
130 .map_err(|_| AppError::Internal("密码处理失败"))?;
131
132 let result = client
135 .query_opt(
136 "INSERT INTO users (username, email, password_hash, role)
137 VALUES ($1, $2, $3, 'admin')
138 ON CONFLICT DO NOTHING
139 RETURNING id",
140 &[&username, &email, &password_hash],
141 )
142 .await
143 .map_err(AppError::query)?;
144
145 if result.is_some() {
146 return Ok(AuthResponse {
147 success: true,
148 message: "注册成功".to_string(),
149 token: None,
150 });
151 }
152
153 let admin_exists: bool = client
155 .query_one(
156 "SELECT EXISTS (SELECT 1 FROM users WHERE role = 'admin')",
157 &[],
158 )
159 .await
160 .map_err(AppError::query)?
161 .get(0);
162
163 let message = if admin_exists {
164 "Registration is closed".to_string()
165 } else {
166 "用户名或邮箱已存在".to_string()
167 };
168
169 Ok(AuthResponse {
170 success: false,
171 message,
172 token: None,
173 })
174}
175
176#[server(Login, "/api")]
182pub async fn login(username: String, password: String) -> Result<AuthResponse, ServerFnError> {
183 #[cfg(feature = "server")]
185 {
186 if let Some(ctx) = dioxus::fullstack::FullstackContext::current() {
187 let headers = ctx.parts_mut().headers.clone();
188 let ip = crate::api::rate_limit::get_client_ip(&headers).await;
189 if let Err(msg) = crate::api::rate_limit::check_strict_limit(&ip) {
190 return Ok(AuthResponse {
191 success: false,
192 message: msg,
193 token: None,
194 });
195 }
196 }
197 }
198
199 let mut client = get_conn().await.map_err(AppError::db_conn)?;
200
201 let row = match client
202 .query_opt(
203 "SELECT id, username, email, password_hash, role, created_at FROM users WHERE username = $1 OR email = $1",
204 &[&username],
205 )
206 .await
207 {
208 Ok(Some(row)) => row,
209 Ok(None) => {
210 const DUMMY_HASH: &str =
214 "$argon2id$v=19$m=19456,t=2,p=1$j3rNaAXzdExYaL94WBWtfg$n1S75LUQKaYJwaRl5bkFF/f/N1tLfRYR/7TuQxKP94c";
215 let dummy_pw = password.clone();
216 let _ = tokio::task::spawn_blocking(move || {
217 crate::auth::password::verify_password(&dummy_pw, DUMMY_HASH)
218 })
219 .await;
220 return Ok(AuthResponse {
221 success: false,
222 message: "Invalid credentials".to_string(),
223 token: None,
224 });
225 }
226 Err(e) => {
227 return Err(AppError::query(e).into());
228 }
229 };
230
231 let password_hash: String = row.get("password_hash");
232 let pw_for_verify = password.clone();
234 let hash_for_verify = password_hash.clone();
235 let valid = tokio::task::spawn_blocking(move || {
236 password::verify_password(&pw_for_verify, &hash_for_verify)
237 })
238 .await
239 .map_err(|_| AppError::Internal("密码处理任务失败"))?
240 .map_err(|_| AppError::Internal("密码处理失败"))?;
241
242 if !valid {
243 return Ok(AuthResponse {
244 success: false,
245 message: "Invalid credentials".to_string(),
246 token: None,
247 });
248 }
249
250 let role_str: String = row.get("role");
253 if !matches!(UserRole::from_str(&role_str), Some(UserRole::Admin)) {
254 return Ok(AuthResponse {
255 success: false,
256 message: "Invalid credentials".to_string(),
257 token: None,
258 });
259 }
260
261 let user_id: i32 = row.get("id");
262 let token = session::generate_token();
263 let token_hash = session::hash_token(&token);
264 let expires_at = session::default_expiry();
265
266 let max_sessions = crate::api::settings::runtime_security_settings()
267 .await
268 .max_sessions_per_user
269 .max(1) as i64;
270
271 let tx = client.transaction().await.map_err(AppError::query)?;
274 tx.execute("SELECT 1 FROM users WHERE id = $1 FOR UPDATE", &[&user_id])
276 .await
277 .map_err(AppError::query)?;
278
279 let session_count: i64 = tx
280 .query_one(
281 "SELECT COUNT(*) FROM sessions WHERE user_id = $1 AND expires_at > NOW()",
282 &[&user_id],
283 )
284 .await
285 .map_err(AppError::query)?
286 .get(0);
287
288 if session_count >= max_sessions {
289 tx.execute(
290 "DELETE FROM sessions WHERE id IN (
291 SELECT id FROM sessions
292 WHERE user_id = $1 AND expires_at > NOW()
293 ORDER BY created_at ASC
294 LIMIT 1
295 )",
296 &[&user_id],
297 )
298 .await
299 .map_err(AppError::query)?;
300 }
301
302 tx.execute(
303 "INSERT INTO sessions (user_id, token_hash, user_agent, expires_at) VALUES ($1, $2, $3, $4)",
304 &[&user_id, &token_hash, &None::<String>, &expires_at],
305 )
306 .await
307 .map_err(AppError::query)?;
308
309 tx.commit().await.map_err(AppError::query)?;
310
311 let cookie = session::session_cookie(
312 &token,
313 session::SESSION_MAX_AGE_SECS,
314 crate::api::settings::runtime_security_settings()
315 .await
316 .cookie_secure,
317 );
318 if let Some(ctx) = dioxus::fullstack::FullstackContext::current() {
320 if let Ok(value) = HeaderValue::try_from(cookie.as_str()) {
321 ctx.add_response_header(SET_COOKIE, value);
322 }
323 }
324
325 Ok(AuthResponse {
326 success: true,
327 message: "登录成功".to_string(),
328 token: None,
329 })
330}
331
332#[server(Logout, "/api")]
337pub async fn logout() -> Result<AuthResponse, ServerFnError> {
338 let token = get_session_from_ctx();
339
340 let client = get_conn().await.map_err(AppError::db_conn)?;
341
342 let cookie = session::session_cookie(
344 "",
345 0,
346 crate::api::settings::runtime_security_settings()
347 .await
348 .cookie_secure,
349 );
350 if let Some(ctx) = dioxus::fullstack::FullstackContext::current() {
351 if let Ok(value) = HeaderValue::try_from(cookie.as_str()) {
352 ctx.add_response_header(SET_COOKIE, value);
353 }
354 }
355
356 if let Some(t) = token {
357 let token_hash = session::hash_token(&t);
358 crate::cache::invalidate_session_user(&token_hash).await;
359 client
360 .execute("DELETE FROM sessions WHERE token_hash = $1", &[&token_hash])
361 .await
362 .map_err(AppError::query)?;
363 }
364
365 Ok(AuthResponse {
366 success: true,
367 message: "登出成功".to_string(),
368 token: None,
369 })
370}
371
372#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
373pub struct CurrentUserResponse {
375 pub user: Option<PublicUser>,
377}
378
379#[cfg(feature = "server")]
380pub async fn get_user_by_token(token: &str) -> Result<Option<SessionUser>, ServerFnError> {
386 let token_hash = session::hash_token(token);
387
388 if let Some(cached) = crate::cache::get_session_user(&token_hash).await {
389 let session_state: Option<(i32, String)> = get_conn()
391 .await
392 .map_err(AppError::db_conn)?
393 .query_opt(
394 "SELECT u.session_generation, u.role
395 FROM sessions s
396 JOIN users u ON s.user_id = u.id
397 WHERE s.token_hash = $1 AND s.expires_at > NOW()",
398 &[&token_hash],
399 )
400 .await
401 .map_err(AppError::query)?
402 .map(|r| (r.get(0), r.get(1)));
403
404 match session_state {
405 Some((generation, role))
406 if generation == cached.session_generation
407 && role == "admin"
408 && matches!(cached.role, UserRole::Admin) =>
409 {
410 return Ok(Some(cached));
411 }
412 _ => {
413 crate::cache::invalidate_session_user(&token_hash).await;
415 return Ok(None);
416 }
417 }
418 }
419
420 let client = get_conn().await.map_err(AppError::db_conn)?;
421
422 let row = client
423 .query_opt(
424 "SELECT u.id, u.username, u.email, u.display_name, u.avatar_url, u.role, u.created_at, u.session_generation
425 FROM sessions s
426 JOIN users u ON s.user_id = u.id
427 WHERE s.token_hash = $1 AND s.expires_at > NOW()",
428 &[&token_hash],
429 )
430 .await
431 .map_err(AppError::query)?;
432
433 let user = match row {
434 Some(row) => {
435 let role_str: String = row.get("role");
436 let role = match UserRole::from_str(&role_str) {
437 Some(UserRole::Admin) => UserRole::Admin,
438 _ => return Ok(None),
439 };
440 Some(SessionUser {
441 id: row.get("id"),
442 username: row.get("username"),
443 email: row.get("email"),
444 display_name: row.get("display_name"),
445 avatar_url: row.get("avatar_url"),
446 role,
447 created_at: row.get("created_at"),
448 session_generation: row.get("session_generation"),
449 })
450 }
451 None => None,
452 };
453
454 if let Some(ref u) = user {
455 crate::cache::set_session_user(&token_hash, u.clone()).await;
456 }
457
458 Ok(user)
459}
460
461#[server(GetCurrentUser, "/api")]
465pub async fn get_current_user() -> Result<CurrentUserResponse, ServerFnError> {
466 let token = match get_session_from_ctx() {
467 Some(t) => t,
468 None => return Ok(CurrentUserResponse { user: None }),
469 };
470
471 let user = get_user_by_token(&token).await?.map(PublicUser::from);
472
473 Ok(CurrentUserResponse { user })
474}
475
476#[cfg(feature = "server")]
477pub async fn get_current_admin_user() -> Result<SessionUser, AppError> {
481 let token = get_session_from_ctx().ok_or(AppError::Unauthorized("未登录"))?;
482
483 let session_user = get_user_by_token(&token)
484 .await
485 .map_err(AppError::query)?
486 .ok_or(AppError::Unauthorized("会话已过期"))?;
487
488 if session_user.role != UserRole::Admin {
489 return Err(AppError::Forbidden("权限不足"));
490 }
491
492 Ok(session_user)
493}
494
495#[cfg(feature = "server")]
499struct AdminEnvConfig {
500 username: String,
501 email: String,
502 password: String,
503}
504
505#[cfg(feature = "server")]
510fn parse_admin_env() -> Option<AdminEnvConfig> {
511 let username = std::env::var("ADMIN_USERNAME").ok()?;
512 let username = username.trim().to_string();
513 if username.is_empty() {
514 return None;
515 }
516 let email = std::env::var("ADMIN_EMAIL").ok()?;
517 let email = email.trim().to_string();
518 if email.is_empty() {
519 return None;
520 }
521 let password = std::env::var("ADMIN_PASSWORD").ok()?;
522 if password.is_empty() {
523 return None;
524 }
525 Some(AdminEnvConfig {
526 username,
527 email,
528 password,
529 })
530}
531
532#[cfg(feature = "server")]
536pub(crate) fn admin_env_active() -> bool {
537 parse_admin_env().is_some()
538}
539
540#[cfg(feature = "server")]
555pub(crate) async fn sync_admin_from_env(client: &tokio_postgres::Client) -> Result<(), AppError> {
556 let Some(cfg) = parse_admin_env() else {
557 return Ok(());
558 };
559 if let Err(e) = validate_username(&cfg.username) {
560 tracing::warn!(
561 "ADMIN_USERNAME={:?} 非法,跳过初始管理员同步: {e}",
562 cfg.username
563 );
564 return Ok(());
565 }
566 if let Err(e) = validate_email(&cfg.email) {
567 tracing::warn!("ADMIN_EMAIL={:?} 非法,跳过初始管理员同步: {e}", cfg.email);
568 return Ok(());
569 }
570 if let Err(e) = validate_password(&cfg.password) {
571 tracing::warn!("ADMIN_PASSWORD 不合法({e}),跳过初始管理员同步");
572 return Ok(());
573 }
574
575 let pw = cfg.password.clone();
577 let password_hash = tokio::task::spawn_blocking(move || password::hash_password(&pw))
578 .await
579 .map_err(|_| AppError::Internal("初始管理员密码哈希任务失败"))?
580 .map_err(|_| AppError::Internal("初始管理员密码哈希失败"))?;
581
582 let row = client
583 .query_opt(
584 "SELECT id, email, role, password_hash FROM users WHERE username = $1",
585 &[&cfg.username],
586 )
587 .await
588 .map_err(AppError::query)?;
589
590 let Some(row) = row else {
591 let created = client
593 .query_opt(
594 "INSERT INTO users (username, email, password_hash, role)
595 VALUES ($1, $2, $3, 'admin')
596 ON CONFLICT DO NOTHING
597 RETURNING id",
598 &[&cfg.username, &cfg.email, &password_hash],
599 )
600 .await
601 .map_err(AppError::query)?;
602 if created.is_some() {
603 tracing::info!(
604 "初始管理员已从环境变量创建: username={}(登录成功后可从 .env 移除 ADMIN_* 变量)",
605 cfg.username
606 );
607 } else {
608 tracing::warn!(
609 "无法创建初始管理员 username={}(用户名/邮箱已被占用,或库中已存在其他 admin——全库仅允许一个 admin)。\
610 可将 ADMIN_USERNAME 改为现有 admin 的用户名以同步其凭据",
611 cfg.username
612 );
613 }
614 return Ok(());
615 };
616
617 let user_id: i32 = row.get("id");
619 let stored_hash: String = row.get("password_hash");
620 let stored_email: String = row.get("email");
621 let role: String = row.get("role");
622
623 let pw_for_verify = cfg.password.clone();
625 let same_password = tokio::task::spawn_blocking(move || {
626 password::verify_password(&pw_for_verify, &stored_hash)
627 })
628 .await
629 .map_err(|_| AppError::Internal("初始管理员密码校验任务失败"))?
630 .unwrap_or(false);
631 if !same_password {
632 client
634 .execute(
635 "UPDATE users SET password_hash = $1, session_generation = session_generation + 1
636 WHERE id = $2",
637 &[&password_hash, &user_id],
638 )
639 .await
640 .map_err(AppError::query)?;
641 tracing::info!("初始管理员密码已从环境变量覆盖: username={}", cfg.username);
642 }
643
644 if role != "admin" {
645 client
647 .execute("UPDATE users SET role = 'admin' WHERE id = $1", &[&user_id])
648 .await
649 .map_err(AppError::query)?;
650 tracing::info!("初始管理员角色已恢复为 admin: username={}", cfg.username);
651 }
652
653 if stored_email != cfg.email {
654 let email_taken: bool = client
655 .query_one(
656 "SELECT EXISTS (SELECT 1 FROM users WHERE email = $1 AND id <> $2)",
657 &[&cfg.email, &user_id],
658 )
659 .await
660 .map_err(AppError::query)?
661 .get(0);
662 if email_taken {
663 tracing::warn!(
664 "ADMIN_EMAIL={:?} 已被其他用户占用,跳过邮箱同步(密码/角色不受影响)",
665 cfg.email
666 );
667 } else {
668 client
669 .execute(
670 "UPDATE users SET email = $1 WHERE id = $2",
671 &[&cfg.email, &user_id],
672 )
673 .await
674 .map_err(AppError::query)?;
675 }
676 }
677
678 Ok(())
679}
680
681#[cfg(all(test, feature = "server"))]
682mod tests {
683 use super::*;
684
685 #[test]
686 fn validate_username_valid() {
687 assert!(validate_username("admin").is_ok());
688 assert!(validate_username("user_123").is_ok());
689 assert!(validate_username("abc").is_ok());
690 }
691
692 #[test]
693 fn validate_username_too_short() {
694 assert!(validate_username("ab").is_err());
695 }
696
697 #[test]
698 fn validate_username_too_long() {
699 assert!(validate_username(&"a".repeat(51)).is_err());
700 }
701
702 #[test]
703 fn validate_username_max_length() {
704 assert!(validate_username(&"a".repeat(50)).is_ok());
705 }
706
707 #[test]
708 fn validate_username_special_chars() {
709 assert!(validate_username("user name").is_err());
710 assert!(validate_username("user@name").is_err());
711 assert!(validate_username("user-name").is_err());
712 }
713
714 #[test]
715 fn validate_username_unicode() {
716 assert!(validate_username("用户名").is_ok());
717 }
718
719 #[test]
720 fn validate_email_valid() {
721 assert!(validate_email("user@example.com").is_ok());
722 assert!(validate_email("a.b+c@domain.co").is_ok());
723 }
724
725 #[test]
726 fn validate_email_invalid() {
727 assert!(validate_email("notanemail").is_err());
728 assert!(validate_email("@domain.com").is_err());
729 assert!(validate_email("user@").is_err());
730 assert!(validate_email("user@.com").is_err());
731 assert!(validate_email("").is_err());
732 }
733
734 #[test]
735 fn validate_password_valid() {
736 assert!(validate_password("12345678").is_ok());
737 assert!(validate_password("a very long password with spaces").is_ok());
738 }
739
740 #[test]
741 fn validate_password_too_short() {
742 assert!(validate_password("1234567").is_err());
743 }
744
745 #[test]
746 fn validate_password_exactly_8() {
747 assert!(validate_password("12345678").is_ok());
748 }
749
750 #[test]
751 fn validate_password_empty() {
752 assert!(validate_password("").is_err());
753 }
754
755 #[test]
756 fn validate_password_too_long() {
757 assert!(validate_password(&"a".repeat(257)).is_err());
758 assert!(validate_password(&"a".repeat(256)).is_ok());
759 }
760
761 fn with_env(vars: &[(&str, Option<&str>)], f: impl FnOnce()) {
763 let saved: Vec<(String, Option<String>)> = vars
764 .iter()
765 .map(|(k, _)| (k.to_string(), std::env::var(k).ok()))
766 .collect();
767 for (k, v) in vars {
768 match v {
769 Some(v) => std::env::set_var(k, v),
770 None => std::env::remove_var(k),
771 }
772 }
773 f();
774 for (k, v) in saved {
775 match v {
776 Some(v) => std::env::set_var(k, v),
777 None => std::env::remove_var(k),
778 }
779 }
780 }
781
782 #[test]
783 #[serial_test::serial]
784 fn parse_admin_env_none_when_all_unset() {
785 with_env(
786 &[
787 ("ADMIN_USERNAME", None),
788 ("ADMIN_EMAIL", None),
789 ("ADMIN_PASSWORD", None),
790 ],
791 || assert!(parse_admin_env().is_none()),
792 );
793 }
794
795 #[test]
796 #[serial_test::serial]
797 fn parse_admin_env_none_when_any_missing() {
798 with_env(
800 &[
801 ("ADMIN_USERNAME", Some("admin")),
802 ("ADMIN_EMAIL", Some("admin@example.com")),
803 ("ADMIN_PASSWORD", None),
804 ],
805 || assert!(parse_admin_env().is_none()),
806 );
807 with_env(
809 &[
810 ("ADMIN_USERNAME", Some("admin")),
811 ("ADMIN_EMAIL", None),
812 ("ADMIN_PASSWORD", Some("strongpass123")),
813 ],
814 || assert!(parse_admin_env().is_none()),
815 );
816 with_env(
818 &[
819 ("ADMIN_USERNAME", None),
820 ("ADMIN_EMAIL", Some("admin@example.com")),
821 ("ADMIN_PASSWORD", Some("strongpass123")),
822 ],
823 || assert!(parse_admin_env().is_none()),
824 );
825 }
826
827 #[test]
828 #[serial_test::serial]
829 fn parse_admin_env_none_on_empty_values() {
830 with_env(
831 &[
832 ("ADMIN_USERNAME", Some("admin")),
833 ("ADMIN_EMAIL", Some("admin@example.com")),
834 ("ADMIN_PASSWORD", Some("")),
835 ],
836 || assert!(parse_admin_env().is_none()),
837 );
838 with_env(
839 &[
840 ("ADMIN_USERNAME", Some(" ")),
841 ("ADMIN_EMAIL", Some("admin@example.com")),
842 ("ADMIN_PASSWORD", Some("strongpass123")),
843 ],
844 || assert!(parse_admin_env().is_none()),
845 );
846 }
847
848 #[test]
849 #[serial_test::serial]
850 fn parse_admin_env_some_when_all_set() {
851 with_env(
852 &[
853 ("ADMIN_USERNAME", Some(" admin ")),
854 ("ADMIN_EMAIL", Some(" admin@example.com ")),
855 ("ADMIN_PASSWORD", Some(" strong pass 123 ")),
856 ],
857 || {
858 let cfg = parse_admin_env().expect("三个变量齐全应返回 Some");
859 assert_eq!(cfg.username, "admin");
861 assert_eq!(cfg.email, "admin@example.com");
862 assert_eq!(cfg.password, " strong pass 123 ");
863 },
864 );
865 }
866}