diff --git a/.env.example b/.env.example index bb5cfe8..2f26ea4 100644 --- a/.env.example +++ b/.env.example @@ -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 diff --git a/Cargo.lock b/Cargo.lock index 44e55af..fa6b436 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -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]] diff --git a/Cargo.toml b/Cargo.toml index 3b3bf76..879d998 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -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] diff --git a/README.md b/README.md index ae9df94..415a5a2 100644 --- a/README.md +++ b/README.md @@ -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 客户端 ``` diff --git a/config/development.toml b/config/development.toml index 5b0be82..f54105c 100644 --- a/config/development.toml +++ b/config/development.toml @@ -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 diff --git a/config/production.toml b/config/production.toml index e67f1b5..4c41097 100644 --- a/config/production.toml +++ b/config/production.toml @@ -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. 已修改数据库密码为强密码 diff --git a/docs/api/api-overview.md b/docs/api/api-overview.md index 2e76b0c..f1a00f6 100644 --- a/docs/api/api-overview.md +++ b/docs/api/api-overview.md @@ -183,3 +183,17 @@ curl http://localhost:3000/health --- **提示**:建议使用 Postman、Insomnia 或类似工具测试 API 接口。 +# 新增通用接口 + +所有受保护接口使用 `Authorization: Bearer `。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` 选择认证错误语言。 diff --git a/docs/deployment/environment-variables.md b/docs/deployment/environment-variables.md index fd103a9..218f74e 100644 --- a/docs/deployment/environment-variables.md +++ b/docs/deployment/environment-variables.md @@ -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` 不允许包含 `*`。 diff --git a/docs/sql/init.sql b/docs/sql/init.sql index 7cdf221..b1a6fbd 100644 --- a/docs/sql/init.sql +++ b/docs/sql/init.sql @@ -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; -- ============================================ diff --git a/docs/sql/mysql.sql b/docs/sql/mysql.sql new file mode 100644 index 0000000..77ff068 --- /dev/null +++ b/docs/sql/mysql.sql @@ -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); diff --git a/docs/sql/postgres.sql b/docs/sql/postgres.sql new file mode 100644 index 0000000..1e5c8a6 --- /dev/null +++ b/docs/sql/postgres.sql @@ -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); diff --git a/docs/sql/sqlite.sql b/docs/sql/sqlite.sql new file mode 100644 index 0000000..f2a91bc --- /dev/null +++ b/docs/sql/sqlite.sql @@ -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); diff --git a/src/cli.rs b/src/cli.rs index 0186986..4f6361a 100644 --- a/src/cli.rs +++ b/src/cli.rs @@ -144,7 +144,10 @@ impl CliArgs { if config.exists() { return Some(config); } - eprintln!("⚠ 警告:环境变量 CONFIG 指定的配置文件不存在: {}", config_path); + eprintln!( + "⚠ 警告:环境变量 CONFIG 指定的配置文件不存在: {}", + config_path + ); eprintln!(" 将仅使用环境变量运行"); return None; } diff --git a/src/config/app.rs b/src/config/app.rs index f85ce4a..aaa0ec3 100644 --- a/src/config/app.rs +++ b/src/config/app.rs @@ -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()); + } +} diff --git a/src/config/database.rs b/src/config/database.rs index 759228e..88469eb 100644 --- a/src/config/database.rs +++ b/src/config/database.rs @@ -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, diff --git a/src/config/email.rs b/src/config/email.rs new file mode 100644 index 0000000..98c7086 --- /dev/null +++ b/src/config/email.rs @@ -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(), + } + } +} diff --git a/src/config/mod.rs b/src/config/mod.rs index 69d1e76..0096ae3 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -1,5 +1,6 @@ pub mod app; pub mod auth; pub mod database; +pub mod email; pub mod redis; pub mod server; diff --git a/src/config/redis.rs b/src/config/redis.rs index c9291bb..830d268 100644 --- a/src/config/redis.rs +++ b/src/config/redis.rs @@ -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, diff --git a/src/config/server.rs b/src/config/server.rs index 3bf66fe..9e84051 100644 --- a/src/config/server.rs +++ b/src/config/server.rs @@ -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, +} + +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 { + vec!["http://localhost:3000".to_string()] } fn default_server_host() -> String { diff --git a/src/db.rs b/src/db.rs index dcaac57..9517979 100644 --- a/src/db.rs +++ b/src/db.rs @@ -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 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!("✅ 数据库表结构检查完成"); diff --git a/src/domain/dto/auth.rs b/src/domain/dto/auth.rs index f5d040d..852dd0a 100644 --- a/src/domain/dto/auth.rs +++ b/src/domain/dto/auth.rs @@ -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 + ) } } diff --git a/src/domain/dto/mod.rs b/src/domain/dto/mod.rs index 0e4a05d..f9bae3d 100644 --- a/src/domain/dto/mod.rs +++ b/src/domain/dto/mod.rs @@ -1 +1,2 @@ pub mod auth; +pub mod user; diff --git a/src/domain/dto/user.rs b/src/domain/dto/user.rs index ffb2e1b..6bbf799 100644 --- a/src/domain/dto/user.rs +++ b/src/domain/dto/user.rs @@ -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, + #[validate(url)] + pub avatar_url: Option, + #[validate(length(max = 500))] + pub bio: Option, +} diff --git a/src/domain/entities/email_logs.rs b/src/domain/entities/email_logs.rs new file mode 100644 index 0000000..daa2c78 --- /dev/null +++ b/src/domain/entities/email_logs.rs @@ -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, + pub recipient: String, + pub kind: String, + pub status: String, + pub error: Option, + 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(self, _: &C, insert: bool) -> Result { + let mut value = self; + if insert { + value.created_at = Set(chrono::Utc::now().naive_utc()); + } + Ok(value) + } +} diff --git a/src/domain/entities/mod.rs b/src/domain/entities/mod.rs index 995a558..b51da07 100644 --- a/src/domain/entities/mod.rs +++ b/src/domain/entities/mod.rs @@ -1,2 +1,3 @@ +pub mod email_logs; +pub mod user_profiles; pub mod users; - diff --git a/src/domain/entities/user_profiles.rs b/src/domain/entities/user_profiles.rs new file mode 100644 index 0000000..e298cd9 --- /dev/null +++ b/src/domain/entities/user_profiles.rs @@ -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, + pub avatar_url: Option, + pub bio: Option, + 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(self, _: &C, insert: bool) -> Result { + 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) + } +} diff --git a/src/domain/entities/users.rs b/src/domain/entities/users.rs index 0bb7e6d..8cd4cc8 100644 --- a/src/domain/entities/users.rs +++ b/src/domain/entities/users.rs @@ -12,6 +12,7 @@ pub struct Model { pub password_hash: String, pub created_at: DateTime, pub updated_at: DateTime, + pub deleted_at: Option, } #[derive(Copy, Clone, Debug, EnumIter, DeriveRelation)] diff --git a/src/domain/mod.rs b/src/domain/mod.rs index 09561d9..4aec842 100644 --- a/src/domain/mod.rs +++ b/src/domain/mod.rs @@ -1,3 +1,3 @@ pub mod dto; -pub mod vo; pub mod entities; +pub mod vo; diff --git a/src/domain/vo/auth.rs b/src/domain/vo/auth.rs index f51f29e..57137f8 100644 --- a/src/domain/vo/auth.rs +++ b/src/domain/vo/auth.rs @@ -4,16 +4,25 @@ use serde::Serialize; #[derive(Debug, Serialize)] pub struct RegisterResult { pub email: String, - pub created_at: String, // ISO 8601 格式 + pub created_at: String, // ISO 8601 格式 pub access_token: String, pub refresh_token: String, } 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, } @@ -25,17 +34,26 @@ impl From<(crate::domain::entities::users::Model, String, String)> for RegisterR pub struct LoginResult { pub id: String, pub email: String, - pub created_at: String, // ISO 8601 格式 + pub created_at: String, // ISO 8601 格式 pub access_token: String, pub refresh_token: String, } 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, } diff --git a/src/domain/vo/mod.rs b/src/domain/vo/mod.rs index 31fb983..4508f95 100644 --- a/src/domain/vo/mod.rs +++ b/src/domain/vo/mod.rs @@ -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 { diff --git a/src/error.rs b/src/error.rs index 4373ebb..eaf432f 100644 --- a/src/error.rs +++ b/src/error.rs @@ -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 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) -> Self { + Self::bad_request(message) + } + pub fn bad_request(message: impl Into) -> Self { Self { status: StatusCode::BAD_REQUEST, message: message.into(), } } - - #[allow(dead_code)] - pub fn not_found(message: impl Into) -> Self { - Self { - status: StatusCode::NOT_FOUND, - message: message.into(), - } - } - - #[allow(dead_code)] pub fn unauthorized(message: impl Into) -> Self { Self { status: StatusCode::UNAUTHORIZED, message: message.into(), } } - - #[allow(dead_code)] + pub fn not_found(message: impl Into) -> Self { + Self { + status: StatusCode::NOT_FOUND, + message: message.into(), + } + } + pub fn too_many_requests(message: impl Into) -> Self { + Self { + status: StatusCode::TOO_MANY_REQUESTS, + message: message.into(), + } + } + pub fn payload_too_large(message: impl Into) -> Self { + Self { + status: StatusCode::PAYLOAD_TOO_LARGE, + message: message.into(), + } + } + pub fn unavailable(message: impl Into) -> Self { + Self { + status: StatusCode::SERVICE_UNAVAILABLE, + message: message.into(), + } + } pub fn internal(message: impl Into) -> 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::<()> { - code: self.status.as_u16(), - message: self.message, - data: None, - }; - - (self.status, Json(body)).into_response() + ( + self.status, + Json(ApiResponse::<()> { + code: self.status.as_u16(), + message: self.message, + data: None, + }), + ) + .into_response() } } diff --git a/src/handlers/auth.rs b/src/handlers/auth.rs index d6753a6..74e0d32 100644 --- a/src/handlers/auth.rs +++ b/src/handlers/auth.rs @@ -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, Json(payload): Json, ) -> Result>, 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, Json(payload): Json, ) -> Result>, 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, State(state): State, - Extension(user_id): Extension, + UserId(user_id): UserId, Json(payload): Json, ) -> Result>, 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, State(state): State, - Extension(user_id): Extension, + UserId(user_id): UserId, ) -> Result>, 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())) } } diff --git a/src/handlers/email.rs b/src/handlers/email.rs new file mode 100644 index 0000000..bbf4700 --- /dev/null +++ b/src/handlers/email.rs @@ -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, + Json(input): Json, +) -> Result>, 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, + Json(input): Json, +) -> Result>, 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, + UserId(user_id): UserId, +) -> Result>, 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, +) -> Result>, 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()); + } +} diff --git a/src/handlers/health.rs b/src/handlers/health.rs index 9180d3b..a0453df 100644 --- a/src/handlers/health.rs +++ b/src/handlers/health.rs @@ -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) -> 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 + }}), + ) } /// 获取服务器信息 diff --git a/src/handlers/mod.rs b/src/handlers/mod.rs index a923656..a52f1d4 100644 --- a/src/handlers/mod.rs +++ b/src/handlers/mod.rs @@ -1,2 +1,4 @@ pub mod auth; +pub mod email; pub mod health; +pub mod user_profile; diff --git a/src/handlers/user_profile.rs b/src/handlers/user_profile.rs new file mode 100644 index 0000000..0708281 --- /dev/null +++ b/src/handlers/user_profile.rs @@ -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, + UserId(user_id): UserId, +) -> Result>, 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, + UserId(user_id): UserId, + Json(input): Json, +) -> Result>, 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, + UserId(user_id): UserId, +) -> Result>, 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", + ))) +} diff --git a/src/infra/mail/mailer.rs b/src/infra/mail/mailer.rs new file mode 100644 index 0000000..41a8b99 --- /dev/null +++ b/src/infra/mail/mailer.rs @@ -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::()?) + .to(recipient.parse::()?) + .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::::starttls_relay(&config.smtp_host)? + .port(config.smtp_port) + .credentials(credentials) + .build(); + mailer.send(email).await?; + Ok(()) +} diff --git a/src/infra/mail/mod.rs b/src/infra/mail/mod.rs new file mode 100644 index 0000000..a38064f --- /dev/null +++ b/src/infra/mail/mod.rs @@ -0,0 +1,2 @@ +pub mod mailer; +pub mod worker; diff --git a/src/infra/mail/worker.rs b/src/infra/mail/worker.rs new file mode 100644 index 0000000..d3370c8 --- /dev/null +++ b/src/infra/mail/worker.rs @@ -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::(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; + } + } + } + }); + } +} diff --git a/src/infra/middleware/auth.rs b/src/infra/middleware/auth.rs index b5b118c..c18451f 100644 --- a/src/infra/middleware/auth.rs +++ b/src/infra/middleware/auth.rs @@ -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, - headers: HeaderMap, mut req: Request, next: Next, -) -> Result { - // 1. 提取 Authorization header - let auth_header = headers +) -> Result { + let language = req + .extensions() + .get::() + .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::( + .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::( 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) } diff --git a/src/infra/middleware/body_limit.rs b/src/infra/middleware/body_limit.rs new file mode 100644 index 0000000..f94779b --- /dev/null +++ b/src/infra/middleware/body_limit.rs @@ -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, 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 +} diff --git a/src/infra/middleware/language.rs b/src/infra/middleware/language.rs new file mode 100644 index 0000000..7d53820 --- /dev/null +++ b/src/infra/middleware/language.rs @@ -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 FromRequestParts for Language { + type Rejection = (StatusCode, &'static str); + async fn from_request_parts(parts: &mut Parts, _: &S) -> Result { + parts.extensions.get::().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 +} diff --git a/src/infra/middleware/logging.rs b/src/infra/middleware/logging.rs index b13d9bd..604d1f7 100644 --- a/src/infra/middleware/logging.rs +++ b/src/infra/middleware/logging.rs @@ -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(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::() + "....." } 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::(&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, 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(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), "中文....."); + } } diff --git a/src/infra/middleware/mod.rs b/src/infra/middleware/mod.rs index 7eaa040..6ee18f4 100644 --- a/src/infra/middleware/mod.rs +++ b/src/infra/middleware/mod.rs @@ -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; diff --git a/src/infra/middleware/rate_limit.rs b/src/infra/middleware/rate_limit.rs new file mode 100644 index 0000000..b7d5573 --- /dev/null +++ b/src/infra/middleware/rate_limit.rs @@ -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>>, +} +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, + req: Request, + next: Next, +) -> Response { + let key = req + .extensions() + .get::>() + .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() + } +} diff --git a/src/infra/middleware/security.rs b/src/infra/middleware/security.rs new file mode 100644 index 0000000..524c029 --- /dev/null +++ b/src/infra/middleware/security.rs @@ -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() +} diff --git a/src/infra/middleware/user_id.rs b/src/infra/middleware/user_id.rs new file mode 100644 index 0000000..f443a28 --- /dev/null +++ b/src/infra/middleware/user_id.rs @@ -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 FromRequestParts for UserId { + type Rejection = ErrorResponse; + async fn from_request_parts( + parts: &mut axum::http::request::Parts, + _: &S, + ) -> Result { + let lang = parts + .extensions + .get::() + .map(|v| v.0.as_str()) + .unwrap_or(ZH_CN); + parts.extensions.get::().cloned().ok_or_else(|| { + ErrorResponse::unauthorized(message(lang, "未找到用户身份", "User identity not found")) + }) + } +} diff --git a/src/infra/mod.rs b/src/infra/mod.rs index fa7f0f6..64735dc 100644 --- a/src/infra/mod.rs +++ b/src/infra/mod.rs @@ -1,2 +1,3 @@ +pub mod mail; pub mod middleware; pub mod redis; diff --git a/src/infra/redis/redis_client.rs b/src/infra/redis/redis_client.rs index 4ed8493..388946e 100644 --- a/src/infra/redis/redis_client.rs +++ b/src/infra/redis/redis_client.rs @@ -52,11 +52,7 @@ impl RedisClient { } /// 使用 RedisKey 设置 JSON 值 - pub async fn set_key( - &self, - key: &RedisKey, - value: &T, - ) -> redis::RedisResult<()> { + pub async fn set_key(&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(&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 serde::Deserialize<'de>>( + &self, + key: &str, + ) -> redis::RedisResult> { + let mut connection = self.conn.lock().await; + let value: Option = 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 { + let mut connection = self.conn.lock().await; + connection.llen(key).await + } } diff --git a/src/main.rs b/src/main.rs index affb886..b463d0d 100644 --- a/src/main.rs +++ b/src/main.rs @@ -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, +} + +fn cors_layer(config: &config::server::ServerConfig) -> anyhow::Result { + let origins = config + .cors_origins + .iter() + .map(|v| HeaderValue::from_str(v)) + .collect::, _>>()?; + 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 { + 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), - ) - // 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?; - + 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::(), + ) + .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")); + } +} diff --git a/src/repositories/email_log_repository.rs b/src/repositories/email_log_repository.rs new file mode 100644 index 0000000..4ec4bab --- /dev/null +++ b/src/repositories/email_log_repository.rs @@ -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, + recipient: String, + status: String, + error: Option, + ) -> 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> { + Ok(email_logs::Entity::find() + .filter(email_logs::Column::UserId.eq(user_id)) + .order_by_desc(email_logs::Column::CreatedAt) + .one(&self.db) + .await?) + } +} diff --git a/src/repositories/mod.rs b/src/repositories/mod.rs index 15ee656..b4efd8b 100644 --- a/src/repositories/mod.rs +++ b/src/repositories/mod.rs @@ -1,2 +1,3 @@ +pub mod email_log_repository; +pub mod user_profile_repository; pub mod user_repository; - diff --git a/src/repositories/user_profile_repository.rs b/src/repositories/user_profile_repository.rs new file mode 100644 index 0000000..4d66417 --- /dev/null +++ b/src/repositories/user_profile_repository.rs @@ -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> { + Ok(user_profiles::Entity::find_by_id(user_id) + .one(&self.db) + .await?) + } + pub async fn upsert( + &self, + user_id: String, + input: UpdateProfileRequest, + ) -> Result { + 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(()) + } +} diff --git a/src/repositories/user_repository.rs b/src/repositories/user_repository.rs index 1eea54e..5f32078 100644 --- a/src/repositories/user_repository.rs +++ b/src/repositories/user_repository.rs @@ -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> { 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> { + Ok(users::Entity::find_by_id(id).one(&self.db).await?) + } + /// 统计邮箱数量 pub async fn count_by_email(&self, email: &str) -> Result { 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> { + 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 { + pub async fn insert( + &self, + id: String, + email: String, + password_hash: String, + ) -> Result { 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(()) + } } diff --git a/src/services/auth_service.rs b/src/services/auth_service.rs index 0e704a0..c212890 100644 --- a/src/services/auth_service.rs +++ b/src/services/auth_service.rs @@ -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, 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, + auth_config: AuthConfig, + ) -> Self { + Self { + user_repo, + redis_client, + auth_config, + } } /// 哈希密码 - pub fn hash_password(&self, password: &str) -> Result { + pub fn hash_password(password: &str) -> Result { 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 = self.redis_client + let redis = self + .redis_client + .as_ref() + .ok_or_else(|| anyhow::anyhow!("Redis 服务不可用"))?; + let token: Option = 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()); + } +} diff --git a/src/utils/i18n.rs b/src/utils/i18n.rs new file mode 100644 index 0000000..450143a --- /dev/null +++ b/src/utils/i18n.rs @@ -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"), "中"); + } +} diff --git a/src/utils/jwt.rs b/src/utils/jwt.rs index 8739815..5d195bd 100644 --- a/src/utils/jwt.rs +++ b/src/utils/jwt.rs @@ -75,6 +75,7 @@ impl TokenService { } /// 从 token 中解码出 user_id + #[allow(dead_code)] pub fn decode_user_id(token: &str, jwt_secret: &str) -> Result { let token_data = decode::( token, @@ -85,6 +86,18 @@ impl TokenService { Ok(token_data.claims.sub) } + + pub fn decode_refresh_user_id(token: &str, jwt_secret: &str) -> Result { + let token_data = decode::( + 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" + ); + } +} diff --git a/src/utils/mod.rs b/src/utils/mod.rs index 417233c..e6645f7 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -1 +1,2 @@ +pub mod i18n; pub mod jwt;