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")]
48pub(crate) fn validate_password(password: &str) -> Result<(), String> {
49 if password.len() < 8 {
50 return Err("密码长度至少 8 位".to_string());
51 }
52 Ok(())
53}
54
55#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
56pub struct AuthResponse {
58 pub success: bool,
60 pub message: String,
62 pub token: Option<String>,
64}
65
66#[server(Register, "/api")]
72pub async fn register(
73 username: String,
74 email: String,
75 password: String,
76) -> Result<AuthResponse, ServerFnError> {
77 #[cfg(feature = "server")]
79 {
80 if let Some(ctx) = dioxus::fullstack::FullstackContext::current() {
81 let headers = ctx.parts_mut().headers.clone();
82 let ip = crate::api::rate_limit::get_client_ip(&headers).await;
83 if let Err(msg) = crate::api::rate_limit::check_strict_limit(&ip) {
84 return Ok(AuthResponse {
85 success: false,
86 message: msg,
87 token: None,
88 });
89 }
90 }
91 }
92
93 if let Err(e) = validate_username(&username) {
94 return Ok(AuthResponse {
95 success: false,
96 message: e,
97 token: None,
98 });
99 }
100 if let Err(e) = validate_email(&email) {
101 return Ok(AuthResponse {
102 success: false,
103 message: e,
104 token: None,
105 });
106 }
107 if let Err(e) = validate_password(&password) {
108 return Ok(AuthResponse {
109 success: false,
110 message: e,
111 token: None,
112 });
113 }
114
115 let client = get_conn().await.map_err(AppError::db_conn)?;
116
117 let pw_for_hash = password.clone();
119 let password_hash = tokio::task::spawn_blocking(move || password::hash_password(&pw_for_hash))
120 .await
121 .map_err(|_| AppError::Internal("密码处理任务失败"))?
122 .map_err(|_| AppError::Internal("密码处理失败"))?;
123
124 let result = client
127 .query_opt(
128 "INSERT INTO users (username, email, password_hash, role)
129 VALUES ($1, $2, $3, 'admin')
130 ON CONFLICT DO NOTHING
131 RETURNING id",
132 &[&username, &email, &password_hash],
133 )
134 .await
135 .map_err(AppError::query)?;
136
137 if result.is_some() {
138 return Ok(AuthResponse {
139 success: true,
140 message: "注册成功".to_string(),
141 token: None,
142 });
143 }
144
145 let admin_exists: bool = client
147 .query_one(
148 "SELECT EXISTS (SELECT 1 FROM users WHERE role = 'admin')",
149 &[],
150 )
151 .await
152 .map_err(AppError::query)?
153 .get(0);
154
155 let message = if admin_exists {
156 "Registration is closed".to_string()
157 } else {
158 "用户名或邮箱已存在".to_string()
159 };
160
161 Ok(AuthResponse {
162 success: false,
163 message,
164 token: None,
165 })
166}
167
168#[server(Login, "/api")]
174pub async fn login(username: String, password: String) -> Result<AuthResponse, ServerFnError> {
175 #[cfg(feature = "server")]
177 {
178 if let Some(ctx) = dioxus::fullstack::FullstackContext::current() {
179 let headers = ctx.parts_mut().headers.clone();
180 let ip = crate::api::rate_limit::get_client_ip(&headers).await;
181 if let Err(msg) = crate::api::rate_limit::check_strict_limit(&ip) {
182 return Ok(AuthResponse {
183 success: false,
184 message: msg,
185 token: None,
186 });
187 }
188 }
189 }
190
191 let mut client = get_conn().await.map_err(AppError::db_conn)?;
192
193 let row = match client
194 .query_opt(
195 "SELECT id, username, email, password_hash, role, created_at FROM users WHERE username = $1 OR email = $1",
196 &[&username],
197 )
198 .await
199 {
200 Ok(Some(row)) => row,
201 Ok(None) => {
202 const DUMMY_HASH: &str =
206 "$argon2id$v=19$m=19456,t=2,p=1$j3rNaAXzdExYaL94WBWtfg$n1S75LUQKaYJwaRl5bkFF/f/N1tLfRYR/7TuQxKP94c";
207 let dummy_pw = password.clone();
208 let _ = tokio::task::spawn_blocking(move || {
209 crate::auth::password::verify_password(&dummy_pw, DUMMY_HASH)
210 })
211 .await;
212 return Ok(AuthResponse {
213 success: false,
214 message: "Invalid credentials".to_string(),
215 token: None,
216 });
217 }
218 Err(e) => {
219 return Err(AppError::query(e).into());
220 }
221 };
222
223 let password_hash: String = row.get("password_hash");
224 let pw_for_verify = password.clone();
226 let hash_for_verify = password_hash.clone();
227 let valid = tokio::task::spawn_blocking(move || {
228 password::verify_password(&pw_for_verify, &hash_for_verify)
229 })
230 .await
231 .map_err(|_| AppError::Internal("密码处理任务失败"))?
232 .map_err(|_| AppError::Internal("密码处理失败"))?;
233
234 if !valid {
235 return Ok(AuthResponse {
236 success: false,
237 message: "Invalid credentials".to_string(),
238 token: None,
239 });
240 }
241
242 let user_id: i32 = row.get("id");
243 let token = session::generate_token();
244 let token_hash = session::hash_token(&token);
245 let expires_at = session::default_expiry();
246
247 let max_sessions = crate::api::settings::runtime_security_settings()
248 .await
249 .max_sessions_per_user
250 .max(1) as i64;
251
252 let tx = client.transaction().await.map_err(AppError::query)?;
255 tx.execute("SELECT 1 FROM users WHERE id = $1 FOR UPDATE", &[&user_id])
257 .await
258 .map_err(AppError::query)?;
259
260 let session_count: i64 = tx
261 .query_one(
262 "SELECT COUNT(*) FROM sessions WHERE user_id = $1 AND expires_at > NOW()",
263 &[&user_id],
264 )
265 .await
266 .map_err(AppError::query)?
267 .get(0);
268
269 if session_count >= max_sessions {
270 tx.execute(
271 "DELETE FROM sessions WHERE id IN (
272 SELECT id FROM sessions
273 WHERE user_id = $1 AND expires_at > NOW()
274 ORDER BY created_at ASC
275 LIMIT 1
276 )",
277 &[&user_id],
278 )
279 .await
280 .map_err(AppError::query)?;
281 }
282
283 tx.execute(
284 "INSERT INTO sessions (user_id, token_hash, user_agent, expires_at) VALUES ($1, $2, $3, $4)",
285 &[&user_id, &token_hash, &None::<String>, &expires_at],
286 )
287 .await
288 .map_err(AppError::query)?;
289
290 tx.commit().await.map_err(AppError::query)?;
291
292 let cookie = session::session_cookie(&token, 30 * 24 * 60 * 60, session::cookie_secure().await);
293 if let Some(ctx) = dioxus::fullstack::FullstackContext::current() {
295 if let Ok(value) = HeaderValue::try_from(cookie.as_str()) {
296 ctx.add_response_header(SET_COOKIE, value);
297 }
298 }
299
300 Ok(AuthResponse {
301 success: true,
302 message: "登录成功".to_string(),
303 token: None,
304 })
305}
306
307#[server(Logout, "/api")]
312pub async fn logout() -> Result<AuthResponse, ServerFnError> {
313 let token = get_session_from_ctx();
314
315 let client = get_conn().await.map_err(AppError::db_conn)?;
316
317 let cookie = session::session_cookie("", 0, session::cookie_secure().await);
319 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 if let Some(t) = token {
326 let token_hash = session::hash_token(&t);
327 crate::cache::invalidate_session_user(&token_hash).await;
328 client
329 .execute("DELETE FROM sessions WHERE token_hash = $1", &[&token_hash])
330 .await
331 .map_err(AppError::query)?;
332 }
333
334 Ok(AuthResponse {
335 success: true,
336 message: "登出成功".to_string(),
337 token: None,
338 })
339}
340
341#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
342pub struct CurrentUserResponse {
344 pub user: Option<PublicUser>,
346}
347
348#[cfg(feature = "server")]
349pub async fn get_user_by_token(token: &str) -> Result<Option<SessionUser>, ServerFnError> {
356 let token_hash = session::hash_token(token);
357
358 if let Some(cached) = crate::cache::get_session_user(&token_hash).await {
359 let current_gen: Option<i32> = get_conn()
362 .await
363 .map_err(AppError::db_conn)?
364 .query_opt(
365 "SELECT session_generation FROM users WHERE id = $1",
366 &[&cached.id],
367 )
368 .await
369 .map_err(AppError::query)?
370 .map(|r| r.get::<_, i32>(0));
371 match current_gen {
372 Some(gen) if gen == cached.session_generation => return Ok(Some(cached)),
373 _ => {
374 crate::cache::invalidate_session_user(&token_hash).await;
376 }
377 }
378 }
379
380 let client = get_conn().await.map_err(AppError::db_conn)?;
381
382 let row = client
383 .query_opt(
384 "SELECT u.id, u.username, u.email, u.display_name, u.avatar_url, u.role, u.created_at, u.session_generation
385 FROM sessions s
386 JOIN users u ON s.user_id = u.id
387 WHERE s.token_hash = $1 AND s.expires_at > NOW()",
388 &[&token_hash],
389 )
390 .await
391 .map_err(AppError::query)?;
392
393 let user = match row {
394 Some(row) => {
395 let role_str: String = row.get("role");
396 let role = UserRole::from_str(&role_str).unwrap_or(UserRole::Blocked);
397 Some(SessionUser {
398 id: row.get("id"),
399 username: row.get("username"),
400 email: row.get("email"),
401 display_name: row.get("display_name"),
402 avatar_url: row.get("avatar_url"),
403 role,
404 created_at: row.get("created_at"),
405 session_generation: row.get("session_generation"),
406 })
407 }
408 None => None,
409 };
410
411 if let Some(ref u) = user {
412 crate::cache::set_session_user(&token_hash, u.clone()).await;
413 }
414
415 Ok(user)
416}
417
418#[server(GetCurrentUser, "/api")]
422pub async fn get_current_user() -> Result<CurrentUserResponse, ServerFnError> {
423 let token = match get_session_from_ctx() {
424 Some(t) => t,
425 None => return Ok(CurrentUserResponse { user: None }),
426 };
427
428 let user = get_user_by_token(&token).await?.map(PublicUser::from);
429
430 Ok(CurrentUserResponse { user })
431}
432
433#[cfg(feature = "server")]
434pub async fn get_current_admin_user() -> Result<SessionUser, AppError> {
438 let token = get_session_from_ctx().ok_or(AppError::Unauthorized("未登录"))?;
439
440 let session_user = get_user_by_token(&token)
441 .await
442 .map_err(AppError::query)?
443 .ok_or(AppError::Unauthorized("会话已过期"))?;
444
445 if session_user.role != UserRole::Admin {
446 return Err(AppError::Forbidden("权限不足"));
447 }
448
449 Ok(session_user)
450}
451
452#[cfg(feature = "server")]
456struct AdminEnvConfig {
457 username: String,
458 email: String,
459 password: String,
460}
461
462#[cfg(feature = "server")]
467fn parse_admin_env() -> Option<AdminEnvConfig> {
468 let username = std::env::var("ADMIN_USERNAME").ok()?;
469 let username = username.trim().to_string();
470 if username.is_empty() {
471 return None;
472 }
473 let email = std::env::var("ADMIN_EMAIL").ok()?;
474 let email = email.trim().to_string();
475 if email.is_empty() {
476 return None;
477 }
478 let password = std::env::var("ADMIN_PASSWORD").ok()?;
479 if password.is_empty() {
480 return None;
481 }
482 Some(AdminEnvConfig {
483 username,
484 email,
485 password,
486 })
487}
488
489#[cfg(feature = "server")]
493pub(crate) fn admin_env_active() -> bool {
494 parse_admin_env().is_some()
495}
496
497#[cfg(feature = "server")]
512pub(crate) async fn sync_admin_from_env(client: &tokio_postgres::Client) -> Result<(), AppError> {
513 let Some(cfg) = parse_admin_env() else {
514 return Ok(());
515 };
516 if let Err(e) = validate_username(&cfg.username) {
517 tracing::warn!(
518 "ADMIN_USERNAME={:?} 非法,跳过初始管理员同步: {e}",
519 cfg.username
520 );
521 return Ok(());
522 }
523 if let Err(e) = validate_email(&cfg.email) {
524 tracing::warn!("ADMIN_EMAIL={:?} 非法,跳过初始管理员同步: {e}", cfg.email);
525 return Ok(());
526 }
527 if let Err(e) = validate_password(&cfg.password) {
528 tracing::warn!("ADMIN_PASSWORD 不合法({e}),跳过初始管理员同步");
529 return Ok(());
530 }
531
532 let pw = cfg.password.clone();
534 let password_hash = tokio::task::spawn_blocking(move || password::hash_password(&pw))
535 .await
536 .map_err(|_| AppError::Internal("初始管理员密码哈希任务失败"))?
537 .map_err(|_| AppError::Internal("初始管理员密码哈希失败"))?;
538
539 let row = client
540 .query_opt(
541 "SELECT id, email, role, password_hash FROM users WHERE username = $1",
542 &[&cfg.username],
543 )
544 .await
545 .map_err(AppError::query)?;
546
547 let Some(row) = row else {
548 let created = client
550 .query_opt(
551 "INSERT INTO users (username, email, password_hash, role)
552 VALUES ($1, $2, $3, 'admin')
553 ON CONFLICT DO NOTHING
554 RETURNING id",
555 &[&cfg.username, &cfg.email, &password_hash],
556 )
557 .await
558 .map_err(AppError::query)?;
559 if created.is_some() {
560 tracing::info!(
561 "初始管理员已从环境变量创建: username={}(登录成功后可从 .env 移除 ADMIN_* 变量)",
562 cfg.username
563 );
564 } else {
565 tracing::warn!(
566 "无法创建初始管理员 username={}(用户名/邮箱已被占用,或库中已存在其他 admin——全库仅允许一个 admin)。\
567 可将 ADMIN_USERNAME 改为现有 admin 的用户名以同步其凭据",
568 cfg.username
569 );
570 }
571 return Ok(());
572 };
573
574 let user_id: i32 = row.get("id");
576 let stored_hash: String = row.get("password_hash");
577 let stored_email: String = row.get("email");
578 let role: String = row.get("role");
579
580 let pw_for_verify = cfg.password.clone();
582 let same_password = tokio::task::spawn_blocking(move || {
583 password::verify_password(&pw_for_verify, &stored_hash)
584 })
585 .await
586 .map_err(|_| AppError::Internal("初始管理员密码校验任务失败"))?
587 .unwrap_or(false);
588 if !same_password {
589 client
591 .execute(
592 "UPDATE users SET password_hash = $1, session_generation = session_generation + 1
593 WHERE id = $2",
594 &[&password_hash, &user_id],
595 )
596 .await
597 .map_err(AppError::query)?;
598 tracing::info!("初始管理员密码已从环境变量覆盖: username={}", cfg.username);
599 }
600
601 if role != "admin" {
602 client
604 .execute("UPDATE users SET role = 'admin' WHERE id = $1", &[&user_id])
605 .await
606 .map_err(AppError::query)?;
607 tracing::info!("初始管理员角色已恢复为 admin: username={}", cfg.username);
608 }
609
610 if stored_email != cfg.email {
611 let email_taken: bool = client
612 .query_one(
613 "SELECT EXISTS (SELECT 1 FROM users WHERE email = $1 AND id <> $2)",
614 &[&cfg.email, &user_id],
615 )
616 .await
617 .map_err(AppError::query)?
618 .get(0);
619 if email_taken {
620 tracing::warn!(
621 "ADMIN_EMAIL={:?} 已被其他用户占用,跳过邮箱同步(密码/角色不受影响)",
622 cfg.email
623 );
624 } else {
625 client
626 .execute(
627 "UPDATE users SET email = $1 WHERE id = $2",
628 &[&cfg.email, &user_id],
629 )
630 .await
631 .map_err(AppError::query)?;
632 }
633 }
634
635 Ok(())
636}
637
638#[cfg(all(test, feature = "server"))]
639mod tests {
640 use super::*;
641
642 #[test]
643 fn validate_username_valid() {
644 assert!(validate_username("admin").is_ok());
645 assert!(validate_username("user_123").is_ok());
646 assert!(validate_username("abc").is_ok());
647 }
648
649 #[test]
650 fn validate_username_too_short() {
651 assert!(validate_username("ab").is_err());
652 }
653
654 #[test]
655 fn validate_username_too_long() {
656 assert!(validate_username(&"a".repeat(51)).is_err());
657 }
658
659 #[test]
660 fn validate_username_max_length() {
661 assert!(validate_username(&"a".repeat(50)).is_ok());
662 }
663
664 #[test]
665 fn validate_username_special_chars() {
666 assert!(validate_username("user name").is_err());
667 assert!(validate_username("user@name").is_err());
668 assert!(validate_username("user-name").is_err());
669 }
670
671 #[test]
672 fn validate_username_unicode() {
673 assert!(validate_username("用户名").is_ok());
674 }
675
676 #[test]
677 fn validate_email_valid() {
678 assert!(validate_email("user@example.com").is_ok());
679 assert!(validate_email("a.b+c@domain.co").is_ok());
680 }
681
682 #[test]
683 fn validate_email_invalid() {
684 assert!(validate_email("notanemail").is_err());
685 assert!(validate_email("@domain.com").is_err());
686 assert!(validate_email("user@").is_err());
687 assert!(validate_email("user@.com").is_err());
688 assert!(validate_email("").is_err());
689 }
690
691 #[test]
692 fn validate_password_valid() {
693 assert!(validate_password("12345678").is_ok());
694 assert!(validate_password("a very long password with spaces").is_ok());
695 }
696
697 #[test]
698 fn validate_password_too_short() {
699 assert!(validate_password("1234567").is_err());
700 }
701
702 #[test]
703 fn validate_password_exactly_8() {
704 assert!(validate_password("12345678").is_ok());
705 }
706
707 #[test]
708 fn validate_password_empty() {
709 assert!(validate_password("").is_err());
710 }
711
712 fn with_env(vars: &[(&str, Option<&str>)], f: impl FnOnce()) {
714 let saved: Vec<(String, Option<String>)> = vars
715 .iter()
716 .map(|(k, _)| (k.to_string(), std::env::var(k).ok()))
717 .collect();
718 for (k, v) in vars {
719 match v {
720 Some(v) => std::env::set_var(k, v),
721 None => std::env::remove_var(k),
722 }
723 }
724 f();
725 for (k, v) in saved {
726 match v {
727 Some(v) => std::env::set_var(k, v),
728 None => std::env::remove_var(k),
729 }
730 }
731 }
732
733 #[test]
734 #[serial_test::serial]
735 fn parse_admin_env_none_when_all_unset() {
736 with_env(
737 &[
738 ("ADMIN_USERNAME", None),
739 ("ADMIN_EMAIL", None),
740 ("ADMIN_PASSWORD", None),
741 ],
742 || assert!(parse_admin_env().is_none()),
743 );
744 }
745
746 #[test]
747 #[serial_test::serial]
748 fn parse_admin_env_none_when_any_missing() {
749 with_env(
751 &[
752 ("ADMIN_USERNAME", Some("admin")),
753 ("ADMIN_EMAIL", Some("admin@example.com")),
754 ("ADMIN_PASSWORD", None),
755 ],
756 || assert!(parse_admin_env().is_none()),
757 );
758 with_env(
760 &[
761 ("ADMIN_USERNAME", Some("admin")),
762 ("ADMIN_EMAIL", None),
763 ("ADMIN_PASSWORD", Some("strongpass123")),
764 ],
765 || assert!(parse_admin_env().is_none()),
766 );
767 with_env(
769 &[
770 ("ADMIN_USERNAME", None),
771 ("ADMIN_EMAIL", Some("admin@example.com")),
772 ("ADMIN_PASSWORD", Some("strongpass123")),
773 ],
774 || assert!(parse_admin_env().is_none()),
775 );
776 }
777
778 #[test]
779 #[serial_test::serial]
780 fn parse_admin_env_none_on_empty_values() {
781 with_env(
782 &[
783 ("ADMIN_USERNAME", Some("admin")),
784 ("ADMIN_EMAIL", Some("admin@example.com")),
785 ("ADMIN_PASSWORD", Some("")),
786 ],
787 || assert!(parse_admin_env().is_none()),
788 );
789 with_env(
790 &[
791 ("ADMIN_USERNAME", Some(" ")),
792 ("ADMIN_EMAIL", Some("admin@example.com")),
793 ("ADMIN_PASSWORD", Some("strongpass123")),
794 ],
795 || assert!(parse_admin_env().is_none()),
796 );
797 }
798
799 #[test]
800 #[serial_test::serial]
801 fn parse_admin_env_some_when_all_set() {
802 with_env(
803 &[
804 ("ADMIN_USERNAME", Some(" admin ")),
805 ("ADMIN_EMAIL", Some(" admin@example.com ")),
806 ("ADMIN_PASSWORD", Some(" strong pass 123 ")),
807 ],
808 || {
809 let cfg = parse_admin_env().expect("三个变量齐全应返回 Some");
810 assert_eq!(cfg.username, "admin");
812 assert_eq!(cfg.email, "admin@example.com");
813 assert_eq!(cfg.password, " strong pass 123 ");
814 },
815 );
816 }
817}