Compare commits

..

2 Commits

58 changed files with 2042 additions and 349 deletions
+16
View File
@@ -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
View File
@@ -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
View File
@@ -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]
+15 -2
View File
@@ -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 客户端
```
+18
View File
@@ -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
+18
View File
@@ -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. 已修改数据库密码为强密码
+14
View File
@@ -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` 选择认证错误语言。
+20
View File
@@ -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` 不允许包含 `*`
+12
View File
@@ -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;
-- ============================================
+5
View File
@@ -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);
+5
View File
@@ -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);
+5
View File
@@ -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
View File
@@ -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
View File
@@ -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());
}
}
+1 -1
View File
@@ -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,
+55
View File
@@ -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
View File
@@ -1,5 +1,6 @@
pub mod app;
pub mod auth;
pub mod database;
pub mod email;
pub mod redis;
pub mod server;
+2
View File
@@ -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,
+26
View File
@@ -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 {
+62 -9
View File
@@ -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
View File
@@ -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
View File
@@ -1 +1,2 @@
pub mod auth;
pub mod user;
+12 -1
View File
@@ -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>,
}
+27
View File
@@ -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)
}
}
+2 -1
View File
@@ -1,2 +1,3 @@
pub mod email_logs;
pub mod user_profiles;
pub mod users;
+29
View File
@@ -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)
}
}
+1
View File
@@ -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
View File
@@ -1,3 +1,3 @@
pub mod dto;
pub mod vo;
pub mod entities;
pub mod vo;
+22 -4
View File
@@ -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 -1
View File
@@ -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
View File
@@ -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
View File
@@ -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()))
}
}
+196
View File
@@ -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());
}
}
+7 -5
View File
@@ -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
}}),
)
}
/// 获取服务器信息
+2
View File
@@ -1,2 +1,4 @@
pub mod auth;
pub mod email;
pub mod health;
pub mod user_profile;
+51
View File
@@ -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",
)))
}
+23
View File
@@ -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(())
}
+2
View File
@@ -0,0 +1,2 @@
pub mod mailer;
pub mod worker;
+50
View File
@@ -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;
}
}
}
});
}
}
+69 -27
View File
@@ -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)
}
+19
View File
@@ -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
}
+43
View File
@@ -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
}
+62 -65
View File
@@ -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), "中文.....");
}
}
+8
View File
@@ -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;
+64
View File
@@ -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()
}
}
+38
View File
@@ -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()
}
+28
View File
@@ -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
View File
@@ -1,2 +1,3 @@
pub mod mail;
pub mod middleware;
pub mod redis;
+37 -5
View File
@@ -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
View File
@@ -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"));
}
}
+39
View File
@@ -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?)
}
}
+2 -1
View File
@@ -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(())
}
}
+42 -7
View File
@@ -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(())
}
}
+93 -28
View File
@@ -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());
}
}
+20
View File
@@ -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
View File
@@ -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
View File
@@ -1 +1,2 @@
pub mod i18n;
pub mod jwt;