feat: 完善生产级 Rust Web 模板 #1
@@ -18,3 +18,19 @@
|
||||
|
||||
# 配置文件路径
|
||||
# CONFIG=config/production.toml
|
||||
|
||||
# 可选能力与 Web 防护
|
||||
REDIS_ENABLED=false
|
||||
EMAIL_ENABLED=false
|
||||
EMAIL_SMTP_HOST=smtp.example.com
|
||||
EMAIL_SMTP_PORT=587
|
||||
EMAIL_SMTP_USERNAME=
|
||||
EMAIL_SMTP_PASSWORD=
|
||||
EMAIL_FROM_EMAIL=noreply@example.com
|
||||
EMAIL_QUEUE_ENABLED=false
|
||||
EMAIL_WORKER_POOL_SIZE=2
|
||||
SERVER_CORS_ORIGINS=["http://localhost:5173"]
|
||||
SERVER_REQUEST_TIMEOUT_SECONDS=30
|
||||
SERVER_MAX_BODY_BYTES=1048576
|
||||
SERVER_CONCURRENCY_LIMIT=256
|
||||
SERVER_RATE_LIMIT_PER_MINUTE=120
|
||||
|
||||
Generated
+276
-15
@@ -2,6 +2,41 @@
|
||||
# It is not intended for manual editing.
|
||||
version = 4
|
||||
|
||||
[[package]]
|
||||
name = "aead"
|
||||
version = "0.5.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d122413f284cf2d62fb1b7db97e02edb8cda96d769b16e443a4f6195e35662b0"
|
||||
dependencies = [
|
||||
"crypto-common",
|
||||
"generic-array",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "aes"
|
||||
version = "0.8.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b169f7a6d4742236a0a00c541b845991d0ac43e546831af1249753ab4c3aa3a0"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"cipher",
|
||||
"cpufeatures",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "aes-gcm"
|
||||
version = "0.10.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "831010a0f742e1209b3bcea8fab6a8e149051ba6099432c8cb2cc117dec3ead1"
|
||||
dependencies = [
|
||||
"aead",
|
||||
"aes",
|
||||
"cipher",
|
||||
"ctr",
|
||||
"ghash",
|
||||
"subtle",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ahash"
|
||||
version = "0.7.8"
|
||||
@@ -208,7 +243,7 @@ dependencies = [
|
||||
"serde_urlencoded",
|
||||
"sync_wrapper",
|
||||
"tokio",
|
||||
"tower 0.5.3",
|
||||
"tower",
|
||||
"tower-layer",
|
||||
"tower-service",
|
||||
"tracing",
|
||||
@@ -420,6 +455,16 @@ dependencies = [
|
||||
"windows-link",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "cipher"
|
||||
version = "0.4.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad"
|
||||
dependencies = [
|
||||
"crypto-common",
|
||||
"inout",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "clap"
|
||||
version = "4.5.58"
|
||||
@@ -498,7 +543,7 @@ dependencies = [
|
||||
"async-trait",
|
||||
"json5",
|
||||
"lazy_static",
|
||||
"nom",
|
||||
"nom 7.1.3",
|
||||
"pathdiff",
|
||||
"ron",
|
||||
"rust-ini",
|
||||
@@ -566,9 +611,19 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a"
|
||||
dependencies = [
|
||||
"generic-array",
|
||||
"rand_core",
|
||||
"typenum",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ctr"
|
||||
version = "0.9.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0369ee1ad671834580515889b80f2ea915f23b8be8d0daa4bbaf2ac5c7590835"
|
||||
dependencies = [
|
||||
"cipher",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "der"
|
||||
version = "0.7.10"
|
||||
@@ -656,6 +711,22 @@ dependencies = [
|
||||
"serde",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "email-encoding"
|
||||
version = "0.4.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "9298e6504d9b9e780ed3f7dfd43a61be8cd0e09eb07f7706a945b0072b6670b6"
|
||||
dependencies = [
|
||||
"base64 0.22.1",
|
||||
"memchr",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "email_address"
|
||||
version = "0.2.9"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e079f19b08ca6239f47f8ba8509c11cf3ea30095831f7fed61441475edd8c449"
|
||||
|
||||
[[package]]
|
||||
name = "equivalent"
|
||||
version = "1.0.2"
|
||||
@@ -873,6 +944,16 @@ dependencies = [
|
||||
"wasip2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ghash"
|
||||
version = "0.5.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f0d8a4362ccb29cb0b265253fb0a2728f592895ee6854fd9bc13f2ffda266ff1"
|
||||
dependencies = [
|
||||
"opaque-debug",
|
||||
"polyval",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "hashbrown"
|
||||
version = "0.12.3"
|
||||
@@ -1139,6 +1220,16 @@ dependencies = [
|
||||
"zerovec",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "idna"
|
||||
version = "0.4.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7d20d6b07bfbc108882d88ed8e37d39636dcc260e15e30c45e6ba089610b917c"
|
||||
dependencies = [
|
||||
"unicode-bidi",
|
||||
"unicode-normalization",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "idna"
|
||||
version = "1.1.0"
|
||||
@@ -1160,6 +1251,12 @@ dependencies = [
|
||||
"icu_properties",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "if_chain"
|
||||
version = "1.0.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "cd62e6b5e86ea8eeeb8db1de02880a6abc01a397b2ebb64b5d74ac255318f5cb"
|
||||
|
||||
[[package]]
|
||||
name = "indexmap"
|
||||
version = "2.13.0"
|
||||
@@ -1181,6 +1278,15 @@ dependencies = [
|
||||
"syn 2.0.115",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "inout"
|
||||
version = "0.1.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "879f10e63c20629ecabbb64a8010319738c66a5cd0c29b02d63d272b03751d01"
|
||||
dependencies = [
|
||||
"generic-array",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "is_terminal_polyfill"
|
||||
version = "1.70.2"
|
||||
@@ -1247,6 +1353,33 @@ dependencies = [
|
||||
"spin",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "lettre"
|
||||
version = "0.11.22"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0da65617f6cb926332d039cb578aad56178da86e128db6a1b09f4c94fa5b3349"
|
||||
dependencies = [
|
||||
"async-trait",
|
||||
"base64 0.22.1",
|
||||
"email-encoding",
|
||||
"email_address",
|
||||
"fastrand",
|
||||
"futures-io",
|
||||
"futures-util",
|
||||
"httpdate",
|
||||
"idna 1.1.0",
|
||||
"mime",
|
||||
"nom 8.0.0",
|
||||
"percent-encoding",
|
||||
"quoted_printable",
|
||||
"rustls",
|
||||
"socket2 0.6.2",
|
||||
"tokio",
|
||||
"tokio-rustls",
|
||||
"url",
|
||||
"webpki-roots 1.0.6",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "libc"
|
||||
version = "0.2.181"
|
||||
@@ -1372,6 +1505,15 @@ dependencies = [
|
||||
"minimal-lexical",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "nom"
|
||||
version = "8.0.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "df9761775871bdef83bee530e60050f7e54b1105350d6884eb0fb4f46c2f9405"
|
||||
dependencies = [
|
||||
"memchr",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "nu-ansi-term"
|
||||
version = "0.50.3"
|
||||
@@ -1455,6 +1597,12 @@ version = "1.70.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe"
|
||||
|
||||
[[package]]
|
||||
name = "opaque-debug"
|
||||
version = "0.3.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381"
|
||||
|
||||
[[package]]
|
||||
name = "ordered-float"
|
||||
version = "4.6.0"
|
||||
@@ -1660,6 +1808,18 @@ version = "0.3.32"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7edddbd0b52d732b21ad9a5fab5c704c14cd949e5e9a1ec5929a24fded1b904c"
|
||||
|
||||
[[package]]
|
||||
name = "polyval"
|
||||
version = "0.6.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "9d1fe60d06143b2430aa532c94cfe9e29783047f06c0d7fd359a9a51b729fa25"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"cpufeatures",
|
||||
"opaque-debug",
|
||||
"universal-hash",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "potential_utf"
|
||||
version = "0.1.4"
|
||||
@@ -1693,6 +1853,30 @@ dependencies = [
|
||||
"toml_edit",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "proc-macro-error"
|
||||
version = "1.0.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "da25490ff9892aab3fcf7c36f08cfb902dd3e71ca0f9f9517bea02a73a5ce38c"
|
||||
dependencies = [
|
||||
"proc-macro-error-attr",
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 1.0.109",
|
||||
"version_check",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "proc-macro-error-attr"
|
||||
version = "1.0.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "a1be40180e52ecc98ad80b184934baf3d0d29f979574e439af5a55274b35f869"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"version_check",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "proc-macro-error-attr2"
|
||||
version = "2.0.0"
|
||||
@@ -1766,6 +1950,12 @@ dependencies = [
|
||||
"proc-macro2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "quoted_printable"
|
||||
version = "0.5.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "478e0585659a122aa407eb7e3c0e1fa51b1d8a870038bd29f0cf4a8551eea972"
|
||||
|
||||
[[package]]
|
||||
name = "r-efi"
|
||||
version = "5.3.0"
|
||||
@@ -1852,6 +2042,18 @@ dependencies = [
|
||||
"bitflags 2.10.0",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "regex"
|
||||
version = "1.12.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e10754a14b9137dd7b1e3e5b0493cc9171fdd105e0ab477f51b72e7f3ac0e276"
|
||||
dependencies = [
|
||||
"aho-corasick",
|
||||
"memchr",
|
||||
"regex-automata",
|
||||
"regex-syntax",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "regex-automata"
|
||||
version = "0.4.14"
|
||||
@@ -1993,6 +2195,7 @@ version = "0.23.36"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c665f33d38cea657d9614f766881e4d510e0eda4239891eea56b4cadcf01801b"
|
||||
dependencies = [
|
||||
"log",
|
||||
"once_cell",
|
||||
"ring",
|
||||
"rustls-pki-types",
|
||||
@@ -2769,6 +2972,16 @@ dependencies = [
|
||||
"syn 2.0.115",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tokio-rustls"
|
||||
version = "0.26.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1729aa945f29d91ba541258c8df89027d5792d85a8841fb65e8bf0f4ede4ef61"
|
||||
dependencies = [
|
||||
"rustls",
|
||||
"tokio",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tokio-stream"
|
||||
version = "0.1.18"
|
||||
@@ -2832,17 +3045,6 @@ dependencies = [
|
||||
"winnow",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tower"
|
||||
version = "0.4.13"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b8fa9be0de6cf49e536ce1851f987bd21a43b771b09473c3549a6c853db37c1c"
|
||||
dependencies = [
|
||||
"tower-layer",
|
||||
"tower-service",
|
||||
"tracing",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tower"
|
||||
version = "0.5.3"
|
||||
@@ -2854,6 +3056,7 @@ dependencies = [
|
||||
"pin-project-lite",
|
||||
"sync_wrapper",
|
||||
"tokio",
|
||||
"tokio-util",
|
||||
"tower-layer",
|
||||
"tower-service",
|
||||
"tracing",
|
||||
@@ -2867,10 +3070,12 @@ checksum = "1e9cd434a998747dd2c4276bc96ee2e0c7a2eadf3cae88e52be55a05fa9053f5"
|
||||
dependencies = [
|
||||
"bitflags 2.10.0",
|
||||
"bytes",
|
||||
"futures-util",
|
||||
"http",
|
||||
"http-body",
|
||||
"http-body-util",
|
||||
"pin-project-lite",
|
||||
"tokio",
|
||||
"tower-layer",
|
||||
"tower-service",
|
||||
"tracing",
|
||||
@@ -2995,6 +3200,16 @@ version = "0.2.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ebc1c04c71510c7f702b52b7c350734c9ff1295c464a03335b00bb84fc54f853"
|
||||
|
||||
[[package]]
|
||||
name = "universal-hash"
|
||||
version = "0.5.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "fc1de2c688dc15305988b563c3854064043356019f97a4b46276fe734c4f07ea"
|
||||
dependencies = [
|
||||
"crypto-common",
|
||||
"subtle",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "untrusted"
|
||||
version = "0.9.0"
|
||||
@@ -3008,7 +3223,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ff67a8a4397373c3ef660812acab3268222035010ab8680ec4215f38ba3d0eed"
|
||||
dependencies = [
|
||||
"form_urlencoded",
|
||||
"idna",
|
||||
"idna 1.1.0",
|
||||
"percent-encoding",
|
||||
"serde",
|
||||
]
|
||||
@@ -3037,6 +3252,48 @@ dependencies = [
|
||||
"wasm-bindgen",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "validator"
|
||||
version = "0.16.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b92f40481c04ff1f4f61f304d61793c7b56ff76ac1469f1beb199b1445b253bd"
|
||||
dependencies = [
|
||||
"idna 0.4.0",
|
||||
"lazy_static",
|
||||
"regex",
|
||||
"serde",
|
||||
"serde_derive",
|
||||
"serde_json",
|
||||
"url",
|
||||
"validator_derive",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "validator_derive"
|
||||
version = "0.16.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "bc44ca3088bb3ba384d9aecf40c6a23a676ce23e09bdaca2073d99c207f864af"
|
||||
dependencies = [
|
||||
"if_chain",
|
||||
"lazy_static",
|
||||
"proc-macro-error",
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"regex",
|
||||
"syn 1.0.109",
|
||||
"validator_types",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "validator_types"
|
||||
version = "0.16.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "111abfe30072511849c5910134e8baf8dc05de4c0e5903d681cbd5c9c4d611e3"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"syn 1.0.109",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "valuable"
|
||||
version = "0.1.1"
|
||||
@@ -3125,6 +3382,7 @@ dependencies = [
|
||||
name = "web-rust-template"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"aes-gcm",
|
||||
"anyhow",
|
||||
"argon2",
|
||||
"async-trait",
|
||||
@@ -3133,7 +3391,9 @@ dependencies = [
|
||||
"chrono",
|
||||
"clap",
|
||||
"config",
|
||||
"http-body-util",
|
||||
"jsonwebtoken",
|
||||
"lettre",
|
||||
"rand",
|
||||
"redis",
|
||||
"sea-orm",
|
||||
@@ -3142,11 +3402,12 @@ dependencies = [
|
||||
"sha2",
|
||||
"thiserror 1.0.69",
|
||||
"tokio",
|
||||
"tower 0.4.13",
|
||||
"tower",
|
||||
"tower-http",
|
||||
"tracing",
|
||||
"tracing-subscriber",
|
||||
"uuid",
|
||||
"validator",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
||||
+8
-2
@@ -7,8 +7,8 @@ edition = "2021"
|
||||
# ===== Web 框架 =====
|
||||
axum = "0.7"
|
||||
tokio = { version = "1", features = ["full"] }
|
||||
tower = "0.4"
|
||||
tower-http = { version = "0.5", features = ["cors", "trace"] }
|
||||
tower = { version = "0.5", features = ["limit", "timeout", "util"] }
|
||||
tower-http = { version = "0.5", features = ["catch-panic", "cors", "limit", "set-header", "timeout", "trace"] }
|
||||
|
||||
# ===== 数据库(支持 MySQL、SQLite、PostgreSQL) =====
|
||||
# SeaORM - 数据库 ORM(替代 SQLX 直接使用)
|
||||
@@ -30,6 +30,10 @@ jsonwebtoken = "9"
|
||||
argon2 = "0.5"
|
||||
sha2 = "0.10"
|
||||
base64 = "0.22"
|
||||
aes-gcm = "0.10"
|
||||
|
||||
# ===== 邮件发送 =====
|
||||
lettre = { version = "0.11", default-features = false, features = ["tokio1", "tokio1-rustls", "builder", "smtp-transport", "webpki-roots", "ring"] }
|
||||
|
||||
# ===== Redis =====
|
||||
redis = { version = "0.27", features = ["tokio-comp", "connection-manager"] }
|
||||
@@ -45,6 +49,8 @@ tracing-subscriber = { version = "0.3", features = ["env-filter"] }
|
||||
config = "0.13"
|
||||
rand = "0.8"
|
||||
clap = { version = "4", features = ["derive", "env"] }
|
||||
validator = { version = "0.16", features = ["derive"] }
|
||||
http-body-util = "0.1"
|
||||
|
||||
# 优化发布版本
|
||||
[profile.release]
|
||||
|
||||
@@ -6,7 +6,7 @@
|
||||
|
||||
### 架构特色
|
||||
- **DDD 分层架构**:领域层、基础设施层、应用层清晰分离
|
||||
- **生产就绪**:JWT 双 Token 认证、Argon2 密码哈希、结构化日志
|
||||
- **生产就绪**:JWT 双 Token 认证、Argon2 密码哈希、可选 Redis/SMTP、结构化日志与请求保护
|
||||
- **多数据库支持**:MySQL / PostgreSQL / SQLite 无缝切换
|
||||
|
||||
### 技术栈
|
||||
@@ -14,7 +14,13 @@
|
||||
- **数据库 ORM**:SeaORM 1.1(支持多数据库)
|
||||
- **认证**:JWT (Access Token 15min + Refresh Token 7天)
|
||||
- **缓存**:Redis 存储 Refresh Token
|
||||
- **安全**:Argon2 密码哈希、CORS 支持
|
||||
- **安全**:Token 类型隔离、软删除校验、配置化 CORS、超时、正文/并发/频率限制、panic 捕获和安全响应头
|
||||
|
||||
### 可选能力
|
||||
|
||||
- Redis 默认关闭;服务无需 Redis 即可启动。Token 轮换、验证码等依赖 Redis 的接口会在不可用时返回 `503`。
|
||||
- 邮件默认关闭;启用后提供验证码、密码重置和邮件发送记录接口,可选择同步 SMTP 或 Redis 队列 Worker。
|
||||
- `Accept-Language` 支持 `zh-CN` 与 `en`,认证错误会按请求语言返回。
|
||||
|
||||
## 快速开始
|
||||
|
||||
@@ -139,11 +145,17 @@ cargo run -- -e production -c config/production.toml
|
||||
- `POST /auth/register` - 用户注册
|
||||
- `POST /auth/login` - 用户登录
|
||||
- `POST /auth/refresh` - 刷新 Token
|
||||
- `POST /auth/request-verification-code` - 发送验证码(邮件启用时)
|
||||
- `POST /auth/reset-password` - 重置密码(邮件启用时)
|
||||
|
||||
### 需要认证的接口
|
||||
|
||||
- `POST /auth/delete` - 删除账号
|
||||
- `POST /auth/delete-refresh-token` - 删除 Refresh Token
|
||||
- `POST /auth/logout` - 注销 Refresh Token
|
||||
- `GET|PUT|DELETE /api/user/profile` - 当前用户资料
|
||||
- `GET /api/email/latest-log` - 最近邮件记录(邮件启用时)
|
||||
- `GET /api/email/queue-status` - 邮件能力状态(邮件启用时)
|
||||
|
||||
> 查看 [完整 API 文档](docs/api/api-overview.md)
|
||||
|
||||
@@ -166,6 +178,7 @@ src/
|
||||
├── repositories/ # 数据访问层
|
||||
└── infra/ # 基础设施层
|
||||
├── middleware/ # 中间件
|
||||
├── mail/ # SMTP 邮件发送
|
||||
└── redis/ # Redis 客户端
|
||||
```
|
||||
|
||||
|
||||
@@ -2,6 +2,11 @@
|
||||
[server]
|
||||
host = "0.0.0.0"
|
||||
port = 3000
|
||||
request_timeout_seconds = 30
|
||||
max_body_bytes = 1048576
|
||||
concurrency_limit = 256
|
||||
rate_limit_per_minute = 120
|
||||
cors_origins = ["http://localhost:3000", "http://localhost:5173"]
|
||||
|
||||
[database]
|
||||
# 数据库类型: mysql, sqlite, postgresql
|
||||
@@ -27,7 +32,20 @@ access_token_expiration_minutes = 15 # access_token 15 分钟
|
||||
refresh_token_expiration_days = 7 # refresh_token 7 天
|
||||
|
||||
[redis]
|
||||
enabled = false
|
||||
host = "localhost"
|
||||
port = 6379
|
||||
password = "" # 可选
|
||||
db = 0
|
||||
|
||||
[email]
|
||||
enabled = false
|
||||
smtp_host = "smtp.example.com"
|
||||
smtp_port = 587
|
||||
smtp_username = ""
|
||||
smtp_password = ""
|
||||
from_email = "noreply@example.com"
|
||||
from_name = "Web Rust Template"
|
||||
verification_code_ttl_seconds = 600
|
||||
queue_enabled = false
|
||||
worker_pool_size = 2
|
||||
|
||||
@@ -3,6 +3,11 @@
|
||||
[server]
|
||||
host = "0.0.0.0" # 服务器监听地址(0.0.0.0=允许所有网络访问)
|
||||
port = 3000 # 服务器监听端口(确保防火墙已开放)
|
||||
request_timeout_seconds = 30
|
||||
max_body_bytes = 1048576
|
||||
concurrency_limit = 256
|
||||
rate_limit_per_minute = 120
|
||||
cors_origins = ["https://example.com"]
|
||||
|
||||
[database]
|
||||
database_type = "postgresql" # 数据库类型:sqlite/mysql/postgresql
|
||||
@@ -19,11 +24,24 @@ access_token_expiration_minutes = 15 # Access Token 过期时间(
|
||||
refresh_token_expiration_days = 7 # Refresh Token 过期时间(天)
|
||||
|
||||
[redis]
|
||||
enabled = true
|
||||
host = "localhost" # Redis 服务器地址
|
||||
port = 6379 # Redis 端口(默认 6379)
|
||||
password = "" # Redis 密码(强烈建议设置密码)
|
||||
db = 0 # Redis 数据库编号(0-15)
|
||||
|
||||
[email]
|
||||
enabled = false
|
||||
smtp_host = "smtp.example.com"
|
||||
smtp_port = 587
|
||||
smtp_username = ""
|
||||
smtp_password = ""
|
||||
from_email = "noreply@example.com"
|
||||
from_name = "Web Rust Template"
|
||||
verification_code_ttl_seconds = 600
|
||||
queue_enabled = true
|
||||
worker_pool_size = 4
|
||||
|
||||
# 安全检查清单:部署前请确认
|
||||
# ✅ 1. 已修改 jwt_secret 为强随机字符串
|
||||
# ✅ 2. 已修改数据库密码为强密码
|
||||
|
||||
@@ -183,3 +183,17 @@ curl http://localhost:3000/health
|
||||
---
|
||||
|
||||
**提示**:建议使用 Postman、Insomnia 或类似工具测试 API 接口。
|
||||
# 新增通用接口
|
||||
|
||||
所有受保护接口使用 `Authorization: Bearer <access_token>`。Refresh Token 不能访问受保护接口。
|
||||
|
||||
| 方法 | 路径 | 认证 | 条件 |
|
||||
|---|---|---|---|
|
||||
| POST | `/auth/logout` | 是 | Redis 可用 |
|
||||
| POST | `/auth/request-verification-code` | 否 | 邮件启用且 Redis 可用 |
|
||||
| POST | `/auth/reset-password` | 否 | 邮件启用且 Redis 可用 |
|
||||
| GET/PUT/DELETE | `/api/user/profile` | 是 | 始终注册 |
|
||||
| GET | `/api/email/latest-log` | 是 | 邮件启用 |
|
||||
| GET | `/api/email/queue-status` | 是 | 邮件启用 |
|
||||
|
||||
请求可通过 `Accept-Language: zh-CN` 或 `Accept-Language: en` 选择认证错误语言。
|
||||
|
||||
@@ -533,3 +533,23 @@ volumes:
|
||||
---
|
||||
|
||||
**提示**:使用 `.env.example` 作为模板,不要提交包含敏感信息的 `.env` 文件到版本控制。
|
||||
# 可选基础设施与 Web 防护
|
||||
|
||||
Redis 和邮件能力默认关闭。常用环境变量:
|
||||
|
||||
```text
|
||||
REDIS_ENABLED=false
|
||||
EMAIL_ENABLED=false
|
||||
EMAIL_SMTP_HOST=smtp.example.com
|
||||
EMAIL_SMTP_PORT=587
|
||||
EMAIL_SMTP_USERNAME=
|
||||
EMAIL_SMTP_PASSWORD=
|
||||
EMAIL_FROM_EMAIL=noreply@example.com
|
||||
SERVER_REQUEST_TIMEOUT_SECONDS=30
|
||||
SERVER_MAX_BODY_BYTES=1048576
|
||||
SERVER_CONCURRENCY_LIMIT=256
|
||||
SERVER_RATE_LIMIT_PER_MINUTE=120
|
||||
SERVER_CORS_ORIGINS=["https://app.example.com"]
|
||||
```
|
||||
|
||||
生产环境必须设置非默认 JWT 密钥,且 `cors_origins` 不允许包含 `*`。
|
||||
|
||||
@@ -14,6 +14,18 @@ CREATE TABLE IF NOT EXISTS users (
|
||||
password_hash VARCHAR(255) NOT NULL,
|
||||
created_at DATETIME NOT NULL COMMENT '创建时间',
|
||||
updated_at DATETIME NOT NULL COMMENT '更新时间'
|
||||
,deleted_at DATETIME NULL COMMENT '软删除时间'
|
||||
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4;
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_profiles (
|
||||
user_id VARCHAR(10) PRIMARY KEY,
|
||||
display_name VARCHAR(80), avatar_url VARCHAR(2048), bio VARCHAR(500),
|
||||
created_at DATETIME NOT NULL, updated_at DATETIME NOT NULL
|
||||
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4;
|
||||
|
||||
CREATE TABLE IF NOT EXISTS email_logs (
|
||||
id BIGINT AUTO_INCREMENT PRIMARY KEY, user_id VARCHAR(10), recipient VARCHAR(255) NOT NULL,
|
||||
kind VARCHAR(32) NOT NULL, status VARCHAR(32) NOT NULL, error TEXT, created_at DATETIME NOT NULL
|
||||
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4;
|
||||
|
||||
-- ============================================
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
CREATE TABLE IF NOT EXISTS users (id VARCHAR(10) PRIMARY KEY, email VARCHAR(255) NOT NULL UNIQUE, password_hash VARCHAR(255) NOT NULL, created_at DATETIME NOT NULL, updated_at DATETIME NOT NULL, deleted_at DATETIME NULL) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4;
|
||||
CREATE INDEX idx_users_email ON users(email);
|
||||
CREATE TABLE IF NOT EXISTS user_profiles (user_id VARCHAR(10) PRIMARY KEY, display_name VARCHAR(80), avatar_url VARCHAR(2048), bio VARCHAR(500), created_at DATETIME NOT NULL, updated_at DATETIME NOT NULL) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4;
|
||||
CREATE TABLE IF NOT EXISTS email_logs (id BIGINT AUTO_INCREMENT PRIMARY KEY, user_id VARCHAR(10), recipient VARCHAR(255) NOT NULL, kind VARCHAR(32) NOT NULL, status VARCHAR(32) NOT NULL, error TEXT, created_at DATETIME NOT NULL) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4;
|
||||
CREATE INDEX idx_email_logs_user_created ON email_logs(user_id, created_at DESC);
|
||||
@@ -0,0 +1,5 @@
|
||||
CREATE TABLE IF NOT EXISTS users (id VARCHAR(10) PRIMARY KEY, email VARCHAR(255) NOT NULL UNIQUE, password_hash VARCHAR(255) NOT NULL, created_at TIMESTAMP NOT NULL, updated_at TIMESTAMP NOT NULL, deleted_at TIMESTAMP NULL);
|
||||
CREATE INDEX IF NOT EXISTS idx_users_email ON users(email);
|
||||
CREATE TABLE IF NOT EXISTS user_profiles (user_id VARCHAR(10) PRIMARY KEY, display_name VARCHAR(80), avatar_url VARCHAR(2048), bio VARCHAR(500), created_at TIMESTAMP NOT NULL, updated_at TIMESTAMP NOT NULL);
|
||||
CREATE TABLE IF NOT EXISTS email_logs (id BIGSERIAL PRIMARY KEY, user_id VARCHAR(10), recipient VARCHAR(255) NOT NULL, kind VARCHAR(32) NOT NULL, status VARCHAR(32) NOT NULL, error TEXT, created_at TIMESTAMP NOT NULL);
|
||||
CREATE INDEX IF NOT EXISTS idx_email_logs_user_created ON email_logs(user_id, created_at DESC);
|
||||
@@ -0,0 +1,5 @@
|
||||
CREATE TABLE IF NOT EXISTS users (id TEXT PRIMARY KEY, email TEXT NOT NULL UNIQUE, password_hash TEXT NOT NULL, created_at DATETIME NOT NULL, updated_at DATETIME NOT NULL, deleted_at DATETIME NULL);
|
||||
CREATE INDEX IF NOT EXISTS idx_users_email ON users(email);
|
||||
CREATE TABLE IF NOT EXISTS user_profiles (user_id TEXT PRIMARY KEY, display_name TEXT, avatar_url TEXT, bio TEXT, created_at DATETIME NOT NULL, updated_at DATETIME NOT NULL);
|
||||
CREATE TABLE IF NOT EXISTS email_logs (id INTEGER PRIMARY KEY AUTOINCREMENT, user_id TEXT, recipient TEXT NOT NULL, kind TEXT NOT NULL, status TEXT NOT NULL, error TEXT, created_at DATETIME NOT NULL);
|
||||
CREATE INDEX IF NOT EXISTS idx_email_logs_user_created ON email_logs(user_id, created_at DESC);
|
||||
+4
-1
@@ -144,7 +144,10 @@ impl CliArgs {
|
||||
if config.exists() {
|
||||
return Some(config);
|
||||
}
|
||||
eprintln!("⚠ 警告:环境变量 CONFIG 指定的配置文件不存在: {}", config_path);
|
||||
eprintln!(
|
||||
"⚠ 警告:环境变量 CONFIG 指定的配置文件不存在: {}",
|
||||
config_path
|
||||
);
|
||||
eprintln!(" 将仅使用环境变量运行");
|
||||
return None;
|
||||
}
|
||||
|
||||
+55
-1
@@ -1,4 +1,7 @@
|
||||
use super::{auth::AuthConfig, database::DatabaseConfig, redis::RedisConfig, server::ServerConfig};
|
||||
use super::{
|
||||
auth::AuthConfig, database::DatabaseConfig, email::EmailConfig, redis::RedisConfig,
|
||||
server::ServerConfig,
|
||||
};
|
||||
use config::{Config, ConfigError, Environment, File};
|
||||
use serde::Deserialize;
|
||||
use std::path::PathBuf;
|
||||
@@ -12,6 +15,8 @@ pub struct AppConfig {
|
||||
pub database: DatabaseConfig,
|
||||
pub auth: AuthConfig,
|
||||
pub redis: RedisConfig,
|
||||
#[serde(default)]
|
||||
pub email: EmailConfig,
|
||||
}
|
||||
|
||||
impl AppConfig {
|
||||
@@ -46,6 +51,11 @@ impl AppConfig {
|
||||
// 注意:这些值会被环境变量覆盖
|
||||
builder = builder.set_default("server.host", default_server_host())?;
|
||||
builder = builder.set_default("server.port", default_server_port())?;
|
||||
builder = builder.set_default("server.request_timeout_seconds", 30)?;
|
||||
builder = builder.set_default("server.max_body_bytes", 1048576)?;
|
||||
builder = builder.set_default("server.concurrency_limit", 256)?;
|
||||
builder = builder.set_default("server.rate_limit_per_minute", 120)?;
|
||||
builder = builder.set_default("server.cors_origins", vec!["http://localhost:3000"])?;
|
||||
|
||||
// 设置 database 默认值(使用 SQLite 作为默认数据库)
|
||||
builder = builder.set_default("database.database_type", "sqlite")?;
|
||||
@@ -59,8 +69,10 @@ impl AppConfig {
|
||||
|
||||
// 设置 redis 默认值
|
||||
builder = builder.set_default("redis.host", default_redis_host())?;
|
||||
builder = builder.set_default("redis.enabled", false)?;
|
||||
builder = builder.set_default("redis.port", 6379)?;
|
||||
builder = builder.set_default("redis.db", 0)?;
|
||||
builder = builder.set_default("email.enabled", false)?;
|
||||
}
|
||||
|
||||
// 添加环境变量源(会覆盖配置文件的值)
|
||||
@@ -77,6 +89,8 @@ impl AppConfig {
|
||||
let settings = builder.build()?;
|
||||
let config: AppConfig = settings.try_deserialize()?;
|
||||
|
||||
validate_security(&config, _environment)?;
|
||||
|
||||
// 安全警告:检查是否使用了默认的 JWT 密钥
|
||||
if config.auth.jwt_secret == "change-this-to-a-strong-secret-key-in-production" {
|
||||
tracing::warn!("⚠️ 警告:正在使用不安全的默认 JWT 密钥!");
|
||||
@@ -121,6 +135,25 @@ impl AppConfig {
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_security(config: &AppConfig, environment: &str) -> Result<(), ConfigError> {
|
||||
if environment.eq_ignore_ascii_case("production") {
|
||||
if config.auth.jwt_secret == default_jwt_secret() {
|
||||
return Err(ConfigError::Message("生产环境禁止使用默认 JWT 密钥".into()));
|
||||
}
|
||||
if config
|
||||
.server
|
||||
.cors_origins
|
||||
.iter()
|
||||
.any(|origin| origin == "*")
|
||||
{
|
||||
return Err(ConfigError::Message(
|
||||
"生产环境禁止使用通配 CORS 来源".into(),
|
||||
));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// 默认值函数(复用)
|
||||
fn default_server_host() -> String {
|
||||
"127.0.0.1".to_string()
|
||||
@@ -133,3 +166,24 @@ fn default_server_port() -> u16 {
|
||||
fn default_jwt_secret() -> String {
|
||||
"change-this-to-a-strong-secret-key-in-production".to_string()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn production_rejects_default_secret_and_wildcard_cors() {
|
||||
let mut config = AppConfig::load_from_path("config/development.toml").unwrap();
|
||||
config.auth.jwt_secret = default_jwt_secret();
|
||||
assert!(validate_security(&config, "production").is_err());
|
||||
config.auth.jwt_secret = "a-production-secret".into();
|
||||
config.server.cors_origins = vec!["*".into()];
|
||||
assert!(validate_security(&config, "production").is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn development_accepts_safe_defaults() {
|
||||
let config = AppConfig::load_from_path("config/development.toml").unwrap();
|
||||
assert!(validate_security(&config, "development").is_ok());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -35,7 +35,7 @@ pub struct DatabaseConfig {
|
||||
impl DatabaseConfig {
|
||||
/// 获取端口号(根据数据库类型返回默认值)
|
||||
pub fn get_port(&self) -> u16 {
|
||||
self.port.unwrap_or_else(|| match self.database_type {
|
||||
self.port.unwrap_or(match self.database_type {
|
||||
DatabaseType::MySQL => 3306,
|
||||
DatabaseType::PostgreSQL => 5432,
|
||||
DatabaseType::SQLite => 0,
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
use serde::Deserialize;
|
||||
|
||||
#[derive(Debug, Deserialize, Clone)]
|
||||
pub struct EmailConfig {
|
||||
#[serde(default)]
|
||||
pub enabled: bool,
|
||||
#[serde(default)]
|
||||
pub smtp_host: String,
|
||||
#[serde(default = "default_smtp_port")]
|
||||
pub smtp_port: u16,
|
||||
#[serde(default)]
|
||||
pub smtp_username: String,
|
||||
#[serde(default)]
|
||||
pub smtp_password: String,
|
||||
#[serde(default)]
|
||||
pub from_email: String,
|
||||
#[serde(default = "default_from_name")]
|
||||
pub from_name: String,
|
||||
#[serde(default = "default_code_ttl")]
|
||||
pub verification_code_ttl_seconds: u64,
|
||||
#[serde(default)]
|
||||
pub queue_enabled: bool,
|
||||
#[serde(default = "default_worker_pool_size")]
|
||||
pub worker_pool_size: usize,
|
||||
}
|
||||
|
||||
fn default_smtp_port() -> u16 {
|
||||
587
|
||||
}
|
||||
fn default_from_name() -> String {
|
||||
"Web Rust Template".into()
|
||||
}
|
||||
fn default_code_ttl() -> u64 {
|
||||
600
|
||||
}
|
||||
fn default_worker_pool_size() -> usize {
|
||||
2
|
||||
}
|
||||
|
||||
impl Default for EmailConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
enabled: false,
|
||||
smtp_host: String::new(),
|
||||
smtp_port: default_smtp_port(),
|
||||
smtp_username: String::new(),
|
||||
smtp_password: String::new(),
|
||||
from_email: String::new(),
|
||||
from_name: default_from_name(),
|
||||
verification_code_ttl_seconds: default_code_ttl(),
|
||||
queue_enabled: false,
|
||||
worker_pool_size: default_worker_pool_size(),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,5 +1,6 @@
|
||||
pub mod app;
|
||||
pub mod auth;
|
||||
pub mod database;
|
||||
pub mod email;
|
||||
pub mod redis;
|
||||
pub mod server;
|
||||
|
||||
@@ -2,6 +2,8 @@ use serde::Deserialize;
|
||||
|
||||
#[derive(Debug, Deserialize, Clone)]
|
||||
pub struct RedisConfig {
|
||||
#[serde(default)]
|
||||
pub enabled: bool,
|
||||
/// Redis 主机地址
|
||||
#[serde(default = "default_redis_host")]
|
||||
pub host: String,
|
||||
|
||||
@@ -6,6 +6,32 @@ pub struct ServerConfig {
|
||||
pub host: String,
|
||||
#[serde(default = "default_server_port")]
|
||||
pub port: u16,
|
||||
#[serde(default = "default_request_timeout_seconds")]
|
||||
pub request_timeout_seconds: u64,
|
||||
#[serde(default = "default_max_body_bytes")]
|
||||
pub max_body_bytes: usize,
|
||||
#[serde(default = "default_concurrency_limit")]
|
||||
pub concurrency_limit: usize,
|
||||
#[serde(default = "default_rate_limit_per_minute")]
|
||||
pub rate_limit_per_minute: u32,
|
||||
#[serde(default = "default_cors_origins")]
|
||||
pub cors_origins: Vec<String>,
|
||||
}
|
||||
|
||||
fn default_request_timeout_seconds() -> u64 {
|
||||
30
|
||||
}
|
||||
fn default_max_body_bytes() -> usize {
|
||||
1024 * 1024
|
||||
}
|
||||
fn default_concurrency_limit() -> usize {
|
||||
256
|
||||
}
|
||||
fn default_rate_limit_per_minute() -> u32 {
|
||||
120
|
||||
}
|
||||
fn default_cors_origins() -> Vec<String> {
|
||||
vec!["http://localhost:3000".to_string()]
|
||||
}
|
||||
|
||||
fn default_server_host() -> String {
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
use crate::config::database::{DatabaseConfig, DatabaseType};
|
||||
use sea_orm::{
|
||||
ConnectionTrait, Database, DatabaseConnection, DbBackend, EntityName, EntityTrait, ConnectOptions, Schema,
|
||||
Statement,
|
||||
ConnectOptions, ConnectionTrait, Database, DatabaseConnection, DbBackend, EntityName,
|
||||
EntityTrait, Schema, Statement,
|
||||
};
|
||||
use std::time::Duration;
|
||||
|
||||
@@ -75,13 +75,48 @@ pub async fn init_database(config: &DatabaseConfig) -> anyhow::Result<DatabaseCo
|
||||
|
||||
// 创建表
|
||||
create_tables(&pool).await?;
|
||||
migrate_existing_tables(&pool).await?;
|
||||
create_indexes(&pool).await?;
|
||||
|
||||
Ok(pool)
|
||||
}
|
||||
|
||||
async fn create_indexes(db: &DatabaseConnection) -> anyhow::Result<()> {
|
||||
let backend = db.get_database_backend();
|
||||
let statements = match backend {
|
||||
DbBackend::MySql => vec!["CREATE INDEX idx_users_email ON users(email)", "CREATE INDEX idx_email_logs_user_created ON email_logs(user_id, created_at)"],
|
||||
_ => vec!["CREATE INDEX IF NOT EXISTS idx_users_email ON users(email)", "CREATE INDEX IF NOT EXISTS idx_email_logs_user_created ON email_logs(user_id, created_at)"],
|
||||
};
|
||||
for statement in statements {
|
||||
if let Err(error) = db.execute(Statement::from_string(backend, statement)).await {
|
||||
let message = error.to_string().to_lowercase();
|
||||
if !message.contains("duplicate") && !message.contains("already exists") {
|
||||
return Err(anyhow::anyhow!("创建索引失败: {error}"));
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn migrate_existing_tables(db: &DatabaseConnection) -> anyhow::Result<()> {
|
||||
let backend = db.get_database_backend();
|
||||
let statement = match backend {
|
||||
DbBackend::MySql => "ALTER TABLE users ADD COLUMN deleted_at DATETIME NULL",
|
||||
DbBackend::Postgres => "ALTER TABLE users ADD COLUMN deleted_at TIMESTAMP NULL",
|
||||
DbBackend::Sqlite => "ALTER TABLE users ADD COLUMN deleted_at DATETIME NULL",
|
||||
};
|
||||
if let Err(error) = db.execute(Statement::from_string(backend, statement)).await {
|
||||
let message = error.to_string().to_lowercase();
|
||||
if !message.contains("duplicate") && !message.contains("already exists") {
|
||||
return Err(anyhow::anyhow!("迁移 users.deleted_at 失败: {error}"));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 获取端口号(根据数据库类型返回默认值)
|
||||
fn get_database_port(config: &DatabaseConfig) -> u16 {
|
||||
config.port.unwrap_or_else(|| match config.database_type {
|
||||
config.port.unwrap_or(match config.database_type {
|
||||
DatabaseType::MySQL => 3306,
|
||||
DatabaseType::PostgreSQL => 5432,
|
||||
DatabaseType::SQLite => 0,
|
||||
@@ -94,6 +129,7 @@ async fn init_mysql_database(config: &DatabaseConfig) -> anyhow::Result<()> {
|
||||
.database
|
||||
.as_ref()
|
||||
.ok_or_else(|| anyhow::anyhow!("MySQL 需要配置 database.database"))?;
|
||||
validate_database_name(database_name)?;
|
||||
|
||||
let host = config
|
||||
.host
|
||||
@@ -150,6 +186,7 @@ async fn init_postgresql_database(config: &DatabaseConfig) -> anyhow::Result<()>
|
||||
.database
|
||||
.as_ref()
|
||||
.ok_or_else(|| anyhow::anyhow!("PostgreSQL 需要配置 database.database"))?;
|
||||
validate_database_name(database_name)?;
|
||||
|
||||
let host = config
|
||||
.host
|
||||
@@ -190,17 +227,18 @@ async fn init_postgresql_database(config: &DatabaseConfig) -> anyhow::Result<()>
|
||||
);
|
||||
|
||||
let result = conn
|
||||
.execute(Statement::from_string(
|
||||
.query_one(Statement::from_string(
|
||||
sea_orm::DatabaseBackend::Postgres,
|
||||
check_query,
|
||||
))
|
||||
.await;
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("检查 PostgreSQL 数据库失败: {e}"))?;
|
||||
|
||||
match result {
|
||||
Ok(_) => {
|
||||
Some(_) => {
|
||||
tracing::info!("PostgreSQL 数据库 '{}' 已存在", database_name);
|
||||
}
|
||||
Err(_) => {
|
||||
None => {
|
||||
// 数据库不存在,创建它
|
||||
let create_query = format!(
|
||||
"CREATE DATABASE {} WITH ENCODING 'UTF8' LC_COLLATE='en_US.UTF-8' LC_CTYPE='en_US.UTF-8'",
|
||||
@@ -221,6 +259,17 @@ async fn init_postgresql_database(config: &DatabaseConfig) -> anyhow::Result<()>
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn validate_database_name(name: &str) -> anyhow::Result<()> {
|
||||
if name.is_empty()
|
||||
|| !name
|
||||
.bytes()
|
||||
.all(|value| value.is_ascii_alphanumeric() || value == b'_')
|
||||
{
|
||||
anyhow::bail!("数据库名称只能包含字母、数字和下划线");
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 为 SQLite 确保数据库文件目录存在
|
||||
async fn init_sqlite_database(config: &DatabaseConfig) -> anyhow::Result<()> {
|
||||
let path = config
|
||||
@@ -299,7 +348,9 @@ where
|
||||
}
|
||||
Err(e) => {
|
||||
let err_msg = e.to_string();
|
||||
if err_msg.contains("already exists") || (err_msg.contains("table") && err_msg.contains("exists")) {
|
||||
if err_msg.contains("already exists")
|
||||
|| (err_msg.contains("table") && err_msg.contains("exists"))
|
||||
{
|
||||
tracing::info!("✅ {}已存在", table_name);
|
||||
} else {
|
||||
return Err(anyhow::anyhow!("创建{}失败: {}", table_name, e));
|
||||
@@ -318,10 +369,12 @@ async fn create_tables(db: &DatabaseConnection) -> anyhow::Result<()> {
|
||||
let schema = Schema::new(builder);
|
||||
|
||||
// 导入所有 entities
|
||||
use crate::domain::entities::users;
|
||||
use crate::domain::entities::{email_logs, user_profiles, users};
|
||||
|
||||
// 创建所有表(添加新表只需一行!)
|
||||
create_single_table(db, &schema, &builder, users::Entity, "用户表").await?;
|
||||
create_single_table(db, &schema, &builder, user_profiles::Entity, "用户资料表").await?;
|
||||
create_single_table(db, &schema, &builder, email_logs::Entity, "邮件日志表").await?;
|
||||
|
||||
tracing::info!("✅ 数据库表结构检查完成");
|
||||
|
||||
|
||||
+17
-4
@@ -1,24 +1,33 @@
|
||||
use serde::Deserialize;
|
||||
use std::fmt;
|
||||
use validator::Validate;
|
||||
|
||||
/// 注册请求
|
||||
#[derive(Deserialize)]
|
||||
#[derive(Deserialize, Validate)]
|
||||
pub struct RegisterRequest {
|
||||
#[validate(email)]
|
||||
pub email: String,
|
||||
#[validate(length(min = 8, max = 128))]
|
||||
pub password: String,
|
||||
}
|
||||
|
||||
// 实现 Debug trait,对密码进行脱敏
|
||||
impl fmt::Debug for RegisterRequest {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
write!(f, "RegisterRequest {{ email: {}, password: *** }}", self.email)
|
||||
write!(
|
||||
f,
|
||||
"RegisterRequest {{ email: {}, password: *** }}",
|
||||
self.email
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
/// 登录请求
|
||||
#[derive(Deserialize)]
|
||||
#[derive(Deserialize, Validate)]
|
||||
pub struct LoginRequest {
|
||||
#[validate(email)]
|
||||
pub email: String,
|
||||
#[validate(length(min = 1, max = 128))]
|
||||
pub password: String,
|
||||
}
|
||||
|
||||
@@ -39,7 +48,11 @@ pub struct DeleteUserRequest {
|
||||
// 实现 Debug trait
|
||||
impl fmt::Debug for DeleteUserRequest {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
write!(f, "DeleteUserRequest {{ user_id: {}, password: *** }}", self.user_id)
|
||||
write!(
|
||||
f,
|
||||
"DeleteUserRequest {{ user_id: {}, password: *** }}",
|
||||
self.user_id
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1 +1,2 @@
|
||||
pub mod auth;
|
||||
pub mod user;
|
||||
|
||||
+12
-1
@@ -1 +1,12 @@
|
||||
// 用户相关 DTO(预留)
|
||||
use serde::Deserialize;
|
||||
use validator::Validate;
|
||||
|
||||
#[derive(Debug, Deserialize, Validate)]
|
||||
pub struct UpdateProfileRequest {
|
||||
#[validate(length(max = 80))]
|
||||
pub display_name: Option<String>,
|
||||
#[validate(url)]
|
||||
pub avatar_url: Option<String>,
|
||||
#[validate(length(max = 500))]
|
||||
pub bio: Option<String>,
|
||||
}
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
use sea_orm::entity::prelude::*;
|
||||
use sea_orm::Set;
|
||||
use serde::{Deserialize, Serialize};
|
||||
#[derive(Clone, Debug, PartialEq, DeriveEntityModel, Serialize, Deserialize)]
|
||||
#[sea_orm(table_name = "email_logs")]
|
||||
pub struct Model {
|
||||
#[sea_orm(primary_key)]
|
||||
pub id: i64,
|
||||
pub user_id: Option<String>,
|
||||
pub recipient: String,
|
||||
pub kind: String,
|
||||
pub status: String,
|
||||
pub error: Option<String>,
|
||||
pub created_at: DateTime,
|
||||
}
|
||||
#[derive(Copy, Clone, Debug, EnumIter, DeriveRelation)]
|
||||
pub enum Relation {}
|
||||
#[async_trait::async_trait]
|
||||
impl ActiveModelBehavior for ActiveModel {
|
||||
async fn before_save<C: ConnectionTrait>(self, _: &C, insert: bool) -> Result<Self, DbErr> {
|
||||
let mut value = self;
|
||||
if insert {
|
||||
value.created_at = Set(chrono::Utc::now().naive_utc());
|
||||
}
|
||||
Ok(value)
|
||||
}
|
||||
}
|
||||
@@ -1,2 +1,3 @@
|
||||
pub mod email_logs;
|
||||
pub mod user_profiles;
|
||||
pub mod users;
|
||||
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
use sea_orm::entity::prelude::*;
|
||||
use sea_orm::Set;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, DeriveEntityModel, Serialize, Deserialize)]
|
||||
#[sea_orm(table_name = "user_profiles")]
|
||||
pub struct Model {
|
||||
#[sea_orm(primary_key, auto_increment = false)]
|
||||
pub user_id: String,
|
||||
pub display_name: Option<String>,
|
||||
pub avatar_url: Option<String>,
|
||||
pub bio: Option<String>,
|
||||
pub created_at: DateTime,
|
||||
pub updated_at: DateTime,
|
||||
}
|
||||
#[derive(Copy, Clone, Debug, EnumIter, DeriveRelation)]
|
||||
pub enum Relation {}
|
||||
#[async_trait::async_trait]
|
||||
impl ActiveModelBehavior for ActiveModel {
|
||||
async fn before_save<C: ConnectionTrait>(self, _: &C, insert: bool) -> Result<Self, DbErr> {
|
||||
let mut value = self;
|
||||
let now = chrono::Utc::now().naive_utc();
|
||||
if insert {
|
||||
value.created_at = Set(now);
|
||||
}
|
||||
value.updated_at = Set(now);
|
||||
Ok(value)
|
||||
}
|
||||
}
|
||||
@@ -12,6 +12,7 @@ pub struct Model {
|
||||
pub password_hash: String,
|
||||
pub created_at: DateTime,
|
||||
pub updated_at: DateTime,
|
||||
pub deleted_at: Option<DateTime>,
|
||||
}
|
||||
|
||||
#[derive(Copy, Clone, Debug, EnumIter, DeriveRelation)]
|
||||
|
||||
+1
-1
@@ -1,3 +1,3 @@
|
||||
pub mod dto;
|
||||
pub mod vo;
|
||||
pub mod entities;
|
||||
pub mod vo;
|
||||
|
||||
+22
-4
@@ -10,10 +10,19 @@ pub struct RegisterResult {
|
||||
}
|
||||
|
||||
impl From<(crate::domain::entities::users::Model, String, String)> for RegisterResult {
|
||||
fn from((user_model, access_token, refresh_token): (crate::domain::entities::users::Model, String, String)) -> Self {
|
||||
fn from(
|
||||
(user_model, access_token, refresh_token): (
|
||||
crate::domain::entities::users::Model,
|
||||
String,
|
||||
String,
|
||||
),
|
||||
) -> Self {
|
||||
Self {
|
||||
email: user_model.email,
|
||||
created_at: user_model.created_at.format("%Y-%m-%dT%H:%M:%S%.3fZ").to_string(),
|
||||
created_at: user_model
|
||||
.created_at
|
||||
.format("%Y-%m-%dT%H:%M:%S%.3fZ")
|
||||
.to_string(),
|
||||
access_token,
|
||||
refresh_token,
|
||||
}
|
||||
@@ -31,11 +40,20 @@ pub struct LoginResult {
|
||||
}
|
||||
|
||||
impl From<(crate::domain::entities::users::Model, String, String)> for LoginResult {
|
||||
fn from((user_model, access_token, refresh_token): (crate::domain::entities::users::Model, String, String)) -> Self {
|
||||
fn from(
|
||||
(user_model, access_token, refresh_token): (
|
||||
crate::domain::entities::users::Model,
|
||||
String,
|
||||
String,
|
||||
),
|
||||
) -> Self {
|
||||
Self {
|
||||
id: user_model.id,
|
||||
email: user_model.email,
|
||||
created_at: user_model.created_at.format("%Y-%m-%dT%H:%M:%S%.3fZ").to_string(),
|
||||
created_at: user_model
|
||||
.created_at
|
||||
.format("%Y-%m-%dT%H:%M:%S%.3fZ")
|
||||
.to_string(),
|
||||
access_token,
|
||||
refresh_token,
|
||||
}
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
pub mod auth;
|
||||
pub mod user;
|
||||
|
||||
use axum::http::StatusCode;
|
||||
/// 统一的 API 响应结构
|
||||
use serde::Serialize;
|
||||
use axum::http::StatusCode;
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
pub struct ApiResponse<T> {
|
||||
|
||||
+33
-48
@@ -5,38 +5,6 @@ use axum::{
|
||||
Json,
|
||||
};
|
||||
|
||||
/// 应用错误类型
|
||||
#[allow(dead_code)]
|
||||
pub struct AppError(pub anyhow::Error);
|
||||
|
||||
impl IntoResponse for AppError {
|
||||
fn into_response(self) -> Response {
|
||||
tracing::error!("Application error: {:?}", self.0);
|
||||
|
||||
let (status, message) = match self.0.downcast_ref::<&str>() {
|
||||
Some(&"not_found") => (StatusCode::NOT_FOUND, "Resource not found"),
|
||||
Some(&"unauthorized") => (StatusCode::UNAUTHORIZED, "Unauthorized"),
|
||||
_ => (StatusCode::INTERNAL_SERVER_ERROR, "Internal server error"),
|
||||
};
|
||||
|
||||
let body = ApiResponse::<()> {
|
||||
code: status.as_u16(),
|
||||
message: message.to_string(),
|
||||
data: None,
|
||||
};
|
||||
|
||||
(status, Json(body)).into_response()
|
||||
}
|
||||
}
|
||||
|
||||
// 为具体类型实现 From
|
||||
impl From<anyhow::Error> for AppError {
|
||||
fn from(err: anyhow::Error) -> Self {
|
||||
Self(err)
|
||||
}
|
||||
}
|
||||
|
||||
/// 统一的 API 错误响应结构
|
||||
#[derive(Debug)]
|
||||
pub struct ErrorResponse {
|
||||
pub status: StatusCode,
|
||||
@@ -45,29 +13,44 @@ pub struct ErrorResponse {
|
||||
|
||||
impl ErrorResponse {
|
||||
pub fn new(message: impl Into<String>) -> Self {
|
||||
Self::bad_request(message)
|
||||
}
|
||||
pub fn bad_request(message: impl Into<String>) -> Self {
|
||||
Self {
|
||||
status: StatusCode::BAD_REQUEST,
|
||||
message: message.into(),
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub fn not_found(message: impl Into<String>) -> Self {
|
||||
Self {
|
||||
status: StatusCode::NOT_FOUND,
|
||||
message: message.into(),
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub fn unauthorized(message: impl Into<String>) -> Self {
|
||||
Self {
|
||||
status: StatusCode::UNAUTHORIZED,
|
||||
message: message.into(),
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub fn not_found(message: impl Into<String>) -> Self {
|
||||
Self {
|
||||
status: StatusCode::NOT_FOUND,
|
||||
message: message.into(),
|
||||
}
|
||||
}
|
||||
pub fn too_many_requests(message: impl Into<String>) -> Self {
|
||||
Self {
|
||||
status: StatusCode::TOO_MANY_REQUESTS,
|
||||
message: message.into(),
|
||||
}
|
||||
}
|
||||
pub fn payload_too_large(message: impl Into<String>) -> Self {
|
||||
Self {
|
||||
status: StatusCode::PAYLOAD_TOO_LARGE,
|
||||
message: message.into(),
|
||||
}
|
||||
}
|
||||
pub fn unavailable(message: impl Into<String>) -> Self {
|
||||
Self {
|
||||
status: StatusCode::SERVICE_UNAVAILABLE,
|
||||
message: message.into(),
|
||||
}
|
||||
}
|
||||
pub fn internal(message: impl Into<String>) -> Self {
|
||||
Self {
|
||||
status: StatusCode::INTERNAL_SERVER_ERROR,
|
||||
@@ -78,12 +61,14 @@ impl ErrorResponse {
|
||||
|
||||
impl IntoResponse for ErrorResponse {
|
||||
fn into_response(self) -> Response {
|
||||
let body = ApiResponse::<()> {
|
||||
(
|
||||
self.status,
|
||||
Json(ApiResponse::<()> {
|
||||
code: self.status.as_u16(),
|
||||
message: self.message,
|
||||
data: None,
|
||||
};
|
||||
|
||||
(self.status, Json(body)).into_response()
|
||||
}),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
}
|
||||
|
||||
+58
-29
@@ -1,8 +1,9 @@
|
||||
use crate::domain::dto::auth::{DeleteUserRequest, LoginRequest, RefreshRequest, RegisterRequest};
|
||||
use crate::domain::vo::auth::{LoginResult, RefreshResult, RegisterResult};
|
||||
use crate::domain::vo::ApiResponse;
|
||||
use crate::error::ErrorResponse;
|
||||
use crate::infra::middleware::logging::{log_info, RequestId};
|
||||
use crate::domain::dto::auth::{RegisterRequest, LoginRequest, RefreshRequest, DeleteUserRequest};
|
||||
use crate::domain::vo::auth::{RegisterResult, LoginResult, RefreshResult};
|
||||
use crate::domain::vo::ApiResponse;
|
||||
use crate::infra::middleware::UserId;
|
||||
use crate::repositories::user_repository::UserRepository;
|
||||
use crate::services::auth_service::AuthService;
|
||||
use crate::AppState;
|
||||
@@ -11,6 +12,7 @@ use axum::{
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
use validator::Validate;
|
||||
|
||||
/// 注册
|
||||
pub async fn register(
|
||||
@@ -18,10 +20,17 @@ pub async fn register(
|
||||
State(state): State<AppState>,
|
||||
Json(payload): Json<RegisterRequest>,
|
||||
) -> Result<Json<ApiResponse<RegisterResult>>, ErrorResponse> {
|
||||
payload
|
||||
.validate()
|
||||
.map_err(|e| ErrorResponse::bad_request(e.to_string()))?;
|
||||
log_info(&request_id, "注册请求参数", &payload);
|
||||
|
||||
let user_repo = UserRepository::new(state.pool.clone());
|
||||
let service = AuthService::new(user_repo, state.redis_client.clone(), state.config.auth.clone());
|
||||
let service = AuthService::new(
|
||||
user_repo,
|
||||
state.redis_client.clone(),
|
||||
state.config.auth.clone(),
|
||||
);
|
||||
|
||||
match service.register(payload).await {
|
||||
Ok((user_model, access_token, refresh_token)) => {
|
||||
@@ -31,7 +40,7 @@ pub async fn register(
|
||||
Ok(Json(response))
|
||||
}
|
||||
Err(e) => {
|
||||
log_info(&request_id, "注册失败", &e.to_string());
|
||||
log_info(&request_id, "注册失败", e.to_string());
|
||||
Err(ErrorResponse::new(e.to_string()))
|
||||
}
|
||||
}
|
||||
@@ -43,10 +52,17 @@ pub async fn login(
|
||||
State(state): State<AppState>,
|
||||
Json(payload): Json<LoginRequest>,
|
||||
) -> Result<Json<ApiResponse<LoginResult>>, ErrorResponse> {
|
||||
payload
|
||||
.validate()
|
||||
.map_err(|e| ErrorResponse::bad_request(e.to_string()))?;
|
||||
log_info(&request_id, "登录请求参数", &payload);
|
||||
|
||||
let user_repo = UserRepository::new(state.pool.clone());
|
||||
let service = AuthService::new(user_repo, state.redis_client.clone(), state.config.auth.clone());
|
||||
let service = AuthService::new(
|
||||
user_repo,
|
||||
state.redis_client.clone(),
|
||||
state.config.auth.clone(),
|
||||
);
|
||||
|
||||
match service.login(payload).await {
|
||||
Ok((user_model, access_token, refresh_token)) => {
|
||||
@@ -56,7 +72,7 @@ pub async fn login(
|
||||
Ok(Json(response))
|
||||
}
|
||||
Err(e) => {
|
||||
log_info(&request_id, "登录失败", &e.to_string());
|
||||
log_info(&request_id, "登录失败", e.to_string());
|
||||
Err(ErrorResponse::new(e.to_string()))
|
||||
}
|
||||
}
|
||||
@@ -71,16 +87,17 @@ pub async fn refresh(
|
||||
log_info(
|
||||
&request_id,
|
||||
"刷新 token 请求",
|
||||
&json!({"device_id": "default"}),
|
||||
json!({"device_id": "default"}),
|
||||
);
|
||||
|
||||
let user_repo = UserRepository::new(state.pool.clone());
|
||||
let service = AuthService::new(user_repo, state.redis_client.clone(), state.config.auth.clone());
|
||||
let service = AuthService::new(
|
||||
user_repo,
|
||||
state.redis_client.clone(),
|
||||
state.config.auth.clone(),
|
||||
);
|
||||
|
||||
match service
|
||||
.refresh_access_token(&payload.refresh_token)
|
||||
.await
|
||||
{
|
||||
match service.refresh_access_token(&payload.refresh_token).await {
|
||||
Ok((access_token, refresh_token)) => {
|
||||
let data = RefreshResult {
|
||||
access_token,
|
||||
@@ -88,15 +105,11 @@ pub async fn refresh(
|
||||
};
|
||||
let response = ApiResponse::success(data);
|
||||
|
||||
log_info(
|
||||
&request_id,
|
||||
"刷新成功",
|
||||
&json!({"access_token": "***"}),
|
||||
);
|
||||
log_info(&request_id, "刷新成功", json!({"access_token": "***"}));
|
||||
Ok(Json(response))
|
||||
}
|
||||
Err(e) => {
|
||||
log_info(&request_id, "刷新失败", &e.to_string());
|
||||
log_info(&request_id, "刷新失败", e.to_string());
|
||||
Err(ErrorResponse::new(e.to_string()))
|
||||
}
|
||||
}
|
||||
@@ -106,13 +119,17 @@ pub async fn refresh(
|
||||
pub async fn delete_account(
|
||||
Extension(request_id): Extension<RequestId>,
|
||||
State(state): State<AppState>,
|
||||
Extension(user_id): Extension<String>,
|
||||
UserId(user_id): UserId,
|
||||
Json(payload): Json<DeleteUserRequest>,
|
||||
) -> Result<Json<ApiResponse<()>>, ErrorResponse> {
|
||||
log_info(&request_id, "删除账号请求", &format!("user_id={}", user_id));
|
||||
log_info(&request_id, "删除账号请求", format!("user_id={}", user_id));
|
||||
|
||||
let user_repo = UserRepository::new(state.pool.clone());
|
||||
let service = AuthService::new(user_repo, state.redis_client.clone(), state.config.auth.clone());
|
||||
let service = AuthService::new(
|
||||
user_repo,
|
||||
state.redis_client.clone(),
|
||||
state.config.auth.clone(),
|
||||
);
|
||||
|
||||
let delete_request = DeleteUserRequest {
|
||||
user_id: user_id.clone(),
|
||||
@@ -121,12 +138,12 @@ pub async fn delete_account(
|
||||
|
||||
match service.delete_user(delete_request).await {
|
||||
Ok(_) => {
|
||||
log_info(&request_id, "账号删除成功", &format!("user_id={}", user_id));
|
||||
log_info(&request_id, "账号删除成功", format!("user_id={}", user_id));
|
||||
let response = ApiResponse::success_with_message((), "账号删除成功");
|
||||
Ok(Json(response))
|
||||
}
|
||||
Err(e) => {
|
||||
log_info(&request_id, "账号删除失败", &e.to_string());
|
||||
log_info(&request_id, "账号删除失败", e.to_string());
|
||||
Err(ErrorResponse::new(e.to_string()))
|
||||
}
|
||||
}
|
||||
@@ -136,21 +153,33 @@ pub async fn delete_account(
|
||||
pub async fn delete_refresh_token(
|
||||
Extension(request_id): Extension<RequestId>,
|
||||
State(state): State<AppState>,
|
||||
Extension(user_id): Extension<String>,
|
||||
UserId(user_id): UserId,
|
||||
) -> Result<Json<ApiResponse<()>>, ErrorResponse> {
|
||||
log_info(&request_id, "删除刷新令牌请求", &format!("user_id={}", user_id));
|
||||
log_info(
|
||||
&request_id,
|
||||
"删除刷新令牌请求",
|
||||
format!("user_id={}", user_id),
|
||||
);
|
||||
|
||||
let user_repo = UserRepository::new(state.pool.clone());
|
||||
let service = AuthService::new(user_repo, state.redis_client.clone(), state.config.auth.clone());
|
||||
let service = AuthService::new(
|
||||
user_repo,
|
||||
state.redis_client.clone(),
|
||||
state.config.auth.clone(),
|
||||
);
|
||||
|
||||
match service.delete_refresh_token(&user_id).await {
|
||||
Ok(_) => {
|
||||
log_info(&request_id, "刷新令牌删除成功", &format!("user_id={}", user_id));
|
||||
log_info(
|
||||
&request_id,
|
||||
"刷新令牌删除成功",
|
||||
format!("user_id={}", user_id),
|
||||
);
|
||||
let response = ApiResponse::success_with_message((), "刷新令牌删除成功");
|
||||
Ok(Json(response))
|
||||
}
|
||||
Err(e) => {
|
||||
log_info(&request_id, "刷新令牌删除失败", &e.to_string());
|
||||
log_info(&request_id, "刷新令牌删除失败", e.to_string());
|
||||
Err(ErrorResponse::new(e.to_string()))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,196 @@
|
||||
use crate::{
|
||||
domain::vo::ApiResponse,
|
||||
error::ErrorResponse,
|
||||
infra::{
|
||||
mail::{
|
||||
mailer,
|
||||
worker::{EmailJob, MAIL_QUEUE},
|
||||
},
|
||||
middleware::UserId,
|
||||
redis::redis_key::{BusinessType, RedisKey},
|
||||
},
|
||||
repositories::{email_log_repository::EmailLogRepository, user_repository::UserRepository},
|
||||
services::auth_service::AuthService,
|
||||
AppState,
|
||||
};
|
||||
use axum::{extract::State, Json};
|
||||
use rand::Rng;
|
||||
use serde::Deserialize;
|
||||
use serde_json::json;
|
||||
use validator::{Validate, ValidationError};
|
||||
|
||||
#[derive(Debug, Deserialize, Validate)]
|
||||
pub struct VerificationRequest {
|
||||
#[validate(email)]
|
||||
pub email: String,
|
||||
}
|
||||
|
||||
fn valid_code(code: &str) -> Result<(), ValidationError> {
|
||||
if code.len() == 6 && code.bytes().all(|v| v.is_ascii_digit()) {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(ValidationError::new("code"))
|
||||
}
|
||||
}
|
||||
#[derive(Debug, Deserialize, Validate)]
|
||||
pub struct ResetPasswordRequest {
|
||||
#[validate(email)]
|
||||
pub email: String,
|
||||
#[validate(custom = "valid_code")]
|
||||
pub code: String,
|
||||
#[validate(length(min = 8, max = 128))]
|
||||
pub new_password: String,
|
||||
}
|
||||
|
||||
fn code_key(email: &str) -> RedisKey {
|
||||
RedisKey::new(BusinessType::Auth)
|
||||
.add_identifier("verify_code")
|
||||
.add_identifier(email)
|
||||
}
|
||||
|
||||
pub async fn send_verification_code(
|
||||
State(state): State<AppState>,
|
||||
Json(input): Json<VerificationRequest>,
|
||||
) -> Result<Json<ApiResponse<()>>, ErrorResponse> {
|
||||
input
|
||||
.validate()
|
||||
.map_err(|e| ErrorResponse::bad_request(e.to_string()))?;
|
||||
let redis = state
|
||||
.redis_client
|
||||
.as_ref()
|
||||
.ok_or_else(|| ErrorResponse::unavailable("Redis service unavailable"))?;
|
||||
let key = code_key(&input.email);
|
||||
if redis
|
||||
.exists_key(&key)
|
||||
.await
|
||||
.map_err(|e| ErrorResponse::unavailable(e.to_string()))?
|
||||
{
|
||||
return Err(ErrorResponse::too_many_requests(
|
||||
"verification code already sent",
|
||||
));
|
||||
}
|
||||
let code = format!("{:06}", rand::thread_rng().gen_range(0..1_000_000));
|
||||
if state.config.email.queue_enabled {
|
||||
redis
|
||||
.set_key_ex(
|
||||
&key,
|
||||
&code,
|
||||
state.config.email.verification_code_ttl_seconds,
|
||||
)
|
||||
.await
|
||||
.map_err(|e| ErrorResponse::unavailable(e.to_string()))?;
|
||||
redis
|
||||
.queue_push(
|
||||
MAIL_QUEUE,
|
||||
&EmailJob {
|
||||
recipient: input.email,
|
||||
code,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.map_err(|e| ErrorResponse::unavailable(e.to_string()))?;
|
||||
return Ok(Json(ApiResponse::success_with_message(
|
||||
(),
|
||||
"verification email queued",
|
||||
)));
|
||||
}
|
||||
let result = mailer::send_verification(&state.config.email, &input.email, &code).await;
|
||||
EmailLogRepository::new(state.pool.clone())
|
||||
.add(
|
||||
None,
|
||||
input.email.clone(),
|
||||
if result.is_ok() { "sent" } else { "failed" }.into(),
|
||||
result.as_ref().err().map(ToString::to_string),
|
||||
)
|
||||
.await
|
||||
.map_err(|e| ErrorResponse::internal(e.to_string()))?;
|
||||
result.map_err(|e| ErrorResponse::unavailable(e.to_string()))?;
|
||||
redis
|
||||
.set_key_ex(
|
||||
&key,
|
||||
&code,
|
||||
state.config.email.verification_code_ttl_seconds,
|
||||
)
|
||||
.await
|
||||
.map_err(|e| ErrorResponse::unavailable(e.to_string()))?;
|
||||
Ok(Json(ApiResponse::success_with_message(
|
||||
(),
|
||||
"verification code sent",
|
||||
)))
|
||||
}
|
||||
|
||||
pub async fn reset_password(
|
||||
State(state): State<AppState>,
|
||||
Json(input): Json<ResetPasswordRequest>,
|
||||
) -> Result<Json<ApiResponse<()>>, ErrorResponse> {
|
||||
input
|
||||
.validate()
|
||||
.map_err(|e| ErrorResponse::bad_request(e.to_string()))?;
|
||||
let redis = state
|
||||
.redis_client
|
||||
.as_ref()
|
||||
.ok_or_else(|| ErrorResponse::unavailable("Redis service unavailable"))?;
|
||||
let key = code_key(&input.email);
|
||||
let stored = redis
|
||||
.get_key(&key)
|
||||
.await
|
||||
.map_err(|e| ErrorResponse::unavailable(e.to_string()))?
|
||||
.ok_or_else(|| ErrorResponse::bad_request("verification code expired"))?;
|
||||
if stored != input.code {
|
||||
return Err(ErrorResponse::bad_request("invalid verification code"));
|
||||
}
|
||||
let repo = UserRepository::new(state.pool.clone());
|
||||
repo.update_password_by_email(
|
||||
&input.email,
|
||||
AuthService::hash_password(&input.new_password)
|
||||
.map_err(|e| ErrorResponse::internal(e.to_string()))?,
|
||||
)
|
||||
.await
|
||||
.map_err(|e| ErrorResponse::bad_request(e.to_string()))?;
|
||||
redis
|
||||
.delete_key(&key)
|
||||
.await
|
||||
.map_err(|e| ErrorResponse::unavailable(e.to_string()))?;
|
||||
Ok(Json(ApiResponse::success_with_message(
|
||||
(),
|
||||
"password reset",
|
||||
)))
|
||||
}
|
||||
|
||||
pub async fn latest_log(
|
||||
State(state): State<AppState>,
|
||||
UserId(user_id): UserId,
|
||||
) -> Result<Json<ApiResponse<serde_json::Value>>, ErrorResponse> {
|
||||
let value = EmailLogRepository::new(state.pool)
|
||||
.latest(&user_id)
|
||||
.await
|
||||
.map_err(|e| ErrorResponse::internal(e.to_string()))?;
|
||||
Ok(Json(ApiResponse::success(
|
||||
serde_json::to_value(value).unwrap_or_default(),
|
||||
)))
|
||||
}
|
||||
pub async fn queue_status(
|
||||
State(state): State<AppState>,
|
||||
) -> Result<Json<ApiResponse<serde_json::Value>>, ErrorResponse> {
|
||||
let pending = if let Some(redis) = &state.redis_client {
|
||||
redis
|
||||
.queue_len(MAIL_QUEUE)
|
||||
.await
|
||||
.map_err(|e| ErrorResponse::unavailable(e.to_string()))?
|
||||
} else {
|
||||
0
|
||||
};
|
||||
Ok(Json(ApiResponse::success(
|
||||
json!({"enabled": state.config.email.queue_enabled, "redis_available": state.redis_client.is_some(), "pending": pending, "workers": state.config.email.worker_pool_size}),
|
||||
)))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
#[test]
|
||||
fn validates_code() {
|
||||
assert!(valid_code("123456").is_ok());
|
||||
assert!(valid_code("12x").is_err());
|
||||
}
|
||||
}
|
||||
@@ -1,5 +1,5 @@
|
||||
use crate::AppState;
|
||||
use crate::db;
|
||||
use crate::AppState;
|
||||
use axum::{
|
||||
extract::State,
|
||||
response::{IntoResponse, Json},
|
||||
@@ -8,10 +8,12 @@ use serde_json::json;
|
||||
|
||||
/// 健康检查端点
|
||||
pub async fn health_check(State(state): State<AppState>) -> impl IntoResponse {
|
||||
match db::health_check(&state.pool).await {
|
||||
Ok(_) => Json(json!({"status": "ok"})),
|
||||
Err(_) => Json(json!({"status": "unavailable"})),
|
||||
}
|
||||
let database = db::health_check(&state.pool).await.is_ok();
|
||||
Json(
|
||||
json!({"status": if database { "ok" } else { "unavailable" }, "capabilities": {
|
||||
"database": database, "redis": state.redis_client.is_some(), "email": state.config.email.enabled
|
||||
}}),
|
||||
)
|
||||
}
|
||||
|
||||
/// 获取服务器信息
|
||||
|
||||
@@ -1,2 +1,4 @@
|
||||
pub mod auth;
|
||||
pub mod email;
|
||||
pub mod health;
|
||||
pub mod user_profile;
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
use crate::{
|
||||
domain::{dto::user::UpdateProfileRequest, vo::ApiResponse},
|
||||
error::ErrorResponse,
|
||||
infra::middleware::UserId,
|
||||
repositories::user_profile_repository::UserProfileRepository,
|
||||
AppState,
|
||||
};
|
||||
use axum::{extract::State, Json};
|
||||
use validator::Validate;
|
||||
|
||||
pub async fn get_profile(
|
||||
State(state): State<AppState>,
|
||||
UserId(user_id): UserId,
|
||||
) -> Result<Json<ApiResponse<serde_json::Value>>, ErrorResponse> {
|
||||
let value = UserProfileRepository::new(state.pool)
|
||||
.get(&user_id)
|
||||
.await
|
||||
.map_err(|e| ErrorResponse::internal(e.to_string()))?;
|
||||
Ok(Json(ApiResponse::success(
|
||||
serde_json::to_value(value).unwrap_or_default(),
|
||||
)))
|
||||
}
|
||||
pub async fn update_profile(
|
||||
State(state): State<AppState>,
|
||||
UserId(user_id): UserId,
|
||||
Json(input): Json<UpdateProfileRequest>,
|
||||
) -> Result<Json<ApiResponse<serde_json::Value>>, ErrorResponse> {
|
||||
input
|
||||
.validate()
|
||||
.map_err(|e| ErrorResponse::bad_request(e.to_string()))?;
|
||||
let value = UserProfileRepository::new(state.pool)
|
||||
.upsert(user_id, input)
|
||||
.await
|
||||
.map_err(|e| ErrorResponse::internal(e.to_string()))?;
|
||||
Ok(Json(ApiResponse::success(
|
||||
serde_json::to_value(value).unwrap_or_default(),
|
||||
)))
|
||||
}
|
||||
pub async fn delete_profile(
|
||||
State(state): State<AppState>,
|
||||
UserId(user_id): UserId,
|
||||
) -> Result<Json<ApiResponse<()>>, ErrorResponse> {
|
||||
UserProfileRepository::new(state.pool)
|
||||
.delete(&user_id)
|
||||
.await
|
||||
.map_err(|e| ErrorResponse::internal(e.to_string()))?;
|
||||
Ok(Json(ApiResponse::success_with_message(
|
||||
(),
|
||||
"profile deleted",
|
||||
)))
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
use crate::config::email::EmailConfig;
|
||||
use anyhow::Result;
|
||||
use lettre::{
|
||||
message::Mailbox, transport::smtp::authentication::Credentials, AsyncSmtpTransport,
|
||||
AsyncTransport, Message, Tokio1Executor,
|
||||
};
|
||||
|
||||
pub async fn send_verification(config: &EmailConfig, recipient: &str, code: &str) -> Result<()> {
|
||||
let email = Message::builder()
|
||||
.from(format!("{} <{}>", config.from_name, config.from_email).parse::<Mailbox>()?)
|
||||
.to(recipient.parse::<Mailbox>()?)
|
||||
.subject("Verification code")
|
||||
.body(format!(
|
||||
"Your verification code is {code}. It expires soon."
|
||||
))?;
|
||||
let credentials = Credentials::new(config.smtp_username.clone(), config.smtp_password.clone());
|
||||
let mailer = AsyncSmtpTransport::<Tokio1Executor>::starttls_relay(&config.smtp_host)?
|
||||
.port(config.smtp_port)
|
||||
.credentials(credentials)
|
||||
.build();
|
||||
mailer.send(email).await?;
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,2 @@
|
||||
pub mod mailer;
|
||||
pub mod worker;
|
||||
@@ -0,0 +1,50 @@
|
||||
use crate::{
|
||||
config::email::EmailConfig,
|
||||
db::DbPool,
|
||||
infra::{mail::mailer, redis::redis_client::RedisClient},
|
||||
repositories::email_log_repository::EmailLogRepository,
|
||||
};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::time::Duration;
|
||||
|
||||
pub const MAIL_QUEUE: &str = "mail:verification:queue";
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize)]
|
||||
pub struct EmailJob {
|
||||
pub recipient: String,
|
||||
pub code: String,
|
||||
}
|
||||
|
||||
pub fn start(redis: RedisClient, config: EmailConfig, pool: DbPool) {
|
||||
for worker_id in 0..config.worker_pool_size.max(1) {
|
||||
let redis = redis.clone();
|
||||
let config = config.clone();
|
||||
let pool = pool.clone();
|
||||
tokio::spawn(async move {
|
||||
loop {
|
||||
match redis.queue_pop::<EmailJob>(MAIL_QUEUE).await {
|
||||
Ok(Some(job)) => {
|
||||
let result =
|
||||
mailer::send_verification(&config, &job.recipient, &job.code).await;
|
||||
if let Err(error) = EmailLogRepository::new(pool.clone())
|
||||
.add(
|
||||
None,
|
||||
job.recipient,
|
||||
if result.is_ok() { "sent" } else { "failed" }.into(),
|
||||
result.err().map(|e| e.to_string()),
|
||||
)
|
||||
.await
|
||||
{
|
||||
tracing::error!(worker_id, %error, "failed to persist email log");
|
||||
}
|
||||
}
|
||||
Ok(None) => tokio::time::sleep(Duration::from_millis(500)).await,
|
||||
Err(error) => {
|
||||
tracing::warn!(worker_id, %error, "mail queue unavailable");
|
||||
tokio::time::sleep(Duration::from_secs(2)).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -1,7 +1,15 @@
|
||||
use crate::AppState;
|
||||
use crate::{
|
||||
error::ErrorResponse,
|
||||
infra::middleware::{Language, UserId},
|
||||
repositories::user_repository::UserRepository,
|
||||
utils::{
|
||||
i18n::{message, ZH_CN},
|
||||
jwt::TokenType,
|
||||
},
|
||||
AppState,
|
||||
};
|
||||
use axum::{
|
||||
extract::{Request, State},
|
||||
http::{HeaderMap, StatusCode},
|
||||
middleware::Next,
|
||||
response::Response,
|
||||
};
|
||||
@@ -10,42 +18,76 @@ use serde::Deserialize;
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub struct Claims {
|
||||
pub sub: String, // user_id
|
||||
pub sub: String,
|
||||
#[allow(dead_code)]
|
||||
pub exp: usize,
|
||||
pub token_type: TokenType,
|
||||
}
|
||||
|
||||
/// JWT 认证中间件
|
||||
pub async fn auth_middleware(
|
||||
State(state): State<AppState>,
|
||||
headers: HeaderMap,
|
||||
mut req: Request,
|
||||
next: Next,
|
||||
) -> Result<Response, StatusCode> {
|
||||
// 1. 提取 Authorization header
|
||||
let auth_header = headers
|
||||
) -> Result<Response, ErrorResponse> {
|
||||
let language = req
|
||||
.extensions()
|
||||
.get::<Language>()
|
||||
.map(|v| v.0.as_str())
|
||||
.unwrap_or(ZH_CN);
|
||||
let header = req
|
||||
.headers()
|
||||
.get("Authorization")
|
||||
.and_then(|h| h.to_str().ok())
|
||||
.ok_or(StatusCode::UNAUTHORIZED)?;
|
||||
|
||||
if !auth_header.starts_with("Bearer ") {
|
||||
return Err(StatusCode::UNAUTHORIZED);
|
||||
}
|
||||
|
||||
let token = &auth_header[7..];
|
||||
|
||||
// 2. 验证 JWT
|
||||
let jwt_secret = &state.config.auth.jwt_secret;
|
||||
|
||||
let token_data = decode::<Claims>(
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.ok_or_else(|| {
|
||||
ErrorResponse::unauthorized(message(
|
||||
language,
|
||||
"缺少认证请求头",
|
||||
"Missing authorization header",
|
||||
))
|
||||
})?;
|
||||
let token = header
|
||||
.strip_prefix("Bearer ")
|
||||
.filter(|v| !v.is_empty())
|
||||
.ok_or_else(|| {
|
||||
ErrorResponse::unauthorized(message(
|
||||
language,
|
||||
"认证格式无效",
|
||||
"Invalid authorization format",
|
||||
))
|
||||
})?;
|
||||
let claims = decode::<Claims>(
|
||||
token,
|
||||
&DecodingKey::from_secret(jwt_secret.as_ref()),
|
||||
&DecodingKey::from_secret(state.config.auth.jwt_secret.as_bytes()),
|
||||
&Validation::default(),
|
||||
)
|
||||
.map_err(|_| StatusCode::UNAUTHORIZED)?;
|
||||
|
||||
// 3. 将 user_id 添加到请求扩展
|
||||
req.extensions_mut().insert(token_data.claims.sub);
|
||||
|
||||
.map_err(|_| {
|
||||
ErrorResponse::unauthorized(message(
|
||||
language,
|
||||
"令牌无效或已过期",
|
||||
"Token is invalid or expired",
|
||||
))
|
||||
})?
|
||||
.claims;
|
||||
if claims.token_type != TokenType::Access {
|
||||
return Err(ErrorResponse::unauthorized(message(
|
||||
language,
|
||||
"令牌类型无效",
|
||||
"Invalid token type",
|
||||
)));
|
||||
}
|
||||
let user = UserRepository::new(state.pool.clone())
|
||||
.find_by_id_raw(&claims.sub)
|
||||
.await
|
||||
.map_err(|_| {
|
||||
ErrorResponse::internal(message(language, "验证用户失败", "Failed to verify user"))
|
||||
})?;
|
||||
if user.map(|v| v.deleted_at.is_some()).unwrap_or(true) {
|
||||
return Err(ErrorResponse::unauthorized(message(
|
||||
language,
|
||||
"用户不存在或已删除",
|
||||
"User not found or deleted",
|
||||
)));
|
||||
}
|
||||
req.extensions_mut().insert(UserId(claims.sub));
|
||||
Ok(next.run(req).await)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,19 @@
|
||||
use crate::error::ErrorResponse;
|
||||
use axum::{
|
||||
body::HttpBody,
|
||||
extract::{Request, State},
|
||||
middleware::Next,
|
||||
response::{IntoResponse, Response},
|
||||
};
|
||||
|
||||
pub async fn enforce_body_limit(State(limit): State<usize>, req: Request, next: Next) -> Response {
|
||||
if req
|
||||
.body()
|
||||
.size_hint()
|
||||
.upper()
|
||||
.is_some_and(|size| size > limit as u64)
|
||||
{
|
||||
return ErrorResponse::payload_too_large("request body too large").into_response();
|
||||
}
|
||||
next.run(req).await
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
use crate::utils::i18n::{EN, ZH_CN};
|
||||
use async_trait::async_trait;
|
||||
use axum::{
|
||||
extract::{FromRequestParts, Request},
|
||||
http::{request::Parts, StatusCode},
|
||||
middleware::Next,
|
||||
response::Response,
|
||||
};
|
||||
use std::ops::Deref;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct Language(pub String);
|
||||
impl Deref for Language {
|
||||
type Target = String;
|
||||
fn deref(&self) -> &Self::Target {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl<S: Send + Sync> FromRequestParts<S> for Language {
|
||||
type Rejection = (StatusCode, &'static str);
|
||||
async fn from_request_parts(parts: &mut Parts, _: &S) -> Result<Self, Self::Rejection> {
|
||||
parts.extensions.get::<Language>().cloned().ok_or((
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
"language context missing",
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn language_middleware(mut req: Request, next: Next) -> Response {
|
||||
let language = req
|
||||
.headers()
|
||||
.get("Accept-Language")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.and_then(|v| v.split(',').next())
|
||||
.map(str::trim)
|
||||
.filter(|v| *v == ZH_CN || *v == EN)
|
||||
.unwrap_or(ZH_CN)
|
||||
.to_string();
|
||||
req.extensions_mut().insert(Language(language));
|
||||
next.run(req).await
|
||||
}
|
||||
@@ -1,71 +1,68 @@
|
||||
use axum::{extract::Request, response::Response};
|
||||
use axum::{
|
||||
body::{to_bytes, Body, Bytes},
|
||||
extract::Request,
|
||||
middleware::Next,
|
||||
response::Response,
|
||||
};
|
||||
use std::time::Instant;
|
||||
|
||||
/// Request ID 标记
|
||||
#[derive(Clone)]
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct RequestId(pub String);
|
||||
|
||||
/// 请求日志中间件
|
||||
pub async fn request_logging_middleware(
|
||||
mut req: Request,
|
||||
next: axum::middleware::Next,
|
||||
) -> Response {
|
||||
let start = Instant::now();
|
||||
|
||||
// 提取请求信息
|
||||
let method = req.method().clone();
|
||||
let path = req.uri().path().to_string();
|
||||
let query = req.uri().query().map(|s| s.to_string());
|
||||
|
||||
// 生成请求 ID
|
||||
let request_id = uuid::Uuid::new_v4().to_string();
|
||||
|
||||
// 将 request_id 存储到请求扩展中
|
||||
req.extensions_mut().insert(RequestId(request_id.clone()));
|
||||
|
||||
// 第1条日志:请求开始
|
||||
let separator = "=".repeat(80);
|
||||
let header = format!("{} {}", method, path);
|
||||
|
||||
tracing::info!("{}", separator);
|
||||
tracing::info!("{}", header);
|
||||
tracing::info!("{}", separator);
|
||||
|
||||
let now_beijing = chrono::Local::now().format("%Y-%m-%d %H:%M:%S%.3f");
|
||||
let query_str = query.as_deref().unwrap_or("无");
|
||||
tracing::info!(
|
||||
"[{}] 📥 查询参数: {} | 时间: {}",
|
||||
request_id,
|
||||
query_str,
|
||||
now_beijing
|
||||
);
|
||||
|
||||
// 调用下一个处理器
|
||||
let response = next.run(req).await;
|
||||
|
||||
// 第3条日志:请求完成
|
||||
let duration = start.elapsed();
|
||||
let status = response.status();
|
||||
tracing::info!(
|
||||
"[{}] ✅ 状态码: {} | 耗时: {}ms",
|
||||
request_id,
|
||||
status.as_u16(),
|
||||
duration.as_millis()
|
||||
);
|
||||
|
||||
tracing::info!("{}", separator);
|
||||
|
||||
response
|
||||
}
|
||||
|
||||
/// 请求日志辅助工具
|
||||
pub fn log_info<T: std::fmt::Debug>(request_id: &RequestId, label: &str, data: T) {
|
||||
let data_str = format!("{:?}", data);
|
||||
let truncated = if data_str.len() > 300 {
|
||||
format!("{}...", &data_str[..300])
|
||||
pub fn truncate_string(value: &str, max: usize) -> String {
|
||||
if value.chars().count() > max {
|
||||
value.chars().take(max).collect::<String>() + "....."
|
||||
} else {
|
||||
data_str
|
||||
};
|
||||
|
||||
tracing::info!("[{}] 🔧 {} | {}", request_id.0, label, truncated);
|
||||
value.to_string()
|
||||
}
|
||||
}
|
||||
fn truncate_json(value: &mut serde_json::Value, max: usize) {
|
||||
match value {
|
||||
serde_json::Value::String(v) => *v = truncate_string(v, max),
|
||||
serde_json::Value::Array(v) => v.iter_mut().for_each(|v| truncate_json(v, max)),
|
||||
serde_json::Value::Object(v) => v.values_mut().for_each(|v| truncate_json(v, max)),
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
fn pretty(bytes: &Bytes) -> String {
|
||||
let raw = String::from_utf8_lossy(bytes);
|
||||
match serde_json::from_str::<serde_json::Value>(&raw) {
|
||||
Ok(mut value) => {
|
||||
truncate_json(&mut value, 50);
|
||||
serde_json::to_string_pretty(&value).unwrap_or_else(|_| raw.into_owned())
|
||||
}
|
||||
Err(_) => truncate_string(&raw, 50),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn request_logging_middleware(mut req: Request<Body>, next: Next) -> Response {
|
||||
let started = Instant::now();
|
||||
let id = uuid::Uuid::new_v4().to_string();
|
||||
let method = req.method().clone();
|
||||
let uri = req.uri().clone();
|
||||
req.extensions_mut().insert(RequestId(id.clone()));
|
||||
let (parts, body) = req.into_parts();
|
||||
let request_bytes = to_bytes(body, usize::MAX).await.unwrap_or_default();
|
||||
tracing::info!(request_id=%id, %method, %uri, body=%pretty(&request_bytes), "request started");
|
||||
let response = next
|
||||
.run(Request::from_parts(parts, Body::from(request_bytes)))
|
||||
.await;
|
||||
let status = response.status();
|
||||
let (parts, body) = response.into_parts();
|
||||
let response_bytes = to_bytes(body, usize::MAX).await.unwrap_or_default();
|
||||
tracing::info!(request_id=%id, status=%status, elapsed_ms=started.elapsed().as_millis(), body=%pretty(&response_bytes), "request completed");
|
||||
Response::from_parts(parts, Body::from(response_bytes))
|
||||
}
|
||||
|
||||
pub fn log_info<T: std::fmt::Debug>(request_id: &RequestId, label: &str, data: T) {
|
||||
tracing::info!(request_id=%request_id.0, %label, data=?data);
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
#[test]
|
||||
fn truncates_at_character_boundary() {
|
||||
assert_eq!(truncate_string("中文测试", 2), "中文.....");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,2 +1,10 @@
|
||||
pub mod auth;
|
||||
pub mod body_limit;
|
||||
pub mod language;
|
||||
pub mod logging;
|
||||
pub mod rate_limit;
|
||||
pub mod security;
|
||||
pub mod user_id;
|
||||
|
||||
pub use language::Language;
|
||||
pub use user_id::UserId;
|
||||
|
||||
@@ -0,0 +1,64 @@
|
||||
use axum::{
|
||||
extract::{ConnectInfo, Request, State},
|
||||
http::StatusCode,
|
||||
middleware::Next,
|
||||
response::{IntoResponse, Response},
|
||||
};
|
||||
use std::{
|
||||
collections::HashMap,
|
||||
net::SocketAddr,
|
||||
sync::Arc,
|
||||
time::{Duration, Instant},
|
||||
};
|
||||
use tokio::sync::Mutex;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct RateLimiter {
|
||||
limit: u32,
|
||||
entries: Arc<Mutex<HashMap<String, (Instant, u32)>>>,
|
||||
}
|
||||
impl RateLimiter {
|
||||
pub fn new(limit: u32) -> Self {
|
||||
Self {
|
||||
limit,
|
||||
entries: Arc::new(Mutex::new(HashMap::new())),
|
||||
}
|
||||
}
|
||||
async fn allow(&self, key: String) -> bool {
|
||||
let mut entries = self.entries.lock().await;
|
||||
let value = entries.entry(key).or_insert((Instant::now(), 0));
|
||||
if value.0.elapsed() >= Duration::from_secs(60) {
|
||||
*value = (Instant::now(), 0);
|
||||
}
|
||||
if value.1 >= self.limit {
|
||||
false
|
||||
} else {
|
||||
value.1 += 1;
|
||||
true
|
||||
}
|
||||
}
|
||||
}
|
||||
pub async fn rate_limit_middleware(
|
||||
State(limiter): State<RateLimiter>,
|
||||
req: Request,
|
||||
next: Next,
|
||||
) -> Response {
|
||||
let key = req
|
||||
.extensions()
|
||||
.get::<ConnectInfo<SocketAddr>>()
|
||||
.map(|v| v.0.ip().to_string())
|
||||
.or_else(|| {
|
||||
req.headers()
|
||||
.get("x-forwarded-for")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.and_then(|v| v.split(',').next())
|
||||
.map(str::trim)
|
||||
.map(str::to_string)
|
||||
})
|
||||
.unwrap_or_else(|| "unknown".into());
|
||||
if limiter.allow(key).await {
|
||||
next.run(req).await
|
||||
} else {
|
||||
(StatusCode::TOO_MANY_REQUESTS, "rate limit exceeded").into_response()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
use crate::error::ErrorResponse;
|
||||
use axum::response::IntoResponse;
|
||||
use axum::{
|
||||
extract::Request,
|
||||
http::{
|
||||
header::{HeaderName, HeaderValue},
|
||||
StatusCode,
|
||||
},
|
||||
middleware::Next,
|
||||
response::Response,
|
||||
};
|
||||
|
||||
pub async fn security_headers(req: Request, next: Next) -> Response {
|
||||
let mut response = next.run(req).await;
|
||||
let headers = response.headers_mut();
|
||||
headers.insert(
|
||||
"x-content-type-options",
|
||||
HeaderValue::from_static("nosniff"),
|
||||
);
|
||||
headers.insert("x-frame-options", HeaderValue::from_static("DENY"));
|
||||
headers.insert("referrer-policy", HeaderValue::from_static("no-referrer"));
|
||||
headers.insert(
|
||||
HeaderName::from_static("permissions-policy"),
|
||||
HeaderValue::from_static("camera=(), microphone=(), geolocation=()"),
|
||||
);
|
||||
response
|
||||
}
|
||||
|
||||
pub async fn fallback_404() -> Response {
|
||||
ErrorResponse::not_found("route not found").into_response()
|
||||
}
|
||||
pub async fn fallback_405() -> Response {
|
||||
ErrorResponse {
|
||||
status: StatusCode::METHOD_NOT_ALLOWED,
|
||||
message: "method not allowed".into(),
|
||||
}
|
||||
.into_response()
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
use crate::{
|
||||
error::ErrorResponse,
|
||||
infra::middleware::Language,
|
||||
utils::i18n::{message, ZH_CN},
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
use axum::extract::FromRequestParts;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct UserId(pub String);
|
||||
|
||||
#[async_trait]
|
||||
impl<S: Send + Sync> FromRequestParts<S> for UserId {
|
||||
type Rejection = ErrorResponse;
|
||||
async fn from_request_parts(
|
||||
parts: &mut axum::http::request::Parts,
|
||||
_: &S,
|
||||
) -> Result<Self, Self::Rejection> {
|
||||
let lang = parts
|
||||
.extensions
|
||||
.get::<Language>()
|
||||
.map(|v| v.0.as_str())
|
||||
.unwrap_or(ZH_CN);
|
||||
parts.extensions.get::<UserId>().cloned().ok_or_else(|| {
|
||||
ErrorResponse::unauthorized(message(lang, "未找到用户身份", "User identity not found"))
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1,2 +1,3 @@
|
||||
pub mod mail;
|
||||
pub mod middleware;
|
||||
pub mod redis;
|
||||
|
||||
@@ -52,11 +52,7 @@ impl RedisClient {
|
||||
}
|
||||
|
||||
/// 使用 RedisKey 设置 JSON 值
|
||||
pub async fn set_key<T: Serialize>(
|
||||
&self,
|
||||
key: &RedisKey,
|
||||
value: &T,
|
||||
) -> redis::RedisResult<()> {
|
||||
pub async fn set_key<T: Serialize>(&self, key: &RedisKey, value: &T) -> redis::RedisResult<()> {
|
||||
let json = serde_json::to_string(value).map_err(|e| {
|
||||
redis::RedisError::from((
|
||||
redis::ErrorKind::TypeError,
|
||||
@@ -132,4 +128,40 @@ impl RedisClient {
|
||||
let mut c = self.conn.lock().await;
|
||||
c.expire(key.build(), seconds as i64).await
|
||||
}
|
||||
|
||||
pub async fn queue_push<T: Serialize>(&self, key: &str, value: &T) -> redis::RedisResult<()> {
|
||||
let json = serde_json::to_string(value).map_err(|e| {
|
||||
redis::RedisError::from((
|
||||
redis::ErrorKind::TypeError,
|
||||
"JSON serialization failed",
|
||||
e.to_string(),
|
||||
))
|
||||
})?;
|
||||
let mut connection = self.conn.lock().await;
|
||||
connection.lpush(key, json).await
|
||||
}
|
||||
|
||||
pub async fn queue_pop<T: for<'de> serde::Deserialize<'de>>(
|
||||
&self,
|
||||
key: &str,
|
||||
) -> redis::RedisResult<Option<T>> {
|
||||
let mut connection = self.conn.lock().await;
|
||||
let value: Option<String> = connection.rpop(key, None).await?;
|
||||
value
|
||||
.map(|json| {
|
||||
serde_json::from_str(&json).map_err(|e| {
|
||||
redis::RedisError::from((
|
||||
redis::ErrorKind::TypeError,
|
||||
"JSON deserialization failed",
|
||||
e.to_string(),
|
||||
))
|
||||
})
|
||||
})
|
||||
.transpose()
|
||||
}
|
||||
|
||||
pub async fn queue_len(&self, key: &str) -> redis::RedisResult<usize> {
|
||||
let mut connection = self.conn.lock().await;
|
||||
connection.llen(key).await
|
||||
}
|
||||
}
|
||||
|
||||
+247
-84
@@ -1,3 +1,4 @@
|
||||
#![recursion_limit = "512"]
|
||||
mod cli;
|
||||
mod config;
|
||||
mod db;
|
||||
@@ -10,28 +11,120 @@ mod services;
|
||||
mod utils;
|
||||
|
||||
use axum::{
|
||||
http::{HeaderValue, Method},
|
||||
middleware,
|
||||
routing::{get, post},
|
||||
Router,
|
||||
};
|
||||
use clap::Parser;
|
||||
use cli::CliArgs;
|
||||
use tower_http::cors::{Any, CorsLayer};
|
||||
use std::time::Duration;
|
||||
use tower::limit::ConcurrencyLimitLayer;
|
||||
use tower_http::{
|
||||
catch_panic::CatchPanicLayer, cors::CorsLayer, limit::RequestBodyLimitLayer,
|
||||
timeout::TimeoutLayer,
|
||||
};
|
||||
use tracing_subscriber::{layer::SubscriberExt, util::SubscriberInitExt};
|
||||
|
||||
/// 应用状态
|
||||
#[derive(Clone)]
|
||||
pub struct AppState {
|
||||
pub pool: db::DbPool,
|
||||
pub config: config::app::AppConfig,
|
||||
pub redis_client: infra::redis::redis_client::RedisClient,
|
||||
pub redis_client: Option<infra::redis::redis_client::RedisClient>,
|
||||
}
|
||||
|
||||
fn cors_layer(config: &config::server::ServerConfig) -> anyhow::Result<CorsLayer> {
|
||||
let origins = config
|
||||
.cors_origins
|
||||
.iter()
|
||||
.map(|v| HeaderValue::from_str(v))
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
Ok(CorsLayer::new()
|
||||
.allow_origin(origins)
|
||||
.allow_methods([Method::GET, Method::POST, Method::PUT, Method::DELETE])
|
||||
.allow_headers(tower_http::cors::Any))
|
||||
}
|
||||
|
||||
fn build_router(state: AppState) -> anyhow::Result<Router> {
|
||||
let mut public = Router::new()
|
||||
.route("/health", get(handlers::health::health_check))
|
||||
.route("/info", get(handlers::health::server_info))
|
||||
.route("/auth/register", post(handlers::auth::register))
|
||||
.route("/auth/login", post(handlers::auth::login))
|
||||
.route("/auth/refresh", post(handlers::auth::refresh));
|
||||
if state.config.email.enabled {
|
||||
public = public
|
||||
.route(
|
||||
"/auth/request-verification-code",
|
||||
post(handlers::email::send_verification_code),
|
||||
)
|
||||
.route(
|
||||
"/auth/reset-password",
|
||||
post(handlers::email::reset_password),
|
||||
);
|
||||
}
|
||||
let mut protected = Router::new()
|
||||
.route("/auth/logout", post(handlers::auth::delete_refresh_token))
|
||||
.route(
|
||||
"/auth/delete-refresh-token",
|
||||
post(handlers::auth::delete_refresh_token),
|
||||
)
|
||||
.route("/auth/delete", post(handlers::auth::delete_account))
|
||||
.route(
|
||||
"/api/user/profile",
|
||||
get(handlers::user_profile::get_profile)
|
||||
.put(handlers::user_profile::update_profile)
|
||||
.delete(handlers::user_profile::delete_profile),
|
||||
);
|
||||
if state.config.email.enabled {
|
||||
protected = protected
|
||||
.route("/api/email/latest-log", get(handlers::email::latest_log))
|
||||
.route(
|
||||
"/api/email/queue-status",
|
||||
get(handlers::email::queue_status),
|
||||
);
|
||||
}
|
||||
protected = protected.route_layer(middleware::from_fn_with_state(
|
||||
state.clone(),
|
||||
infra::middleware::auth::auth_middleware,
|
||||
));
|
||||
let limiter =
|
||||
infra::middleware::rate_limit::RateLimiter::new(state.config.server.rate_limit_per_minute);
|
||||
let server = state.config.server.clone();
|
||||
Ok(public
|
||||
.merge(protected)
|
||||
.fallback(infra::middleware::security::fallback_404)
|
||||
.method_not_allowed_fallback(infra::middleware::security::fallback_405)
|
||||
.layer(middleware::from_fn(
|
||||
infra::middleware::security::security_headers,
|
||||
))
|
||||
.layer(middleware::from_fn(
|
||||
infra::middleware::language::language_middleware,
|
||||
))
|
||||
.layer(cors_layer(&server)?)
|
||||
.layer(middleware::from_fn(
|
||||
infra::middleware::logging::request_logging_middleware,
|
||||
))
|
||||
.layer(middleware::from_fn_with_state(
|
||||
limiter,
|
||||
infra::middleware::rate_limit::rate_limit_middleware,
|
||||
))
|
||||
.layer(TimeoutLayer::new(Duration::from_secs(
|
||||
server.request_timeout_seconds,
|
||||
)))
|
||||
.layer(CatchPanicLayer::new())
|
||||
.layer(ConcurrencyLimitLayer::new(server.concurrency_limit))
|
||||
.layer(RequestBodyLimitLayer::new(server.max_body_bytes))
|
||||
.layer(middleware::from_fn_with_state(
|
||||
server.max_body_bytes,
|
||||
infra::middleware::body_limit::enforce_body_limit,
|
||||
))
|
||||
.with_state(state))
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> anyhow::Result<()> {
|
||||
// 解析命令行参数
|
||||
let args = CliArgs::parse();
|
||||
|
||||
// 初始化日志
|
||||
tracing_subscriber::registry()
|
||||
.with(
|
||||
tracing_subscriber::EnvFilter::try_from_default_env()
|
||||
@@ -39,93 +132,163 @@ async fn main() -> anyhow::Result<()> {
|
||||
)
|
||||
.with(tracing_subscriber::fmt::layer())
|
||||
.init();
|
||||
|
||||
// 打印启动信息
|
||||
args.print_startup_info();
|
||||
|
||||
// 设置工作目录(如果指定)
|
||||
if let Some(ref work_dir) = args.work_dir {
|
||||
std::env::set_current_dir(work_dir).ok();
|
||||
println!("Working directory set to: {}", work_dir.display());
|
||||
if let Some(ref dir) = args.work_dir {
|
||||
std::env::set_current_dir(dir)?;
|
||||
}
|
||||
|
||||
// 解析配置文件路径(可选)
|
||||
let config_path = args.resolve_config_path();
|
||||
|
||||
// 加载配置(支持 CLI 覆盖)
|
||||
// 如果没有配置文件,将仅使用环境变量和默认值
|
||||
let config = config::app::AppConfig::load_with_overrides(
|
||||
config_path,
|
||||
args.resolve_config_path(),
|
||||
args.get_overrides(),
|
||||
args.env.as_str(),
|
||||
)?;
|
||||
|
||||
tracing::info!("Configuration loaded successfully");
|
||||
tracing::info!("Environment: {}", args.env.as_str());
|
||||
tracing::info!("Debug mode: {}", args.is_debug_enabled());
|
||||
|
||||
// 初始化数据库(自动创建数据库和表)
|
||||
let pool = db::init_database(&config.database).await?;
|
||||
|
||||
// 初始化 Redis 客户端
|
||||
let redis_client = infra::redis::redis_client::RedisClient::new(&config.redis.build_url())
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("Redis 初始化失败: {}", e))?;
|
||||
|
||||
tracing::info!("Redis 连接池初始化成功");
|
||||
|
||||
// 创建应用状态
|
||||
let app_state = AppState {
|
||||
pool: pool.clone(),
|
||||
config: config.clone(),
|
||||
let redis_client = if config.redis.enabled {
|
||||
match infra::redis::redis_client::RedisClient::new(&config.redis.build_url()).await {
|
||||
Ok(client) => {
|
||||
tracing::info!("Redis connected");
|
||||
Some(client)
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(%error, "Redis unavailable; dependent capabilities are disabled");
|
||||
None
|
||||
}
|
||||
}
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let address = format!("{}:{}", config.server.host, config.server.port);
|
||||
let state = AppState {
|
||||
pool,
|
||||
config,
|
||||
redis_client,
|
||||
};
|
||||
|
||||
// ========== 公开路由 ==========
|
||||
let public_routes = Router::new()
|
||||
.route("/health", get(handlers::health::health_check))
|
||||
.route("/info", get(handlers::health::server_info))
|
||||
.route("/auth/register", post(handlers::auth::register))
|
||||
.route("/auth/login", post(handlers::auth::login))
|
||||
.route("/auth/refresh", post(handlers::auth::refresh));
|
||||
|
||||
// ========== 受保护路由 ==========
|
||||
let protected_routes = Router::new()
|
||||
.route("/auth/delete", post(handlers::auth::delete_account))
|
||||
.route(
|
||||
"/auth/delete-refresh-token",
|
||||
post(handlers::auth::delete_refresh_token),
|
||||
if state.config.email.enabled && state.config.email.queue_enabled {
|
||||
if let Some(redis) = state.redis_client.clone() {
|
||||
infra::mail::worker::start(redis, state.config.email.clone(), state.pool.clone());
|
||||
} else {
|
||||
tracing::warn!("mail queue requested but Redis is unavailable");
|
||||
}
|
||||
}
|
||||
let listener = tokio::net::TcpListener::bind(&address).await?;
|
||||
tracing::info!(%address, "server listening");
|
||||
axum::serve(
|
||||
listener,
|
||||
build_router(state)?.into_make_service_with_connect_info::<std::net::SocketAddr>(),
|
||||
)
|
||||
// JWT 认证中间件(仅应用于受保护路由)
|
||||
.route_layer(axum::middleware::from_fn_with_state(
|
||||
app_state.clone(),
|
||||
infra::middleware::auth::auth_middleware,
|
||||
));
|
||||
|
||||
// ========== 合并路由 ==========
|
||||
let app = public_routes
|
||||
.merge(protected_routes)
|
||||
// CORS(应用于所有路由)
|
||||
.layer(
|
||||
CorsLayer::new()
|
||||
.allow_origin(Any)
|
||||
.allow_methods(Any)
|
||||
.allow_headers(Any)
|
||||
)
|
||||
// 日志中间件(应用于所有路由)
|
||||
.layer(axum::middleware::from_fn_with_state(
|
||||
app_state.clone(),
|
||||
infra::middleware::logging::request_logging_middleware,
|
||||
))
|
||||
.with_state(app_state);
|
||||
|
||||
// 启动服务器
|
||||
let addr = format!("{}:{}", config.server.host, config.server.port);
|
||||
let listener = tokio::net::TcpListener::bind(&addr).await?;
|
||||
tracing::info!("Server listening on {}", addr);
|
||||
tracing::info!("Press Ctrl+C to stop");
|
||||
|
||||
axum::serve(listener, app).await?;
|
||||
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod route_tests {
|
||||
use super::*;
|
||||
use axum::{
|
||||
body::{to_bytes, Body},
|
||||
http::{Request, StatusCode},
|
||||
};
|
||||
use tower::ServiceExt;
|
||||
|
||||
async fn test_app(max_body_bytes: usize) -> Router {
|
||||
let mut config = config::app::AppConfig::load_from_path("config/development.toml").unwrap();
|
||||
let db_path =
|
||||
std::env::temp_dir().join(format!("web-rust-template-{}.sqlite", uuid::Uuid::new_v4()));
|
||||
config.database.path = Some(db_path);
|
||||
config.redis.enabled = false;
|
||||
config.email.enabled = false;
|
||||
config.server.max_body_bytes = max_body_bytes;
|
||||
let pool = db::init_database(&config.database).await.unwrap();
|
||||
build_router(AppState {
|
||||
pool,
|
||||
config,
|
||||
redis_client: None,
|
||||
})
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn health_works_without_redis() {
|
||||
let response = test_app(1024)
|
||||
.await
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.uri("/health")
|
||||
.body(Body::empty())
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let body = to_bytes(response.into_body(), 4096).await.unwrap();
|
||||
assert!(String::from_utf8_lossy(&body).contains("\"redis\":false"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn fallback_is_structured_json() {
|
||||
let response = test_app(1024)
|
||||
.await
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.uri("/missing")
|
||||
.body(Body::empty())
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::NOT_FOUND);
|
||||
let body = to_bytes(response.into_body(), 4096).await.unwrap();
|
||||
assert!(String::from_utf8_lossy(&body).contains("\"code\":404"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn rejects_oversized_body() {
|
||||
let response = test_app(8)
|
||||
.await
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.method("POST")
|
||||
.uri("/auth/login")
|
||||
.header("content-type", "application/json")
|
||||
.body(Body::from("0123456789"))
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::PAYLOAD_TOO_LARGE);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn protected_route_requires_token_and_uses_language() {
|
||||
let response = test_app(1024)
|
||||
.await
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.uri("/api/user/profile")
|
||||
.header("accept-language", "en")
|
||||
.body(Body::empty())
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
|
||||
let body = to_bytes(response.into_body(), 4096).await.unwrap();
|
||||
assert!(String::from_utf8_lossy(&body).contains("Missing authorization header"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn method_not_allowed_is_structured() {
|
||||
let response = test_app(1024)
|
||||
.await
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.method("PATCH")
|
||||
.uri("/health")
|
||||
.body(Body::empty())
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::METHOD_NOT_ALLOWED);
|
||||
let body = to_bytes(response.into_body(), 4096).await.unwrap();
|
||||
assert!(String::from_utf8_lossy(&body).contains("\"code\":405"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,39 @@
|
||||
use crate::domain::entities::email_logs;
|
||||
use anyhow::Result;
|
||||
use sea_orm::{
|
||||
ActiveModelTrait, ColumnTrait, DatabaseConnection, EntityTrait, QueryFilter, QueryOrder, Set,
|
||||
};
|
||||
pub struct EmailLogRepository {
|
||||
db: DatabaseConnection,
|
||||
}
|
||||
impl EmailLogRepository {
|
||||
pub fn new(db: DatabaseConnection) -> Self {
|
||||
Self { db }
|
||||
}
|
||||
pub async fn add(
|
||||
&self,
|
||||
user_id: Option<String>,
|
||||
recipient: String,
|
||||
status: String,
|
||||
error: Option<String>,
|
||||
) -> Result<()> {
|
||||
email_logs::ActiveModel {
|
||||
user_id: Set(user_id),
|
||||
recipient: Set(recipient),
|
||||
kind: Set("verification".into()),
|
||||
status: Set(status),
|
||||
error: Set(error),
|
||||
..Default::default()
|
||||
}
|
||||
.insert(&self.db)
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
pub async fn latest(&self, user_id: &str) -> Result<Option<email_logs::Model>> {
|
||||
Ok(email_logs::Entity::find()
|
||||
.filter(email_logs::Column::UserId.eq(user_id))
|
||||
.order_by_desc(email_logs::Column::CreatedAt)
|
||||
.one(&self.db)
|
||||
.await?)
|
||||
}
|
||||
}
|
||||
@@ -1,2 +1,3 @@
|
||||
pub mod email_log_repository;
|
||||
pub mod user_profile_repository;
|
||||
pub mod user_repository;
|
||||
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
use crate::domain::{dto::user::UpdateProfileRequest, entities::user_profiles};
|
||||
use anyhow::Result;
|
||||
use sea_orm::{ActiveModelTrait, DatabaseConnection, EntityTrait, Set, TryIntoModel};
|
||||
|
||||
pub struct UserProfileRepository {
|
||||
db: DatabaseConnection,
|
||||
}
|
||||
impl UserProfileRepository {
|
||||
pub fn new(db: DatabaseConnection) -> Self {
|
||||
Self { db }
|
||||
}
|
||||
pub async fn get(&self, user_id: &str) -> Result<Option<user_profiles::Model>> {
|
||||
Ok(user_profiles::Entity::find_by_id(user_id)
|
||||
.one(&self.db)
|
||||
.await?)
|
||||
}
|
||||
pub async fn upsert(
|
||||
&self,
|
||||
user_id: String,
|
||||
input: UpdateProfileRequest,
|
||||
) -> Result<user_profiles::Model> {
|
||||
let existing = self.get(&user_id).await?;
|
||||
let mut model = existing
|
||||
.map(Into::into)
|
||||
.unwrap_or_else(|| user_profiles::ActiveModel {
|
||||
user_id: Set(user_id),
|
||||
..Default::default()
|
||||
});
|
||||
model.display_name = Set(input.display_name);
|
||||
model.avatar_url = Set(input.avatar_url);
|
||||
model.bio = Set(input.bio);
|
||||
Ok(model.save(&self.db).await?.try_into_model()?)
|
||||
}
|
||||
pub async fn delete(&self, user_id: &str) -> Result<()> {
|
||||
user_profiles::Entity::delete_by_id(user_id)
|
||||
.exec(&self.db)
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,9 @@
|
||||
use sea_orm::{EntityTrait, QueryFilter, ColumnTrait, DatabaseConnection, Set, ActiveModelTrait, PaginatorTrait};
|
||||
use crate::domain::entities::users;
|
||||
use anyhow::Result;
|
||||
use sea_orm::{
|
||||
ActiveModelTrait, ColumnTrait, DatabaseConnection, EntityTrait, PaginatorTrait, QueryFilter,
|
||||
Set,
|
||||
};
|
||||
|
||||
/// 用户数据访问仓库
|
||||
pub struct UserRepository {
|
||||
@@ -16,6 +19,7 @@ impl UserRepository {
|
||||
pub async fn find_by_email(&self, email: &str) -> Result<Option<users::Model>> {
|
||||
let user = users::Entity::find()
|
||||
.filter(users::Column::Email.eq(email))
|
||||
.filter(users::Column::DeletedAt.is_null())
|
||||
.one(&self.db)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("查询失败: {}", e))?;
|
||||
@@ -23,6 +27,10 @@ impl UserRepository {
|
||||
Ok(user)
|
||||
}
|
||||
|
||||
pub async fn find_by_id_raw(&self, id: &str) -> Result<Option<users::Model>> {
|
||||
Ok(users::Entity::find_by_id(id).one(&self.db).await?)
|
||||
}
|
||||
|
||||
/// 统计邮箱数量
|
||||
pub async fn count_by_email(&self, email: &str) -> Result<i64> {
|
||||
let count = users::Entity::find()
|
||||
@@ -56,8 +64,20 @@ impl UserRepository {
|
||||
Ok(user.map(|u| u.password_hash))
|
||||
}
|
||||
|
||||
pub async fn get_password_hash_by_id(&self, id: &str) -> Result<Option<String>> {
|
||||
Ok(users::Entity::find_by_id(id)
|
||||
.one(&self.db)
|
||||
.await?
|
||||
.map(|user| user.password_hash))
|
||||
}
|
||||
|
||||
/// 插入用户(created_at 和 updated_at 会自动填充),返回插入后的用户对象
|
||||
pub async fn insert(&self, id: String, email: String, password_hash: String) -> Result<users::Model> {
|
||||
pub async fn insert(
|
||||
&self,
|
||||
id: String,
|
||||
email: String,
|
||||
password_hash: String,
|
||||
) -> Result<users::Model> {
|
||||
let user_model = users::ActiveModel {
|
||||
id: Set(id),
|
||||
email: Set(email),
|
||||
@@ -66,7 +86,8 @@ impl UserRepository {
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let inserted_user = user_model.insert(&self.db)
|
||||
let inserted_user = user_model
|
||||
.insert(&self.db)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("插入失败: {}", e))?;
|
||||
|
||||
@@ -75,11 +96,25 @@ impl UserRepository {
|
||||
|
||||
/// 根据 ID 删除用户
|
||||
pub async fn delete_by_id(&self, id: &str) -> Result<()> {
|
||||
users::Entity::delete_by_id(id)
|
||||
.exec(&self.db)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("删除失败: {}", e))?;
|
||||
let user = users::Entity::find_by_id(id)
|
||||
.one(&self.db)
|
||||
.await?
|
||||
.ok_or_else(|| anyhow::anyhow!("用户不存在"))?;
|
||||
let mut active: users::ActiveModel = user.into();
|
||||
active.deleted_at = Set(Some(chrono::Utc::now().naive_utc()));
|
||||
active.update(&self.db).await?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn update_password_by_email(&self, email: &str, password_hash: String) -> Result<()> {
|
||||
let user = self
|
||||
.find_by_email(email)
|
||||
.await?
|
||||
.ok_or_else(|| anyhow::anyhow!("用户不存在"))?;
|
||||
let mut active: users::ActiveModel = user.into();
|
||||
active.password_hash = Set(password_hash);
|
||||
active.update(&self.db).await?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,26 +5,37 @@ use argon2::{
|
||||
};
|
||||
use rand::Rng;
|
||||
|
||||
use crate::utils::jwt::TokenService;
|
||||
use crate::domain::dto::auth::{RegisterRequest, LoginRequest, DeleteUserRequest};
|
||||
use crate::domain::entities::users;
|
||||
use crate::config::auth::AuthConfig;
|
||||
use crate::infra::redis::{redis_client::RedisClient, redis_key::{BusinessType, RedisKey}};
|
||||
use crate::domain::dto::auth::{DeleteUserRequest, LoginRequest, RegisterRequest};
|
||||
use crate::domain::entities::users;
|
||||
use crate::infra::redis::{
|
||||
redis_client::RedisClient,
|
||||
redis_key::{BusinessType, RedisKey},
|
||||
};
|
||||
use crate::repositories::user_repository::UserRepository;
|
||||
use crate::utils::jwt::TokenService;
|
||||
|
||||
pub struct AuthService {
|
||||
user_repo: UserRepository,
|
||||
redis_client: RedisClient,
|
||||
redis_client: Option<RedisClient>,
|
||||
auth_config: AuthConfig,
|
||||
}
|
||||
|
||||
impl AuthService {
|
||||
pub fn new(user_repo: UserRepository, redis_client: RedisClient, auth_config: AuthConfig) -> Self {
|
||||
Self { user_repo, redis_client, auth_config }
|
||||
pub fn new(
|
||||
user_repo: UserRepository,
|
||||
redis_client: Option<RedisClient>,
|
||||
auth_config: AuthConfig,
|
||||
) -> Self {
|
||||
Self {
|
||||
user_repo,
|
||||
redis_client,
|
||||
auth_config,
|
||||
}
|
||||
}
|
||||
|
||||
/// 哈希密码
|
||||
pub fn hash_password(&self, password: &str) -> Result<String> {
|
||||
pub fn hash_password(password: &str) -> Result<String> {
|
||||
let salt = SaltString::generate(&mut OsRng);
|
||||
let argon2 = Argon2::default();
|
||||
let password_hash = argon2
|
||||
@@ -62,7 +73,12 @@ impl AuthService {
|
||||
}
|
||||
|
||||
/// 保存 refresh_token 到 Redis
|
||||
async fn save_refresh_token(&self, user_id: &str, refresh_token: &str, expiration_days: i64) -> Result<()> {
|
||||
async fn save_refresh_token(
|
||||
&self,
|
||||
user_id: &str,
|
||||
refresh_token: &str,
|
||||
expiration_days: i64,
|
||||
) -> Result<()> {
|
||||
let key = RedisKey::new(BusinessType::Auth)
|
||||
.add_identifier("refresh_token")
|
||||
.add_identifier(user_id);
|
||||
@@ -70,6 +86,8 @@ impl AuthService {
|
||||
let expiration_seconds = expiration_days * 24 * 3600;
|
||||
|
||||
self.redis_client
|
||||
.as_ref()
|
||||
.ok_or_else(|| anyhow::anyhow!("Redis 服务不可用"))?
|
||||
.set_ex(&key.build(), refresh_token, expiration_seconds as u64)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("Redis 保存失败: {}", e))?;
|
||||
@@ -83,13 +101,17 @@ impl AuthService {
|
||||
.add_identifier("refresh_token")
|
||||
.add_identifier(user_id);
|
||||
|
||||
let token: Option<String> = self.redis_client
|
||||
let redis = self
|
||||
.redis_client
|
||||
.as_ref()
|
||||
.ok_or_else(|| anyhow::anyhow!("Redis 服务不可用"))?;
|
||||
let token: Option<String> = redis
|
||||
.get(&key.build())
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("Redis 查询失败: {}", e))?;
|
||||
|
||||
if token.is_some() {
|
||||
self.redis_client
|
||||
redis
|
||||
.delete_key(&key)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("Redis 删除失败: {}", e))?;
|
||||
@@ -105,6 +127,8 @@ impl AuthService {
|
||||
.add_identifier(user_id);
|
||||
|
||||
self.redis_client
|
||||
.as_ref()
|
||||
.ok_or_else(|| anyhow::anyhow!("Redis 服务不可用"))?
|
||||
.delete_key(&key)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("Redis 删除失败: {}", e))?;
|
||||
@@ -125,13 +149,16 @@ impl AuthService {
|
||||
}
|
||||
|
||||
// 2. 哈希密码
|
||||
let password_hash = self.hash_password(&request.password)?;
|
||||
let password_hash = Self::hash_password(&request.password)?;
|
||||
|
||||
// 3. 生成用户 ID
|
||||
let user_id = self.generate_unique_user_id().await?;
|
||||
|
||||
// 4. 插入数据库并获取包含真实 created_at 的用户对象
|
||||
let user = self.user_repo.insert(user_id.clone(), request.email, password_hash).await?;
|
||||
let user = self
|
||||
.user_repo
|
||||
.insert(user_id.clone(), request.email, password_hash)
|
||||
.await?;
|
||||
|
||||
// 5. 生成 token
|
||||
let (access_token, refresh_token) = TokenService::generate_token_pair(
|
||||
@@ -142,22 +169,32 @@ impl AuthService {
|
||||
)?;
|
||||
|
||||
// 6. 保存 refresh_token
|
||||
self.save_refresh_token(&user_id, &refresh_token, self.auth_config.refresh_token_expiration_days as i64).await?;
|
||||
if self.redis_client.is_some() {
|
||||
self.save_refresh_token(
|
||||
&user_id,
|
||||
&refresh_token,
|
||||
self.auth_config.refresh_token_expiration_days,
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
|
||||
Ok((user, access_token, refresh_token))
|
||||
}
|
||||
|
||||
/// 登录
|
||||
pub async fn login(
|
||||
&self,
|
||||
request: LoginRequest,
|
||||
) -> Result<(users::Model, String, String)> {
|
||||
pub async fn login(&self, request: LoginRequest) -> Result<(users::Model, String, String)> {
|
||||
// 1. 查询用户
|
||||
let user = self.user_repo.find_by_email(&request.email).await?
|
||||
let user = self
|
||||
.user_repo
|
||||
.find_by_email(&request.email)
|
||||
.await?
|
||||
.ok_or_else(|| anyhow::anyhow!("邮箱或密码错误"))?;
|
||||
|
||||
// 2. 验证密码
|
||||
let password_hash = self.user_repo.get_password_hash(&request.email).await?
|
||||
let password_hash = self
|
||||
.user_repo
|
||||
.get_password_hash(&request.email)
|
||||
.await?
|
||||
.ok_or_else(|| anyhow::anyhow!("邮箱或密码错误"))?;
|
||||
|
||||
let parsed_hash = PasswordHash::new(&password_hash)
|
||||
@@ -177,18 +214,23 @@ impl AuthService {
|
||||
)?;
|
||||
|
||||
// 4. 保存 refresh_token
|
||||
self.save_refresh_token(&user.id, &refresh_token, self.auth_config.refresh_token_expiration_days as i64).await?;
|
||||
if self.redis_client.is_some() {
|
||||
self.save_refresh_token(
|
||||
&user.id,
|
||||
&refresh_token,
|
||||
self.auth_config.refresh_token_expiration_days,
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
|
||||
Ok((user, access_token, refresh_token))
|
||||
}
|
||||
|
||||
/// 使用 refresh_token 刷新 access_token
|
||||
pub async fn refresh_access_token(
|
||||
&self,
|
||||
refresh_token: &str,
|
||||
) -> Result<(String, String)> {
|
||||
pub async fn refresh_access_token(&self, refresh_token: &str) -> Result<(String, String)> {
|
||||
// 1. 从 refresh_token 中解码出 user_id
|
||||
let user_id = TokenService::decode_user_id(refresh_token, &self.auth_config.jwt_secret)?;
|
||||
let user_id =
|
||||
TokenService::decode_refresh_user_id(refresh_token, &self.auth_config.jwt_secret)?;
|
||||
|
||||
// 2. 从 Redis 获取存储的 token 并删除
|
||||
let stored_token = self.get_and_delete_refresh_token(&user_id).await?;
|
||||
@@ -207,14 +249,22 @@ impl AuthService {
|
||||
)?;
|
||||
|
||||
// 5. 保存新的 refresh_token
|
||||
self.save_refresh_token(&user_id, &new_refresh_token, self.auth_config.refresh_token_expiration_days as i64).await?;
|
||||
self.save_refresh_token(
|
||||
&user_id,
|
||||
&new_refresh_token,
|
||||
self.auth_config.refresh_token_expiration_days,
|
||||
)
|
||||
.await?;
|
||||
|
||||
Ok((new_access_token, new_refresh_token))
|
||||
}
|
||||
|
||||
/// 删除用户
|
||||
pub async fn delete_user(&self, request: DeleteUserRequest) -> Result<()> {
|
||||
let password_hash = self.user_repo.get_password_hash(&request.user_id).await?
|
||||
let password_hash = self
|
||||
.user_repo
|
||||
.get_password_hash_by_id(&request.user_id)
|
||||
.await?
|
||||
.ok_or_else(|| anyhow::anyhow!("用户不存在"))?;
|
||||
|
||||
let parsed_hash = PasswordHash::new(&password_hash)
|
||||
@@ -230,3 +280,18 @@ impl AuthService {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
#[test]
|
||||
fn hashes_password_with_random_salt() {
|
||||
let first = AuthService::hash_password("password123").unwrap();
|
||||
let second = AuthService::hash_password("password123").unwrap();
|
||||
assert_ne!(first, second);
|
||||
let parsed = PasswordHash::new(&first).unwrap();
|
||||
assert!(Argon2::default()
|
||||
.verify_password(b"password123", &parsed)
|
||||
.is_ok());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
pub const ZH_CN: &str = "zh-CN";
|
||||
pub const EN: &str = "en";
|
||||
|
||||
pub fn message(language: &str, zh: &'static str, en: &'static str) -> &'static str {
|
||||
if language == EN {
|
||||
en
|
||||
} else {
|
||||
zh
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
#[test]
|
||||
fn selects_language() {
|
||||
assert_eq!(message(EN, "中", "en"), "en");
|
||||
assert_eq!(message(ZH_CN, "中", "en"), "中");
|
||||
}
|
||||
}
|
||||
+30
-1
@@ -75,6 +75,7 @@ impl TokenService {
|
||||
}
|
||||
|
||||
/// 从 token 中解码出 user_id
|
||||
#[allow(dead_code)]
|
||||
pub fn decode_user_id(token: &str, jwt_secret: &str) -> Result<String> {
|
||||
let token_data = decode::<Claims>(
|
||||
token,
|
||||
@@ -85,6 +86,18 @@ impl TokenService {
|
||||
|
||||
Ok(token_data.claims.sub)
|
||||
}
|
||||
|
||||
pub fn decode_refresh_user_id(token: &str, jwt_secret: &str) -> Result<String> {
|
||||
let token_data = decode::<Claims>(
|
||||
token,
|
||||
&DecodingKey::from_secret(jwt_secret.as_ref()),
|
||||
&Validation::default(),
|
||||
)?;
|
||||
if token_data.claims.token_type != TokenType::Refresh {
|
||||
anyhow::bail!("Token 类型无效");
|
||||
}
|
||||
Ok(token_data.claims.sub)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
@@ -94,8 +107,24 @@ pub struct Claims {
|
||||
pub token_type: TokenType,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub enum TokenType {
|
||||
Access,
|
||||
Refresh,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
const SECRET: &str = "a-test-secret-that-is-long-enough";
|
||||
#[test]
|
||||
fn token_pair_keeps_subject_and_type() {
|
||||
let (access, refresh) = TokenService::generate_token_pair("42", 5, 1, SECRET).unwrap();
|
||||
assert_eq!(TokenService::decode_user_id(&access, SECRET).unwrap(), "42");
|
||||
assert!(TokenService::decode_refresh_user_id(&access, SECRET).is_err());
|
||||
assert_eq!(
|
||||
TokenService::decode_refresh_user_id(&refresh, SECRET).unwrap(),
|
||||
"42"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1 +1,2 @@
|
||||
pub mod i18n;
|
||||
pub mod jwt;
|
||||
|
||||
Reference in New Issue
Block a user