From b91f1e4fa5a906590991853466aab243eb16e9c8 Mon Sep 17 00:00:00 2001 From: Asfmq <2696428814@qq.com> Date: Tue, 28 Jul 2026 21:54:02 +0800 Subject: [PATCH] =?UTF-8?q?feat(server,dashboard):=20=E5=BC=95=E5=85=A5?= =?UTF-8?q?=E5=A4=9A=E5=B7=A5=E4=BD=9C=E6=B5=81=E6=95=B0=E6=8D=AE=E9=9A=94?= =?UTF-8?q?=E7=A6=BB=E3=80=81=E5=AE=89=E5=85=A8=E4=B8=AD=E9=97=B4=E4=BB=B6?= =?UTF-8?q?=E4=B8=8E=E5=89=8D=E7=AB=AF=20ESM=20=E6=A8=A1=E5=9D=97=E5=8C=96?= =?UTF-8?q?=E9=87=8D=E6=9E=84?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - server: 实现按 workflow_name 的多工作流数据隔离与旧数据库平滑迁移机制 - server: 新增 API Key 认证(auth)、限流中间件(rate_limit)与运维备份接口(admin) - server: 统一 AppError 错误处理体系,重构调度器 scheduler 支持工作流级重置与抢占 - node: 节点 ID 缺失时自动生成随机 UUID,原生支持 `docker compose --scale node=N` 动态扩容 - dashboard: 前端模块化重构(state/api/components),升级 CSS 变量设计系统与 Toast 通知 - docker/docs: 更新 /healthz 健康检查、部署脚本 IP 配置及数据库设计文档 --- .env.example | 23 +- .gitignore | 4 + Cargo.lock | 431 +++---- Cargo.toml | 8 +- Dockerfile.node | 12 +- Dockerfile.server | 32 +- README.md | 65 +- config_dense.yaml | 43 - crates/common/src/config.rs | 140 ++- crates/common/src/conv_check.rs | 35 +- crates/common/src/embedded.rs | 33 +- crates/common/src/fort55_writer.rs | 10 +- crates/common/src/gen_input5.rs | 189 ++- crates/common/src/lib.rs | 1 - crates/common/src/logging.rs | 8 +- crates/common/src/models.rs | 13 +- crates/common/src/nst_writer.rs | 1 - crates/common/src/runner.rs | 88 +- crates/common/src/seed_finder.rs | 9 +- crates/mq/src/lib.rs | 1 - crates/mq/src/sqlite_queue.rs | 246 +++- crates/node/src/executor.rs | 141 ++- crates/node/src/main.rs | 83 +- crates/node/src/reporter.rs | 69 +- crates/node/src/worker.rs | 183 ++- crates/server/Cargo.toml | 3 + crates/server/src/api/admin.rs | 154 +++ crates/server/src/api/auth.rs | 121 ++ crates/server/src/api/data.rs | 93 +- crates/server/src/api/error.rs | 51 + crates/server/src/api/mod.rs | 290 ++++- crates/server/src/api/node.rs | 143 ++- crates/server/src/api/rate_limit.rs | 196 ++++ crates/server/src/api/seed.rs | 42 +- crates/server/src/api/status.rs | 28 +- crates/server/src/api/task.rs | 179 ++- crates/server/src/api/workflow.rs | 279 ++--- crates/server/src/cors.rs | 58 + crates/server/src/db.rs | 1460 ++++++++++++++++++++++-- crates/server/src/lib.rs | 2 +- crates/server/src/main.rs | 323 +++++- crates/server/src/scheduler.rs | 385 ++++++- crates/server/tests/api_tests.rs | 1192 ++++++++++++++++++- dashboard/index.html | 156 ++- dashboard/src/api.js | 109 ++ dashboard/src/components/modal.js | 83 ++ dashboard/src/components/nodesTable.js | 122 ++ dashboard/src/components/toast.js | 20 + dashboard/src/components/workflows.js | 95 ++ dashboard/src/main.js | 565 +++++---- dashboard/src/state.js | 134 +++ dashboard/src/style.css | 813 ++++++++++--- docker-compose.yml | 15 +- docs/api.md | 98 +- docs/architecture.md | 9 +- docs/database.md | 18 +- docs/deployment.md | 7 +- scripts/deploy.sh | 2 +- tools/sync_seeds/src/main.rs | 74 +- 59 files changed, 7697 insertions(+), 1490 deletions(-) delete mode 100644 config_dense.yaml create mode 100644 crates/server/src/api/admin.rs create mode 100644 crates/server/src/api/auth.rs create mode 100644 crates/server/src/api/error.rs create mode 100644 crates/server/src/api/rate_limit.rs create mode 100644 crates/server/src/cors.rs create mode 100644 dashboard/src/api.js create mode 100644 dashboard/src/components/modal.js create mode 100644 dashboard/src/components/nodesTable.js create mode 100644 dashboard/src/components/toast.js create mode 100644 dashboard/src/components/workflows.js create mode 100644 dashboard/src/state.js diff --git a/.env.example b/.env.example index bad892d..0e1f45f 100644 --- a/.env.example +++ b/.env.example @@ -6,8 +6,23 @@ # 日志输出级别过滤 (可选格式: info, debug, warn 等) DCTS_LOG=info,server=debug,node=debug -# 共享 API 身份鉴权 Token (留空或不配置则默认使用内网无鉴权模式) -# DCTS_AUTH_TOKEN=your_secure_secret_token_here +# ===== 鉴权凭据 (公网部署务必配置 DCTS_ADMIN_TOKEN) ===== +# +# DCTS_ADMIN_TOKEN 管理员登录密码 / 工作流 CRUD / 起停计算 / 节点审批授权。 +# 支持配置人类易记的短密码(如 admin123)或强随机字符串。 +# Web Dashboard 界面将提供登录弹窗,内置 5 分钟 5 次暴破锁死保护。 +# +# 节点准入模式: +# 计算节点 (Node Worker) 物理部署时【无需配置任何凭据/Token】(零凭据部署)。 +# 节点启动后会自动提交申请,管理员登录 Web Dashboard 界面在“节点管理”中 +# 点击【同意接入】即可自动下发专属身份 Token 授权加入集群。 +# +# 示例配置: +# DCTS_ADMIN_TOKEN=admin123 +# +# 应急:临时关闭全部鉴权(仅本地调试,切勿生产使用) +# DCTS_AUTH_DISABLE=0 + # 静态资源与数据目录路径 (配分函数、谱线列表文件所在目录) # DCTS_ASSETS_DIR=assets @@ -26,6 +41,10 @@ DCTS_PORT=8090 # 网格模型计算结果文件保存根目录 # DCTS_RESULTS_DIR=data/results +# 数据库自动备份目录(每日备份 + 7 天保留期自动清理)。默认 data/backups。 +# 注意:生产部署建议与 DCTS_DB_PATH 位于同一持久化卷,避免备份落到临时层。 +# DCTS_BACKUP_DIR=data/backups + # --- 计算节点 (Node Worker) 专用配置 --- # 计算节点固定身份 ID (若留空则自动生成随机 UUID node-) diff --git a/.gitignore b/.gitignore index b6d8569..85bdda3 100644 --- a/.gitignore +++ b/.gitignore @@ -38,3 +38,7 @@ fort.84 *.swp *~ assets/data/ + + +# agent +.zcode/ \ No newline at end of file diff --git a/Cargo.lock b/Cargo.lock index acf2000..cbbb487 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -121,10 +121,10 @@ dependencies = [ "axum-core", "bytes", "futures-util", - "http 1.4.2", - "http-body 1.1.0", + "http", + "http-body", "http-body-util", - "hyper 1.11.0", + "hyper", "hyper-util", "itoa", "matchit", @@ -138,7 +138,7 @@ dependencies = [ "serde_json", "serde_path_to_error", "serde_urlencoded", - "sync_wrapper 1.0.2", + "sync_wrapper", "tokio", "tower 0.5.3", "tower-layer", @@ -155,13 +155,13 @@ dependencies = [ "async-trait", "bytes", "futures-util", - "http 1.4.2", - "http-body 1.1.0", + "http", + "http-body", "http-body-util", "mime", "pin-project-lite", "rustversion", - "sync_wrapper 1.0.2", + "sync_wrapper", "tower-layer", "tower-service", "tracing", @@ -169,15 +169,9 @@ dependencies = [ [[package]] name = "base64" -version = "0.21.7" +version = "0.22.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9d297deb1925b89f2ccc13d7635fa0714f12c87adce1c75356b39ca9b7178567" - -[[package]] -name = "bitflags" -version = "1.3.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bef38d45163c2f1dde094a7dfd33ccf595c92905c8f8f4fdc18d06fb1037718a" +checksum = "72b3254f16251a8381aa12e40e3c4d2f0199f8c6508fbecb9d91f575e0fbb8c6" [[package]] name = "bitflags" @@ -576,7 +570,18 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "dc3655aa6818d65bc620d6911f05aa7b6aeb596291e1e9f79e52df85583d1e30" dependencies = [ "rustix 0.38.44", - "windows-targets 0.52.6", + "windows-targets", +] + +[[package]] +name = "getrandom" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ff2abc00be7fca6ebc474524697ae276ad847ad0a6b3faa4bcb027e9a4614ad0" +dependencies = [ + "cfg-if", + "libc", + "wasi", ] [[package]] @@ -593,16 +598,16 @@ dependencies = [ [[package]] name = "h2" -version = "0.3.27" +version = "0.4.15" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0beca50380b1fc32983fc1cb4587bfa4bb9e78fc259aad4a0032d2080309222d" +checksum = "6cb093c84e8bd9b188d4c4a8cb6579fc016968d14c99882163cd3ff402a4f155" dependencies = [ + "atomic-waker", "bytes", "fnv", "futures-core", "futures-sink", - "futures-util", - "http 0.2.12", + "http", "indexmap", "slab", "tokio", @@ -646,17 +651,6 @@ version = "0.4.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70" -[[package]] -name = "http" -version = "0.2.12" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "601cbb57e577e2f5ef5be8e7b83f0f63994f25aa94d673e54a92d5c516d101f1" -dependencies = [ - "bytes", - "fnv", - "itoa", -] - [[package]] name = "http" version = "1.4.2" @@ -667,17 +661,6 @@ dependencies = [ "itoa", ] -[[package]] -name = "http-body" -version = "0.4.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7ceab25649e9960c0311ea418d17bee82c0dcec1bd053b5f9a66e265a693bed2" -dependencies = [ - "bytes", - "http 0.2.12", - "pin-project-lite", -] - [[package]] name = "http-body" version = "1.1.0" @@ -685,7 +668,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ca2a8f2913ee65f60facd6a5905613afaa448497a0230cc41ce022d93290bc2c" dependencies = [ "bytes", - "http 1.4.2", + "http", ] [[package]] @@ -696,8 +679,8 @@ checksum = "e9f41fd6a08e4d4ec69df65976da761afd5ad5e58a9d4acb46bd1c953a9e3ff2" dependencies = [ "bytes", "futures-core", - "http 1.4.2", - "http-body 1.1.0", + "http", + "http-body", "pin-project-lite", ] @@ -719,30 +702,6 @@ version = "1.0.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9" -[[package]] -name = "hyper" -version = "0.14.32" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "41dfc780fdec9373c01bae43289ea34c972e40ee3c9f6b3c8801a35f35586ce7" -dependencies = [ - "bytes", - "futures-channel", - "futures-core", - "futures-util", - "h2", - "http 0.2.12", - "http-body 0.4.6", - "httparse", - "httpdate", - "itoa", - "pin-project-lite", - "socket2 0.4.10", - "tokio", - "tower-service", - "tracing", - "want", -] - [[package]] name = "hyper" version = "1.11.0" @@ -753,27 +712,47 @@ dependencies = [ "bytes", "futures-channel", "futures-core", - "http 1.4.2", - "http-body 1.1.0", + "h2", + "http", + "http-body", "httparse", "httpdate", "itoa", "pin-project-lite", "smallvec", "tokio", + "want", +] + +[[package]] +name = "hyper-rustls" +version = "0.27.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "33ca68d021ef39cf6463ab54c1d0f5daf03377b70561305bb89a8f83aab66e0f" +dependencies = [ + "http", + "hyper", + "hyper-util", + "rustls", + "tokio", + "tokio-rustls", + "tower-service", ] [[package]] name = "hyper-tls" -version = "0.5.0" +version = "0.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d6183ddfa99b85da61a140bea0efc93fdf56ceaa041b37d553518030827f9905" +checksum = "70206fc6890eaca9fde8a0bf71caa2ddfc9fe045ac9e5c70df101a7dbde866e0" dependencies = [ "bytes", - "hyper 0.14.32", + "http-body-util", + "hyper", + "hyper-util", "native-tls", "tokio", "tokio-native-tls", + "tower-service", ] [[package]] @@ -782,13 +761,23 @@ version = "0.1.20" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "96547c2556ec9d12fb1578c4eaf448b04993e7fb79cbaad930a656880a6bdfa0" dependencies = [ + "base64", "bytes", - "http 1.4.2", - "http-body 1.1.0", - "hyper 1.11.0", + "futures-channel", + "futures-util", + "http", + "http-body", + "hyper", + "ipnet", + "libc", + "percent-encoding", "pin-project-lite", + "socket2", + "system-configuration", "tokio", "tower-service", + "tracing", + "windows-registry", ] [[package]] @@ -1088,7 +1077,7 @@ dependencies = [ "bytes", "encoding_rs", "futures-util", - "http 1.4.2", + "http", "httparse", "memchr", "mime", @@ -1184,7 +1173,7 @@ version = "0.10.81" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "77823a27f0babb03091cb9ed9ef80af3b39dbc82f97e8fa530374b7dafd87a45" dependencies = [ - "bitflags 2.13.1", + "bitflags", "cfg-if", "foreign-types", "libc", @@ -1350,7 +1339,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c7f5fa3a058cd35567ef9bfa5e75732bee0f9e4c55fa90477bef2dfcdbc4be80" dependencies = [ "chacha20", - "getrandom", + "getrandom 0.4.3", "rand_core", ] @@ -1386,7 +1375,7 @@ version = "0.5.18" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ed2bf2547551a7053d6fdfafda3f938979645c44812fbfcda098faae3f1a362d" dependencies = [ - "bitflags 2.13.1", + "bitflags", ] [[package]] @@ -1420,9 +1409,9 @@ checksum = "d6f6ff9a378485b298a5286656da665ba74413d36db0979633275d2e708145d4" [[package]] name = "reqwest" -version = "0.11.27" +version = "0.12.28" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dd67538700a17451e7cba03ac727fb961abb7607553461627b97de0b89cf4a62" +checksum = "eddd3ca559203180a307f12d114c268abf583f59b03cb906fd0b3ff8646c1147" dependencies = [ "base64", "bytes", @@ -1430,33 +1419,48 @@ dependencies = [ "futures-core", "futures-util", "h2", - "http 0.2.12", - "http-body 0.4.6", - "hyper 0.14.32", + "http", + "http-body", + "http-body-util", + "hyper", + "hyper-rustls", "hyper-tls", - "ipnet", + "hyper-util", "js-sys", "log", "mime", "mime_guess", "native-tls", - "once_cell", "percent-encoding", "pin-project-lite", - "rustls-pemfile", + "rustls-pki-types", "serde", "serde_json", "serde_urlencoded", - "sync_wrapper 0.1.2", - "system-configuration", + "sync_wrapper", "tokio", "tokio-native-tls", + "tower 0.5.3", + "tower-http 0.6.11", "tower-service", "url", "wasm-bindgen", "wasm-bindgen-futures", "web-sys", - "winreg", +] + +[[package]] +name = "ring" +version = "0.17.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a4689e6c2294d81e88dc6261c768b63bc4fcdb852be6d1352498b114f61383b7" +dependencies = [ + "cc", + "cfg-if", + "getrandom 0.2.17", + "libc", + "untrusted", + "windows-sys 0.52.0", ] [[package]] @@ -1465,7 +1469,7 @@ version = "0.31.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b838eba278d213a8beaf485bd313fd580ca4505a00d5871caeb1457c55322cae" dependencies = [ - "bitflags 2.13.1", + "bitflags", "fallible-iterator", "fallible-streaming-iterator", "hashlink", @@ -1479,7 +1483,7 @@ version = "0.38.44" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fdb5bc1ae2baa591800df16c9ca78619bf65c0488b41b96ccec5d11220d8c154" dependencies = [ - "bitflags 2.13.1", + "bitflags", "errno", "libc", "linux-raw-sys 0.4.15", @@ -1492,7 +1496,7 @@ version = "1.1.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6fe4565b9518b83ef4f91bb47ce29620ca828bd32cb7e408f0062e9930ba190" dependencies = [ - "bitflags 2.13.1", + "bitflags", "errno", "libc", "linux-raw-sys 0.12.1", @@ -1500,12 +1504,36 @@ dependencies = [ ] [[package]] -name = "rustls-pemfile" -version = "1.0.4" +name = "rustls" +version = "0.23.42" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1c74cae0a4cf6ccbbf5f359f08efdf8ee7e1dc532573bf0db71968cb56b1448c" +checksum = "3c54fcab019b409d04215d3a17cb438fd7fbf192ee61461f20f4fe18704bc138" dependencies = [ - "base64", + "once_cell", + "rustls-pki-types", + "rustls-webpki", + "subtle", + "zeroize", +] + +[[package]] +name = "rustls-pki-types" +version = "1.15.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2f4925028c7eb5d1fcdaf196971378ed9d2c1c4efc7dc5d011256f76c99c0a96" +dependencies = [ + "zeroize", +] + +[[package]] +name = "rustls-webpki" +version = "0.103.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "61c429a8649f110dddef65e2a5ad240f747e85f7758a6bccc7e5777bd33f756e" +dependencies = [ + "ring", + "rustls-pki-types", + "untrusted", ] [[package]] @@ -1550,7 +1578,7 @@ version = "3.7.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b7f4bc775c73d9a02cde8bf7b2ec4c9d12743edf609006c7facc23998404cd1d" dependencies = [ - "bitflags 2.13.1", + "bitflags", "core-foundation 0.10.1", "core-foundation-sys", "libc", @@ -1656,6 +1684,7 @@ dependencies = [ "clap", "common", "dotenvy", + "hex", "mq", "r2d2", "r2d2_sqlite", @@ -1663,11 +1692,13 @@ dependencies = [ "serde", "serde_json", "serde_yaml", + "sha2", + "subtle", "tempfile", "tokio", "tokio-util", "tower 0.4.13", - "tower-http", + "tower-http 0.5.2", "tracing", "tracing-subscriber", "uuid", @@ -1721,16 +1752,6 @@ version = "1.15.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8ed6a63f02c8539c91a8685a86f4099661ba3da017932f6ebbea6de3f0fa7c90" -[[package]] -name = "socket2" -version = "0.4.10" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9f7916fc008ca5542385b89a3d3ce689953c143e9304a9bf8beec1de48994c0d" -dependencies = [ - "libc", - "winapi", -] - [[package]] name = "socket2" version = "0.6.5" @@ -1759,6 +1780,12 @@ version = "0.11.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7da8b5736845d9f2fcb837ea5d9e2628564b3b043a70948a3f0b778838c5fb4f" +[[package]] +name = "subtle" +version = "2.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" + [[package]] name = "symlink" version = "0.1.0" @@ -1803,17 +1830,14 @@ dependencies = [ "uuid", ] -[[package]] -name = "sync_wrapper" -version = "0.1.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2047c6ded9c721764247e62cd3b03c09ffc529b2ba5b10ec482ae507a4a70160" - [[package]] name = "sync_wrapper" version = "1.0.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0bf256ce5efdfa370213c1dabab5935a12e49f2c58d15e9eac2870d3b4f27263" +dependencies = [ + "futures-core", +] [[package]] name = "synstructure" @@ -1843,20 +1867,20 @@ dependencies = [ [[package]] name = "system-configuration" -version = "0.5.1" +version = "0.7.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ba3a3adc5c275d719af8cb4272ea1c4a6d668a777f37e115f6d11ddbc1c8e0e7" +checksum = "a13f3d0daba03132c0aa9767f98351b3488edc2c100cda2d2ec2b04f3d8d3c8b" dependencies = [ - "bitflags 1.3.2", + "bitflags", "core-foundation 0.9.4", "system-configuration-sys", ] [[package]] name = "system-configuration-sys" -version = "0.5.0" +version = "0.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a75fb188eb626b924683e3b95e3a48e63551fcfb51949de2f06a9d91dbee93c9" +checksum = "8e1d1b10ced5ca923a1fcb8d03e96b8d3268065d724548c0211415ff6ac6bac4" dependencies = [ "core-foundation-sys", "libc", @@ -1869,7 +1893,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "32497e9a4c7b38532efcdebeef879707aa9f794296a4f0244f6f69e9bc8574bd" dependencies = [ "fastrand", - "getrandom", + "getrandom 0.4.3", "once_cell", "rustix 1.1.4", "windows-sys 0.61.2", @@ -1956,7 +1980,7 @@ dependencies = [ "parking_lot", "pin-project-lite", "signal-hook-registry", - "socket2 0.6.5", + "socket2", "tokio-macros", "windows-sys 0.61.2", ] @@ -1982,6 +2006,16 @@ dependencies = [ "tokio", ] +[[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-util" version = "0.7.19" @@ -2006,6 +2040,8 @@ dependencies = [ "futures-util", "pin-project", "pin-project-lite", + "tokio", + "tokio-util", "tower-layer", "tower-service", "tracing", @@ -2020,7 +2056,7 @@ dependencies = [ "futures-core", "futures-util", "pin-project-lite", - "sync_wrapper 1.0.2", + "sync_wrapper", "tokio", "tower-layer", "tower-service", @@ -2033,11 +2069,11 @@ version = "0.5.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1e9cd434a998747dd2c4276bc96ee2e0c7a2eadf3cae88e52be55a05fa9053f5" dependencies = [ - "bitflags 2.13.1", + "bitflags", "bytes", "futures-util", - "http 1.4.2", - "http-body 1.1.0", + "http", + "http-body", "http-body-util", "http-range-header", "httpdate", @@ -2052,6 +2088,24 @@ dependencies = [ "tracing", ] +[[package]] +name = "tower-http" +version = "0.6.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4cfcf7e2740e6fc6d4d688b4ef00650406bb94adf4731e43c096c3a19fe40840" +dependencies = [ + "bitflags", + "bytes", + "futures-util", + "http", + "http-body", + "pin-project-lite", + "tower 0.5.3", + "tower-layer", + "tower-service", + "url", +] + [[package]] name = "tower-layer" version = "0.3.3" @@ -2182,6 +2236,12 @@ version = "0.2.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "673aac59facbab8a9007c7f6108d11f63b603f7cabff99fabf650fea5c32b861" +[[package]] +name = "untrusted" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ecb6da28b8a351d773b68d5825ac39017e680750f980f3a1a85cd8dd28a47c1" + [[package]] name = "url" version = "2.5.8" @@ -2212,7 +2272,7 @@ version = "1.24.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "bf3923a6f5c4c6382e0b653c4117f48d631ea17f38ed86e2a828e6f7412f5239" dependencies = [ - "getrandom", + "getrandom 0.4.3", "js-sys", "rand", "serde_core", @@ -2346,7 +2406,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e48a53791691ab099e5e2ad123536d0fff50652600abaf43bbf952894110d0be" dependencies = [ "windows-core 0.52.0", - "windows-targets 0.52.6", + "windows-targets", ] [[package]] @@ -2355,7 +2415,7 @@ version = "0.52.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "33ab640c8d7e35bf8ba19b884ba838ceb4fba93a4e8c65a9059d08afcfc683d9" dependencies = [ - "windows-targets 0.52.6", + "windows-targets", ] [[package]] @@ -2399,6 +2459,17 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" +[[package]] +name = "windows-registry" +version = "0.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "02752bf7fbdcce7f2a27a742f798510f3e5ad88dbe84871e5168e2120c3d5720" +dependencies = [ + "windows-link", + "windows-result", + "windows-strings", +] + [[package]] name = "windows-result" version = "0.4.1" @@ -2417,22 +2488,13 @@ dependencies = [ "windows-link", ] -[[package]] -name = "windows-sys" -version = "0.48.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "677d2418bec65e3338edb076e806bc1ec15693c5d0104683f2efe857f61056a9" -dependencies = [ - "windows-targets 0.48.5", -] - [[package]] name = "windows-sys" version = "0.52.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "282be5f36a8ce781fad8c8ae18fa3f9beff57ec1b52cb3de0789201425d9a33d" dependencies = [ - "windows-targets 0.52.6", + "windows-targets", ] [[package]] @@ -2444,67 +2506,34 @@ dependencies = [ "windows-link", ] -[[package]] -name = "windows-targets" -version = "0.48.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9a2fa6e2155d7247be68c096456083145c183cbbbc2764150dda45a87197940c" -dependencies = [ - "windows_aarch64_gnullvm 0.48.5", - "windows_aarch64_msvc 0.48.5", - "windows_i686_gnu 0.48.5", - "windows_i686_msvc 0.48.5", - "windows_x86_64_gnu 0.48.5", - "windows_x86_64_gnullvm 0.48.5", - "windows_x86_64_msvc 0.48.5", -] - [[package]] name = "windows-targets" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9b724f72796e036ab90c1021d4780d4d3d648aca59e491e6b98e725b84e99973" dependencies = [ - "windows_aarch64_gnullvm 0.52.6", - "windows_aarch64_msvc 0.52.6", - "windows_i686_gnu 0.52.6", + "windows_aarch64_gnullvm", + "windows_aarch64_msvc", + "windows_i686_gnu", "windows_i686_gnullvm", - "windows_i686_msvc 0.52.6", - "windows_x86_64_gnu 0.52.6", - "windows_x86_64_gnullvm 0.52.6", - "windows_x86_64_msvc 0.52.6", + "windows_i686_msvc", + "windows_x86_64_gnu", + "windows_x86_64_gnullvm", + "windows_x86_64_msvc", ] -[[package]] -name = "windows_aarch64_gnullvm" -version = "0.48.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2b38e32f0abccf9987a4e3079dfb67dcd799fb61361e53e2882c3cbaf0d905d8" - [[package]] name = "windows_aarch64_gnullvm" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "32a4622180e7a0ec044bb555404c800bc9fd9ec262ec147edd5989ccd0c02cd3" -[[package]] -name = "windows_aarch64_msvc" -version = "0.48.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dc35310971f3b2dbbf3f0690a219f40e2d9afcf64f9ab7cc1be722937c26b4bc" - [[package]] name = "windows_aarch64_msvc" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "09ec2a7bb152e2252b53fa7803150007879548bc709c039df7627cabbd05d469" -[[package]] -name = "windows_i686_gnu" -version = "0.48.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a75915e7def60c94dcef72200b9a8e58e5091744960da64ec734a6c6e9b3743e" - [[package]] name = "windows_i686_gnu" version = "0.52.6" @@ -2517,64 +2546,30 @@ version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0eee52d38c090b3caa76c563b86c3a4bd71ef1a819287c19d586d7334ae8ed66" -[[package]] -name = "windows_i686_msvc" -version = "0.48.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8f55c233f70c4b27f66c523580f78f1004e8b5a8b659e05a4eb49d4166cca406" - [[package]] name = "windows_i686_msvc" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "240948bc05c5e7c6dabba28bf89d89ffce3e303022809e73deaefe4f6ec56c66" -[[package]] -name = "windows_x86_64_gnu" -version = "0.48.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "53d40abd2583d23e4718fddf1ebec84dbff8381c07cae67ff7768bbf19c6718e" - [[package]] name = "windows_x86_64_gnu" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "147a5c80aabfbf0c7d901cb5895d1de30ef2907eb21fbbab29ca94c5b08b1a78" -[[package]] -name = "windows_x86_64_gnullvm" -version = "0.48.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0b7b52767868a23d5bab768e390dc5f5c55825b6d30b86c844ff2dc7414044cc" - [[package]] name = "windows_x86_64_gnullvm" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "24d5b23dc417412679681396f2b49f3de8c1473deb516bd34410872eff51ed0d" -[[package]] -name = "windows_x86_64_msvc" -version = "0.48.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ed94fce61571a4006852b7389a063ab983c02eb1bb37b47f8272ce92d06d9538" - [[package]] name = "windows_x86_64_msvc" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" -[[package]] -name = "winreg" -version = "0.50.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "524e57b2c537c0f9b1e69f1965311ec12182b4122e45035b1508cd24d2adadb1" -dependencies = [ - "cfg-if", - "windows-sys 0.48.0", -] - [[package]] name = "writeable" version = "0.6.3" @@ -2645,6 +2640,12 @@ dependencies = [ "synstructure", ] +[[package]] +name = "zeroize" +version = "1.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e13c156562582aa81c60cb29407084cdb54c4164760106ab78e6c5b0858cf64e" + [[package]] name = "zerotrie" version = "0.2.4" diff --git a/Cargo.toml b/Cargo.toml index 5a5356d..3ca06d1 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -11,7 +11,9 @@ members = [ [workspace.dependencies] serde = { version = "1.0", features = ["derive"] } serde_json = "1.0" -serde_yaml = "0.9" +# 注:官方 serde_yaml (dtolnay/serde-yaml) 已于 2024 年 3 月归档停止维护。 +# 本项目选择固定使用官方最终稳定版本 0.9.34,因其经过多年大规模生产验证,功能完备且极度稳定无隐患。 +serde_yaml = "0.9.34" tokio = { version = "1.35", features = ["full"] } tracing = "0.1" tracing-subscriber = { version = "0.3", features = ["env-filter", "json"] } @@ -21,13 +23,13 @@ tempfile = "3.8" uuid = { version = "1.6", features = ["v4", "serde"] } chrono = { version = "0.4", features = ["serde"] } regex = "1.10" -reqwest = { version = "0.11", features = ["json", "multipart"] } +reqwest = { version = "0.12", features = ["json", "multipart"] } rusqlite = { version = "0.31", features = ["bundled"] } async-trait = "0.1" axum = { version = "0.7", features = ["multipart"] } tokio-util = { version = "0.7", features = ["io"] } tower-http = { version = "0.5", features = ["cors", "trace", "fs"] } -tower = { version = "0.4", features = ["util"] } +tower = { version = "0.4", features = ["util", "limit"] } clap = { version = "4.4", features = ["derive"] } sysinfo = "0.30" gethostname = "0.5" diff --git a/Dockerfile.node b/Dockerfile.node index d1f4ee6..e1820d5 100644 --- a/Dockerfile.node +++ b/Dockerfile.node @@ -3,12 +3,20 @@ # ============================================================================= # ─── Stage 1: Rust Node Worker 二进制编译 ──────────────────────────────────── -FROM rust:1.80-alpine AS node-builder -RUN apk add --no-cache musl-dev g++ make pkgconfig sqlite-dev +FROM rust:1.94-alpine AS node-builder + +ARG USE_MIRRORS=1 + +RUN if [ "$USE_MIRRORS" = "1" ]; then \ + sed -i 's|dl-cdn.alpinelinux.org|mirrors.aliyun.com|g' /etc/apk/repositories; \ + fi + +RUN apk add --no-cache musl-dev g++ make cmake pkgconfig sqlite-dev openssl-dev openssl-libs-static WORKDIR /app COPY Cargo.toml Cargo.lock ./ COPY crates/ ./crates/ COPY tools/ ./tools/ +COPY assets/tlusty_static assets/synspec_static ./assets/ RUN cargo build --release -p node && \ cp /app/target/release/node /usr/local/bin/dcts-node diff --git a/Dockerfile.server b/Dockerfile.server index 4ddacc3..0967775 100644 --- a/Dockerfile.server +++ b/Dockerfile.server @@ -4,16 +4,29 @@ # ─── Stage 1: 前端静态资源构建 ──────────────────────────────────────────────── FROM node:22-alpine AS frontend-builder + +ARG USE_MIRRORS=1 WORKDIR /app/dashboard + +RUN if [ "$USE_MIRRORS" = "1" ]; then \ + npm config set registry https://registry.npmmirror.com; \ + fi + COPY dashboard/package.json dashboard/package-lock.json* ./ RUN npm install COPY dashboard/ ./ RUN npm run build # ─── Stage 2: Rust 服务端编译 (Alpine/musl 静态编译) ───────────────────────── -FROM rust:1.80-alpine AS backend-builder +FROM rust:1.94-alpine AS backend-builder -RUN apk add --no-cache musl-dev g++ make pkgconfig sqlite-dev +ARG USE_MIRRORS=1 + +RUN if [ "$USE_MIRRORS" = "1" ]; then \ + sed -i 's|dl-cdn.alpinelinux.org|mirrors.aliyun.com|g' /etc/apk/repositories; \ + fi + +RUN apk add --no-cache musl-dev g++ make cmake pkgconfig sqlite-dev openssl-dev openssl-libs-static WORKDIR /app ENV SKIP_DASHBOARD_BUILD=1 @@ -21,13 +34,20 @@ ENV SKIP_DASHBOARD_BUILD=1 COPY Cargo.toml Cargo.lock ./ COPY crates/ ./crates/ COPY tools/ ./tools/ +COPY assets/tlusty_static assets/synspec_static ./assets/ COPY --from=frontend-builder /app/dashboard/dist ./dashboard/dist RUN cargo build --release -p server && \ cp /app/target/release/server /usr/local/bin/dcts-server # ─── Stage 3: 最小化生产运行镜像 ───────────────────────────────────────────── -FROM alpine:3.20 +FROM alpine:3.21 + +ARG USE_MIRRORS=1 + +RUN if [ "$USE_MIRRORS" = "1" ]; then \ + sed -i 's|dl-cdn.alpinelinux.org|mirrors.aliyun.com|g' /etc/apk/repositories; \ + fi RUN apk add --no-cache ca-certificates tzdata sqlite @@ -38,7 +58,6 @@ WORKDIR /app COPY --from=backend-builder /usr/local/bin/dcts-server /app/ COPY --from=frontend-builder /app/dashboard/dist ./dashboard/dist -COPY config_dense.yaml ./ RUN mkdir -p /app/data /app/data/results /app/logs /app/assets && \ chown -R dcts:dcts /app @@ -51,11 +70,14 @@ ENV DCTS_PORT=8090 ENV DCTS_DB_PATH=/app/data/dcts.db ENV DCTS_QUEUE_DB_PATH=/app/data/dcts_queue.db ENV DCTS_RESULTS_DIR=/app/data/results +ENV DCTS_BACKUP_DIR=/app/data/backups ENV DCTS_ASSETS_DIR=/app/assets VOLUME ["/app/data", "/app/logs"] +# 健康检查走不走鉴权的 /healthz(启用 admin/enrollment token 后 /api/status 返回 401, +# 会令容器被误判不健康而反复重启)。与 docker-compose.yml 的 healthcheck 保持一致。 HEALTHCHECK --interval=30s --timeout=10s --start-period=15s --retries=3 \ - CMD wget -q --spider http://localhost:8090/api/status || exit 1 + CMD wget -q --spider http://localhost:8090/healthz || exit 1 ENTRYPOINT ["/app/dcts-server"] diff --git a/README.md b/README.md index 0ede410..307bde9 100644 --- a/README.md +++ b/README.md @@ -50,6 +50,69 @@ DCTS_MAX_SLOTS=4 --- +## 🔐 安全与鉴权 (Security & Auth) + +DCTS 采用**分层鉴权**模型,公网部署务必按下表配置凭据。 + +### 鉴权主体 + +| 主体 | 环境变量 | 用途 | 持有方式 | +| :--- | :--- | :--- | :--- | +| **Admin** | `DCTS_ADMIN_TOKEN` | Dashboard 登录、工作流 CRUD、起停计算 | 人工,Dashboard 输入 | +| **Enrollment** | `DCTS_ENROLLMENT_TOKEN` | 节点首次注册领取专属 token | 部署脚本/人工 | +| **Node** | _(服务端自动颁发)_ | 心跳、领任务、上报、下载数据 | 节点本地 `.node_token` 文件(权限 600) | + +> **兼容**:旧变量 `DCTS_AUTH_TOKEN` 仍有效,自动同时充当 Admin + Enrollment 凭据(建议迁移到上面两个独立变量)。 + +### 节点注册流程(L2) + +1. 启动时优先读取本地 `runtime/.node_token`;不存在则用 `DCTS_ENROLLMENT_TOKEN` 调 `/api/node/register`。 +2. 服务端注册成功后**颁发该节点专属 token**(仅返回一次,DB 只存 SHA-256 hash),节点持久化到 `.node_token`。 +3. 后续所有请求携带专属 token;服务端按 token 反查 `node_id` 鉴权。 +4. **吊销/重发**:通过 Dashboard「节点凭据管理」面板或下方管理 API 操作。 + +### 节点凭据管理 API(Admin 角色) + +| 方法 | 路径 | 说明 | +| :--- | :--- | :--- | +| GET | `/api/admin/nodes` | 列出全部节点及凭据状态(在线/token 有效/吊销/颁发时间) | +| POST | `/api/admin/nodes/:node_id/revoke` | 吊销指定节点 token(立即失效,幂等) | +| POST | `/api/admin/nodes/:node_id/reissue` | 重新颁发 token,返回新明文(旧 token 失效) | + +所有端点要求 Admin token(`Authorization: Bearer `)。被攻陷节点持有的 node token 无权访问这些端点,因此吊销/重发始终是管理员主动行为。 + +### 公网部署清单 + +```bash +# 生成强随机 token +openssl rand -hex 32 +``` + +```env +# .env(服务端 + 节点共享此文件时各自读取所需变量) +DCTS_ADMIN_TOKEN=<强随机值> +DCTS_ENROLLMENT_TOKEN=<强随机值> +``` + +**TLS 反代**(推荐 Caddy,自动 HTTPS): + +```bash +# 1. 编辑 Caddyfile,把 dcts.example.com 改为真实域名 +# 2. 启用 public profile 拉起反代 +docker compose --profile public up -d --build +# 3. 节点的 DCTS_SERVER_URL 改为 https://你的域名 +``` + +### 默认安全策略 + +- **CORS**:仅允许同源或本地 Origin(localhost / 127.0.0.1 / [::1])。 +- **请求体限制**:普通 API 10MB,任务上报 256MB。 +- **安全响应头**:CSP / `X-Content-Type-Options` / `X-Frame-Options` / `Referrer-Policy` 默认开启。 +- **审计日志**:所有写操作(POST/PUT/DELETE)记录 `subject + method + path`(不记请求体)。 +- **应急调试**:`DCTS_AUTH_DISABLE=1` 跳过全部鉴权(仅本地,切勿生产)。 + +--- + ## 🏛️ Workspace 核心模块 | Crate / Tool | 类型 | 职责说明 | 详细文档 | @@ -68,7 +131,7 @@ DCTS_MAX_SLOTS=4 系统技术细节按以下主题组织: - 📐 **[系统架构 (Architecture)](docs/architecture.md)**:Master-Worker 拓扑结构、任务生命周期与心跳机制。 -- 🔗 **[API 参考 (API Reference)](docs/api_reference.md)**:Axum RESTful 接口规格明细与鉴权方式。 +- 🔗 **[API 参考 (API Reference)](docs/api.md)**:Axum RESTful 接口规格明细与鉴权方式。 - 💾 **[数据库设计 (Database)](docs/database.md)**:SQLite 数据表结构模式与队列状态机设计。 - ⚙️ **[物理链设计 (Design)](docs/design.md)**:4 阶段 TLUSTY/SYNSPEC 计算链、冷启动与种子步进(Seed Step)降级重试逻辑。 - 📦 **[部署运维指南 (Deployment)](docs/deployment.md)**:生产环境部署、Systemd 服务配置、安全令牌与日志管理。 diff --git a/config_dense.yaml b/config_dense.yaml deleted file mode 100644 index bf84fc2..0000000 --- a/config_dense.yaml +++ /dev/null @@ -1,43 +0,0 @@ -# 6 维 CNO NLTE 热亚矮星网格 —— 加密版配置 -# -# 相比 config.yaml 的改动: -# 1. CNO 各维加 -3:[-4,-2,-1] → [-4,-3,-2,-1] -# 消除 -4→-2 的 100× 丰度跳跃,每步均匀 10×,种子步进更稳 -# 2. Teff 加 50000:消除 40K→60K 的跨度,中间点有助种子传递 -# 3. 配合 run_grid.py 的 wave scheduling(按 CNO 总量分批提交) -# -# 总点数: 5*4*4*4*4*4 = 5120 -# 预计耗时: 5120/16 * 12min ≈ 64 小时 - -# ---- 网格轴 ---- -grid: - teff: [20000, 30000, 40000, 50000, 60000] - logg: [5.0, 5.5, 6.0, 6.5] - loghe: [-4, -2, 0, 2] - logc: [-4, -3, -2, -1] - logn: [-4, -3, -2, -1] - logo: [-4, -3, -2, -1] - # 共 5*4*4*4*4*4 = 5120 个点 - # CNO 每步丰度跳跃: 10× (均匀) - -# ---- 收敛链(同 config.yaml)---- -chain: - - {label: lte, lte: T, ltgray: T, ilvlin: 0, require_converged: false, niter: 0} - - {label: nc, lte: F, ltgray: F, ilvlin: 0, require_converged: false, niter: 10} - - {label: nl, lte: F, ltgray: F, ilvlin: 100, require_converged: true, niter: 100} -itek_fallback: [] -niter: 100 - -# ---- 种子步进回退 ---- -seed_step_fallback: true - -# ---- 执行参数 ---- -nworkers: 20 -timeout_sec: 3600 -resume: true - -# ---- 路径 ---- -template: templates/cno_atmos.5.tpl -fort55: templates/fort.55.lin -linelist: data/gfVIS99.dat -results: results diff --git a/crates/common/src/config.rs b/crates/common/src/config.rs index 000b0c3..e523488 100644 --- a/crates/common/src/config.rs +++ b/crates/common/src/config.rs @@ -135,19 +135,56 @@ impl GridConfig { } } -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Clone, Serialize, Deserialize)] pub struct ServerConfig { pub bind_addr: String, pub db_path: String, pub queue_db_path: String, pub results_dir: String, + /// 数据库备份目录(每日自动备份落盘位置)。默认 "data/backups",可经 DCTS_BACKUP_DIR 覆盖。 + pub backup_dir: String, pub grid_config: String, pub stale_sec: u64, #[serde(default = "default_node_stale_sec")] pub node_stale_sec: u64, pub mq_type: String, // "sqlite" or "rabbitmq" pub rabbitmq_url: Option, + /// Admin 凭据(Dashboard 登录用)。优先 DCTS_ADMIN_TOKEN,回退旧变量 DCTS_AUTH_TOKEN。 + #[serde(default)] + pub admin_token: Option, + /// 兼容字段:保留以判断「是否启用鉴权」与旧中间件逻辑。取 admin_token 的值。 + #[serde(default)] pub auth_token: Option, + /// 应急开关:DCTS_AUTH_DISABLE=1 时跳过全部鉴权(仅本地调试,默认关闭)。 + #[serde(default)] + pub auth_disabled: bool, +} + +// 手写 Debug:token 类字段脱敏为 ***REDACTED***,防止日志/错误链泄露明文凭据。 +impl std::fmt::Debug for ServerConfig { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("ServerConfig") + .field("bind_addr", &self.bind_addr) + .field("db_path", &self.db_path) + .field("queue_db_path", &self.queue_db_path) + .field("results_dir", &self.results_dir) + .field("backup_dir", &self.backup_dir) + .field("grid_config", &self.grid_config) + .field("stale_sec", &self.stale_sec) + .field("node_stale_sec", &self.node_stale_sec) + .field("mq_type", &self.mq_type) + .field("rabbitmq_url", &self.rabbitmq_url) + .field( + "admin_token", + &self.admin_token.as_ref().map(|_| "***REDACTED***"), + ) + .field( + "auth_token", + &self.auth_token.as_ref().map(|_| "***REDACTED***"), + ) + .field("auth_disabled", &self.auth_disabled) + .finish() + } } fn default_node_stale_sec() -> u64 { @@ -160,12 +197,13 @@ impl Default for ServerConfig { .or_else(|_| std::env::var("CNO_PORT")) .or_else(|_| std::env::var("PORT")) .unwrap_or_else(|_| "8090".to_string()); - let db_path = std::env::var("DCTS_DB_PATH") - .unwrap_or_else(|_| "data/dcts.db".to_string()); + let db_path = std::env::var("DCTS_DB_PATH").unwrap_or_else(|_| "data/dcts.db".to_string()); let queue_db_path = std::env::var("DCTS_QUEUE_DB_PATH") .unwrap_or_else(|_| "data/dcts_queue.db".to_string()); - let results_dir = std::env::var("DCTS_RESULTS_DIR") - .unwrap_or_else(|_| "data/results".to_string()); + let results_dir = + std::env::var("DCTS_RESULTS_DIR").unwrap_or_else(|_| "data/results".to_string()); + let backup_dir = + std::env::var("DCTS_BACKUP_DIR").unwrap_or_else(|_| "data/backups".to_string()); let grid_config = std::env::var("DCTS_GRID_CONFIG") .unwrap_or_else(|_| "workflows/sdB_cno.yaml".to_string()); // 默认设置为 7800 秒,比计算任务默认超时(7200 秒)高 600 秒缓冲,避免两边的超时检测同时触发冲突 @@ -177,26 +215,63 @@ impl Default for ServerConfig { .ok() .and_then(|v| v.parse::().ok()) .unwrap_or(60); - let mq_type = std::env::var("DCTS_MQ_TYPE") - .unwrap_or_else(|_| "sqlite".to_string()); + let mq_type = std::env::var("DCTS_MQ_TYPE").unwrap_or_else(|_| "sqlite".to_string()); let rabbitmq_url = std::env::var("DCTS_RABBITMQ_URL").ok(); - let auth_token = std::env::var("DCTS_AUTH_TOKEN").ok(); + + // ── 鉴权凭据解析 ── + let legacy_token = std::env::var("DCTS_AUTH_TOKEN") + .ok() + .filter(|s| !s.is_empty()); + let admin_token = std::env::var("DCTS_ADMIN_TOKEN") + .ok() + .filter(|s| !s.is_empty()) + .or_else(|| legacy_token.clone()); + if legacy_token.is_some() + && (std::env::var("DCTS_ADMIN_TOKEN").is_err() + || std::env::var("DCTS_ENROLLMENT_TOKEN").is_err()) + { + tracing::warn!( + "检测到旧的 DCTS_AUTH_TOKEN,已自动用作 admin/enrollment 凭据。\ + 建议迁移到 DCTS_ADMIN_TOKEN(管理)与 DCTS_ENROLLMENT_TOKEN(节点注册)" + ); + } + + let auth_disabled = std::env::var("DCTS_AUTH_DISABLE") + .map(|v| v == "1" || v.eq_ignore_ascii_case("true")) + .unwrap_or(false); + if auth_disabled { + tracing::warn!( + "⚠️ DCTS_AUTH_DISABLE=1 已生效:全部鉴权被跳过,仅供本地调试,切勿用于生产!" + ); + } + + // auth_token 兼容字段:用于 main.rs 判断「是否启用鉴权中间件」。 + // 启用条件 = 显式配置了 admin 或 enrollment 凭据,且未应急关闭。 + let auth_token = if auth_disabled { + None + } else { + admin_token.clone() + }; + Self { bind_addr: format!("0.0.0.0:{}", port), db_path, queue_db_path, results_dir, + backup_dir, grid_config, stale_sec, node_stale_sec, mq_type, rabbitmq_url, + admin_token, auth_token, + auth_disabled, } } } -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Clone, Serialize, Deserialize)] pub struct NodeConfig { pub node_id: String, pub server_url: String, @@ -204,7 +279,19 @@ pub struct NodeConfig { pub runtime_dir: String, pub work_dir: String, pub heartbeat_sec: u64, - pub auth_token: Option, +} + +impl std::fmt::Debug for NodeConfig { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("NodeConfig") + .field("node_id", &self.node_id) + .field("server_url", &self.server_url) + .field("max_slots", &self.max_slots) + .field("runtime_dir", &self.runtime_dir) + .field("work_dir", &self.work_dir) + .field("heartbeat_sec", &self.heartbeat_sec) + .finish() + } } impl Default for NodeConfig { @@ -215,22 +302,27 @@ impl Default for NodeConfig { .unwrap_or_else(|_| "http://127.0.0.1:8090".to_string()); let node_id = std::env::var("DCTS_NODE_ID") .or_else(|_| std::env::var("NODE_ID")) - .and_then(|v| if v.trim().is_empty() { Err(std::env::VarError::NotPresent) } else { Ok(v) }) + .and_then(|v| { + if v.trim().is_empty() { + Err(std::env::VarError::NotPresent) + } else { + Ok(v) + } + }) .unwrap_or_else(|_| format!("node-{}", uuid::Uuid::new_v4().simple())); let max_slots = std::env::var("DCTS_MAX_SLOTS") .or_else(|_| std::env::var("MAX_SLOTS")) .ok() .and_then(|v| v.parse::().ok()) .unwrap_or(4); - let runtime_dir = std::env::var("DCTS_RUNTIME_DIR") - .unwrap_or_else(|_| "data/runtime".to_string()); - let work_dir = std::env::var("DCTS_WORK_DIR") - .unwrap_or_else(|_| "data/work".to_string()); + let runtime_dir = + std::env::var("DCTS_RUNTIME_DIR").unwrap_or_else(|_| "data/runtime".to_string()); + let work_dir = std::env::var("DCTS_WORK_DIR").unwrap_or_else(|_| "data/work".to_string()); let heartbeat_sec = std::env::var("DCTS_HEARTBEAT_SEC") .ok() .and_then(|v| v.parse::().ok()) .unwrap_or(15); - let auth_token = std::env::var("DCTS_AUTH_TOKEN").ok(); + Self { node_id, server_url, @@ -238,7 +330,6 @@ impl Default for NodeConfig { runtime_dir, work_dir, heartbeat_sec, - auth_token, } } } @@ -250,21 +341,14 @@ mod tests { #[test] fn test_load_real_grid_configs() { let root = std::path::Path::new(env!("CARGO_MANIFEST_DIR")).join("../.."); - + let sdb_path = root.join("workflows/sdB_cno.yaml"); if sdb_path.exists() { - let cfg = GridConfig::load_from_file(&sdb_path).expect("解析 workflows/sdB_cno.yaml 发生失败"); + let cfg = GridConfig::load_from_file(&sdb_path) + .expect("解析 workflows/sdB_cno.yaml 发生失败"); assert_eq!(cfg.nworkers, 16); assert_eq!(cfg.niter, Some(100)); - } - - let dense_path = root.join("config_dense.yaml"); - if dense_path.exists() { - let cfg = GridConfig::load_from_file(&dense_path).expect("解析 config_dense.yaml 发生失败"); - assert_eq!(cfg.template.as_deref(), Some("templates/cno_atmos.5.tpl")); - assert_eq!(cfg.linelist.as_deref(), Some("data/gfVIS99.dat")); + assert!(cfg.seed_step_fallback); } } } - - diff --git a/crates/common/src/conv_check.rs b/crates/common/src/conv_check.rs index 1506ec2..875c0e8 100644 --- a/crates/common/src/conv_check.rs +++ b/crates/common/src/conv_check.rs @@ -80,15 +80,12 @@ pub fn check_fort9(path: &Path, chmax: f64) -> ConvCheckResult { } // Safely find depth with maximum absolute change without unwrap panic on NaN - let worst = match cur_rows - .iter() - .max_by(|a, b| { - a.maximum - .abs() - .partial_cmp(&b.maximum.abs()) - .unwrap_or(std::cmp::Ordering::Equal) - }) - { + let worst = match cur_rows.iter().max_by(|a, b| { + a.maximum + .abs() + .partial_cmp(&b.maximum.abs()) + .unwrap_or(std::cmp::Ordering::Equal) + }) { Some(row) => row, None => { return ConvCheckResult { @@ -98,7 +95,9 @@ pub fn check_fort9(path: &Path, chmax: f64) -> ConvCheckResult { last_iter, n_depths: 0, chmax, - error: Some("No valid iteration rows found when calculating maximum change".to_string()), + error: Some( + "No valid iteration rows found when calculating maximum change".to_string(), + ), }; } }; @@ -113,15 +112,22 @@ pub fn check_fort9(path: &Path, chmax: f64) -> ConvCheckResult { last_iter, n_depths: cur_rows.len(), chmax, - error: if is_valid_num { None } else { Some("Convergence value is NaN or Inf".to_string()) }, + error: if is_valid_num { + None + } else { + Some("Convergence value is NaN or Inf".to_string()) + }, } } /// Checks if an atmosphere file (.7) contains NaN lines (>10% NaN lines = invalid) using exact word boundary +/// +/// 文件缺失时返回 `false`(语义:不存在 NaN 内容)。这与“含 NaN 导致无效”是不同语义; +/// 调用方需先自行确认文件存在性,不应将“缺失”与“含 NaN”混为一谈。 pub fn atmosphere_has_nan(path: &Path) -> bool { let file = match File::open(path) { Ok(f) => f, - Err(_) => return true, + Err(_) => return false, }; let reader = BufReader::new(file); let mut total_lines = 0; @@ -142,7 +148,6 @@ pub fn atmosphere_has_nan(path: &Path) -> bool { (nan_lines as f64) > (total_lines as f64 * 0.1) } - #[cfg(test)] mod tests { use super::*; @@ -162,5 +167,9 @@ mod tests { let banana_file_path = dir.path().join("banana.7"); std::fs::write(&banana_file_path, "banana 2 3\nbanana 5 6\n7 8 9\n").unwrap(); assert!(!atmosphere_has_nan(&banana_file_path)); + + // Missing file returns false (absence != contains NaN) + let missing_path = dir.path().join("missing.7"); + assert!(!atmosphere_has_nan(&missing_path)); } } diff --git a/crates/common/src/embedded.rs b/crates/common/src/embedded.rs index 5b433bd..ad7a9f5 100644 --- a/crates/common/src/embedded.rs +++ b/crates/common/src/embedded.rs @@ -54,9 +54,13 @@ pub async fn ensure_runtime( write_if_changed(&synspec_exe, SYNSPEC_BIN, true)?; } - // 2. Fetch baseline equation of state partition function tables if missing locally - let common_files = &["irwin_bc.dat", "irwin_orig.dat", "tsuji.molec_bc2", "tsuji.molec_orig"]; + let common_files = &[ + "irwin_bc.dat", + "irwin_orig.dat", + "tsuji.molec_bc2", + "tsuji.molec_orig", + ]; ensure_specific_data_files(&data_dir, server_url, client, common_files).await?; // 3. Check gfVIS99.dat @@ -69,11 +73,15 @@ pub async fn ensure_runtime( fs::write(&linelist, &bytes)?; info!("成功下载并保存主谱线库 gfVIS99.dat"); } else { - anyhow::bail!("从服务端下载主谱线库 gfVIS99.dat 失败,HTTP 状态码: {}", resp.status()); + anyhow::bail!( + "从服务端下载主谱线库 gfVIS99.dat 失败,HTTP 状态码: {}", + resp.status() + ); } } - let abs_runtime_dir = fs::canonicalize(runtime_dir).unwrap_or_else(|_| runtime_dir.to_path_buf()); + let abs_runtime_dir = + fs::canonicalize(runtime_dir).unwrap_or_else(|_| runtime_dir.to_path_buf()); let tlusty_exe = abs_runtime_dir.join("tlusty_static"); let synspec_exe = abs_runtime_dir.join("synspec_static"); let data_dir = abs_runtime_dir.join("data"); @@ -103,17 +111,28 @@ pub async fn ensure_specific_data_files( let local_file = data_dir.join(filename); if !local_file.exists() { let file_url = format!("{}/api/data/file/{}", server_url, filename); - info!("本地缺失数据文件 {},开始从服务端拉取: {}", filename, file_url); + info!( + "本地缺失数据文件 {},开始从服务端拉取: {}", + filename, file_url + ); let resp = client.get(&file_url).send().await?; if resp.status().is_success() { let bytes = resp.bytes().await?; - let tmp_file = data_dir.join(format!("{}.{}.tmp", filename, uuid::Uuid::new_v4().simple())); + let tmp_file = data_dir.join(format!( + "{}.{}.tmp", + filename, + uuid::Uuid::new_v4().simple() + )); tokio::fs::write(&tmp_file, &bytes).await?; tokio::fs::rename(&tmp_file, &local_file).await?; info!("成功保存数据文件: {}", filename); } else { - anyhow::bail!("服务端返回 HTTP {} 错误,数据文件: {}", resp.status(), filename); + anyhow::bail!( + "服务端返回 HTTP {} 错误,数据文件: {}", + resp.status(), + filename + ); } } } diff --git a/crates/common/src/fort55_writer.rs b/crates/common/src/fort55_writer.rs index 59bd4a3..e5e6289 100644 --- a/crates/common/src/fort55_writer.rs +++ b/crates/common/src/fort55_writer.rs @@ -7,10 +7,16 @@ pub fn generate_fort55_content(cfg: &SynspecConfig) -> String { let line3 = " 0 0 0 0 0"; let line4 = " 1 1 0 0 0"; let line5 = " 0 0 0"; - let line6 = format!(" {:.1} {:.1} 10 0 {} {}", cfg.wstart, cfg.wend, cfg.rel_cutoff, cfg.abs_cutoff); + let line6 = format!( + " {:.1} {:.1} 10 0 {} {}", + cfg.wstart, cfg.wend, cfg.rel_cutoff, cfg.abs_cutoff + ); let line7 = " 0 0"; - format!("{}\n{}\n{}\n{}\n{}\n{}\n{}\n", line1, line2, line3, line4, line5, line6, line7) + format!( + "{}\n{}\n{}\n{}\n{}\n{}\n{}\n", + line1, line2, line3, line4, line5, line6, line7 + ) } #[cfg(test)] diff --git a/crates/common/src/gen_input5.rs b/crates/common/src/gen_input5.rs index 66d16ce..26a0e3a 100644 --- a/crates/common/src/gen_input5.rs +++ b/crates/common/src/gen_input5.rs @@ -9,40 +9,172 @@ struct IonDef { } const IONS_H: &[IonDef] = &[ - IonDef { iat: 1, iz: 0, nlevs: 9, typion: " H 1", filei: "data/h1.dat" }, - IonDef { iat: 1, iz: 1, nlevs: 1, typion: " H 2", filei: " " }, + IonDef { + iat: 1, + iz: 0, + nlevs: 9, + typion: " H 1", + filei: "data/h1.dat", + }, + IonDef { + iat: 1, + iz: 1, + nlevs: 1, + typion: " H 2", + filei: " ", + }, ]; const IONS_HE: &[IonDef] = &[ - IonDef { iat: 2, iz: 0, nlevs: 14, typion: "He 1", filei: "data/he1.dat" }, - IonDef { iat: 2, iz: 1, nlevs: 14, typion: "He 2", filei: "data/he2.dat" }, - IonDef { iat: 2, iz: 2, nlevs: 1, typion: "He 3", filei: " " }, + IonDef { + iat: 2, + iz: 0, + nlevs: 14, + typion: "He 1", + filei: "data/he1.dat", + }, + IonDef { + iat: 2, + iz: 1, + nlevs: 14, + typion: "He 2", + filei: "data/he2.dat", + }, + IonDef { + iat: 2, + iz: 2, + nlevs: 1, + typion: "He 3", + filei: " ", + }, ]; const IONS_C: &[IonDef] = &[ - IonDef { iat: 6, iz: 0, nlevs: 40, typion: " C 1", filei: "data/c1.dat" }, - IonDef { iat: 6, iz: 1, nlevs: 22, typion: " C 2", filei: "data/c2.dat" }, - IonDef { iat: 6, iz: 2, nlevs: 46, typion: " C 3", filei: "data/c3_34+12lev.dat" }, - IonDef { iat: 6, iz: 3, nlevs: 25, typion: " C 4", filei: "data/c4.dat" }, - IonDef { iat: 6, iz: 4, nlevs: 1, typion: " C 5", filei: " " }, + IonDef { + iat: 6, + iz: 0, + nlevs: 40, + typion: " C 1", + filei: "data/c1.dat", + }, + IonDef { + iat: 6, + iz: 1, + nlevs: 22, + typion: " C 2", + filei: "data/c2.dat", + }, + IonDef { + iat: 6, + iz: 2, + nlevs: 46, + typion: " C 3", + filei: "data/c3_34+12lev.dat", + }, + IonDef { + iat: 6, + iz: 3, + nlevs: 25, + typion: " C 4", + filei: "data/c4.dat", + }, + IonDef { + iat: 6, + iz: 4, + nlevs: 1, + typion: " C 5", + filei: " ", + }, ]; const IONS_N: &[IonDef] = &[ - IonDef { iat: 7, iz: 0, nlevs: 34, typion: " N 1", filei: "data/n1.dat" }, - IonDef { iat: 7, iz: 1, nlevs: 42, typion: " N 2", filei: "data/n2_32+10lev.dat" }, - IonDef { iat: 7, iz: 2, nlevs: 32, typion: " N 3", filei: "data/n3.dat" }, - IonDef { iat: 7, iz: 3, nlevs: 48, typion: " N 4", filei: "data/n4_34+14lev.dat" }, - IonDef { iat: 7, iz: 4, nlevs: 16, typion: " N 5", filei: "data/n5.dat" }, - IonDef { iat: 7, iz: 5, nlevs: 1, typion: " N 6", filei: " " }, + IonDef { + iat: 7, + iz: 0, + nlevs: 34, + typion: " N 1", + filei: "data/n1.dat", + }, + IonDef { + iat: 7, + iz: 1, + nlevs: 42, + typion: " N 2", + filei: "data/n2_32+10lev.dat", + }, + IonDef { + iat: 7, + iz: 2, + nlevs: 32, + typion: " N 3", + filei: "data/n3.dat", + }, + IonDef { + iat: 7, + iz: 3, + nlevs: 48, + typion: " N 4", + filei: "data/n4_34+14lev.dat", + }, + IonDef { + iat: 7, + iz: 4, + nlevs: 16, + typion: " N 5", + filei: "data/n5.dat", + }, + IonDef { + iat: 7, + iz: 5, + nlevs: 1, + typion: " N 6", + filei: " ", + }, ]; const IONS_O: &[IonDef] = &[ - IonDef { iat: 8, iz: 0, nlevs: 33, typion: " O 1", filei: "data/o1_23+10lev.dat" }, - IonDef { iat: 8, iz: 1, nlevs: 48, typion: " O 2", filei: "data/o2_36+12lev.dat" }, - IonDef { iat: 8, iz: 2, nlevs: 41, typion: " O 3", filei: "data/o3_28+13lev.dat" }, - IonDef { iat: 8, iz: 3, nlevs: 39, typion: " O 4", filei: "data/o4.dat" }, - IonDef { iat: 8, iz: 4, nlevs: 6, typion: " O 5", filei: "data/o5.dat" }, - IonDef { iat: 8, iz: 5, nlevs: 1, typion: " O 6", filei: " " }, + IonDef { + iat: 8, + iz: 0, + nlevs: 33, + typion: " O 1", + filei: "data/o1_23+10lev.dat", + }, + IonDef { + iat: 8, + iz: 1, + nlevs: 48, + typion: " O 2", + filei: "data/o2_36+12lev.dat", + }, + IonDef { + iat: 8, + iz: 2, + nlevs: 41, + typion: " O 3", + filei: "data/o3_28+13lev.dat", + }, + IonDef { + iat: 8, + iz: 3, + nlevs: 39, + typion: " O 4", + filei: "data/o4.dat", + }, + IonDef { + iat: 8, + iz: 4, + nlevs: 6, + typion: " O 5", + filei: "data/o5.dat", + }, + IonDef { + iat: 8, + iz: 5, + nlevs: 1, + typion: " O 6", + filei: " ", + }, ]; fn fmt_abn(logx: f64) -> String { @@ -64,11 +196,11 @@ pub fn make_input5( // Atoms block let mut atom_rows: Vec<(i32, String)> = vec![ - (2, "0.".to_string()), // 1 H - (2, fmt_abn(params.loghe)), // 2 He - (0, "0.".to_string()), // 3 Li - (0, "0.".to_string()), // 4 Be - (0, "0.".to_string()), // 5 B + (2, "0.".to_string()), // 1 H + (2, fmt_abn(params.loghe)), // 2 He + (0, "0.".to_string()), // 3 Li + (0, "0.".to_string()), // 4 Be + (0, "0.".to_string()), // 5 B ]; if has_c { @@ -81,7 +213,8 @@ pub fn make_input5( atom_rows.push((2, fmt_abn(params.logo))); // 8 O } - let natoms = 5 + (if has_c { 1 } else { 0 }) + (if has_n { 1 } else { 0 }) + (if has_o { 1 } else { 0 }); + let natoms = + 5 + (if has_c { 1 } else { 0 }) + (if has_n { 1 } else { 0 }) + (if has_o { 1 } else { 0 }); let mut atoms_block = format!(" {}\n* mode abn modpf\n", natoms); for (mode, abn) in &atom_rows { diff --git a/crates/common/src/lib.rs b/crates/common/src/lib.rs index 9b7006e..28f87f2 100644 --- a/crates/common/src/lib.rs +++ b/crates/common/src/lib.rs @@ -8,4 +8,3 @@ pub mod models; pub mod nst_writer; pub mod runner; pub mod seed_finder; - diff --git a/crates/common/src/logging.rs b/crates/common/src/logging.rs index b2583fd..29c1d80 100644 --- a/crates/common/src/logging.rs +++ b/crates/common/src/logging.rs @@ -17,7 +17,6 @@ impl FormatTime for LocalTimeFormatter { } } - /// Initializes high-performance, non-blocking structured logging for DCTS applications pub fn init_logging(app_name: &str, default_filter: &str) -> Result> { let mut guards = Vec::new(); @@ -29,7 +28,8 @@ pub fn init_logging(app_name: &str, default_filter: &str) -> Result + Send + Sync>> = Vec::new(); @@ -39,7 +39,9 @@ pub fn init_logging(app_name: &str, default_filter: &str) -> Result for GridPointStatus { fn from(s: &str) -> Self { match s { @@ -106,6 +105,10 @@ pub struct TaskSpec { pub task_type: TaskType, pub seed_point_name: Option, pub timeout_sec: u64, + /// 所属工作流名称,用于按工作流隔离队列清理(stop_workflow 只清当前工作流的任务)。 + /// 旧数据反序列化时缺省为 None。 + #[serde(default)] + pub workflow_name: Option, } #[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] @@ -251,10 +254,12 @@ mod tests { assert_eq!(GridPointStatus::Failed.to_string(), "failed"); assert_eq!(GridPointStatus::from("queued"), GridPointStatus::Queued); - assert_eq!(GridPointStatus::from("converged"), GridPointStatus::Converged); + assert_eq!( + GridPointStatus::from("converged"), + GridPointStatus::Converged + ); assert_eq!(GridPointStatus::from("done"), GridPointStatus::Converged); assert_eq!(GridPointStatus::from("failed"), GridPointStatus::Failed); assert_eq!(GridPointStatus::from("unknown"), GridPointStatus::Pending); } } - diff --git a/crates/common/src/nst_writer.rs b/crates/common/src/nst_writer.rs index 5d53dd9..9379bd1 100644 --- a/crates/common/src/nst_writer.rs +++ b/crates/common/src/nst_writer.rs @@ -63,4 +63,3 @@ mod tests { assert!(content.contains("IELCOR=-1")); } } - diff --git a/crates/common/src/runner.rs b/crates/common/src/runner.rs index 06c7925..b6ddcba 100644 --- a/crates/common/src/runner.rs +++ b/crates/common/src/runner.rs @@ -6,11 +6,11 @@ use crate::gen_input5::make_input5; use crate::models::{GridPointParams, ModelSummary, StageSummary, TaskType}; use crate::nst_writer::generate_nst_content; use anyhow::Result; -use tokio::fs::File; use std::path::{Path, PathBuf}; use std::process::Stdio; -use tokio::process::Command as AsyncCommand; use std::time::Instant; +use tokio::fs::File; +use tokio::process::Command as AsyncCommand; use tracing::{info, warn}; pub fn default_cold_chain() -> Vec { @@ -130,7 +130,15 @@ impl<'a> ExecutionRunner<'a> { seed_atmos: Option<&Path>, synspec_cfg: Option<&SynspecConfig>, ) -> Result { - self.run_model_with_timeout(params, task_type, custom_chain, seed_atmos, synspec_cfg, 7200).await + self.run_model_with_timeout( + params, + task_type, + custom_chain, + seed_atmos, + synspec_cfg, + 7200, + ) + .await } pub async fn run_model_with_timeout( @@ -157,7 +165,9 @@ impl<'a> ExecutionRunner<'a> { #[cfg(unix)] { - let abs_data_dir = tokio::fs::canonicalize(&self.runtime.data_dir).await.unwrap_or_else(|_| self.runtime.data_dir.clone()); + let abs_data_dir = tokio::fs::canonicalize(&self.runtime.data_dir) + .await + .unwrap_or_else(|_| self.runtime.data_dir.clone()); if let Err(e) = std::os::unix::fs::symlink(&abs_data_dir, &link_data) { warn!("构建 data 数据集软链时发生提示性告警: {}", e); } @@ -171,7 +181,10 @@ impl<'a> ExecutionRunner<'a> { if let Some(seed_path) = seed_atmos { if seed_path.is_file() { if let Err(e) = tokio::fs::copy(seed_path, &fort8).await { - warn!("向工作沙盒引导填载首期收敛模型种子 fort.8 发生复制错误: {}", e); + warn!( + "向工作沙盒引导填载首期收敛模型种子 fort.8 发生复制错误: {}", + e + ); } } } @@ -219,15 +232,24 @@ impl<'a> ExecutionRunner<'a> { } else if let Some(ref s_path) = current_seed { if s_path.is_file() { if let Err(e) = tokio::fs::copy(s_path, &fort8).await { - warn!("阶段 {} 重载候选近邻推算种子模型期间发生文件复制异常: {}", stage_def.label, e); + warn!( + "阶段 {} 重载候选近邻推算种子模型期间发生文件复制异常: {}", + stage_def.label, e + ); } } } // Run tlusty.exe let fin = File::open(&input5_path).await?.into_std().await; - let fout = File::create(model_dir.join(format!("{}.6", name))).await?.into_std().await; - let ferr = File::create(model_dir.join(format!("{}.err", name))).await?.into_std().await; + let fout = File::create(model_dir.join(format!("{}.6", name))) + .await? + .into_std() + .await; + let ferr = File::create(model_dir.join(format!("{}.err", name))) + .await? + .into_std() + .await; let child = AsyncCommand::new(&self.runtime.tlusty_exe) .current_dir(&model_dir) @@ -246,7 +268,6 @@ impl<'a> ExecutionRunner<'a> { } }; - let fort9 = model_dir.join("fort.9"); let fort7 = model_dir.join("fort.7"); @@ -293,7 +314,10 @@ impl<'a> ExecutionRunner<'a> { stage_summaries.push(stage_summary); if !final_converged && stage_def.require_converged { - warn!("阶段 {} 要求收敛但未达标,中止后续收敛链阶段", stage_def.label); + warn!( + "阶段 {} 要求收敛但未达标,中止后续收敛链阶段", + stage_def.label + ); break; } } @@ -344,14 +368,19 @@ impl<'a> ExecutionRunner<'a> { #[cfg(unix)] { - let abs_linelist = tokio::fs::canonicalize(&self.runtime.linelist).await.unwrap_or_else(|_| self.runtime.linelist.clone()); + let abs_linelist = tokio::fs::canonicalize(&self.runtime.linelist) + .await + .unwrap_or_else(|_| self.runtime.linelist.clone()); let _ = std::os::unix::fs::symlink(&abs_linelist, &fort19_path); } let input5_path = model_dir.join(format!("{}.5", name)); if input5_path.is_file() { let fin = File::open(&input5_path).await?.into_std().await; - let fout = File::create(model_dir.join(format!("{}.log", name))).await?.into_std().await; + let fout = File::create(model_dir.join(format!("{}.log", name))) + .await? + .into_std() + .await; let child = AsyncCommand::new(&self.runtime.synspec_exe) .current_dir(&model_dir) @@ -361,7 +390,8 @@ impl<'a> ExecutionRunner<'a> { .kill_on_drop(true) .spawn()?; - let status_res = run_child_async_with_timeout(child, timeout_sec).await; + let synspec_timeout_sec = 600_u64.min(timeout_sec); + let status_res = run_child_async_with_timeout(child, synspec_timeout_sec).await; let rc = match status_res { Ok(st) => st.code().unwrap_or(-1), Err(e) => { @@ -374,13 +404,25 @@ impl<'a> ExecutionRunner<'a> { // Copy/move outputs: fort.7 (Synspec spectrum) -> .spec, fort.17 -> .cont, fort.12 -> .iden if model_dir.join("fort.7").is_file() { - let _ = tokio::fs::rename(model_dir.join("fort.7"), model_dir.join(format!("{}.spec", name))).await; + let _ = tokio::fs::rename( + model_dir.join("fort.7"), + model_dir.join(format!("{}.spec", name)), + ) + .await; } if model_dir.join("fort.17").is_file() { - let _ = tokio::fs::copy(model_dir.join("fort.17"), model_dir.join(format!("{}.cont", name))).await; + let _ = tokio::fs::copy( + model_dir.join("fort.17"), + model_dir.join(format!("{}.cont", name)), + ) + .await; } if model_dir.join("fort.12").is_file() { - let _ = tokio::fs::copy(model_dir.join("fort.12"), model_dir.join(format!("{}.iden", name))).await; + let _ = tokio::fs::copy( + model_dir.join("fort.12"), + model_dir.join(format!("{}.iden", name)), + ) + .await; } } } else { @@ -415,3 +457,17 @@ impl<'a> ExecutionRunner<'a> { Ok(summary) } } + +#[cfg(test)] +mod tests { + #[test] + fn test_synspec_timeout_calculation() { + let long_tlusty_timeout: u64 = 7200; + let synspec_timeout = 600_u64.min(long_tlusty_timeout); + assert_eq!(synspec_timeout, 600); + + let short_tlusty_timeout: u64 = 300; + let synspec_timeout_short = 600_u64.min(short_tlusty_timeout); + assert_eq!(synspec_timeout_short, 300); + } +} diff --git a/crates/common/src/seed_finder.rs b/crates/common/src/seed_finder.rs index 6858bb6..ada937c 100644 --- a/crates/common/src/seed_finder.rs +++ b/crates/common/src/seed_finder.rs @@ -18,7 +18,11 @@ pub fn calculate_seed_distance(cand: &GridPointParams, target: &GridPointParams) + (cand.logn - target.logn).abs() + (cand.logo - target.logo).abs(); - if d_teff < 1.0 && d_logg < 0.01 && d_loghe < 0.01 { + // exact family 判定:Teff/logg/logHe 视为“同物理族”,仅 CNO 丰度不同。 + // Teff 容忍度取半步 5000K:实际网格 Teff 档位通常为整数千(20000/30000/.../60000), + // 半步既能覆盖 config_dense 等 10000K 步长的相邻档互作种子, + // 又避免跨过大 Teff 间距导致 sdB 高温模型用低温种子而不收敛(sdB_cno 步长 40000K 仍不命中 exact)。 + if d_teff < 5000.0 && d_logg < 0.01 && d_loghe < 0.01 { (true, d_cno) } else { // 距离公式物理意义与标定阐释: @@ -32,6 +36,3 @@ pub fn calculate_seed_distance(cand: &GridPointParams, target: &GridPointParams) (false, global_d) } } - - - diff --git a/crates/mq/src/lib.rs b/crates/mq/src/lib.rs index 7f6bf63..0ee4d14 100644 --- a/crates/mq/src/lib.rs +++ b/crates/mq/src/lib.rs @@ -1,2 +1 @@ pub mod sqlite_queue; - diff --git a/crates/mq/src/sqlite_queue.rs b/crates/mq/src/sqlite_queue.rs index 0a5dc3c..9cd7f53 100644 --- a/crates/mq/src/sqlite_queue.rs +++ b/crates/mq/src/sqlite_queue.rs @@ -10,11 +10,32 @@ struct SqliteCustomizer; impl r2d2::CustomizeConnection for SqliteCustomizer { fn on_acquire(&self, conn: &mut rusqlite::Connection) -> Result<(), rusqlite::Error> { - conn.pragma_update(None, "busy_timeout", 5000)?; + // 与主库一致:高并发 claim/report 下给 SQLITE_BUSY 足够重试窗口。 + conn.pragma_update(None, "busy_timeout", 15000)?; + conn.pragma_update(None, "wal_autocheckpoint", 1000)?; Ok(()) } } +/// 将 SQLite db 文件及其 WAL/SHM 侧车文件权限收紧为 0600(仅 owner 读写)。 +/// 与主库 dcts.db 的口径一致,作为纵深防御(队列库不含 token,但含任务 payload)。 +#[cfg(unix)] +fn restrict_db_file_perms(db_path: &str) { + use std::os::unix::fs::PermissionsExt; + let candidates = [ + std::path::PathBuf::from(db_path), + std::path::PathBuf::from(format!("{}-wal", db_path)), + std::path::PathBuf::from(format!("{}-shm", db_path)), + ]; + for p in candidates { + if let Ok(meta) = std::fs::metadata(&p) { + let mut perms = meta.permissions(); + perms.set_mode(0o600); + let _ = std::fs::set_permissions(&p, perms); + } + } +} + #[derive(Clone)] pub struct SqliteTaskQueue { pool: Pool, @@ -29,11 +50,15 @@ impl SqliteTaskQueue { } let manager = SqliteConnectionManager::file(&db_path_owned); let pool = Pool::builder() - .max_size(4) + .max_size(8) .connection_customizer(Box::new(SqliteCustomizer)) .build(manager) .context("Failed to build SQLite queue connection pool")?; + // 收紧队列 db 文件权限为 0600(与主库口径一致,纵深防御) + #[cfg(unix)] + restrict_db_file_perms(&db_path_owned); + let conn = pool.get()?; let _: String = conn.pragma_update_and_check(None, "journal_mode", "WAL", |r| r.get(0))?; conn.execute( @@ -42,14 +67,39 @@ impl SqliteTaskQueue { payload TEXT NOT NULL, status TEXT NOT NULL, created_at DATETIME NOT NULL, - claimed_at DATETIME + claimed_at DATETIME, + workflow_name TEXT, + claimed_by_node_id TEXT )", [], )?; + // 兼容旧库:若 task_queue 表已存在但缺少 workflow_name / claimed_by_node_id 列,则补列。 + // SQLite 的 ALTER TABLE ADD COLUMN 是在线操作,旧数据该列默认 NULL。 + // PRAGMA table_info 检测列是否存在以实现幂等 migration。 + let has_col = |conn: &rusqlite::Connection, col: &str| -> rusqlite::Result { + let mut stmt = conn.prepare("PRAGMA table_info(task_queue)")?; + let rows = stmt.query_map([], |r| r.get::<_, String>(1))?; + for r in rows { + if r.map(|name| name == col).unwrap_or(false) { + return Ok(true); + } + } + Ok(false) + }; + if !has_col(&conn, "workflow_name")? { + conn.execute("ALTER TABLE task_queue ADD COLUMN workflow_name TEXT", [])?; + } + if !has_col(&conn, "claimed_by_node_id")? { + conn.execute("ALTER TABLE task_queue ADD COLUMN claimed_by_node_id TEXT", [])?; + } conn.execute( "CREATE INDEX IF NOT EXISTS idx_task_queue_status_created ON task_queue(status, created_at)", [], )?; + conn.execute( + "CREATE INDEX IF NOT EXISTS idx_task_queue_workflow ON task_queue(workflow_name)", + [], + )?; Ok(pool) }) .await??; @@ -61,14 +111,15 @@ impl SqliteTaskQueue { pub async fn push_task(&self, task: &TaskSpec) -> Result<()> { let payload = serde_json::to_string(task)?; let task_id_str = task.task_id.to_string(); + let workflow_name = task.workflow_name.clone(); let pool = self.pool.clone(); tokio::task::spawn_blocking(move || -> Result<()> { let conn = pool.get().map_err(|e| anyhow::anyhow!("Queue DB pool error: {}", e))?; conn.execute( - "INSERT OR REPLACE INTO task_queue (task_id, payload, status, created_at) - VALUES (?1, ?2, 'pending', datetime('now'))", - params![task_id_str, payload], + "INSERT OR REPLACE INTO task_queue (task_id, payload, status, created_at, workflow_name) + VALUES (?1, ?2, 'pending', datetime('now'), ?3)", + params![task_id_str, payload, workflow_name], )?; Ok(()) }) @@ -77,8 +128,9 @@ impl SqliteTaskQueue { Ok(()) } - pub async fn pop_task(&self) -> Result> { + pub async fn pop_task(&self, claimant_node_id: &str) -> Result> { let pool = self.pool.clone(); + let claimant = claimant_node_id.to_string(); tokio::task::spawn_blocking(move || -> Result> { let mut attempts = 0; @@ -109,9 +161,11 @@ impl SqliteTaskQueue { let task: TaskSpec = serde_json::from_str(&payload)?; + // 记录任务归属:claim 时写入领用方 node_id,供 report 阶段校验, + // 杜绝「节点 A 领用、节点 B 上报」的跨节点伪造结果投毒。 tx.execute( - "UPDATE task_queue SET status = 'claimed', claimed_at = datetime('now') WHERE task_id = ?1", - params![task_id], + "UPDATE task_queue SET status = 'claimed', claimed_at = datetime('now'), claimed_by_node_id = ?2 WHERE task_id = ?1", + params![task_id, claimant], )?; tx.commit()?; @@ -139,8 +193,13 @@ impl SqliteTaskQueue { let id_owned = task_id.to_string(); tokio::task::spawn_blocking(move || -> Result<()> { - let conn = pool.get().map_err(|e| anyhow::anyhow!("Queue DB pool error: {}", e))?; - conn.execute("DELETE FROM task_queue WHERE task_id = ?1", params![id_owned])?; + let conn = pool + .get() + .map_err(|e| anyhow::anyhow!("Queue DB pool error: {}", e))?; + conn.execute( + "DELETE FROM task_queue WHERE task_id = ?1", + params![id_owned], + )?; Ok(()) }) .await??; @@ -148,13 +207,60 @@ impl SqliteTaskQueue { Ok(()) } - pub async fn requeue_stale_tasks(&self, stale_sec: u64) -> Result> { + /// 校验指定 task 是否由指定 node 领用(claimed 态且 claimed_by_node_id 匹配)。 + /// + /// 用于 report_task 阶段防止跨节点伪造结果:只有真正领用该 task 的 node 才能上报结果。 + /// 返回 (point_name, workflow_name):匹配时附带二者供 report 进一步校验「上报的点与领用的 + /// task 一致」并把 workflow_name 传给 record_task_report 以定向更新对应工作流的 grid_points + /// (多工作流分区:避免按 name 全局更新误改其他工作流同名点)。 + pub async fn verify_task_claim( + &self, + task_id: &str, + claimant_node_id: &str, + ) -> Result)>> { + let pool = self.pool.clone(); + let task_id = task_id.to_string(); + let claimant = claimant_node_id.to_string(); + + let res = + tokio::task::spawn_blocking(move || -> Result)>> { + let conn = pool + .get() + .map_err(|e| anyhow::anyhow!("Queue DB pool error: {}", e))?; + // 仅 claimed 态(尚未被 report 清理)且归属匹配才算有效领用 + let mut stmt = conn.prepare( + "SELECT payload FROM task_queue + WHERE task_id = ?1 AND claimed_by_node_id = ?2 AND status = 'claimed' LIMIT 1", + )?; + let row = stmt.query_row(params![task_id, claimant], |r| r.get::<_, String>(0)); + match row { + Ok(payload) => { + // 解析 payload 取出 point_name + workflow_name,供调用方校验与定向更新 + let task: TaskSpec = serde_json::from_str(&payload)?; + Ok(Some((task.point_name, task.workflow_name))) + } + Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None), + Err(e) => Err(e.into()), + } + }) + .await??; + Ok(res) + } + + /// 重投超时 claimed 任务回 pending,返回每个被重投任务的 (point_name, workflow_name)。 + /// + /// 返回 workflow_name 供调用方(main.rs 后台循环)按工作流分组调用 + /// reset_specific_grid_points_to_pending,避免跨工作流误改同名点(多工作流分区)。 + pub async fn requeue_stale_tasks( + &self, + stale_sec: u64, + ) -> Result)>> { let pool = self.pool.clone(); - let names = tokio::task::spawn_blocking(move || -> Result> { + let entries = tokio::task::spawn_blocking(move || -> Result)>> { let mut conn = pool.get().map_err(|e| anyhow::anyhow!("Queue DB pool error: {}", e))?; let tx = conn.transaction()?; - let mut point_names = Vec::new(); + let mut entries = Vec::new(); { // 改写为单一原子更新带 RETURNING 返回语句,消弭 TOCTOU (Time-Of-Check-To-Time-Of-Use) 竞态问题 @@ -165,32 +271,56 @@ impl SqliteTaskQueue { )?; let rows = stmt.query_map(params![stale_sec as i64], |row| row.get::<_, String>(0))?; for r in rows { - if let Ok(payload) = r { - if let Ok(task) = serde_json::from_str::(&payload) { - point_names.push(task.point_name); - } + let payload = match r { + Ok(p) => p, + Err(_) => continue, + }; + if let Ok(task) = serde_json::from_str::(&payload) { + entries.push((task.point_name, task.workflow_name)); } } } tx.commit()?; - Ok(point_names) + Ok(entries) }) .await??; - Ok(names) + Ok(entries) } pub async fn clear_queue(&self) -> Result<()> { let pool = self.pool.clone(); tokio::task::spawn_blocking(move || -> Result<()> { - let conn = pool.get().map_err(|e| anyhow::anyhow!("Queue DB pool error: {}", e))?; + let conn = pool + .get() + .map_err(|e| anyhow::anyhow!("Queue DB pool error: {}", e))?; conn.execute("DELETE FROM task_queue", [])?; Ok(()) }) .await??; Ok(()) } + + /// 仅清理指定工作流的排队任务。 + /// + /// 用于 stop_workflow 按工作流隔离清理,避免在多工作流场景下误清其他工作流的任务。 + pub async fn clear_queue_by_workflow(&self, workflow_name: &str) -> Result<()> { + let pool = self.pool.clone(); + let wf_owned = workflow_name.to_string(); + tokio::task::spawn_blocking(move || -> Result<()> { + let conn = pool + .get() + .map_err(|e| anyhow::anyhow!("Queue DB pool error: {}", e))?; + conn.execute( + "DELETE FROM task_queue WHERE workflow_name = ?1", + params![wf_owned], + )?; + Ok(()) + }) + .await??; + Ok(()) + } } #[cfg(test)] @@ -203,9 +333,11 @@ mod tests { async fn test_sqlite_task_queue_operations() { let temp_dir = tempfile::tempdir().unwrap(); let db_path = temp_dir.path().join("test_queue.db"); - let queue = SqliteTaskQueue::new(&db_path.to_string_lossy()).await.unwrap(); + let queue = SqliteTaskQueue::new(&db_path.to_string_lossy()) + .await + .unwrap(); - assert!(queue.pop_task().await.unwrap().is_none()); + assert!(queue.pop_task("test-node").await.unwrap().is_none()); let task_id = Uuid::new_v4(); let task = TaskSpec { @@ -222,24 +354,84 @@ mod tests { task_type: TaskType::ColdRun, seed_point_name: None, timeout_sec: 3600, + workflow_name: Some("test_wf".to_string()), }; queue.push_task(&task).await.unwrap(); - let popped = queue.pop_task().await.unwrap(); + let popped = queue.pop_task("test-node").await.unwrap(); assert!(popped.is_some()); let popped_task = popped.unwrap(); assert_eq!(popped_task.task_id, task_id); assert_eq!(popped_task.point_name, task.point_name); - assert!(queue.pop_task().await.unwrap().is_none()); + assert!(queue.pop_task("test-node").await.unwrap().is_none()); let requeued = queue.requeue_stale_tasks(0).await.unwrap(); assert_eq!(requeued.len(), 1); - let popped2 = queue.pop_task().await.unwrap(); + let popped2 = queue.pop_task("test-node").await.unwrap(); assert!(popped2.is_some()); queue.remove_task(&task_id.to_string()).await.unwrap(); - assert!(queue.pop_task().await.unwrap().is_none()); + assert!(queue.pop_task("test-node").await.unwrap().is_none()); + } + + /// 任务归属校验:领用方 node 匹配才放行,其他 node 校验失败(防跨节点伪造结果)。 + #[tokio::test] + async fn test_verify_task_claim_ownership() { + let temp_dir = tempfile::tempdir().unwrap(); + let db_path = temp_dir.path().join("claim_test.db"); + let queue = SqliteTaskQueue::new(&db_path.to_string_lossy()) + .await + .unwrap(); + + let task_id = Uuid::new_v4(); + let params = GridPointParams { + teff: 35000.0, + logg: 5.5, + loghe: -1.0, + logc: -2.0, + logn: -2.0, + logo: -2.0, + }; + let task = TaskSpec { + task_id, + point_name: params.model_name(), + params: params.clone(), + task_type: TaskType::ColdRun, + seed_point_name: None, + timeout_sec: 60, + workflow_name: None, + }; + queue.push_task(&task).await.unwrap(); + + // node-A 领用 + let popped = queue.pop_task("node-A").await.unwrap(); + assert!(popped.is_some()); + + // node-A 校验:匹配,返回绑定的 (point_name, workflow_name) + let claim = queue + .verify_task_claim(&task_id.to_string(), "node-A") + .await + .unwrap(); + assert_eq!( + claim.map(|(p, _)| p).as_deref(), + Some(params.model_name().as_str()) + ); + + // node-B 校验:非领用方,返回 None + let claim_b = queue + .verify_task_claim(&task_id.to_string(), "node-B") + .await + .unwrap(); + assert!(claim_b.is_none()); + + // 任务被清理(remove)后,任何 node 校验都失败 + queue.remove_task(&task_id.to_string()).await.unwrap(); + let claim_after = queue + .verify_task_claim(&task_id.to_string(), "node-A") + .await + .unwrap(); + assert!(claim_after.is_none()); } } diff --git a/crates/node/src/executor.rs b/crates/node/src/executor.rs index 61fb350..e60219a 100644 --- a/crates/node/src/executor.rs +++ b/crates/node/src/executor.rs @@ -13,17 +13,35 @@ pub async fn execute_task( work_dir: &Path, task: &TaskSpec, ) -> Result<(ModelSummary, Option>)> { - info!("开始执行计算任务 {} (网格点: {})", task.task_id, task.point_name); + info!( + "开始执行计算任务 {} (网格点: {})", + task.task_id, task.point_name + ); // 1. Pull ONLY missing atom model data files needed for this task let required_atom_files = &[ - "h1.dat", "he1.dat", "he2.dat", - "c1.dat", "c2.dat", "c3_34+12lev.dat", "c4.dat", - "n1.dat", "n2_32+10lev.dat", "n3.dat", "n4_34+14lev.dat", "n5.dat", - "o1_23+10lev.dat", "o2_36+12lev.dat", "o3_28+13lev.dat", "o4.dat", "o5.dat", + "h1.dat", + "he1.dat", + "he2.dat", + "c1.dat", + "c2.dat", + "c3_34+12lev.dat", + "c4.dat", + "n1.dat", + "n2_32+10lev.dat", + "n3.dat", + "n4_34+14lev.dat", + "n5.dat", + "o1_23+10lev.dat", + "o2_36+12lev.dat", + "o3_28+13lev.dat", + "o4.dat", + "o5.dat", ]; - if let Err(e) = ensure_specific_data_files(&runtime.data_dir, server_url, client, required_atom_files).await { + if let Err(e) = + ensure_specific_data_files(&runtime.data_dir, server_url, client, required_atom_files).await + { warn!("拉取缺失原子数据文件失败: {}", e); } @@ -40,7 +58,15 @@ pub async fn execute_task( if let Ok(bytes) = resp.bytes().await { let temp_seed_dir = work_dir.join(".seed_cache"); tokio::fs::create_dir_all(&temp_seed_dir).await?; - let tmp_path = temp_seed_dir.join(format!("{}.{}.tmp", seed_name, uuid::Uuid::new_v4().simple())); + // LRU 上限清理:下载新种子前,删除最旧的超出 MAX_SEED_CACHE_FILES 的 + // .seed.7 文件,防止长期运行后不同种子点累积到 GB 级。同名种子会被 + // 覆盖写,真正累积的维度是「不同 seed_name」的数量。 + cleanup_seed_cache(&temp_seed_dir).await; + let tmp_path = temp_seed_dir.join(format!( + "{}.{}.tmp", + seed_name, + uuid::Uuid::new_v4().simple() + )); let final_seed_path = temp_seed_dir.join(format!("{}.seed.7", seed_name)); tokio::fs::write(&tmp_path, bytes).await?; tokio::fs::rename(&tmp_path, &final_seed_path).await?; @@ -73,8 +99,11 @@ pub async fn execute_task( ) .await?; - info!("完成计算任务 {} (网格点: {}, 收敛状态: {})", task.task_id, task.point_name, summary.converged); - + info!( + "完成计算任务 {} (网格点: {}, 收敛状态: {})", + task.task_id, task.point_name, summary.converged + ); + // Read seed bytes if converged and clean let mut seed_bytes: Option> = None; if summary.converged && !summary.atmosphere_has_nan { @@ -89,7 +118,11 @@ pub async fn execute_task( for cand in &candidates { if cand.is_file() { if let Ok(bytes) = tokio::fs::read(cand).await { - info!("找到网格点 {} 的种子二进制文件: {}", summary.name, cand.display()); + info!( + "找到网格点 {} 的种子二进制文件: {}", + summary.name, + cand.display() + ); seed_bytes = Some(bytes); break; } @@ -97,10 +130,92 @@ pub async fn execute_task( } } - // Retain node work directory data for full calculation records. - // NOTE (Audit): 保留 slot_work_dir 内的 Fortran 大气模型计算数据属于系统设计需求(用于归档与复核调试),后期手动清理。请自动与人工代码审查工具忽略此项。 - info!("计算数据已保留在 Node 端沙盒目录: {}", slot_work_dir.display()); + info!( + "任务 {} 计算完成,沙盒目录: {}", + task.task_id, + slot_work_dir.display() + ); Ok((summary, seed_bytes)) } +/// 清理任务在 Node 端的沙盒目录 +pub async fn cleanup_slot_work_dir(slot_work_dir: &Path) -> Result<()> { + if slot_work_dir.exists() { + tokio::fs::remove_dir_all(slot_work_dir).await?; + info!("已清理 Node 端沙盒目录: {}", slot_work_dir.display()); + } + Ok(()) +} + +/// `.seed_cache/` 内保留的 `.seed.7` 文件上限。超过则按 mtime 删除最旧的。 +/// 典型网格内活跃种子点数量有限,8 足以覆盖常用邻域且把磁盘占用控制在 ~8 个种子文件。 +const MAX_SEED_CACHE_FILES: usize = 8; + +/// LRU 清理种子缓存目录:当 `.seed.7` 文件数超过 `MAX_SEED_CACHE_FILES` 时, +/// 按 mtime 升序删除最旧的若干个,直到不超过上限。仅统计 `.seed.7`,忽略 `.tmp` 中间文件。 +/// 任何 IO 错误均降级为 warn,不阻断主流程。 +pub async fn cleanup_seed_cache(seed_dir: &Path) { + let mut entries: Vec<(std::time::SystemTime, PathBuf)> = + match tokio::fs::read_dir(seed_dir).await { + Ok(mut rd) => { + let mut v = Vec::new(); + while let Ok(Some(entry)) = rd.next_entry().await { + let path = entry.path(); + // 仅纳入 .seed.7 文件(最终产物),跳过 .tmp 中间文件 + if path.extension().and_then(|e| e.to_str()) != Some("7") { + continue; + } + let file_name = match path.file_name().and_then(|n| n.to_str()) { + Some(n) => n, + None => continue, + }; + if !file_name.ends_with(".seed.7") { + continue; + } + let meta = match entry.metadata().await { + Ok(m) => m, + Err(_) => continue, + }; + let mtime = meta.modified().unwrap_or(std::time::SystemTime::UNIX_EPOCH); + v.push((mtime, path)); + } + v + } + Err(_) => return, + }; + + if entries.len() <= MAX_SEED_CACHE_FILES { + return; + } + + // 按 mtime 升序(最旧在前),删除超出上限的最旧文件 + entries.sort_by_key(|(mtime, _)| *mtime); + let to_remove = entries.len().saturating_sub(MAX_SEED_CACHE_FILES); + for (_, path) in entries.into_iter().take(to_remove) { + if let Err(e) = tokio::fs::remove_file(&path).await { + warn!("清理种子缓存文件 {} 失败: {}", path.display(), e); + } else { + info!("LRU 清理种子缓存文件: {}", path.display()); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn test_cleanup_slot_work_dir() { + let temp_dir = + std::env::temp_dir().join(format!("test_slot_work_dir_{}", uuid::Uuid::new_v4())); + tokio::fs::create_dir_all(&temp_dir).await.unwrap(); + tokio::fs::write(temp_dir.join("dummy.txt"), "content") + .await + .unwrap(); + + assert!(temp_dir.exists()); + cleanup_slot_work_dir(&temp_dir).await.unwrap(); + assert!(!temp_dir.exists()); + } +} diff --git a/crates/node/src/main.rs b/crates/node/src/main.rs index ba64134..8044822 100644 --- a/crates/node/src/main.rs +++ b/crates/node/src/main.rs @@ -8,7 +8,7 @@ use common::embedded::ensure_runtime; use common::logging::init_logging; use reqwest::Client; use std::path::Path; -use tracing::info; +use tracing::{info, warn}; use worker::NodeWorker; #[tokio::main] @@ -19,17 +19,36 @@ async fn main() -> Result<()> { info!("启动 DCTS 计算节点 (Distributed Computing TLUSTY/SYNSPEC Worker Node)..."); let node_cfg = NodeConfig::default(); - let runtime_dir = Path::new(&node_cfg.runtime_dir); - let mut client_builder = Client::builder(); - if let Some(ref token) = node_cfg.auth_token { - let mut headers = reqwest::header::HeaderMap::new(); - if let Ok(val) = reqwest::header::HeaderValue::from_str(&format!("Bearer {}", token)) { - headers.insert(reqwest::header::AUTHORIZATION, val); + + // ── Node 凭据获取 ── + // 1. 优先读取本地持久化的 node 专属 token(`.node_token`,权限 600)。 + // 2. 若不存在,调 /node/register 提交注册申请并轮询等待 Dashboard 管理员审批授权。 + let token_path = runtime_dir.join(".node_token"); + let node_token = match read_node_token(&token_path) { + Some(t) => { + info!("已加载本地持久化的 node 专属 token"); + t } - client_builder = client_builder.default_headers(headers); - } - let client = client_builder.build().unwrap_or_else(|_| Client::new()); + None => { + info!("本地未发现 node token,准备向服务端提交注册申请并等待管理员审批..."); + let public_client = build_client_with_token(None); + let issued = NodeWorker::register_and_fetch_token( + &public_client, + &node_cfg.server_url, + &node_cfg.node_id, + ) + .await + .context("向服务端提交申请或获取专属 token 失败")?; + + write_node_token(&token_path, &issued)?; + info!("已持久化获批的专属 node token 到 {}", token_path.display()); + issued + } + }; + + // 用 node 专属 token 构造后续所有请求的 client + let client = build_client_with_token(Some(&node_token)); info!("检查本地运行时二进制与基础数据文件,必要时从服务端拉取..."); let runtime = ensure_runtime(runtime_dir, &node_cfg.server_url, &client) @@ -41,3 +60,47 @@ async fn main() -> Result<()> { Ok(()) } + +/// 读取本地持久化的 node token;文件须存在且非空。 +fn read_node_token(path: &Path) -> Option { + let content = std::fs::read_to_string(path).ok()?; + let t = content.trim().to_string(); + if t.is_empty() { + None + } else { + Some(t) + } +} + +/// 持久化 node token 到本地文件,并设权限 600(仅 owner 可读写)。 +fn write_node_token(path: &Path, token: &str) -> Result<()> { + if let Some(parent) = path.parent() { + std::fs::create_dir_all(parent) + .with_context(|| format!("创建 token 目录失败: {}", parent.display()))?; + } + std::fs::write(path, token) + .with_context(|| format!("写入 token 文件失败: {}", path.display()))?; + #[cfg(unix)] + { + use std::os::unix::fs::PermissionsExt; + let mut perms = std::fs::metadata(path)?.permissions(); + perms.set_mode(0o600); + std::fs::set_permissions(path, perms)?; + } + Ok(()) +} + +/// 构造带 Authorization: Bearer 头的 reqwest client。 +fn build_client_with_token(token: Option<&str>) -> Client { + let mut builder = Client::builder(); + if let Some(t) = token { + let mut headers = reqwest::header::HeaderMap::new(); + if let Ok(val) = reqwest::header::HeaderValue::from_str(&format!("Bearer {}", t)) { + headers.insert(reqwest::header::AUTHORIZATION, val); + } else { + warn!("node token 含非法 HTTP 头字符,已忽略鉴权头"); + } + builder = builder.default_headers(headers); + } + builder.build().unwrap_or_else(|_| Client::new()) +} diff --git a/crates/node/src/reporter.rs b/crates/node/src/reporter.rs index a1bbe47..fdfbdf2 100644 --- a/crates/node/src/reporter.rs +++ b/crates/node/src/reporter.rs @@ -4,7 +4,6 @@ use reqwest::multipart::{Form, Part}; use reqwest::Client; use tracing::{info, warn}; - pub async fn report_result( client: &Client, server_url: &str, @@ -14,32 +13,33 @@ pub async fn report_result( ) -> Result<()> { let report_url = format!("{}/api/task/report", server_url); - let (status, converged, max_relc, atmo_has_nan, elapsed_sec, err_msg, summary_json, seed_bytes) = match exec_res { - Ok((s, s_bytes)) => ( - if s.converged { - TaskStatus::Completed - } else { - TaskStatus::Failed - }, - s.converged, - s.final_max_relc, - s.atmosphere_has_nan, - s.elapsed_sec, - s.note.clone(), - serde_json::to_string(&s).unwrap_or_default(), - s_bytes, - ), - Err(e) => ( - TaskStatus::Failed, - false, - None, - false, - 0.0, - Some(e.clone()), - serde_json::json!({"error": e}).to_string(), - None, - ), - }; + let (status, converged, max_relc, atmo_has_nan, elapsed_sec, err_msg, summary_json, seed_bytes) = + match exec_res { + Ok((s, s_bytes)) => ( + if s.converged { + TaskStatus::Completed + } else { + TaskStatus::Failed + }, + s.converged, + s.final_max_relc, + s.atmosphere_has_nan, + s.elapsed_sec, + s.note.clone(), + serde_json::to_string(&s).unwrap_or_default(), + s_bytes, + ), + Err(e) => ( + TaskStatus::Failed, + false, + None, + false, + 0.0, + Some(e.clone()), + serde_json::json!({"error": e}).to_string(), + None, + ), + }; let report = TaskReport { task_id: task.task_id, @@ -83,9 +83,19 @@ pub async fn report_result( return Ok(()); } Ok(resp) => { + let status = resp.status(); + // 401/403 表明 node token 已失效/被吊销(非临时故障),重试无意义且会丢结果。 + // 立即 bail 并打 error,与 claim_task 侧口径统一,提示运维介入。 + if status.as_u16() == 401 || status.as_u16() == 403 { + tracing::error!( + "上报任务 {} 被服务端拒绝 (HTTP {}):node token 可能已失效或被吊销,请检查并清理 .node_token 文件后重启节点以重新向服务端发起注册审批,停止重试", + task.task_id, status + ); + anyhow::bail!("node token 失效或被吊销 (HTTP {}),结果未上报", status); + } warn!( "向服务端上报任务 {} 结果失败 (尝试 {}/{}): HTTP {}", - task.task_id, attempt, max_attempts, resp.status() + task.task_id, attempt, max_attempts, status ); } Err(e) => { @@ -97,7 +107,7 @@ pub async fn report_result( } if attempt < max_attempts { - let backoff_secs = (1 << (attempt - 1)).min(60); + let backoff_secs = (1u64 << (attempt - 1).min(6)).min(60); let backoff = std::time::Duration::from_secs(backoff_secs); tokio::time::sleep(backoff).await; } @@ -109,4 +119,3 @@ pub async fn report_result( task.task_id ) } - diff --git a/crates/node/src/worker.rs b/crates/node/src/worker.rs index 1c9b317..7f2ce15 100644 --- a/crates/node/src/worker.rs +++ b/crates/node/src/worker.rs @@ -29,8 +29,96 @@ impl NodeWorker { } } + /// 仅注册并领取专属 token(供 main.rs 在本地无 token 时调用)。 + /// 支持免凭据申请注册并轮询等待管理员在 Web Dashboard 上点击同意。 + pub async fn register_and_fetch_token( + client: &Client, + server_url: &str, + node_id: &str, + ) -> Result { + info!( + "正在向服务端 {} 提交计算节点 {} 的注册申请...", + server_url, node_id + ); + + let req = NodeRegisterRequest { + node_id: node_id.to_string(), + host_name: gethostname::gethostname().to_string_lossy().to_string(), + max_slots: 0, + }; + + let resp = client + .post(format!("{}/api/node/register", server_url)) + .json(&req) + .send() + .await?; + + if !resp.status().is_success() { + anyhow::bail!("向服务端提交注册申请失败,HTTP 状态码: {}", resp.status()); + } + + let json: Value = resp.json().await?; + let status = json.get("status").and_then(|v| v.as_str()).unwrap_or(""); + + if status == "approved" { + if let Some(t) = json.get("node_token").and_then(|v| v.as_str()) { + return Ok(t.to_string()); + } + } + + info!( + "⏳ 节点 {} 的注册申请已提交!等待管理员在管理 Dashboard 上点击【同意接入】...", + node_id + ); + + // 轮询等待管理员在 Dashboard 上的 Approve + loop { + sleep(Duration::from_secs(5)).await; + + let check_req = serde_json::json!({ "node_id": node_id }); + let resp = match client + .post(format!("{}/api/node/check_status", server_url)) + .json(&check_req) + .send() + .await + { + Ok(r) => r, + Err(e) => { + warn!("轮询节点审批状态网络异常: {}", e); + continue; + } + }; + + if !resp.status().is_success() { + continue; + } + + let body: Value = match resp.json().await { + Ok(b) => b, + Err(_) => continue, + }; + + let check_status = body.get("status").and_then(|v| v.as_str()).unwrap_or(""); + if check_status == "approved" { + if let Some(token) = body.get("node_token").and_then(|v| v.as_str()) { + info!( + "🎉 节点 {} 已成功获取管理员授权!专属访问 Token 接收完成。", + node_id + ); + return Ok(token.to_string()); + } + } else if check_status == "rejected" { + anyhow::bail!("节点 {} 的注册申请已被管理员拒绝或清理", node_id); + } + } + } + + /// 正式注册(带真实 slot 数),供 run() 启动时刷新节点信息用。 pub async fn register(&self) -> Result<()> { - info!("正在向服务端 {} 注册计算节点 {}...", self.config.server_url, self.config.node_id); + info!( + "正在向服务端 {} 刷新节点 {} 注册信息...", + self.config.server_url, self.config.node_id + ); let req = NodeRegisterRequest { node_id: self.config.node_id.clone(), @@ -54,7 +142,10 @@ impl NodeWorker { pub async fn run(&self) -> Result<()> { self.register().await?; - info!("计算节点已激活,最大并行 Slot 槽位数: {}", self.config.max_slots); + info!( + "计算节点已激活,最大并行 Slot 槽位数: {}", + self.config.max_slots + ); // Start background heartbeat loop let hb_client = self.client.clone(); @@ -71,7 +162,8 @@ impl NodeWorker { if let Ok(mut sys) = s.lock() { sys.refresh_cpu(); } - }).await; + }) + .await; } sleep(Duration::from_millis(200)).await; { @@ -80,7 +172,8 @@ impl NodeWorker { if let Ok(mut sys) = s.lock() { sys.refresh_cpu(); } - }).await; + }) + .await; } loop { @@ -96,13 +189,17 @@ impl NodeWorker { let cpu_usage = sys.global_cpu_info().cpu_usage(); let mem_total = sys.total_memory() as f32; let mem_used = sys.used_memory() as f32; - let memory_usage = if mem_total > 0.0 { (mem_used / mem_total) * 100.0 } else { 0.0 }; + let memory_usage = if mem_total > 0.0 { + (mem_used / mem_total) * 100.0 + } else { + 0.0 + }; (cpu_usage, memory_usage) }) .await .unwrap_or((0.0, 0.0)); - let active = hb_slots.load(Ordering::Relaxed); + let active = hb_slots.load(Ordering::Acquire); let req = NodeHeartbeatRequest { node_id: hb_node_id.clone(), active_slots: active, @@ -110,7 +207,22 @@ impl NodeWorker { memory_usage, }; - let _ = hb_client.post(&hb_url).json(&req).send().await; + match hb_client.post(&hb_url).json(&req).send().await { + Ok(resp) => { + let status = resp.status(); + // 401/403:token 失效或被吊销。与 claim_task 口径统一:直接退出进程, + // 避免心跳线程持续发被拒请求刷日志、占用服务端限流计数。心跳通常比 + // claim 更高频,往往先于 claim_task 发现吊销。 + if status.as_u16() == 401 || status.as_u16() == 403 { + tracing::error!( + "节点 {} 心跳被服务端拒绝 (HTTP {}):node token 已失效或被吊销。请清理 .node_token 文件后重启节点以重新发起注册审批。进程将退出,依赖编排系统重启。", + hb_node_id, status + ); + std::process::exit(1); + } + } + Err(e) => warn!("节点 {} 心跳上报失败: {}", hb_node_id, e), + } } }); @@ -123,7 +235,7 @@ impl NodeWorker { tokio::spawn(async move { if tokio::signal::ctrl_c().await.is_ok() { info!("收到 Ctrl+C 终止信号,停止领用新任务,准备优雅退出 (再次按 Ctrl+C 可强制立即退出)..."); - shutdown_signal.store(true, Ordering::SeqCst); + shutdown_signal.store(true, Ordering::Release); // 二次 Ctrl+C 强行立即退出 if tokio::signal::ctrl_c().await.is_ok() { @@ -137,11 +249,11 @@ impl NodeWorker { // 带有优雅退出信号响应的任务领用主循环 loop { - if shutting_down.load(Ordering::Relaxed) { + if shutting_down.load(Ordering::Acquire) { break; } - let active = self.active_slots.load(Ordering::Relaxed); + let active = self.active_slots.load(Ordering::Acquire); if (active as usize) < self.config.max_slots { match self.claim_task().await { Ok(Some(task)) => { @@ -149,7 +261,7 @@ impl NodeWorker { info!("与服务端恢复网络连接,已自动重新上线并开始领用计算任务!"); was_disconnected = false; } - self.active_slots.fetch_add(1, Ordering::SeqCst); + self.active_slots.fetch_add(1, Ordering::AcqRel); let client = self.client.clone(); let server_url = self.config.server_url.clone(); let node_id = self.config.node_id.clone(); @@ -158,14 +270,29 @@ impl NodeWorker { let slots_counter = self.active_slots.clone(); tokio::spawn(async move { - let res = execute_task(&client, &server_url, &runtime, &work_dir, &task) - .await - .map_err(|e| e.to_string()); + let slot_work_dir = work_dir.join(format!("task_{}", task.task_id)); + let res = + execute_task(&client, &server_url, &runtime, &work_dir, &task) + .await + .map_err(|e| e.to_string()); - if let Err(e) = report_result(&client, &server_url, &node_id, &task, res).await { + let report_res = + report_result(&client, &server_url, &node_id, &task, res).await; + if report_res.is_ok() { + if let Err(e) = + crate::executor::cleanup_slot_work_dir(&slot_work_dir).await + { + warn!( + "清理任务 {} 的沙盒目录 {} 失败: {}", + task.task_id, + slot_work_dir.display(), + e + ); + } + } else if let Err(ref e) = report_res { warn!("向服务端上报任务 {} 计算结果失败: {}", task.task_id, e); } - slots_counter.fetch_sub(1, Ordering::SeqCst); + slots_counter.fetch_sub(1, Ordering::AcqRel); }); } Ok(None) => { @@ -187,17 +314,17 @@ impl NodeWorker { } // 等待在途任务完结(最多等待 30 秒) - if self.active_slots.load(Ordering::SeqCst) > 0 { + if self.active_slots.load(Ordering::Acquire) > 0 { info!( "正在等待 {} 个在途计算任务优雅完结 (上限 30 秒,按二次 Ctrl+C 可强行中断)...", - self.active_slots.load(Ordering::SeqCst) + self.active_slots.load(Ordering::Acquire) ); } let start_wait = std::time::Instant::now(); let mut last_log_time = std::time::Instant::now(); - while self.active_slots.load(Ordering::SeqCst) > 0 { + while self.active_slots.load(Ordering::Acquire) > 0 { if start_wait.elapsed().as_secs() >= 30 { warn!("在途任务等待超时 (30s),强制退出节点"); break; @@ -205,7 +332,7 @@ impl NodeWorker { if last_log_time.elapsed().as_secs() >= 5 { info!( "仍在等待 {} 个在途计算任务完结...", - self.active_slots.load(Ordering::SeqCst) + self.active_slots.load(Ordering::Acquire) ); last_log_time = std::time::Instant::now(); } @@ -220,7 +347,21 @@ impl NodeWorker { let claim_url = format!("{}/api/task/claim", self.config.server_url); let resp = self.client.post(&claim_url).send().await?; - if !resp.status().is_success() { + let status = resp.status(); + // 401/403 表明 node token 已被吊销或失效(区别于「暂无任务」与服务端 5xx 故障)。 + // 服务端故障返回 5xx 会走 !is_success() 的 Ok(None) 分支,仅在网络层/鉴权层拒绝时 + // 才是真正的吊销。此时继续轮询只会持续产生被拒请求并刷日志,故直接退出进程, + // 由编排系统(Docker restart / systemd / k8s)拉起;新进程发现 .node_token 失效后 + // 会自动走注册审批流程重新申请。 + if status.as_u16() == 401 || status.as_u16() == 403 { + tracing::error!( + "领用任务被服务端拒绝 (HTTP {}):node token 已失效或被吊销。请清理 .node_token 文件后重启节点以重新发起注册审批。进程将退出,依赖编排系统重启。", + status + ); + std::process::exit(1); + } + + if !status.is_success() { return Ok(None); } diff --git a/crates/server/Cargo.toml b/crates/server/Cargo.toml index 3fef921..0b1495c 100644 --- a/crates/server/Cargo.toml +++ b/crates/server/Cargo.toml @@ -25,5 +25,8 @@ chrono.workspace = true uuid.workspace = true tempfile.workspace = true dotenvy.workspace = true +sha2.workspace = true +hex.workspace = true +subtle = "2" diff --git a/crates/server/src/api/admin.rs b/crates/server/src/api/admin.rs new file mode 100644 index 0000000..8380855 --- /dev/null +++ b/crates/server/src/api/admin.rs @@ -0,0 +1,154 @@ +//! 管理 API(Admin 角色)。 +//! +//! 提供 node 凭据的可视化与运维操作,供 Dashboard 管理界面调用: +//! - 列出所有节点及其凭据状态(在线/token 是否有效/吊销/颁发时间) +//! - 吊销指定节点的专属 token(立即失效,不影响其他节点) +//! - 重新颁发指定节点的专属 token(返回新明文,旧 token 失效) +//! +//! 这些端点均要求 Admin 角色(见 mod.rs 授权矩阵),node 自身无权操作他人或自身凭据, +//! 从而保证「吊销/重发」是管理员主动行为,避免被攻陷节点篡改凭据体系。 + +use super::{is_valid_node_id, AppState}; +use axum::{ + extract::{Path as AxumPath, State}, + http::StatusCode, + response::IntoResponse, + Json, +}; +use serde_json::json; +use tracing::{info, warn}; + +/// GET /api/admin/nodes — 列出全部节点及凭据状态。 +pub async fn list_nodes( + State(state): State, +) -> Result { + match state.db.list_nodes_with_credentials().await { + Ok(list) => Ok(( + StatusCode::OK, + Json(json!({ "success": true, "message": "成功获取节点列表", "data": list })), + )), + Err(e) => Err(e.into()), + } +} + +/// POST /api/admin/nodes/:node_id/revoke — 吊销指定节点的专属 token。 +/// +/// 吊销后该 node 的现有 token 立即失效,须重新走注册流程领取新 token。 +/// 操作幂等:对无凭据记录或已吊销的节点调用不会报错。 +pub async fn revoke_node( + State(state): State, + AxumPath(node_id): AxumPath, +) -> Result { + // node_id 白名单校验,防止注入或异常输入(与 register_node 的 node_id 来源口径一致) + if !is_valid_node_id(&node_id) { + return Err(crate::api::AppError::BadRequest( + "非法的节点 ID 参数".to_string(), + )); + } + match state.db.revoke_node_token(&node_id).await { + Ok(_) => { + info!("管理员已吊销节点 {} 的专属 token", node_id); + Ok(( + StatusCode::OK, + Json( + json!({ "success": true, "message": format!("节点 '{}' 的 token 已吊销", node_id) }), + ), + )) + } + Err(e) => { + warn!("吊销节点 {} token 失败: {}", node_id, e); + Err(e.into()) + } + } +} + +/// POST /api/admin/nodes/:node_id/reissue — 重新颁发指定节点的专属 token。 +/// +/// 旧 token 立即失效,返回新 token 明文(仅此一次,DB 只存 hash)。 +/// 节点需用新 token 重新注册或由管理员手动同步到节点本地 `.node_token`。 +pub async fn reissue_node( + State(state): State, + AxumPath(node_id): AxumPath, +) -> Result { + if !is_valid_node_id(&node_id) { + return Err(crate::api::AppError::BadRequest( + "非法的节点 ID 参数".to_string(), + )); + } + // 仅允许对已注册的节点重发 token(防止凭据表被写入幽灵 node_id) + match state.db.get_node_exists(&node_id).await { + Ok(false) => { + return Err(crate::api::AppError::NotFound(format!( + "节点 '{}' 不存在,请先注册", + node_id + ))); + } + Ok(true) => {} + Err(e) => return Err(e.into()), + } + match state.db.issue_node_token(&node_id).await { + Ok(new_token) => { + info!("管理员已为节点 {} 重新颁发专属 token", node_id); + Ok(( + StatusCode::OK, + Json(json!({ + "success": true, + "message": format!("节点 '{}' 的 token 已重新颁发,请将新 token 同步到该节点", node_id), + "node_token": new_token, + })), + )) + } + Err(e) => { + warn!("为节点 {} 重新颁发 token 失败: {}", node_id, e); + Err(e.into()) + } + } +} + +/// POST /api/admin/nodes/:node_id/approve — 管理员同意节点接入申请。 +pub async fn approve_node( + State(state): State, + AxumPath(node_id): AxumPath, +) -> Result { + if !is_valid_node_id(&node_id) { + return Err(crate::api::AppError::BadRequest( + "非法的节点 ID 参数".to_string(), + )); + } + match state.db.approve_node(&node_id).await { + Ok(_token) => { + info!("管理员已同意节点 {} 的接入申请并生成专属 Token", node_id); + Ok(( + StatusCode::OK, + Json( + json!({ "success": true, "message": format!("节点 '{}' 已授权加入集群", node_id) }), + ), + )) + } + Err(e) => Err(e.into()), + } +} + +/// POST /api/admin/nodes/:node_id/reject — 管理员拒绝节点接入申请。 +pub async fn reject_node( + State(state): State, + AxumPath(node_id): AxumPath, +) -> Result { + if !is_valid_node_id(&node_id) { + return Err(crate::api::AppError::BadRequest( + "非法的节点 ID 参数".to_string(), + )); + } + match state.db.reject_node(&node_id).await { + Ok(_) => { + info!("管理员已拒绝节点 {} 的接入申请并移除", node_id); + Ok(( + StatusCode::OK, + Json( + json!({ "success": true, "message": format!("已拒绝节点 '{}' 的接入申请", node_id) }), + ), + )) + } + Err(e) => Err(e.into()), + } +} diff --git a/crates/server/src/api/auth.rs b/crates/server/src/api/auth.rs new file mode 100644 index 0000000..feb8899 --- /dev/null +++ b/crates/server/src/api/auth.rs @@ -0,0 +1,121 @@ +//! 管理员表单登录与凭据校验 API。 +//! +//! 提供基于短密码的身份认证服务: +//! - POST /api/login:校验管理员密码,成功后返回 Admin Token,并记录 IP 错误次数防止暴力破解。 +//! - GET /api/auth/check:由 auth_middleware 保护,供前端初始化时检测当前保存的 Token 是否有效。 + +use super::{ct_eq_str, AppState}; +use axum::{ + extract::{ConnectInfo, State}, + http::StatusCode, + response::IntoResponse, + Json, +}; +use serde::{Deserialize, Serialize}; +use std::net::SocketAddr; +use tracing::{info, warn}; + +#[derive(Debug, Deserialize)] +pub struct LoginRequest { + pub password: String, +} + +#[derive(Debug, Serialize)] +pub struct LoginResponse { + pub success: bool, + pub message: String, + pub token: Option, +} + +/// POST /api/login — 管理员密码登录端点。 +pub async fn login( + State(state): State, + ConnectInfo(addr): ConnectInfo, + Json(req): Json, +) -> Result { + let client_ip = addr.ip(); + + // 限流检查:5 分钟内最多允许 5 次失败尝试(基于 RateLimiter 防暴力破解) + if state.rate_limiter.is_rate_limited(client_ip) { + warn!("客户端 IP {} 登录失败次数过多,已临时封禁锁定", client_ip); + return Err(crate::api::AppError::TooManyRequests( + "登录失败次数过多,已被临时锁定,请 5 分钟后再试".to_string(), + )); + } + + let admin_token = match state.admin_token.as_deref() { + Some(t) if !t.is_empty() => t, + _ => { + warn!("系统未配置 admin_token 且鉴权未禁用,拒绝登录"); + return Err(crate::api::AppError::Forbidden( + "服务端未配置管理员凭据,请检查配置文件".to_string(), + )); + } + }; + + // 恒定时间密码比对(防时序旁路攻击) + if ct_eq_str(&req.password, admin_token) { + // 生成随机 64 位 Session Token + let session_token = format!( + "{}{}", + uuid::Uuid::new_v4().simple(), + uuid::Uuid::new_v4().simple() + ); + let expiry = std::time::Instant::now() + std::time::Duration::from_secs(24 * 3600); + + // 存储 Token 到内存中(带容量上限清理) + { + let mut sessions = state.admin_sessions.write().await; + let now = std::time::Instant::now(); + // 1. 清理已过期的 session + sessions.retain(|_, exp| *exp > now); + // 2. 若超出容量限制,淘汰最老/最快过期的 session + while sessions.len() >= crate::api::MAX_ADMIN_SESSIONS { + if let Some(oldest_key) = sessions + .iter() + .min_by_key(|(_, exp)| **exp) + .map(|(k, _)| k.clone()) + { + sessions.remove(&oldest_key); + } else { + break; + } + } + sessions.insert(session_token.clone(), expiry); + } + + info!( + "客户端 IP {} 密码验证成功,已颁发 Admin Session Token", + client_ip + ); + Ok(( + StatusCode::OK, + Json(LoginResponse { + success: true, + message: "登录成功".to_string(), + token: Some(session_token), + }), + )) + } else { + warn!("客户端 IP {} 登录密码校验失败", client_ip); + // 记录一次失败 + state.rate_limiter.record_failure(client_ip); + Err(crate::api::AppError::Unauthorized( + "管理员密码错误,请重新输入".to_string(), + )) + } +} + +/// GET /api/auth/check — 校验当前 Admin Token 是否有效。 +/// +/// 放在 auth_middleware(Role::Admin)之后,只要到达此 handler 说明 Token 校验必定成功。 +pub async fn check_auth() -> impl IntoResponse { + ( + StatusCode::OK, + Json(serde_json::json!({ + "success": true, + "message": "Token 验证有效", + "authenticated": true + })), + ) +} diff --git a/crates/server/src/api/data.rs b/crates/server/src/api/data.rs index 7156b7a..4b7243c 100644 --- a/crates/server/src/api/data.rs +++ b/crates/server/src/api/data.rs @@ -1,37 +1,81 @@ -use axum::{ - body::Body, - extract::Path as AxumPath, - http::{header, StatusCode}, - response::IntoResponse, -}; +use axum::{body::Body, extract::Path as AxumPath, http::header, response::IntoResponse}; use std::path::{Path, PathBuf}; use tokio::fs::File; use tokio_util::io::ReaderStream; -pub async fn download_single_data_file(AxumPath(filename): AxumPath) -> axum::response::Response { +pub async fn download_single_data_file( + AxumPath(filename): AxumPath, +) -> Result { let safe_name = Path::new(&filename) .file_name() .map(|s| s.to_string_lossy().to_string()) .unwrap_or_default(); if safe_name.is_empty() || safe_name.starts_with('.') { - return (StatusCode::BAD_REQUEST, "无效的数据文件名").into_response(); + return Err(crate::api::AppError::BadRequest( + "无效的数据文件名".to_string(), + )); } // 严苛白名单过滤:严防 `..`、特殊符号注入及路径穿透攻击 - if !safe_name.chars().all(|c| c.is_ascii_alphanumeric() || c == '.' || c == '_' || c == '-' || c == '+' || c == '@') { - tracing::warn!("拦截到疑似非法字符构造的敏感及越界资源抓取行为: {}", safe_name); - return (StatusCode::BAD_REQUEST, "参数非法,请求的文件包含系统不许可的危险专属占位或路径重定向字符").into_response(); + if !safe_name.chars().all(|c| { + c.is_ascii_alphanumeric() || c == '.' || c == '_' || c == '-' || c == '+' || c == '@' + }) { + tracing::warn!( + "拦截到疑似非法字符构造的敏感及越界资源抓取行为: {}", + safe_name + ); + return Err(crate::api::AppError::BadRequest( + "参数非法,请求的文件包含系统不许可的危险专属占位或路径重定向字符".to_string(), + )); } let rel_path = format!("assets/data/{}", safe_name); tracing::debug!("服务端处理数据文件下载请求: {}", safe_name); - stream_asset_file(&rel_path, "application/octet-stream").await.into_response() + stream_asset_file(&rel_path, "application/octet-stream").await } -pub async fn download_linelist() -> axum::response::Response { - let linelist_path = std::env::var("DCTS_LINELIST_PATH").unwrap_or_else(|_| "assets/gfVIS99.dat".to_string()); - stream_asset_file(&linelist_path, "application/octet-stream").await.into_response() +pub async fn download_linelist() -> Result { + let linelist_path = + std::env::var("DCTS_LINELIST_PATH").unwrap_or_else(|_| "assets/gfVIS99.dat".to_string()); + + // 路径规约校验:DCTS_LINELIST_PATH 解析后的绝对路径必须落在 assets 根目录内, + // 防止环境变量被设为 ../../etc/passwd 之类导致任意文件流出。 + // assets 根目录优先取 DCTS_ASSETS_DIR,回退到相对路径 assets。 + let assets_root = std::env::var("DCTS_ASSETS_DIR").unwrap_or_else(|_| "assets".to_string()); + if !is_path_within_assets(&linelist_path, &assets_root) { + tracing::warn!( + "DCTS_LINELIST_PATH '{}' 不在 assets 根目录 '{}' 内,拒绝下载", + linelist_path, + assets_root + ); + return Err(crate::api::AppError::Forbidden( + "请求的谱线文件路径越界,已被拒绝".to_string(), + )); + } + + stream_asset_file(&linelist_path, "application/octet-stream").await +} + +/// 校验 target 路径(经 canonicalize 后)是否落在 assets 根目录之内。 +/// 对不存在的路径(canonicalize 失败)回退到 starts_with 的词法比较,宁可偏严。 +fn is_path_within_assets(target: &str, assets_root: &str) -> bool { + // 严防 `..` 词法穿透 + if target.contains("..") { + return false; + } + + let target_path = Path::new(target); + let root_path = Path::new(assets_root); + + let target_abs = std::fs::canonicalize(target_path).ok(); + let root_abs = std::fs::canonicalize(root_path).ok(); + + match (target_abs, root_abs) { + (Some(t), Some(r)) => t.starts_with(&r), + // 路径尚未存在时用词法前缀比较(canonicalize 需要文件存在) + _ => target_path.starts_with(root_path), + } } fn resolve_asset(rel_path: &str) -> Option { @@ -65,10 +109,17 @@ fn resolve_asset(rel_path: &str) -> Option { None } -async fn stream_asset_file(rel_path: &str, content_type: &'static str) -> impl IntoResponse { +async fn stream_asset_file( + rel_path: &str, + content_type: &'static str, +) -> Result { let resolved_path = match resolve_asset(rel_path) { Some(p) => p, - None => return (StatusCode::NOT_FOUND, "资源数据文件不存在").into_response(), + None => { + return Err(crate::api::AppError::NotFound( + "资源数据文件不存在".to_string(), + )) + } }; match File::open(&resolved_path).await { @@ -89,9 +140,11 @@ async fn stream_asset_file(rel_path: &str, content_type: &'static str) -> impl I (header::CONTENT_DISPOSITION, disposition), ]; - (headers, body).into_response() + Ok((headers, body).into_response()) + } + Err(e) => { + let boxed_err: anyhow::Error = e.into(); + Err(crate::api::AppError::Internal(boxed_err)) } - Err(_) => (StatusCode::INTERNAL_SERVER_ERROR, "无法读取资源数据文件").into_response(), } } - diff --git a/crates/server/src/api/error.rs b/crates/server/src/api/error.rs new file mode 100644 index 0000000..18070c1 --- /dev/null +++ b/crates/server/src/api/error.rs @@ -0,0 +1,51 @@ +use axum::{ + http::StatusCode, + response::{IntoResponse, Response}, + Json, +}; +use serde_json::json; +use tracing::error; + +pub enum AppError { + BadRequest(String), + Unauthorized(String), + Forbidden(String), + NotFound(String), + Conflict(String), + TooManyRequests(String), + Internal(anyhow::Error), +} + +impl IntoResponse for AppError { + fn into_response(self) -> Response { + let (status, error_message) = match self { + AppError::BadRequest(msg) => (StatusCode::BAD_REQUEST, msg), + AppError::Unauthorized(msg) => (StatusCode::UNAUTHORIZED, msg), + AppError::Forbidden(msg) => (StatusCode::FORBIDDEN, msg), + AppError::NotFound(msg) => (StatusCode::NOT_FOUND, msg), + AppError::Conflict(msg) => (StatusCode::CONFLICT, msg), + AppError::TooManyRequests(msg) => (StatusCode::TOO_MANY_REQUESTS, msg), + AppError::Internal(err) => { + error!("Internal server error: {:?}", err); + ( + StatusCode::INTERNAL_SERVER_ERROR, + "Internal server error".to_string(), + ) + } + }; + + let body = Json(json!({ + "success": false, + "message": error_message, + "data": serde_json::Value::Null + })); + + (status, body).into_response() + } +} + +impl From for AppError { + fn from(inner: anyhow::Error) -> Self { + AppError::Internal(inner) + } +} diff --git a/crates/server/src/api/mod.rs b/crates/server/src/api/mod.rs index 8166c41..3b10686 100644 --- a/crates/server/src/api/mod.rs +++ b/crates/server/src/api/mod.rs @@ -1,20 +1,28 @@ +pub mod admin; +pub mod auth; pub mod data; +pub mod error; pub mod node; +pub mod rate_limit; pub mod seed; pub mod status; pub mod task; pub mod workflow; +pub use error::AppError; + +use crate::db::Database; +use crate::scheduler::GridScheduler; use axum::{ extract::State, http::{header, Request, StatusCode}, middleware::Next, response::IntoResponse, }; -use crate::db::Database; -use crate::scheduler::GridScheduler; use mq::sqlite_queue::SqliteTaskQueue; +use sha2::{Digest, Sha256}; use std::sync::Arc; +use subtle::ConstantTimeEq; #[derive(Clone)] pub struct AppState { @@ -22,53 +30,253 @@ pub struct AppState { pub queue: Arc, pub scheduler: Arc, pub results_dir: String, + /// 限流与密码防暴破限速器 + pub rate_limiter: rate_limit::RateLimiter, + /// 兼容字段:Some 表示「已启用某种鉴权」,用于 main.rs 决定是否挂载鉴权中间件。 pub auth_token: Option, + /// Admin 凭据(管理 Dashboard / workflow 写操作)。 + pub admin_token: Option, + /// 应急开关:跳过全部鉴权(仅本地调试)。 + pub auth_disabled: bool, + /// 动态 Session Tokens(登录后发放),设置 24 小时过期 + pub admin_sessions: + Arc>>, } -/// 固定时间敏感字符串一致性核验函数,彻底消解时序测信道猜测危险 -fn constant_time_eq(a: &str, b: &str) -> bool { - let a_bytes = a.as_bytes(); - let b_bytes = b.as_bytes(); - let mut diff = (a_bytes.len() ^ b_bytes.len()) as u64; - // 遍历目标 secret (b_bytes) 的完整长度,使耗时仅受 server 预期 token 长度决定 - for (i, &y) in b_bytes.iter().enumerate() { - let x = if i < a_bytes.len() { a_bytes[i] } else { 0 }; - diff |= (x ^ y) as u64; +/// Admin Session 最大保存上限 +pub const MAX_ADMIN_SESSIONS: usize = 100; + +/// 已认证的 Node 身份(中间件校验 node token 通过后注入 request extension)。 +/// +/// 下游 handler(heartbeat / claim / report)通过 `Extension` 取出, +/// 用于校验请求体里声称的 node_id 与 token 绑定的 node_id 一致,杜绝跨节点冒充。 +#[derive(Clone)] +pub struct AuthenticatedNode { + pub node_id: String, +} + +/// 授权角色:决定某条路径需要哪类主体才能访问。 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum Role { + /// 公开免鉴权端点:登录 /login,节点注册申请 /node/register,审批状态查询 /node/check_status。 + Public, + /// Node 运行态:心跳/领任务/上报/下载种子与数据。需 node 专属 token。 + Node, + /// 管理操作:workflow CRUD、起停、查看 status、审批节点。需 admin token。 + Admin, +} + +/// 路径 → 角色授权矩阵。 +/// +/// 设计依据(最小权限): +/// - Admin 写操作(workflow CRUD / start / stop / status / approve / reject)只对 admin token 开放。 +/// - Node 运行态接口只认 node 专属 token(管理员在 Dashboard 审批后颁发,绑定 node_id,可吊销)。 +/// - 注册端点 /node/register 和状态轮询 /node/check_status 为 Public 免凭据(提交申请 ➔ 待管理员审批)。 +/// +/// 注意:路径已去掉 `/api` 前缀(nest 挂载后中间件看到的 path 不含 nest 前缀)。 +fn required_role(path: &str, method: &axum::http::Method) -> Option { + use axum::http::Method; + // 公开免鉴权端点 + if (path == "/login" || path == "/node/register" || path == "/node/check_status") + && method == Method::POST + { + return Some(Role::Public); } - diff == 0 + // 校验身份与状态 -> Admin + if path == "/auth/check" && method == Method::GET { + return Some(Role::Admin); + } + // 写操作 → Admin + if path == "/workflows" && (method == Method::POST || method == Method::GET) { + return Some(Role::Admin); + } + if path.starts_with("/workflows/") { + // GET/PUT/DELETE /workflows/:name, POST /start|stop → Admin + return Some(Role::Admin); + } + if path == "/status" && method == Method::GET { + return Some(Role::Admin); + } + // 管理 API(节点凭据查看/审批/吊销/重发)→ Admin + if path.starts_with("/admin/") { + return Some(Role::Admin); + } + // Node 运行态 → Node + if path == "/node/heartbeat" && method == Method::POST { + return Some(Role::Node); + } + if path == "/task/claim" && method == Method::POST { + return Some(Role::Node); + } + if path == "/task/report" && method == Method::POST { + return Some(Role::Node); + } + if path.starts_with("/seed/") && method == Method::GET { + return Some(Role::Node); + } + if (path.starts_with("/data/file/") || path == "/data/linelist") && method == Method::GET { + return Some(Role::Node); + } + None } -/// Axum 鉴权中间件:若 AppState 中配置了 auth_token 则强制校验 Bearer Token 或 X-API-Key -pub async fn auth_middleware( - State(state): State, - req: Request, - next: Next, -) -> impl IntoResponse { - if let Some(ref expected_token) = state.auth_token { - let auth_header = req - .headers() - .get(header::AUTHORIZATION) - .and_then(|v| v.to_str().ok()); - let api_key_header = req - .headers() - .get("x-api-key") - .and_then(|v| v.to_str().ok()); +/// 恒定时间字符串比较。 +/// +/// 先对两个串各自做 SHA-256,再比较等长摘要(32 字节),彻底消除长度时序旁路—— +/// 任意长度的输入都产生相同长度的摘要,比较耗时固定,攻击者无法通过响应时间探得 token 长度。 +fn ct_eq_str(a: &str, b: &str) -> bool { + let ha = { + let mut h = Sha256::new(); + h.update(a.as_bytes()); + h.finalize() + }; + let hb = { + let mut h = Sha256::new(); + h.update(b.as_bytes()); + h.finalize() + }; + ha.ct_eq(&hb).into() +} - let token_valid = match (auth_header, api_key_header) { - (Some(auth), _) if auth.starts_with("Bearer ") => constant_time_eq(&auth[7..], expected_token), - (Some(auth), _) => constant_time_eq(auth, expected_token), - (_, Some(key)) => constant_time_eq(key, expected_token), - _ => false, - }; +/// node_id 白名单:字母、数字、点、下划线、连字符,长度 1-128。 +/// 用于 register_node / admin revoke / reissue 统一入口校验,与 Dashboard XSS 防护口径一致。 +pub(crate) fn is_valid_node_id(id: &str) -> bool { + !id.is_empty() + && id.len() <= 128 + && id + .chars() + .all(|c| c.is_ascii_alphanumeric() || c == '.' || c == '_' || c == '-') +} - if !token_valid { - return ( - StatusCode::UNAUTHORIZED, - "Unauthorized: Invalid or missing authentication token", - ) - .into_response(); +/// host_name 白名单:可打印 ASCII(排除控制字符),长度 1-128。 +/// 防止 host_name 携带 HTML/控制字符进入管理 Dashboard 触发存储型 XSS 或污染显示。 +pub(crate) fn is_valid_host_name(name: &str) -> bool { + !name.is_empty() + && name.len() <= 128 + && name.chars().all(|c| c.is_ascii() && !c.is_ascii_control()) +} + +/// 从请求头提取凭据原文(支持 `Authorization: Bearer ` 与 `X-API-Key: `)。 +/// +/// 安全:非 `Bearer ` 前缀的 Authorization 一律视为无 token(不再回退为裸头值比较), +/// 避免 `Authorization: Basic ...` 之类的上游代理头被误送入 token 比对。 +fn extract_token(req: &Request) -> Option { + if let Some(auth) = req + .headers() + .get(header::AUTHORIZATION) + .and_then(|v| v.to_str().ok()) + { + if let Some(rest) = auth.strip_prefix("Bearer ") { + if !rest.is_empty() { + return Some(rest.to_string()); + } + } + // 非 Bearer 前缀或空值:不作为 token + } + if let Some(key) = req.headers().get("x-api-key").and_then(|v| v.to_str().ok()) { + if !key.is_empty() { + return Some(key.to_string()); + } + } + None +} + +/// Axum 鉴权中间件(L2)。 +/// +/// 流程: +/// 1. 应急关闭(auth_disabled)→ 直接放行。 +/// 2. 路径不在授权矩阵 → 视为未公开接口,拒绝(401)。 +/// 3. 按角色校验对应凭据: +/// - Admin: admin token 恒定时间比对。 +/// - Node: node 专属 token 经 DB 反查 node_id(token 只存 hash)。 +/// 4. Node 角色额外校验:请求声称的 node_id 须与 token 绑定的 node_id 一致 +/// (防 A 节点用 B 节点的 token 越权操作)。claim_task / data 下载无 node_id +/// 输入,则仅校验 token 有效即可。 +pub async fn auth_middleware( + State(state): State, + mut req: Request, + next: Next, +) -> impl IntoResponse { + // 应急关闭:本地调试专用,绕过全部校验 + if state.auth_disabled { + return next.run(req).await.into_response(); + } + + let path = req.uri().path().to_string(); + let method = req.method().clone(); + let role = match required_role(&path, &method) { + Some(Role::Public) => { + return next.run(req).await.into_response(); + } + Some(r) => r, + None => { + // 未在矩阵中的路径一律拒绝(默认拒绝原则) + return (StatusCode::UNAUTHORIZED, "Unauthorized").into_response(); + } + }; + + let token = match extract_token(&req) { + Some(t) => t, + None => { + return (StatusCode::UNAUTHORIZED, "Unauthorized: missing token").into_response(); + } + }; + + // 审计日志:仅记录写操作(POST/PUT/DELETE)的「谁、做了什么」,不记请求体(防泄露)。 + // 在校验通过后记录 subject;校验失败由 401 分支体现,不单独审计。 + use axum::http::Method; + let is_write = matches!(method, Method::POST | Method::PUT | Method::DELETE); + + match role { + Role::Public => unreachable!(), + Role::Admin => { + let mut valid = false; + if let Some(ref admin) = state.admin_token { + if ct_eq_str(&token, admin) { + valid = true; + } + } + if !valid { + let mut sessions = state.admin_sessions.write().await; + let now = std::time::Instant::now(); + sessions.retain(|_, expiry| *expiry > now); + if sessions.contains_key(&token) { + valid = true; + } + } + + if valid { + if is_write { + tracing::info!(target: "dcts_audit", "AUDIT subject=admin method={} path={}", method, path); + } + return next.run(req).await.into_response(); + } + ( + StatusCode::UNAUTHORIZED, + "Unauthorized: invalid admin token", + ) + .into_response() + } + Role::Node => { + // 用 token 反查所属 node_id(DB 只存 hash,明文不落库) + match state.db.find_node_by_token(&token).await { + Some(token_node_id) => { + // 仅对非例行高频请求(如任务结果上报 /task/report)记录 AUDIT 审计日志, + // 成功的例行心跳 (/node/heartbeat) 与空闲领任务 (/task/claim) 静默跳过。 + if is_write && path != "/node/heartbeat" && path != "/task/claim" { + tracing::info!(target: "dcts_audit", "AUDIT subject=node:{} method={} path={}", token_node_id, method, path); + } + req.extensions_mut().insert(AuthenticatedNode { + node_id: token_node_id, + }); + next.run(req).await.into_response() + } + None => ( + StatusCode::UNAUTHORIZED, + "Unauthorized: invalid or revoked node token", + ) + .into_response(), + } } } - - next.run(req).await.into_response() } diff --git a/crates/server/src/api/node.rs b/crates/server/src/api/node.rs index f40d524..d175e01 100644 --- a/crates/server/src/api/node.rs +++ b/crates/server/src/api/node.rs @@ -1,25 +1,152 @@ -use super::AppState; -use axum::{extract::State, response::IntoResponse, Json}; +use super::{is_valid_host_name, is_valid_node_id, AppState, AuthenticatedNode}; +use axum::{ + extract::{Extension, State}, + http::StatusCode, + response::IntoResponse, + Json, +}; use common::models::{NodeHeartbeatRequest, NodeRegisterRequest}; use serde_json::json; +use tracing::{info, warn}; pub async fn register_node( State(state): State, + auth_node: Option>, Json(req): Json, -) -> impl IntoResponse { +) -> Result { + // 入口白名单校验 + if !is_valid_node_id(&req.node_id) { + return Err(crate::api::AppError::BadRequest( + "非法的节点 ID(仅允许字母、数字、点、下划线、连字符,长度 1-128)".to_string(), + )); + } + if !is_valid_host_name(&req.host_name) { + return Err(crate::api::AppError::BadRequest( + "非法的主机名(仅允许可打印 ASCII,长度 1-128)".to_string(), + )); + } + + // 已认证已拿到 Token 的节点刷新元数据配置 + if let Some(Extension(ref auth)) = auth_node { + if auth.node_id == req.node_id { + let _ = state.db.register_node(&req).await; + info!("已授权节点 {} 刷新配置成功", req.node_id); + return Ok(( + StatusCode::OK, + Json(json!({ + "status": "approved", + "message": "节点配置更新成功", + "node_token": null, + })), + )); + } + } + + // 申请注册新节点(免凭据提交申请,进入 pending_approval 状态) match state.db.register_node(&req).await { - Ok(_) => Json(json!({"status": "ok", "message": "节点注册成功"})), - Err(e) => Json(json!({"status": "error", "message": e.to_string()})), + Ok(true) => { + info!( + "接收到新节点 {} 的注册申请,已加入待审批 (pending_approval) 队列", + req.node_id + ); + Ok(( + StatusCode::OK, + Json(json!({ + "status": "pending_approval", + "message": "节点注册申请已成功提交!请在管理 Dashboard 控制台上点击【同意接入】授权该节点", + "node_token": null, + })), + )) + } + Ok(false) => { + // 节点已处于待审批或已存在列表 + Ok(( + StatusCode::OK, + Json(json!({ + "status": "pending_approval", + "message": "节点注册申请等待管理员审批中", + "node_token": null, + })), + )) + } + Err(e) => Err(e.into()), } } +#[derive(serde::Deserialize)] +pub struct CheckNodeStatusRequest { + pub node_id: String, +} + +/// POST /api/node/check_status — Node 端轮询检查审批结果。 +pub async fn check_node_status( + State(state): State, + Json(req): Json, +) -> Result { + if !is_valid_node_id(&req.node_id) { + return Err(crate::api::AppError::BadRequest( + "非法的节点 ID 参数".to_string(), + )); + } + + // 尝试拉取取走即焚的暂存明文 Token + match state.db.take_pending_node_token(&req.node_id).await { + Ok(Some(raw_token)) => { + info!( + "节点 {} 的注册申请已被管理员审批同意,下发专属 Token", + req.node_id + ); + Ok(( + StatusCode::OK, + Json(json!({ + "status": "approved", + "message": "节点已通过审批授权", + "node_token": raw_token, + })), + )) + } + Ok(None) | Err(_) => { + // 查节点表状态 + match state.db.get_node_exists(&req.node_id).await { + Ok(true) => Ok(( + StatusCode::OK, + Json(json!({ + "status": "pending_approval", + "message": "等待管理员在控制台点击同意", + "node_token": null, + })), + )), + _ => Ok(( + StatusCode::OK, + Json(json!({ + "status": "rejected", + "message": "节点注册申请未通过或已被移除", + "node_token": null, + })), + )), + } + } + } +} pub async fn heartbeat_node( State(state): State, + Extension(auth_node): Extension, Json(req): Json, -) -> impl IntoResponse { +) -> Result { + // 身份绑定校验:请求体声称的 node_id 必须与 token 绑定的 node_id 一致, + // 杜绝「持有 A 节点 token 却冒充 B 节点发心跳」的跨节点越权。 + if req.node_id != auth_node.node_id { + warn!( + "节点心跳身份校验失败:token 绑定 node={},但请求体声称 node_id={}", + auth_node.node_id, req.node_id + ); + return Err(crate::api::AppError::Forbidden( + "node_id 与凭据不匹配".to_string(), + )); + } match state.db.heartbeat_node(&req).await { - Ok(_) => Json(json!({"status": "ok"})), - Err(e) => Json(json!({"status": "error", "message": e.to_string()})), + Ok(_) => Ok(Json(json!({"status": "ok"}))), + Err(e) => Err(e.into()), } } diff --git a/crates/server/src/api/rate_limit.rs b/crates/server/src/api/rate_limit.rs new file mode 100644 index 0000000..8c29b99 --- /dev/null +++ b/crates/server/src/api/rate_limit.rs @@ -0,0 +1,196 @@ +//! 鉴权失败速率限制中间件(防 token 在线暴力)。 +//! +//! 设计:对返回 401 的请求按客户端 IP 维护滑动窗口失败计数。当某 IP 在窗口内 +//! 累计失败超过阈值,后续请求直接返回 429(持续到窗口内计数回落)。 +//! +//! 仅作用于鉴权路径(与 auth_middleware 叠加),不影响已认证的正常业务流。 +//! 已认证请求返回 2xx,不计入失败窗口,因此合法节点/管理员的高频调用不受影响。 +//! +//! IP 来源:优先取 `X-Forwarded-For` 首段(反代场景),回退到连接的 `ConnectInfo` +//! (需 main.rs 用 `into_make_service_with_connect_info` 启动)。两者都拿不到时按"未知 IP"聚合。 + +use axum::{ + extract::{ConnectInfo, State}, + http::Request, + middleware::Next, + response::{IntoResponse, Response}, +}; +use std::collections::{HashMap, VecDeque}; +use std::net::{IpAddr, SocketAddr}; +use std::sync::{Arc, Mutex}; +use std::time::{Duration, Instant}; +use tracing::warn; + +/// 限流状态:按 IP 维护近窗口内的失败时间戳队列。 +#[derive(Clone)] +pub struct RateLimiter { + inner: Arc>>>, + window: Duration, + max_failures: usize, + /// 计数策略: + /// - `false`(默认,通用 API 限流器):仅对鉴权失败(400/401/403)的响应计数。 + /// - `true`(注册端点专用限流器):对匹配路径(如 `/node/register`)的**所有**响应计数, + /// 无论成败——这是对注册接口的独立节流设计,防止恶意频繁注册。 + /// + /// 历史问题:此前中间件对所有 `/node/register` 请求无条件计数,导致该 limiter 若复用为 + /// 通用 API 限流器时,20 次成功注册会把整个 IP 锁出所有 `/api/*` 端点(跨端点连锁)。 + /// 引入此标志把两种语义显式分离。 + count_all: bool, +} + +impl RateLimiter { + /// 构造通用限流器:仅在鉴权失败(400/401/403)时计数。 + pub fn new(max_failures: usize, window: Duration) -> Self { + Self { + inner: Arc::new(Mutex::new(HashMap::new())), + window, + max_failures, + count_all: false, + } + } + + /// 构造「全量计数」限流器:对匹配路径的所有响应(无论成败)计数。 + /// 用于注册端点专用节流。 + pub fn new_count_all(max_failures: usize, window: Duration) -> Self { + Self { + inner: Arc::new(Mutex::new(HashMap::new())), + window, + max_failures, + count_all: true, + } + } + + /// 检查该 IP 是否已被限流(窗口内失败次数超阈值)。不修改计数。 + pub(crate) fn is_rate_limited(&self, ip: IpAddr) -> bool { + let now = Instant::now(); + let mut map = match self.inner.lock() { + Ok(g) => g, + Err(e) => e.into_inner(), // poisoned:仍尽力返回判断,避免鉴权因锁中毒全部放行 + }; + if let Some(queue) = map.get_mut(&ip) { + // 清理过期时间戳 + while let Some(front) = queue.front() { + if now.duration_since(*front) > self.window { + queue.pop_front(); + } else { + break; + } + } + if queue.is_empty() { + map.remove(&ip); + return false; + } + return queue.len() >= self.max_failures; + } + false + } + + /// 记录一次失败(追加时间戳)。 + pub(crate) fn record_failure(&self, ip: IpAddr) { + let now = Instant::now(); + let mut map = match self.inner.lock() { + Ok(g) => g, + Err(e) => e.into_inner(), + }; + let queue = map.entry(ip).or_default(); + queue.push_back(now); + // 顺带清理,防止队列无限增长 + while let Some(front) = queue.front() { + if now.duration_since(*front) > self.window { + queue.pop_front(); + } else { + break; + } + } + if queue.is_empty() { + map.remove(&ip); + } + } +} + +/// 判断 IP 是否为本地环回或私有网段 IP。 +fn is_private_or_loopback_ip(ip: IpAddr) -> bool { + match ip { + IpAddr::V4(v4) => v4.is_loopback() || v4.is_private(), + IpAddr::V6(v6) => v6.is_loopback(), + } +} + +/// 从请求中提取客户端 IP。 +/// 仅当底层连接 (ConnectInfo) 为本地环回或私有网段时才信任反向代理传递的 X-Forwarded-For / X-Real-IP。 +fn extract_client_ip(req: &Request) -> Option { + let direct_ip = req + .extensions() + .get::>() + .map(|ci| ci.0.ip()); + + // 如果有直连 IP 且不是私有/环回地址,说明未经过可信反代,直接返回直连 IP 拒绝盲信 X-Forwarded-For + if let Some(ip) = direct_ip { + if !is_private_or_loopback_ip(ip) { + return Some(ip); + } + } + + // 只有处于本地/私有网络反代之后时,才尝试提取 X-Forwarded-For + if let Some(xff) = req + .headers() + .get("x-forwarded-for") + .and_then(|v| v.to_str().ok()) + { + if let Some(first) = xff.split(',').map(|s| s.trim()).next() { + if !first.is_empty() { + if let Ok(ip) = first.parse::() { + return Some(ip); + } + } + } + } + // 回退:X-Real-IP + if let Some(xri) = req.headers().get("x-real-ip").and_then(|v| v.to_str().ok()) { + if let Ok(ip) = xri.parse::() { + return Some(ip); + } + } + // 回退:直连 IP + direct_ip +} + +/// 限流中间件:在鉴权之前检查该 IP 是否已被限流。 +/// +/// 放在 auth_middleware **之前**(外层):被限流的 IP 直接 429,不进鉴权逻辑。 +/// 是否记入失败窗口,由 auth_middleware 的结果决定——为此 auth 中间件会把 401 的 IP +/// 通过本 limiter 记录。但为避免跨中间件传参的复杂性,这里采用「先放行让 auth 判定, +/// 若返回 401 再记录」的方式:见下方包装函数 `rate_limit_with_auth`。 +pub async fn rate_limit_middleware( + State(limiter): State, + req: Request, + next: Next, +) -> Response { + let ip = extract_client_ip(&req).unwrap_or(IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED)); + let is_register = req.uri().path().ends_with("/node/register"); + + if limiter.is_rate_limited(ip) { + warn!("客户端 IP {} 鉴权失败次数过多,已限流(429)", ip); + return ( + axum::http::StatusCode::TOO_MANY_REQUESTS, + "鉴权失败次数过多,请稍后重试", + ) + .into_response(); + } + + let resp = next.run(req).await; + + // 计入速率窗口的条件: + // - 鉴权失败(401/403/400):通用与专用限流器都计; + // - 或 limiter 配置为 count_all 且请求落在专用节流路径(如 /node/register): + // 这种情况下成功响应也计,作为对注册接口本身的独立节流(防恶意频繁注册)。 + // 通用 API 限流器(count_all=false)不会因 is_register 把成功请求计入, + // 避免了「成功注册连锁锁出整个 /api/*」的历史缺陷。 + let status = resp.status().as_u16(); + let auth_failed = status == 401 || status == 403 || status == 400; + if auth_failed || (limiter.count_all && is_register) { + limiter.record_failure(ip); + } + + resp +} diff --git a/crates/server/src/api/seed.rs b/crates/server/src/api/seed.rs index 9f64faa..e1d5757 100644 --- a/crates/server/src/api/seed.rs +++ b/crates/server/src/api/seed.rs @@ -2,7 +2,7 @@ use super::AppState; use axum::{ body::Body, extract::{Path as AxumPath, State}, - http::{header, StatusCode}, + http::header, response::Response, }; use tokio::fs::File; @@ -12,13 +12,20 @@ use tracing::warn; pub async fn download_seed( State(state): State, AxumPath(name): AxumPath, -) -> Response { - if name.is_empty() || name.starts_with('.') || !name.chars().all(|c| c.is_ascii_alphanumeric() || c == '.' || c == '_' || c == '-' || c == '+' || c == '@') { - warn!("拒绝可能包含路径穿越或特别注入序列的非法种子下载请求: {}", name); - return Response::builder() - .status(StatusCode::BAD_REQUEST) - .body(Body::from("非法的种子名称参数")) - .unwrap(); +) -> Result { + if name.is_empty() + || name.starts_with('.') + || !name.chars().all(|c| { + c.is_ascii_alphanumeric() || c == '.' || c == '_' || c == '-' || c == '+' || c == '@' + }) + { + warn!( + "拒绝可能包含路径穿越或特别注入序列的非法种子下载请求: {}", + name + ); + return Err(crate::api::AppError::BadRequest( + "非法的种子名称参数".to_string(), + )); } let seed_file_path = std::path::Path::new(&state.results_dir) @@ -27,32 +34,27 @@ pub async fn download_seed( if !seed_file_path.is_file() { warn!("客户端请求的种子文件不存在: {}", seed_file_path.display()); - return Response::builder() - .status(StatusCode::NOT_FOUND) - .body(Body::from("请求的种子文件不存在")) - .unwrap(); + return Err(crate::api::AppError::NotFound( + "请求的种子文件不存在".to_string(), + )); } let file = match File::open(&seed_file_path).await { Ok(file) => file, - Err(_) => { - return Response::builder() - .status(StatusCode::INTERNAL_SERVER_ERROR) - .body(Body::from("无法打开种子文件")) - .unwrap(); + Err(e) => { + return Err(crate::api::AppError::Internal(e.into())); } }; - let stream = ReaderStream::new(file); let body = Body::from_stream(stream); - Response::builder() + Ok(Response::builder() .header(header::CONTENT_TYPE, "application/octet-stream") .header( header::CONTENT_DISPOSITION, format!("attachment; filename=\"{}.7\"", name), ) .body(body) - .unwrap() + .unwrap()) } diff --git a/crates/server/src/api/status.rs b/crates/server/src/api/status.rs index e5147a7..75754dd 100644 --- a/crates/server/src/api/status.rs +++ b/crates/server/src/api/status.rs @@ -2,21 +2,37 @@ use super::AppState; use axum::{extract::State, response::IntoResponse, Json}; use serde_json::json; -pub async fn get_status(State(state): State) -> impl IntoResponse { +/// 轻量健康检查端点(不走鉴权)。 +/// +/// 供 docker healthcheck、负载均衡、外部监控探测。刻意只返回固定 ok, +/// 不触碰数据库或调度器,避免健康检查本身拖累系统或因 DB 瞬时锁导致误判不健康。 +pub async fn healthz() -> Result { + Ok(Json(json!({ "status": "ok" }))) +} + +pub async fn get_status( + State(state): State, +) -> Result { let nodes = state.db.get_active_nodes().await.unwrap_or_default(); let total_active_slots: i32 = nodes.iter().map(|n| n.active_slots).sum(); let total_max_slots: i32 = nodes.iter().map(|n| n.max_slots).sum(); - let grid_stats = state.db.get_grid_summary_stats().await.unwrap_or(serde_json::json!({ - "total": 0, "pending": 0, "running": 0, "converged": 0, "failed": 0 - })); + // dashboard 全局概览:聚合全部工作流的 grid_points(多工作流分区后仍提供全局合计)。 + // 若需单工作流进度,可扩展为按 workflow 查询参数分别聚合。 + let grid_stats = state + .db + .get_grid_summary_stats(None) + .await + .unwrap_or(serde_json::json!({ + "total": 0, "pending": 0, "running": 0, "converged": 0, "failed": 0 + })); - Json(json!({ + Ok(Json(json!({ "status": "online", "nodes_online": nodes.len(), "total_active_slots": total_active_slots, "total_max_slots": total_max_slots, "nodes": nodes, "grid_stats": grid_stats, - })) + }))) } diff --git a/crates/server/src/api/task.rs b/crates/server/src/api/task.rs index 021ea8a..6899e06 100644 --- a/crates/server/src/api/task.rs +++ b/crates/server/src/api/task.rs @@ -1,6 +1,6 @@ -use super::AppState; +use super::{AppState, AuthenticatedNode}; use axum::{ - extract::{Multipart, State}, + extract::{Extension, Multipart, State}, response::IntoResponse, Json, }; @@ -12,27 +12,38 @@ use tracing::{info, warn}; use axum::http::StatusCode; -pub async fn claim_task(State(state): State) -> impl IntoResponse { - match state.queue.pop_task().await { +pub async fn claim_task( + State(state): State, + Extension(auth_node): Extension, +) -> Result { + // 领用时记录任务归属:pop_task 写入 claimed_by_node_id, + // report 阶段据此校验「上报者确为领用者」,杜绝跨节点伪造结果。 + match state.queue.pop_task(&auth_node.node_id).await { Ok(Some(task)) => { - if let Err(e) = state.db.mark_grid_point_running(&task.point_name).await { + // 多工作流分区:mark_grid_point_running 须带 workflow_name,避免按 name 全局更新 + // 误改其他工作流的同名点。TaskSpec.workflow_name 在调度时已绑定。 + let wf = task.workflow_name.as_deref().unwrap_or(""); + if let Err(e) = state.db.mark_grid_point_running(&task.point_name, wf).await { warn!("领用任务 {} 后同步变更为 running 状态遇到异常: {}. 后置 stale 定时自取检索引索将介入修复维护", task.task_id, e); } - (StatusCode::OK, Json(json!({"status": "ok", "task": task}))).into_response() + Ok((StatusCode::OK, Json(json!({"status": "ok", "task": task})))) + } + Ok(None) => Ok(( + StatusCode::OK, + Json(json!({"status": "empty", "task": null})), + )), + Err(e) => { + tracing::error!("领用任务数据库异常: {}", e); + Err(crate::api::AppError::Internal(e)) } - Ok(None) => (StatusCode::OK, Json(json!({"status": "empty", "task": null}))).into_response(), - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(json!({"status": "error", "message": format!("领用任务失败: {}", e)})), - ) - .into_response(), } } pub async fn report_task( State(state): State, + Extension(auth_node): Extension, mut multipart: Multipart, -) -> impl IntoResponse { +) -> Result { let mut report_json: Option = None; let mut seed_file_data: Option> = None; @@ -51,53 +62,98 @@ pub async fn report_task( } } - let report = match report_json { + let mut report = match report_json { Some(r) => r, None => { - return ( - StatusCode::BAD_REQUEST, - Json(json!({"status": "error", "message": "请求中缺少 report 字段"})), - ) - .into_response(); + return Err(crate::api::AppError::BadRequest( + "请求中缺少 report 字段".to_string(), + )); } }; + // ── 任务归属校验(S1 核心,防跨节点伪造结果投毒)── + // 1. 该 task_id 必须由当前鉴权 node 领用(claim 时记录的 claimed_by_node_id 匹配)。 + // 2. 上报的 point_name 必须与该 task 绑定的 point_name 一致(防跨点上报)。 + // 3. 忽略 body 里声称的 node_id,统一以鉴权 node_id 写库(修复审计归因断裂)。 + // 4. 取 task 绑定的 workflow_name,用于定向更新该工作流的 grid_points(多工作流分区)。 + let (claimed_point, claimed_workflow) = match state + .queue + .verify_task_claim(&report.task_id.to_string(), &auth_node.node_id) + .await + { + Ok(Some((p, w))) => (p, w), + Ok(None) => { + warn!( + "任务归属校验失败:node={} 上报 task_id={} 但未领用或已被清理", + auth_node.node_id, report.task_id + ); + return Err(crate::api::AppError::Forbidden( + "任务未由本节点领用或已上报过".to_string(), + )); + } + Err(e) => { + tracing::error!("校验任务归属数据库异常: {}", e); + return Err(crate::api::AppError::Internal(e)); + } + }; + if claimed_point != report.point_name { + warn!( + "任务点名校验失败:task_id={} 领用 point={} 但上报 point={}", + report.task_id, claimed_point, report.point_name + ); + return Err(crate::api::AppError::Forbidden( + "上报的网格点与领用任务不匹配".to_string(), + )); + } + // 统一以鉴权 node_id 覆盖 body 里的 node_id,保证归因可信 + report.node_id = auth_node.node_id.clone(); + // workflow_name 以领用记录为准(claim 时从 TaskSpec 落库),body 无权声称。 + let workflow_name = claimed_workflow.unwrap_or_default(); + let name = report.point_name.clone(); - if name.is_empty() || name.starts_with('.') || !name.chars().all(|c| c.is_ascii_alphanumeric() || c == '.' || c == '_' || c == '-' || c == '+' || c == '@') { - warn!("拒绝可能包含路径穿越或特殊非常规编码号攻击的网格点名称请求: {}", name); - return ( - StatusCode::BAD_REQUEST, - Json(json!({"status": "error", "message": "非法的网格点名称参数"})), - ) - .into_response(); + if name.is_empty() + || name.starts_with('.') + || !name.chars().all(|c| { + c.is_ascii_alphanumeric() || c == '.' || c == '_' || c == '-' || c == '+' || c == '@' + }) + { + warn!( + "拒绝可能包含路径穿越或特殊非常规编码号攻击的网格点名称请求: {}", + name + ); + return Err(crate::api::AppError::BadRequest( + "非法的网格点名称参数".to_string(), + )); } let params = match extract_params(&report) { Some(p) => p, None => { - warn!("网格点 {} 汇报数据解析失败: 无法解析 params 或 summary_json", name); - return ( - StatusCode::BAD_REQUEST, - Json(json!({"status": "error", "message": "无法解析 params 或 summary_json"})), - ) - .into_response(); + warn!( + "网格点 {} 汇报数据解析失败: 无法解析 params 或 summary_json", + name + ); + return Err(crate::api::AppError::BadRequest( + "无法解析 params 或 summary_json".to_string(), + )); } }; - // Record in DB - if let Err(e) = state.db.record_task_report(&report).await { - warn!("记录网格点 {} 任务结果到数据库失败: {}", name, e); - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(json!({"status": "error", "message": format!("记录数据库失败: {}", e)})), - ) - .into_response(); + // Record in DB(带 workflow_name 定向更新该工作流的 grid_points) + if let Err(e) = state.db.record_task_report(&report, &workflow_name).await { + // DB 错误细节进日志,对客户端只返回通用消息(避免泄露表结构/内部错误给未授权方) + tracing::error!("记录网格点 {} 任务结果到数据库失败: {}", name, e); + return Err(crate::api::AppError::Internal(e)); } // Clean up task from task_queue table to prevent queue DB bloat if let Err(e) = state.queue.remove_task(&report.task_id.to_string()).await { - tracing::warn!("从任务队列中清理已上报任务记录 {} 失败: {}", report.task_id, e); + tracing::warn!( + "从任务队列中清理已上报任务记录 {} 失败: {}", + report.task_id, + e + ); } // 采用原子写入模式保持 conv.json 与核心二进制数据完整落地后才揭晓真实文件名 @@ -112,30 +168,46 @@ pub async fn report_task( // Save seed file .7 using atomic temporary writing strategy if report.converged && !report.atmosphere_has_nan { if let Some(bytes) = seed_file_data { - let seed_tmp = model_dir.join(format!("{}.7.{}.tmp", name, uuid::Uuid::new_v4().simple())); + let seed_tmp = + model_dir.join(format!("{}.7.{}.tmp", name, uuid::Uuid::new_v4().simple())); let seed_path = model_dir.join(format!("{}.7", name)); - if fs::write(&seed_tmp, bytes).await.is_ok() { - if fs::rename(&seed_tmp, &seed_path).await.is_ok() { - info!("成功保持原子写入落地并保存网格点 {} 的收敛种子文件: {}", name, seed_path.display()); - let _ = state - .db - .insert_seed(¶ms, &seed_path.to_string_lossy()) - .await; - } + if fs::write(&seed_tmp, bytes).await.is_ok() + && fs::rename(&seed_tmp, &seed_path).await.is_ok() + { + info!( + "成功保持原子写入落地并保存网格点 {} 的收敛种子文件: {}", + name, + seed_path.display() + ); + let _ = state + .db + .insert_seed(¶ms, &seed_path.to_string_lossy()) + .await; } } } } - if report.status == TaskStatus::Failed || report.status == TaskStatus::Timeout || report.atmosphere_has_nan { + if !report.converged + || report.atmosphere_has_nan + || report.status == TaskStatus::Failed + || report.status == TaskStatus::Timeout + { // Task did not succeed -> check if seed_step fallback should be triggered info!("网格点 {} 计算未成功完成,检查种子回退机制...", name); - if let Err(e) = state.scheduler.trigger_seed_step_fallback(¶ms).await { + if let Err(e) = state + .scheduler + .trigger_seed_step_fallback(¶ms, &workflow_name) + .await + { warn!("网格点 {} 触发种子回退机制失败: {}", name, e); } } - (StatusCode::OK, Json(json!({"status": "ok", "message": "上报成功"}))).into_response() + Ok(( + StatusCode::OK, + Json(json!({"status": "ok", "message": "上报成功"})), + )) } fn extract_params(report: &TaskReport) -> Option { @@ -146,4 +218,3 @@ fn extract_params(report: &TaskReport) -> Option { .ok() .map(|summary| summary.params) } - diff --git a/crates/server/src/api/workflow.rs b/crates/server/src/api/workflow.rs index fe0eebb..5c4e85a 100644 --- a/crates/server/src/api/workflow.rs +++ b/crates/server/src/api/workflow.rs @@ -23,162 +23,166 @@ pub struct ApiResponse { pub data: Option, } -pub async fn list_workflows(State(state): State) -> impl IntoResponse { +pub async fn list_workflows( + State(state): State, +) -> Result { match state.db.list_workflows().await { - Ok(list) => (StatusCode::OK, Json(ApiResponse { success: true, message: "成功获取工作流列表".to_string(), data: Some(list) })), - Err(e) => (StatusCode::INTERNAL_SERVER_ERROR, Json(ApiResponse { success: false, message: format!("获取工作流列表失败: {}", e), data: None })), + Ok(list) => Ok(( + StatusCode::OK, + Json(ApiResponse { + success: true, + message: "成功获取工作流列表".to_string(), + data: Some(list), + }), + )), + Err(e) => Err(e.into()), } } pub async fn get_workflow( State(state): State, AxumPath(name): AxumPath, -) -> impl IntoResponse { +) -> Result { match state.db.get_workflow(&name).await { - Ok(Some(item)) => (StatusCode::OK, Json(ApiResponse { success: true, message: "成功获取工作流详情".to_string(), data: Some(item) })), - Ok(None) => (StatusCode::NOT_FOUND, Json(ApiResponse { success: false, message: format!("工作流 '{}' 未找到", name), data: None })), - Err(e) => (StatusCode::INTERNAL_SERVER_ERROR, Json(ApiResponse { success: false, message: format!("获取工作流详情失败: {}", e), data: None })), + Ok(Some(item)) => Ok(( + StatusCode::OK, + Json(ApiResponse { + success: true, + message: "成功获取工作流详情".to_string(), + data: Some(item), + }), + )), + Ok(None) => Err(crate::api::AppError::NotFound(format!( + "工作流 '{}' 未找到", + name + ))), + Err(e) => Err(e.into()), } } +/// 工作流名称白名单:仅允许字母、数字、点、下划线、连字符,长度 1-64。 +/// 与 report_task/download_seed 的网格点名校验口径保持一致,从源头阻止 +/// 名称携带 HTML/JS 特殊字符进入 Dashboard 渲染(存储型 XSS 根因之一)。 +fn is_valid_workflow_name(name: &str) -> bool { + !name.is_empty() + && name.len() <= 64 + && name + .chars() + .all(|c| c.is_ascii_alphanumeric() || c == '.' || c == '_' || c == '-') +} + pub async fn save_workflow( State(state): State, Json(req): Json, -) -> impl IntoResponse { +) -> Result { + // 名称白名单校验(优先于 YAML 校验,拒绝携带特殊字符的名称) + if !is_valid_workflow_name(&req.name) { + return Err(crate::api::AppError::BadRequest( + "工作流名称仅允许字母、数字、点(.)、下划线(_)、连字符(-),长度 1-64".to_string(), + )); + } + // Validate YAML config string if let Err(e) = serde_yaml::from_str::(&req.config_yaml) { - return ( - StatusCode::BAD_REQUEST, - Json(ApiResponse::<()> { - success: false, - message: format!("无效的 YAML 配置: {}", e), - data: None, - }), - ); + return Err(crate::api::AppError::BadRequest(format!( + "无效的 YAML 配置: {}", + e + ))); } // 检查被编辑的工作流是否正处于激活运行中 if let Ok(Some(existing)) = state.db.get_workflow(&req.name).await { if existing.status == "running" || existing.status == "initializing" { - return ( - StatusCode::BAD_REQUEST, - Json(ApiResponse::<()> { - success: false, - message: format!("工作流 '{}' 正处在运行或初始加载流程中,严禁原地覆写参数重设至 IDLE;如待变更参数请先调 API 显式触发停止后再保存", req.name), - data: None, - }), - ); + return Err(crate::api::AppError::BadRequest( + format!("工作流 '{}' 正处在运行或初始加载流程中,严禁原地覆写参数重设至 IDLE;如待变更参数请先调 API 显式触发停止后再保存", req.name) + )); } } - match state.db.upsert_workflow(&req.name, req.description.as_deref(), &req.config_yaml, "idle").await { + match state + .db + .upsert_workflow( + &req.name, + req.description.as_deref(), + &req.config_yaml, + "idle", + ) + .await + { Ok(_) => { info!("成功注册/更新工作流配置: {}", req.name); - ( + Ok(( StatusCode::OK, Json(ApiResponse::<()> { success: true, message: format!("工作流 '{}' 保存成功", req.name), data: None, }), - ) + )) } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(ApiResponse::<()> { - success: false, - message: format!("保存工作流失败: {}", e), - data: None, - }), - ), + Err(e) => Err(e.into()), } } pub async fn delete_workflow( State(state): State, AxumPath(name): AxumPath, -) -> impl IntoResponse { +) -> Result { + // 拦截正在运行或初始加载中的工作流删除请求 + if let Ok(Some(existing)) = state.db.get_workflow(&name).await { + if existing.status == "running" || existing.status == "initializing" { + return Err(crate::api::AppError::BadRequest(format!( + "工作流 '{}' 当前处于 '{}' 状态,无法直接删除。请先显式暂停/停止该工作流。", + name, existing.status + ))); + } + } + + let _ = state.queue.clear_queue_by_workflow(&name).await; match state.db.delete_workflow(&name).await { - Ok(_) => ( + Ok(_) => Ok(( StatusCode::OK, Json(ApiResponse::<()> { success: true, message: format!("工作流 '{}' 已删除", name), data: None, }), - ), - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(ApiResponse::<()> { - success: false, - message: format!("删除工作流失败: {}", e), - data: None, - }), - ), + )), + Err(e) => Err(e.into()), } } pub async fn start_workflow( State(state): State, AxumPath(name): AxumPath, -) -> impl IntoResponse { +) -> Result { let item = match state.db.get_workflow(&name).await { Ok(Some(item)) => item, Ok(None) => { - return ( - StatusCode::NOT_FOUND, - Json(ApiResponse::<()> { - success: false, - message: format!("工作流 '{}' 未找到", name), - data: None, - }), - ) - } - Err(e) => { - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(ApiResponse::<()> { - success: false, - message: format!("获取工作流失败: {}", e), - data: None, - }), - ) + return Err(crate::api::AppError::NotFound(format!( + "工作流 '{}' 未找到", + name + ))) } + Err(e) => return Err(e.into()), }; if item.status == "running" || item.status == "initializing" { - return ( - StatusCode::BAD_REQUEST, - Json(ApiResponse::<()> { - success: false, - message: format!("工作流 '{}' 已处在初始建立状态中或者已处于运行状态,无需且不允许进行并行重置启动", name), - data: None, - }), - ); + return Err(crate::api::AppError::BadRequest(format!( + "工作流 '{}' 已处在初始建立状态中或者已处于运行状态,无需且不允许进行并行重置启动", + name + ))); } // 通过原子性抢占更新将状态切换为 initializing,拦截同名流上的多并发调用导致的双重加载破坏性竞态 match state.db.transition_workflow_to_initializing(&name).await { Ok(false) => { - return ( - StatusCode::CONFLICT, - Json(ApiResponse::<()> { - success: false, - message: format!("工作流 '{}' 初始化抢占挂起异常,表明已在另一会话上下文中顺利推入启动通道", name), - data: None, - }), - ); - } - Err(e) => { - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(ApiResponse::<()> { - success: false, - message: format!("原子化抢占和迁移工作流状态发生异常: {}", e), - data: None, - }), - ); + return Err(crate::api::AppError::Conflict(format!( + "工作流 '{}' 初始化抢占挂起异常,表明已在另一会话上下文中顺利推入启动通道", + name + ))); } + Err(e) => return Err(e.into()), Ok(true) => {} } @@ -186,70 +190,69 @@ pub async fn start_workflow( Ok(cfg) => cfg, Err(e) => { let _ = state.db.update_workflow_status(&name, "idle").await; - return ( - StatusCode::BAD_REQUEST, - Json(ApiResponse::<()> { - success: false, - message: format!("解析工作流 YAML 发生语法或参数解析异常: {}", e), - data: None, - }), - ); + return Err(crate::api::AppError::BadRequest(format!( + "解析工作流 YAML 发生语法或参数解析异常: {}", + e + ))); } }; - info!("成功占据独享启动权,开始启动工作流 '{}',系统进行 64/32 维深度平展开网格结构计算化推列并推送队列...", name); - match state.scheduler.initialize_grid(&grid_cfg).await { - Ok(_) => { - let _ = state.db.update_workflow_status(&name, "running").await; - let _ = state.scheduler.schedule_pending_tasks().await; - ( - StatusCode::OK, - Json(ApiResponse::<()> { - success: true, - message: format!("工作流 '{}' 建立与挂载成功并已接续排班", name), - data: None, - }), - ) + info!("成功占据独享启动权,开始异步启动工作流 '{}',系统将在后台进行 64/32 维深度平展开网格结构计算化推列并推送队列...", name); + + let bg_state = state.clone(); + let bg_name = name.clone(); + let bg_grid_cfg = grid_cfg; + + tokio::spawn(async move { + match bg_state + .scheduler + .initialize_grid(&bg_grid_cfg, &bg_name) + .await + { + Ok(_) => { + let _ = bg_state + .db + .update_workflow_status(&bg_name, "running") + .await; + let _ = bg_state.scheduler.schedule_pending_tasks().await; + } + Err(e) => { + tracing::warn!("工作流 {} 网格初始化中途失败,已回退为 idle;已写入的点保留,重新启动会幂等补齐: {}", bg_name, e); + let _ = bg_state.db.update_workflow_status(&bg_name, "idle").await; + } } - Err(e) => { - let _ = state.db.update_workflow_status(&name, "idle").await; - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(ApiResponse::<()> { - success: false, - message: format!("展开与挂载初始化任务点到系统队列失败: {}", e), - data: None, - }), - ) - } - } + }); + + Ok(( + StatusCode::OK, + Json(ApiResponse::<()> { + success: true, + message: format!("工作流 '{}' 已进入后台异步建立与挂载流程", name), + data: None, + }), + )) } pub async fn stop_workflow( State(state): State, AxumPath(name): AxumPath, -) -> impl IntoResponse { +) -> Result { match state.db.update_workflow_status(&name, "paused").await { Ok(_) => { - let _ = state.queue.clear_queue().await; - let _ = state.db.reset_queued_grid_points_to_pending().await; - ( + // 多工作流分区:清理与重置都限定在本工作流内,避免误伤其他并发运行的工作流。 + // - clear_queue_by_workflow:只删本工作流的排队任务。 + // - reset_queued_grid_points_to_pending(&name):只把本工作流的 queued 点打回 pending。 + let _ = state.queue.clear_queue_by_workflow(&name).await; + let _ = state.db.reset_queued_grid_points_to_pending(&name).await; + Ok(( StatusCode::OK, Json(ApiResponse::<()> { success: true, message: format!("工作流 '{}' 已暂停,排队任务已暂停调度", name), data: None, }), - ) + )) } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(ApiResponse::<()> { - success: false, - message: format!("暂停工作流失败: {}", e), - data: None, - }), - ), + Err(e) => Err(e.into()), } } - diff --git a/crates/server/src/cors.rs b/crates/server/src/cors.rs new file mode 100644 index 0000000..5445291 --- /dev/null +++ b/crates/server/src/cors.rs @@ -0,0 +1,58 @@ +use axum::http::HeaderValue; +use tower_http::cors::{AllowOrigin, CorsLayer}; +use tracing::info; + +/// 构建 CORS 中间件层。 +/// +/// 严格安全策略:仅允许**同源**(Origin 匹配请求头的 Host)或**本地 Origin**(localhost / 127.0.0.1 / [::1])。 +pub fn build_cors_layer() -> CorsLayer { + info!("CORS 策略:仅允许同源或本地 Origin(localhost / 127.0.0.1 / [::1])"); + + CorsLayer::new() + .allow_origin(AllowOrigin::predicate( + |origin: &HeaderValue, head: &axum::http::request::Parts| { + let Ok(origin_str) = origin.to_str() else { + return false; + }; + let Ok(uri) = origin_str.parse::() else { + return false; + }; + let Some(host) = uri.host() else { + return false; + }; + let clean_host = host.trim_start_matches('[').trim_end_matches(']'); + + // 1. 本地来源 (localhost / 127.0.0.1 / [::1]) + if clean_host == "localhost" + || clean_host == "127.0.0.1" + || clean_host == "::1" + || clean_host.starts_with("127.") + { + return true; + } + + // 2. 同源来源 (Origin 匹配请求头的 Host) + if let Some(host_header) = head.headers.get(axum::http::header::HOST) { + if let Ok(host_str) = host_header.to_str() { + if let Some(authority) = uri.authority() { + if authority.as_str().eq_ignore_ascii_case(host_str) { + return true; + } + } + } + } + + false + }, + )) + .allow_methods([ + axum::http::Method::GET, + axum::http::Method::POST, + axum::http::Method::PUT, + axum::http::Method::DELETE, + ]) + .allow_headers([ + axum::http::header::AUTHORIZATION, + axum::http::header::CONTENT_TYPE, + ]) +} diff --git a/crates/server/src/db.rs b/crates/server/src/db.rs index 6350831..6c34279 100644 --- a/crates/server/src/db.rs +++ b/crates/server/src/db.rs @@ -1,19 +1,122 @@ use anyhow::{Context, Result}; use common::models::{ - GridPointParams, GridPointStatus, NodeHeartbeatRequest, NodeInfo, - NodeRegisterRequest, TaskReport, TaskStatus, + GridPointParams, GridPointStatus, NodeHeartbeatRequest, NodeInfo, NodeRegisterRequest, + TaskReport, TaskStatus, }; use r2d2::Pool; use r2d2_sqlite::SqliteConnectionManager; use rusqlite::params; +use sha2::{Digest, Sha256}; use tracing::info; #[derive(Debug)] struct SqliteCustomizer; +/// 计算 token 的 SHA-256 hex hash。凭据表只存 hash,不存明文 token。 +fn hash_token(token: &str) -> String { + let mut hasher = Sha256::new(); + hasher.update(token.as_bytes()); + hex::encode(hasher.finalize()) +} + +/// 多工作流分区迁移:把旧版 grid_points 表(仅 name UNIQUE,无 workflow_name 列) +/// 重建为带 workflow_name 列、(workflow_name, name) 复合唯一的新结构。 +/// +/// 幂等:新库(CREATE TABLE 已含 workflow_name)经 PRAGMA 检测后直接跳过。 +/// 旧库重建步骤: +/// 1. 把旧表重命名为 grid_points_legacy; +/// 2. 重建 grid_points(新 schema,已由 CREATE TABLE IF NOT EXISTS 建好——这里需先 DROP 再建); +/// 3. 从 legacy 复制数据,workflow_name 回填 '__legacy__' 兜底; +/// 4. 删除 legacy 表。 +/// +/// 兜底标记 '__legacy__' 的意义:新工作流的查询恒带 WHERE workflow_name=<真实名>, +/// 不会命中 '__legacy__' 行;这些历史残留行既不干扰新调度,也保留下来供人工排查。 +fn migrate_grid_points_for_workflow_partition(conn: &mut rusqlite::Connection) -> Result<()> { + // 若已是新结构(含 workflow_name 列),无需迁移。 + if has_grid_points_column(conn, "workflow_name")? { + return Ok(()); + } + + tracing::info!("检测到旧版 grid_points 表(无 workflow_name 列),执行多工作流分区迁移..."); + + let tx = conn.transaction()?; + + // 兼容:若存在遗留的迁移中间表(上次迁移被中断),先清理。 + tx.execute("DROP TABLE IF EXISTS grid_points_legacy", [])?; + + // 旧表改名 → 重建新表(按最新 CREATE TABLE 形态)→ 回填数据 → 删旧表 + tx.execute("ALTER TABLE grid_points RENAME TO grid_points_legacy", [])?; + tx.execute( + "CREATE TABLE grid_points ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL, + workflow_name TEXT NOT NULL, + teff REAL NOT NULL, + logg REAL NOT NULL, + loghe REAL NOT NULL, + logc REAL NOT NULL, + logn REAL NOT NULL, + logo REAL NOT NULL, + cno_sum REAL NOT NULL, + wave INTEGER NOT NULL DEFAULT 0, + status TEXT NOT NULL DEFAULT 'pending', + attempt_count INTEGER NOT NULL DEFAULT 0, + success_method TEXT + )", + [], + )?; + // 历史数据 workflow_name 兜底为 '__legacy__',不污染新工作流查询。 + tx.execute( + "INSERT INTO grid_points (name, workflow_name, teff, logg, loghe, logc, logn, logo, cno_sum, wave, status, attempt_count, success_method) + SELECT name, '__legacy__', teff, logg, loghe, logc, logn, logo, cno_sum, wave, status, attempt_count, success_method + FROM grid_points_legacy", + [], + )?; + tx.execute("DROP TABLE grid_points_legacy", [])?; + + tx.commit()?; + + tracing::info!("grid_points 多工作流分区迁移完成,历史数据 workflow_name 标记为 '__legacy__'"); + Ok(()) +} + +/// 检测 grid_points 表是否已含指定列(基于 PRAGMA table_info,与 sqlite_queue.rs 的迁移惯用法一致)。 +fn has_grid_points_column(conn: &rusqlite::Connection, col: &str) -> Result { + let mut stmt = conn.prepare("PRAGMA table_info(grid_points)")?; + let rows = stmt.query_map([], |r| r.get::<_, String>(1))?; + for r in rows { + if r.map(|name| name == col).unwrap_or(false) { + return Ok(true); + } + } + Ok(false) +} + +/// 将 SQLite db 文件及其 WAL/SHM 侧车文件权限收紧为 0600(仅 owner 读写)。 +/// 文件不存在或设置失败时静默忽略(不阻断启动,仅作加固)。 +#[cfg(unix)] +fn restrict_db_file_perms(db_path: &str) { + use std::os::unix::fs::PermissionsExt; + let candidates = [ + std::path::PathBuf::from(db_path), + std::path::PathBuf::from(format!("{}-wal", db_path)), + std::path::PathBuf::from(format!("{}-shm", db_path)), + ]; + for p in candidates { + if let Ok(meta) = std::fs::metadata(&p) { + let mut perms = meta.permissions(); + perms.set_mode(0o600); + let _ = std::fs::set_permissions(&p, perms); + } + } +} + impl r2d2::CustomizeConnection for SqliteCustomizer { fn on_acquire(&self, conn: &mut rusqlite::Connection) -> Result<(), rusqlite::Error> { - conn.pragma_update(None, "busy_timeout", 5000)?; + // 多节点心跳/claim + 后台调度 + dashboard 查询并发时,5s 易触发 SQLITE_BUSY 直接 bail。 + // 调大到 15s 给重试足够窗口,配合 IMMEDIATE 事务退避。 + conn.pragma_update(None, "busy_timeout", 15000)?; + conn.pragma_update(None, "wal_autocheckpoint", 1000)?; Ok(()) } } @@ -25,12 +128,64 @@ pub struct SeedCacheItem { pub file_path: String, } +/// exact_family 种子索引的桶键。 +/// +/// exact_family 判定(seed_finder.rs):`d_teff < 5000 && d_logg < 0.01 && d_loghe < 0.01`。 +/// 把这三个轴量化到桶: +/// - teff 按 5000K 量化为整数(floor),查询时同时查 floor 和 floor+1 两个桶即可覆盖 +/// [floor*5000, (floor+2)*5000) 区间(跨度 10000K),足以容纳双侧 < 5000 的邻域; +/// - logg/loghe 按 0.01 精度量化(×100 四舍五入为整数),相同量化值即满足 d < 0.01。 +/// +/// 命中桶后仍在桶内做精确 distance 计算取最优,故量化只用于缩小候选集,不影响正确性。 +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +struct SeedBucketKey { + teff_bucket: i64, + logg_q: i64, + loghe_q: i64, +} + +impl SeedBucketKey { + fn from_params(params: &GridPointParams) -> [Self; 2] { + // 返回该参数应落入的两个桶(floor 与 floor+1),供插入时双写、查询时双查。 + // teff 量化:floor(teff/5000)。如 teff=35000 → 7;teff=37499 → 7;teff=37500 → 8。 + let teff_floor = (params.teff / 5000.0).floor() as i64; + let logg_q = (params.logg * 100.0).round() as i64; + let loghe_q = (params.loghe * 100.0).round() as i64; + [ + SeedBucketKey { + teff_bucket: teff_floor, + logg_q, + loghe_q, + }, + SeedBucketKey { + teff_bucket: teff_floor + 1, + logg_q, + loghe_q, + }, + ] + } +} + #[derive(Clone)] pub struct Database { pool: Pool, seed_cache: std::sync::Arc>>, + /// exact_family 快速索引:桶键 → 该桶全部种子。命中 exact_family 的查询走 O(1)~O(小), + /// 未命中才退化到 seed_cache 全量 global 扫描。 + seed_index: std::sync::Arc< + tokio::sync::RwLock>>, + >, + /// node token 反查缓存:token_hash → (node_id, 插入时间)。 + /// 鉴权中间件每个 Node 请求都查 find_node_by_token,此缓存把高频心跳/领用请求 + /// 的 DB 查询降为内存读。TTL 由 `TOKEN_CACHE_TTL` 控制;issue/revoke 时整体失效。 + token_cache: std::sync::Arc< + tokio::sync::RwLock>, + >, } +/// token 反查缓存的单条存活时长(秒)。issue/revoke 会立即整体失效,TTL 仅兜底。 +const TOKEN_CACHE_TTL: std::time::Duration = std::time::Duration::from_secs(60); + impl Database { pub async fn new(db_path: &str) -> Result { let db_path_owned = db_path.to_string(); @@ -40,11 +195,16 @@ impl Database { } let manager = SqliteConnectionManager::file(&db_path_owned); let pool = Pool::builder() - .max_size(8) + .max_size(16) .connection_customizer(Box::new(SqliteCustomizer)) .build(manager) .context("Failed to build SQLite main DB connection pool")?; + // 收紧 db 文件权限为 0600(仅 owner 读写),防止裸机部署时其他用户读取节点/凭据信息。 + // 容器内以非 root 运行,此设置仅作加固;WAL/SHM 侧车文件一并处理。 + #[cfg(unix)] + restrict_db_file_perms(&db_path_owned); + Ok(pool) }) .await??; @@ -52,6 +212,12 @@ impl Database { let db = Self { pool, seed_cache: std::sync::Arc::new(tokio::sync::RwLock::new(Vec::new())), + seed_index: std::sync::Arc::new(tokio::sync::RwLock::new( + std::collections::HashMap::new(), + )), + token_cache: std::sync::Arc::new(tokio::sync::RwLock::new( + std::collections::HashMap::new(), + )), }; db.init_tables().await?; db.reload_seed_cache().await?; @@ -61,7 +227,7 @@ impl Database { async fn init_tables(&self) -> Result<()> { let pool = self.pool.clone(); tokio::task::spawn_blocking(move || -> Result<()> { - let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + let mut conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; let _: String = conn.pragma_update_and_check(None, "journal_mode", "WAL", |r| r.get(0))?; conn.execute( "CREATE TABLE IF NOT EXISTS nodes ( @@ -80,7 +246,8 @@ impl Database { conn.execute( "CREATE TABLE IF NOT EXISTS grid_points ( id INTEGER PRIMARY KEY AUTOINCREMENT, - name TEXT UNIQUE NOT NULL, + name TEXT NOT NULL, + workflow_name TEXT NOT NULL, teff REAL NOT NULL, logg REAL NOT NULL, loghe REAL NOT NULL, @@ -95,6 +262,12 @@ impl Database { );", [], )?; + // 多工作流分区迁移:旧库的 grid_points 表只有 name UNIQUE(无 workflow_name), + // 无法支撑「同一物理点属于多个工作流」。SQLite 不能原地删除 CREATE TABLE 内联的 UNIQUE + // 约束,故用 PRAGMA table_info 检测旧表形态:若缺 workflow_name 列,则重建表为 + // (workflow_name, name) 复合唯一。旧数据 workflow_name 回填为 '__legacy__' 兜底, + // 避免新工作流查询 WHERE workflow_name=? 误命中历史残留行。 + migrate_grid_points_for_workflow_partition(&mut conn)?; conn.execute( "CREATE TABLE IF NOT EXISTS tasks ( @@ -110,10 +283,18 @@ impl Database { created_at DATETIME NOT NULL, started_at DATETIME, completed_at DATETIME, - error_message TEXT + error_message TEXT, + workflow_name TEXT );", [], )?; + let has_tasks_wf = conn + .prepare("PRAGMA table_info(tasks)")? + .query_map([], |r| r.get::<_, String>(1))? + .any(|r| r.map(|n| n == "workflow_name").unwrap_or(false)); + if !has_tasks_wf { + let _ = conn.execute("ALTER TABLE tasks ADD COLUMN workflow_name TEXT", []); + } conn.execute( "CREATE TABLE IF NOT EXISTS seeds ( @@ -144,7 +325,13 @@ impl Database { )?; conn.execute( - "CREATE INDEX IF NOT EXISTS idx_grid_points_status_wave ON grid_points(status, wave, cno_sum, teff);", + // 多工作流分区后,调度查询恒带 WHERE workflow_name=?,故索引前置 workflow_name。 + "CREATE INDEX IF NOT EXISTS idx_grid_points_wf_status ON grid_points(workflow_name, status, wave, cno_sum, teff);", + [], + )?; + conn.execute( + // 复合唯一约束:(workflow_name, name) 唯一,支撑 ON CONFLICT(workflow_name, name)。 + "CREATE UNIQUE INDEX IF NOT EXISTS idx_grid_points_wf_name ON grid_points(workflow_name, name);", [], )?; @@ -153,6 +340,34 @@ impl Database { [], )?; + // Node 专属凭据表(L2 鉴权):存储每个 node 颁发的 token 的 SHA-256 hash(不存明文), + // 支持 per-node 吊销。明文 token 仅在注册时返回一次。 + conn.execute( + "CREATE TABLE IF NOT EXISTS node_credentials ( + node_id TEXT PRIMARY KEY, + token_hash TEXT NOT NULL, + issued_at DATETIME NOT NULL, + revoked INTEGER NOT NULL DEFAULT 0, + raw_token_pending TEXT + );", + [], + )?; + let has_raw_pending = conn + .prepare("PRAGMA table_info(node_credentials)")? + .query_map([], |r| r.get::<_, String>(1))? + .any(|r| r.map(|n| n == "raw_token_pending").unwrap_or(false)); + if !has_raw_pending { + let _ = conn.execute("ALTER TABLE node_credentials ADD COLUMN raw_token_pending TEXT", []); + } + + conn.execute( + "CREATE UNIQUE INDEX IF NOT EXISTS idx_node_credentials_token_hash ON node_credentials(token_hash);", + [], + )?; + + // 清理过期或处理过的条目 + let _ = conn.execute("UPDATE node_credentials SET raw_token_pending = NULL WHERE revoked = 1 OR datetime('now', '-1 day') >= issued_at", []); + info!("成功初始化 dcts.db 数据库结构表及索引"); Ok(()) }) @@ -162,21 +377,76 @@ impl Database { } // --- Node operations --- - pub async fn register_node(&self, req: &NodeRegisterRequest) -> Result<()> { + pub async fn register_node(&self, req: &NodeRegisterRequest) -> Result { let pool = self.pool.clone(); let req_cloned = req.clone(); + let is_new = tokio::task::spawn_blocking(move || -> Result { + let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + let mut stmt = conn.prepare("SELECT status FROM nodes WHERE node_id = ?1")?; + let existing_status: Option = stmt.query_row(params![req_cloned.node_id], |r| r.get(0)).ok(); + + match existing_status { + Some(_st) => { + // 已存在的节点:更新配置,保持既有状态 + conn.execute( + "UPDATE nodes SET host_name = ?1, max_slots = ?2, last_heartbeat = datetime('now') WHERE node_id = ?3", + params![req_cloned.host_name, req_cloned.max_slots, req_cloned.node_id], + )?; + Ok(false) + } + None => { + // 新申请节点:插入待审批状态 (pending_approval) + conn.execute( + "INSERT INTO nodes (node_id, host_name, max_slots, status, last_heartbeat) + VALUES (?1, ?2, ?3, 'pending_approval', datetime('now'))", + params![req_cloned.node_id, req_cloned.host_name, req_cloned.max_slots], + )?; + Ok(true) + } + } + }) + .await??; + + Ok(is_new) + } + + /// 管理员审批同意节点接入:将节点状态切为 online 并生成专属 node_token(返回明文 token)。 + pub async fn approve_node(&self, node_id: &str) -> Result { + let pool = self.pool.clone(); + let node_id_owned = node_id.to_string(); + tokio::task::spawn_blocking(move || -> Result<()> { let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; conn.execute( - "INSERT INTO nodes (node_id, host_name, max_slots, status, last_heartbeat) - VALUES (?1, ?2, ?3, 'online', datetime('now')) - ON CONFLICT(node_id) DO UPDATE SET - host_name = excluded.host_name, - max_slots = excluded.max_slots, - status = 'online', - last_heartbeat = datetime('now')", - params![req_cloned.node_id, req_cloned.host_name, req_cloned.max_slots], + "UPDATE nodes SET status = 'online', last_heartbeat = datetime('now') WHERE node_id = ?1", + params![node_id_owned], + )?; + Ok(()) + }) + .await??; + + // 颁发专属 node_token + let new_token = self.issue_node_token(node_id).await?; + Ok(new_token) + } + + /// 管理员拒绝节点接入:彻底清理该节点的注册申请记录。 + pub async fn reject_node(&self, node_id: &str) -> Result<()> { + let pool = self.pool.clone(); + let node_id_owned = node_id.to_string(); + + tokio::task::spawn_blocking(move || -> Result<()> { + let conn = pool + .get() + .map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + conn.execute( + "DELETE FROM nodes WHERE node_id = ?1", + params![node_id_owned], + )?; + conn.execute( + "DELETE FROM node_credentials WHERE node_id = ?1", + params![node_id_owned], )?; Ok(()) }) @@ -193,7 +463,7 @@ impl Database { let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; conn.execute( "UPDATE nodes SET active_slots = ?1, cpu_usage = ?2, memory_usage = ?3, status = 'online', last_heartbeat = datetime('now') - WHERE node_id = ?4", + WHERE node_id = ?4 AND status IN ('online', 'offline')", params![ req_cloned.active_slots, req_cloned.cpu_usage, @@ -208,6 +478,239 @@ impl Database { Ok(()) } + // --- Node credentials (L2 鉴权) --- + + /// 为指定 node 颁发专属 token:生成随机明文 token,DB 存其 SHA-256 hash。 + /// 返回明文 token(仅此一次,由调用方转交 node 持久化)。 + /// 若该 node 已有凭据则覆盖(重新颁发)。 + pub async fn issue_node_token(&self, node_id: &str) -> Result { + let pool = self.pool.clone(); + let node_id_owned = node_id.to_string(); + // 两个 v4 UUID(各 16 字节随机)拼接 → 各 32 hex 字符 = 64 字符 token + let token = + uuid::Uuid::new_v4().simple().to_string() + &uuid::Uuid::new_v4().simple().to_string(); + let token_hash = hash_token(&token); + let token_for_ret = token.clone(); + let token_to_db = token_for_ret.clone(); + + tokio::task::spawn_blocking(move || -> Result<()> { + let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + conn.execute( + "INSERT INTO node_credentials (node_id, token_hash, issued_at, revoked, raw_token_pending) + VALUES (?1, ?2, datetime('now'), 0, ?3) + ON CONFLICT(node_id) DO UPDATE SET + token_hash = excluded.token_hash, + issued_at = datetime('now'), + revoked = 0, + raw_token_pending = excluded.raw_token_pending", + params![node_id_owned, token_hash, token_to_db], + )?; + Ok(()) + }) + .await??; + + // token 轮换:旧 token_hash 已失效,新 token_hash 即将生效。整体清空缓存最稳妥 + // (issue 是低频运维动作,全清代价可忽略)。 + self.invalidate_token_cache().await; + + Ok(token_for_ret) + } + + /// 一次性拉取并清除暂存的明文 node_token(取走即焚安全策略)。 + /// + /// 在单个 IMMEDIATE 事务内:先 SELECT 读出明文,再 UPDATE 置 NULL。IMMEDIATE 事务在 + /// BEGIN 时即获取写锁,保证 SELECT 与 UPDATE 之间不会被其它调用方插入,从而只有一个 + /// 调用方能取到 token(原子语义)。 + /// + /// 注:SQLite 的 `UPDATE ... RETURNING` 返回的是列的**新值**(SET 之后),故清空后 + /// RETURNING 该列只会得到 NULL,无法用于读旧值;因此这里用显式 SELECT + UPDATE。 + pub async fn take_pending_node_token(&self, node_id: &str) -> Result> { + let pool = self.pool.clone(); + let node_id_owned = node_id.to_string(); + + let token = tokio::task::spawn_blocking(move || -> Result> { + let mut conn = pool + .get() + .map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + let tx = conn.transaction_with_behavior(rusqlite::TransactionBehavior::Immediate)?; + let raw_token: Option = { + let mut select_stmt = tx.prepare( + "SELECT raw_token_pending FROM node_credentials + WHERE node_id = ?1 AND revoked = 0 AND raw_token_pending IS NOT NULL", + )?; + select_stmt + .query_row(params![node_id_owned], |r| r.get::<_, String>(0)) + .ok() + }; + if raw_token.is_some() { + tx.execute( + "UPDATE node_credentials SET raw_token_pending = NULL + WHERE node_id = ?1 AND revoked = 0 AND raw_token_pending IS NOT NULL", + params![node_id_owned], + )?; + } + tx.commit()?; + Ok(raw_token) + }) + .await??; + + Ok(token) + } + + /// 按 token(明文)反查所属 node_id;仅当 token 有效且未被吊销时返回 Some。 + /// 用于中间件:请求带来 node token,由此确定调用方身份。 + /// + /// 高频路径(每个 Node 请求一次):先查内存 token_cache,命中且未过期直接返回; + /// miss 才落 DB,并回填缓存。issue/revoke 会主动清空整个缓存。 + pub async fn find_node_by_token(&self, token: &str) -> Option { + let token_hash = hash_token(token); + + // 1) 先查内存缓存 + { + let cache = self.token_cache.read().await; + if let Some((node_id, inserted)) = cache.get(&token_hash) { + if inserted.elapsed() < TOKEN_CACHE_TTL { + return Some(node_id.clone()); + } + } + } + + // 2) miss 落 DB + let pool = self.pool.clone(); + let hash_for_db = token_hash.clone(); + let node_id: Option = tokio::task::spawn_blocking(move || -> Result> { + let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + let mut stmt = conn.prepare( + "SELECT node_id FROM node_credentials WHERE token_hash = ?1 AND revoked = 0 LIMIT 1", + )?; + let res = stmt.query_row(params![hash_for_db], |r| r.get::<_, String>(0)); + match res { + Ok(id) => Ok(Some(id)), + Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None), + Err(e) => Err(e.into()), + } + }) + .await + .ok() + .and_then(|r| r.ok()) + .flatten(); + + // 3) 命中则回填缓存(None 不缓存,避免吊销态被短暂缓存) + if let Some(id) = &node_id { + let mut cache = self.token_cache.write().await; + cache.insert(token_hash, (id.clone(), std::time::Instant::now())); + } + node_id + } + + /// 清空全部 token 反查缓存。在 issue(token 轮换使旧 token 失效)与 revoke 时调用。 + async fn invalidate_token_cache(&self) { + self.token_cache.write().await.clear(); + } + + /// 吊销指定 node 的凭据(置 revoked=1)。吊销后该 node 的 token 立即失效, + /// 需重新走注册流程获取新 token。不影响 nodes 表中的节点记录本身。 + pub async fn revoke_node_token(&self, node_id: &str) -> Result<()> { + let pool = self.pool.clone(); + let node_id_owned = node_id.to_string(); + tokio::task::spawn_blocking(move || -> Result<()> { + let conn = pool + .get() + .map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + conn.execute( + "UPDATE node_credentials SET revoked = 1 WHERE node_id = ?1", + params![node_id_owned], + )?; + Ok(()) + }) + .await??; + // 吊销后旧 token 必须立即失效,整体清缓存确保不残留可用映射。 + self.invalidate_token_cache().await; + Ok(()) + } + + /// 判断指定 node_id 是否已存在于 nodes 表(重发 token 前置校验,防幽灵 node_id)。 + pub async fn get_node_exists(&self, node_id: &str) -> Result { + let pool = self.pool.clone(); + let node_id_owned = node_id.to_string(); + let exists = tokio::task::spawn_blocking(move || -> Result { + let conn = pool + .get() + .map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + let mut stmt = conn.prepare("SELECT 1 FROM nodes WHERE node_id = ?1 LIMIT 1")?; + Ok(stmt.exists(params![node_id_owned])?) + }) + .await??; + Ok(exists) + } + + /// 统计已颁发且未吊销的 node 凭据数量(用于启动期半配置告警判断)。 + pub async fn node_credentials_count(&self) -> Result { + let pool = self.pool.clone(); + let count = tokio::task::spawn_blocking(move || -> Result { + let conn = pool + .get() + .map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + let n: i64 = conn.query_row( + "SELECT COUNT(*) FROM node_credentials WHERE revoked = 0", + [], + |r| r.get(0), + )?; + Ok(n) + }) + .await??; + Ok(count) + } + + /// 列出全部节点及其凭据状态(LEFT JOIN node_credentials)。 + /// 用于管理 API:admin 可查看每个节点的在线状态、是否已颁发 token、是否被吊销、颁发时间。 + /// 尚未注册凭据的节点(如旧数据迁移)token_revoked/token_issued_at 为 None。 + pub async fn list_nodes_with_credentials(&self) -> Result> { + let pool = self.pool.clone(); + tokio::task::spawn_blocking(move || -> Result> { + let conn = pool + .get() + .map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + let mut stmt = conn.prepare( + "SELECT n.node_id, n.host_name, n.max_slots, n.active_slots, n.status, + n.cpu_usage, n.memory_usage, + strftime('%Y-%m-%dT%H:%M:%SZ', n.last_heartbeat), + c.revoked, strftime('%Y-%m-%dT%H:%M:%SZ', c.issued_at) + FROM nodes n + LEFT JOIN node_credentials c ON c.node_id = n.node_id + ORDER BY n.status ASC, n.node_id ASC", + )?; + let rows = stmt.query_map([], |r| { + let hb_str: String = r.get::<_, String>(7)?; + Ok(NodeCredentialView { + node_id: r.get(0)?, + host_name: r.get(1)?, + max_slots: r.get(2)?, + active_slots: r.get(3)?, + status: r.get(4)?, + cpu_usage: r.get(5)?, + memory_usage: r.get(6)?, + last_heartbeat: chrono::DateTime::parse_from_rfc3339(&hb_str) + .map(|d| d.with_timezone(&chrono::Utc)) + .unwrap_or_else(|_| chrono::DateTime::UNIX_EPOCH), + // c.revoked 为 NULL 表示该节点无凭据记录;0=有效,1=已吊销 + token_status: match r.get::<_, Option>(8)? { + None => "none".to_string(), + Some(0) => "active".to_string(), + Some(_) => "revoked".to_string(), + }, + token_issued_at: r.get::<_, Option>(9)?, + }) + })?; + let mut list = Vec::new(); + for row in rows { + list.push(row?); + } + Ok(list) + }) + .await? + } + pub async fn get_active_nodes(&self) -> Result> { let pool = self.pool.clone(); @@ -243,19 +746,25 @@ impl Database { } // --- Grid Point & Task operations --- - pub async fn upsert_grid_point(&self, params_in: &GridPointParams, wave: i32) -> Result<()> { + pub async fn upsert_grid_point( + &self, + params_in: &GridPointParams, + wave: i32, + workflow_name: &str, + ) -> Result<()> { let pool = self.pool.clone(); let p = params_in.clone(); let name = p.model_name(); let cno_sum = p.cno_sum(); + let wf = workflow_name.to_string(); tokio::task::spawn_blocking(move || -> Result<()> { let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; conn.execute( - "INSERT INTO grid_points (name, teff, logg, loghe, logc, logn, logo, cno_sum, wave) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9) - ON CONFLICT(name) DO NOTHING", - params![name, p.teff, p.logg, p.loghe, p.logc, p.logn, p.logo, cno_sum, wave], + "INSERT INTO grid_points (name, workflow_name, teff, logg, loghe, logc, logn, logo, cno_sum, wave) + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10) + ON CONFLICT(workflow_name, name) DO NOTHING", + params![name, wf, p.teff, p.logg, p.loghe, p.logc, p.logn, p.logo, cno_sum, wave], )?; Ok(()) }) @@ -264,21 +773,39 @@ impl Database { Ok(()) } - pub async fn get_pending_grid_points(&self) -> Result> { - self.get_pending_grid_points_limit(usize::MAX).await + /// 仅用于测试/诊断:取指定工作流的全部 pending 点(无 LIMIT)。生产调度走 _limit 版本。 + pub async fn get_pending_grid_points( + &self, + workflow_name: &str, + ) -> Result> { + self.get_pending_grid_points_limit(usize::MAX, workflow_name) + .await } - pub async fn get_pending_grid_points_limit(&self, limit: usize) -> Result> { + pub async fn get_pending_grid_points_limit( + &self, + limit: usize, + workflow_name: &str, + ) -> Result> { let pool = self.pool.clone(); + let wf = workflow_name.to_string(); tokio::task::spawn_blocking(move || -> Result> { - let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + let conn = pool + .get() + .map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; let mut stmt = conn.prepare( - "SELECT name, teff, logg, loghe, logc, logn, logo, wave FROM grid_points WHERE status = 'pending' ORDER BY wave ASC, cno_sum ASC, teff ASC LIMIT ?1" + "SELECT name, teff, logg, loghe, logc, logn, logo, wave FROM grid_points + WHERE status = 'pending' AND workflow_name = ?1 + ORDER BY wave ASC, cno_sum ASC, teff ASC LIMIT ?2", )?; - let limit_param = if limit == usize::MAX { -1i64 } else { limit as i64 }; - let rows_iter = stmt.query_map([limit_param], |r| { + let limit_param = if limit == usize::MAX { + -1i64 + } else { + limit as i64 + }; + let rows_iter = stmt.query_map(params![wf, limit_param], |r| { Ok(( r.get(0)?, GridPointParams { @@ -302,34 +829,44 @@ impl Database { .await? } - pub async fn reset_queued_grid_points_to_pending(&self) -> Result { + /// 重置指定工作流的 queued 点为 pending(系统重启/工作流启动时使用)。 + /// 按 workflow 隔离,避免误伤其他工作流(多工作流分区修复点)。 + pub async fn reset_queued_grid_points_to_pending(&self, workflow_name: &str) -> Result { let pool = self.pool.clone(); + let wf = workflow_name.to_string(); tokio::task::spawn_blocking(move || -> Result { let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; // 只重置 queued 状态的任务为 pending。对于 running (正由 Worker 处理的项目),不可在系统重启或初始化时粗暴清零,让 Worker 正常完成汇报或触发心跳/超时自动逐回 let rows = conn.execute( - "UPDATE grid_points SET status = 'pending' WHERE status = 'queued'", - [], + "UPDATE grid_points SET status = 'pending' WHERE status = 'queued' AND workflow_name = ?1", + params![wf], )?; Ok(rows) }) .await? } - pub async fn reset_specific_grid_points_to_pending(&self, names: &[String]) -> Result { + /// 重置指定工作流内一批点(按 name)为 pending。 + /// 按 workflow 隔离,避免跨工作流误改同名点(多工作流分区修复点)。 + pub async fn reset_specific_grid_points_to_pending( + &self, + names: &[String], + workflow_name: &str, + ) -> Result { if names.is_empty() { return Ok(0); } let pool = self.pool.clone(); let names_owned = names.to_vec(); + let wf = workflow_name.to_string(); tokio::task::spawn_blocking(move || -> Result { let mut conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; let tx = conn.transaction()?; let mut count = 0; for name in &names_owned { count += tx.execute( - "UPDATE grid_points SET status = 'pending' WHERE name = ?1 AND status IN ('queued', 'running')", - params![name], + "UPDATE grid_points SET status = 'pending' WHERE name = ?1 AND workflow_name = ?2 AND status IN ('queued', 'running')", + params![name, wf], )?; } tx.commit()?; @@ -338,16 +875,25 @@ impl Database { .await? } - pub async fn update_grid_status(&self, name: &str, status: GridPointStatus) -> Result<()> { + /// 更新指定工作流内某点的状态。按 workflow 隔离,防跨工作流误改同名点。 + pub async fn update_grid_status( + &self, + name: &str, + status: GridPointStatus, + workflow_name: &str, + ) -> Result<()> { let pool = self.pool.clone(); let name_owned = name.to_string(); let status_str = status.to_string(); + let wf = workflow_name.to_string(); tokio::task::spawn_blocking(move || -> Result<()> { - let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + let conn = pool + .get() + .map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; conn.execute( - "UPDATE grid_points SET status = ?1 WHERE name = ?2", - params![status_str, name_owned], + "UPDATE grid_points SET status = ?1 WHERE name = ?2 AND workflow_name = ?3", + params![status_str, name_owned, wf], )?; Ok(()) }) @@ -356,17 +902,23 @@ impl Database { Ok(()) } - pub async fn mark_grid_point_running(&self, name: &str) -> Result<()> { - self.update_grid_status(name, GridPointStatus::Running).await + pub async fn mark_grid_point_running(&self, name: &str, workflow_name: &str) -> Result<()> { + self.update_grid_status(name, GridPointStatus::Running, workflow_name) + .await } - pub async fn get_grid_point_status(&self, name: &str) -> Result> { + pub async fn get_grid_point_status( + &self, + name: &str, + workflow_name: &str, + ) -> Result> { let pool = self.pool.clone(); let name_owned = name.to_string(); + let wf = workflow_name.to_string(); tokio::task::spawn_blocking(move || -> Result> { let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; - let mut stmt = conn.prepare("SELECT status, attempt_count FROM grid_points WHERE name = ?1")?; - let res = stmt.query_row(params![name_owned], |r| Ok((r.get(0)?, r.get(1)?))); + let mut stmt = conn.prepare("SELECT status, attempt_count FROM grid_points WHERE name = ?1 AND workflow_name = ?2")?; + let res = stmt.query_row(params![name_owned, wf], |r| Ok((r.get(0)?, r.get(1)?))); match res { Ok(tuple) => Ok(Some(tuple)), Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None), @@ -376,6 +928,26 @@ impl Database { .await? } + /// 判断某个网格点在特定工作流内是否已经派发过种子步进 (seed_step) 任务。 + /// + /// 用于“种子回退仅一次”语义:一旦该点在该工作流中已经存在过 task_type='seed_step' 的任务记录, + /// 再次失败就不再触发新的种子回退,直接保持 failed 终态。 + pub async fn has_seed_step_attempt(&self, name: &str, workflow_name: &str) -> Result { + let pool = self.pool.clone(); + let name_owned = name.to_string(); + let wf_owned = workflow_name.to_string(); + let exists = tokio::task::spawn_blocking(move || -> Result { + let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + let mut stmt = conn.prepare( + "SELECT 1 FROM tasks WHERE point_name = ?1 AND task_type = 'seed_step' AND (workflow_name = ?2 OR workflow_name IS NULL OR workflow_name = '') LIMIT 1", + )?; + let res = stmt.exists(params![name_owned, wf_owned])?; + Ok(res) + }) + .await??; + Ok(exists) + } + pub async fn insert_task(&self, spec: &common::models::TaskSpec) -> Result<()> { let pool = self.pool.clone(); let spec = spec.clone(); @@ -386,16 +958,18 @@ impl Database { common::models::TaskType::SeedStep => "seed_step", }; conn.execute( - "INSERT INTO tasks (task_id, point_name, task_type, seed_point_name, status, created_at) - VALUES (?1, ?2, ?3, ?4, 'pending', datetime('now')) + "INSERT INTO tasks (task_id, point_name, task_type, seed_point_name, status, created_at, workflow_name) + VALUES (?1, ?2, ?3, ?4, 'pending', datetime('now'), ?5) ON CONFLICT(task_id) DO UPDATE SET status = 'pending', - seed_point_name = excluded.seed_point_name", + seed_point_name = excluded.seed_point_name, + workflow_name = excluded.workflow_name", params![ spec.task_id.to_string(), spec.point_name, task_type_str, - spec.seed_point_name + spec.seed_point_name, + spec.workflow_name, ], )?; Ok(()) @@ -404,10 +978,11 @@ impl Database { Ok(()) } - pub async fn record_task_report(&self, report: &TaskReport) -> Result<()> { + pub async fn record_task_report(&self, report: &TaskReport, workflow_name: &str) -> Result<()> { let pool = self.pool.clone(); let report_cloned = report.clone(); let point_name = report.point_name.clone(); + let wf = workflow_name.to_string(); let converged = report.converged; let atmo_has_nan = report.atmosphere_has_nan; @@ -434,36 +1009,51 @@ impl Database { ], )?; - // 合并重试次数 +1 与查值操作至一条 atomic sql UPDATE RETURNING 语句,彻底杜绝多事务并发下的读写竞态;若失败则返回 i32::MAX 强制定向至 failed 回退保护 - let current_attempts: i32 = tx + // 失败次数计数自增(attempt_count 仅作观测/统计用途,保留原子 UPDATE 避免并发竞态)。 + // 注意:当前“种子回退仅一次”语义下,状态迁移不再依赖该计数值(失败统一置 failed, + // 是否复活由 trigger_seed_step_fallback 按 has_seed_step_attempt 决定)。 + // 三条 UPDATE grid_points 均带 workflow_name 过滤,避免跨工作流误改同名点。 + let _: i32 = tx .query_row( - "UPDATE grid_points SET attempt_count = attempt_count + 1 WHERE name = ?1 RETURNING attempt_count", - params![point_name], + "UPDATE grid_points SET attempt_count = attempt_count + 1 WHERE name = ?1 AND workflow_name = ?2 RETURNING attempt_count", + params![point_name, wf], |r| r.get(0), ) - .unwrap_or(i32::MAX); + .unwrap_or(0); if report_cloned.status == TaskStatus::Completed && converged && !atmo_has_nan { tx.execute( - "UPDATE grid_points SET status = 'converged', success_method = (SELECT task_type FROM tasks WHERE task_id = ?1) WHERE name = ?2", - params![report_cloned.task_id.to_string(), point_name], + "UPDATE grid_points SET status = 'converged', success_method = (SELECT task_type FROM tasks WHERE task_id = ?1) WHERE name = ?2 AND workflow_name = ?3", + params![report_cloned.task_id.to_string(), point_name, wf], )?; } else { - let max_attempts = 3; - if current_attempts >= max_attempts { - tx.execute( - "UPDATE grid_points SET status = 'failed' WHERE name = ?1", - params![point_name], - )?; - } else { - tx.execute( - "UPDATE grid_points SET status = 'pending' WHERE name = ?1", - params![point_name], - )?; - } + // 失败处理(种子回退仅一次语义): + // 不再把失败点置回 'pending'(否则后台 30s 调度会盲目重投 ColdRun,与 + // report_task 的种子回退路径形成无互斥的双重重试)。失败直接置 'failed', + // 由上层 trigger_seed_step_fallback 决定是否复活为 queued 做一次种子热启动。 + tx.execute( + "UPDATE grid_points SET status = 'failed' WHERE name = ?1 AND workflow_name = ?2", + params![point_name, wf], + )?; } tx.commit()?; + + // 自动判断工作流完成度:如果该工作流下的所有网格点均已到达终态(即无 pending, queued, running 状态的点),自动更新工作流状态为 'completed' + let _ = conn.execute( + "UPDATE workflows + SET status = 'completed', updated_at = datetime('now') + WHERE name = ?1 + AND status = 'running' + AND EXISTS (SELECT 1 FROM grid_points WHERE workflow_name = workflows.name) + AND NOT EXISTS ( + SELECT 1 FROM grid_points + WHERE workflow_name = workflows.name + AND status IN ('pending', 'queued', 'running') + )", + params![wf], + ); + Ok(()) }) .await??; @@ -515,8 +1105,20 @@ impl Database { }) .await??; + // 同步重建 exact_family 索引(每个种子写入其 floor/floor+1 两个桶)。 + let mut index: std::collections::HashMap> = + std::collections::HashMap::new(); + for item in &items { + for key in SeedBucketKey::from_params(&item.params) { + index.entry(key).or_default().push(item.clone()); + } + } + let mut lock = self.seed_cache.write().await; *lock = items; + drop(lock); + let mut idx_lock = self.seed_index.write().await; + *idx_lock = index; Ok(()) } @@ -530,27 +1132,46 @@ impl Database { let path_db = path_owned.clone(); let p_db = p.clone(); tokio::task::spawn_blocking(move || -> Result<()> { - let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + let conn = pool + .get() + .map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; conn.execute( "INSERT INTO seeds (point_name, teff, logg, loghe, logc, logn, logo, file_path) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8) ON CONFLICT(point_name) DO UPDATE SET file_path = excluded.file_path", - params![name_db, p_db.teff, p_db.logg, p_db.loghe, p_db.logc, p_db.logn, p_db.logo, path_db], + params![ + name_db, p_db.teff, p_db.logg, p_db.loghe, p_db.logc, p_db.logn, p_db.logo, + path_db + ], )?; Ok(()) }) .await??; let item = SeedCacheItem { - point_name: name, - params: p, + point_name: name.clone(), + params: p.clone(), file_path: path_owned, }; let mut lock = self.seed_cache.write().await; + let is_new; if let Some(pos) = lock.iter().position(|x| x.point_name == item.point_name) { - lock[pos] = item; + // 已存在:seeds 表 ON CONFLICT 只更新 file_path,point_name/物理参数不变, + // 故 exact_family 桶键不变,索引无需重写,仅同步 Vec 里的 file_path。 + lock[pos].file_path = item.file_path.clone(); + is_new = false; } else { - lock.push(item); + lock.push(item.clone()); + is_new = true; + } + drop(lock); + + // 新种子才需写入索引(已存在的种子 params 不变,桶键未变)。 + if is_new { + let mut idx_lock = self.seed_index.write().await; + for key in SeedBucketKey::from_params(&p) { + idx_lock.entry(key).or_default().push(item.clone()); + } } Ok(()) @@ -560,27 +1181,71 @@ impl Database { &self, target: &GridPointParams, ) -> Result> { - let lock = self.seed_cache.read().await; + // 优先走 exact_family 索引(O(1)~O(小)):取出 target 的两个候选桶的全部种子快照后 + // 立即释放读锁,避免阻塞 insert_seed 写。exact_family 是绝大多数命中的路径。 + let exact_candidates: Vec = { + let idx_lock = self.seed_index.read().await; + let keys = SeedBucketKey::from_params(target); + let mut out = Vec::new(); + for key in keys { + if let Some(bucket) = idx_lock.get(&key) { + out.extend(bucket.iter().cloned()); + } + } + out + }; + let mut exact_family: Option<(String, std::path::PathBuf, f64)> = None; - let mut global_closest: Option<(String, std::path::PathBuf, f64)> = None; - - for item in lock.iter() { - let path = std::path::PathBuf::from(&item.file_path); + for item in &exact_candidates { let (is_exact, d) = common::seed_finder::calculate_seed_distance(&item.params, target); - if is_exact { + let path = std::path::PathBuf::from(&item.file_path); if exact_family.is_none() || d < exact_family.as_ref().unwrap().2 { exact_family = Some((item.point_name.clone(), path, d)); } - } else if d <= common::seed_finder::MAX_GLOBAL_SEED_DISTANCE && (global_closest.is_none() || d < global_closest.as_ref().unwrap().2) { - global_closest = Some((item.point_name.clone(), path, d)); + } + } + if let Some((name, path, d)) = exact_family { + return Ok(Some(common::seed_finder::SeedMatch { + name, + path, + distance: d, + })); + } + + // exact_family 未命中:退化到全量 global 扫描。克隆参数缩小读锁持有范围。 + let snapshot: Vec<_> = { + let lock = self.seed_cache.read().await; + lock.iter() + .map(|item| { + ( + item.point_name.clone(), + item.params.clone(), + item.file_path.clone(), + ) + }) + .collect() + }; + + let mut global_closest: Option<(String, std::path::PathBuf, f64)> = None; + for (point_name, params, file_path) in snapshot { + let (is_exact, d) = common::seed_finder::calculate_seed_distance(¶ms, target); + // exact_family 路径已在上面处理过(索引已覆盖),这里只关心 global 候选。 + if !is_exact + && d <= common::seed_finder::MAX_GLOBAL_SEED_DISTANCE + && (global_closest.is_none() || d < global_closest.as_ref().unwrap().2) + { + let path = std::path::PathBuf::from(&file_path); + global_closest = Some((point_name, path, d)); } } - if let Some((name, path, d)) = exact_family { - Ok(Some(common::seed_finder::SeedMatch { name, path, distance: d })) - } else if let Some((name, path, d)) = global_closest { - Ok(Some(common::seed_finder::SeedMatch { name, path, distance: d })) + if let Some((name, path, d)) = global_closest { + Ok(Some(common::seed_finder::SeedMatch { + name, + path, + distance: d, + })) } else { Ok(None) } @@ -618,6 +1283,32 @@ impl Database { Ok(()) } + /// 自动巡检所有处于 running 状态的工作流: + /// 若某个工作流下的所有网格点均已到达终态(无 pending/queued/running 点), + /// 则自动将该工作流的数据库 status 翻转为 'completed'。 + pub async fn sync_all_running_workflows_completion(&self) -> Result<()> { + let pool = self.pool.clone(); + tokio::task::spawn_blocking(move || -> Result<()> { + let conn = pool + .get() + .map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + let _ = conn.execute( + "UPDATE workflows + SET status = 'completed', updated_at = datetime('now') + WHERE status = 'running' + AND EXISTS (SELECT 1 FROM grid_points WHERE workflow_name = workflows.name) + AND NOT EXISTS ( + SELECT 1 FROM grid_points + WHERE workflow_name = workflows.name + AND status IN ('pending', 'queued', 'running') + )", + [], + ); + Ok(()) + }) + .await? + } + pub async fn list_workflows(&self) -> Result> { let pool = self.pool.clone(); tokio::task::spawn_blocking(move || -> Result> { @@ -678,7 +1369,9 @@ impl Database { let status_owned = status.to_string(); tokio::task::spawn_blocking(move || -> Result<()> { - let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + let conn = pool + .get() + .map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; conn.execute( "UPDATE workflows SET status = ?1, updated_at = datetime('now') WHERE name = ?2", params![status_owned, name_owned], @@ -695,8 +1388,20 @@ impl Database { let name_owned = name.to_string(); tokio::task::spawn_blocking(move || -> Result<()> { - let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; - conn.execute("DELETE FROM workflows WHERE name = ?1", params![name_owned])?; + let mut conn = pool + .get() + .map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + let tx = conn.transaction()?; + tx.execute("DELETE FROM workflows WHERE name = ?1", params![name_owned])?; + tx.execute( + "DELETE FROM grid_points WHERE workflow_name = ?1", + params![name_owned], + )?; + tx.execute( + "DELETE FROM tasks WHERE workflow_name = ?1", + params![name_owned], + )?; + tx.commit()?; Ok(()) }) .await??; @@ -707,7 +1412,9 @@ impl Database { pub async fn has_running_workflow(&self) -> Result { let pool = self.pool.clone(); tokio::task::spawn_blocking(move || -> Result { - let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + let conn = pool + .get() + .map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; let count: i64 = conn.query_row( "SELECT COUNT(*) FROM workflows WHERE status = 'running'", [], @@ -718,6 +1425,49 @@ impl Database { .await? } + /// 返回当前处于 running / initializing 状态的首个工作流名称。 + /// + /// 用于在派发任务时给 TaskSpec 填充 workflow_name(按工作流隔离队列清理)。 + /// 与 get_running_workflow_config_yamls 的状态口径保持一致。 + pub async fn get_running_workflow_name(&self) -> Result> { + let pool = self.pool.clone(); + tokio::task::spawn_blocking(move || -> Result> { + let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + let res = conn.query_row( + "SELECT name FROM workflows WHERE status IN ('running', 'initializing') ORDER BY updated_at DESC LIMIT 1", + [], + |r| r.get::<_, String>(0), + ); + match res { + Ok(name) => Ok(Some(name)), + Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None), + Err(e) => Err(e.into()), + } + }) + .await? + } + + /// 返回当前处于 running / initializing 状态的**全部**工作流名称。 + /// + /// 多工作流并发分区:后台调度需对每个 running 工作流分别派发任务, + /// 替代原来「全局只有一个 running workflow」的 LIMIT 1 假设。 + pub async fn get_running_workflow_names(&self) -> Result> { + let pool = self.pool.clone(); + tokio::task::spawn_blocking(move || -> Result> { + let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + let mut stmt = conn.prepare( + "SELECT name FROM workflows WHERE status IN ('running', 'initializing') ORDER BY updated_at ASC", + )?; + let rows = stmt.query_map([], |r| r.get::<_, String>(0))?; + let mut list = Vec::new(); + for r in rows { + list.push(r?); + } + Ok(list) + }) + .await? + } + /// 原子切转工作流至 initializing 预占启动状态,杜绝高并发 POST /start 触发双重全量排队与重置网格竞态 pub async fn transition_workflow_to_initializing(&self, name: &str) -> Result { let pool = self.pool.clone(); @@ -738,8 +1488,12 @@ impl Database { pub async fn get_running_workflow_config_yamls(&self) -> Result> { let pool = self.pool.clone(); tokio::task::spawn_blocking(move || -> Result> { - let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; - let mut stmt = conn.prepare("SELECT config_yaml FROM workflows WHERE status IN ('running', 'initializing')")?; + let conn = pool + .get() + .map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + let mut stmt = conn.prepare( + "SELECT config_yaml FROM workflows WHERE status IN ('running', 'initializing')", + )?; let rows = stmt.query_map([], |row| row.get(0))?; let mut list = Vec::new(); for r in rows { @@ -750,17 +1504,56 @@ impl Database { .await? } - pub async fn get_grid_summary_stats(&self) -> Result { + /// 网格汇总统计。 + /// + /// `workflow_filter`: + /// - `None`:聚合全部工作流的 grid_points(dashboard 全局概览用)。 + /// - `Some(wf)`:仅聚合指定工作流(按工作流隔离的进度统计)。 + pub async fn get_grid_summary_stats( + &self, + workflow_filter: Option<&str>, + ) -> Result { let pool = self.pool.clone(); + let wf = workflow_filter.map(|s| s.to_string()); tokio::task::spawn_blocking(move || -> Result { let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; - let total: i64 = conn.query_row("SELECT COUNT(*) FROM grid_points", [], |r| r.get(0)).unwrap_or(0); - let pending: i64 = conn.query_row("SELECT COUNT(*) FROM grid_points WHERE status IN ('pending', 'queued')", [], |r| r.get(0)).unwrap_or(0); - let running: i64 = conn.query_row("SELECT COUNT(*) FROM grid_points WHERE status = 'running'", [], |r| r.get(0)).unwrap_or(0); - let converged: i64 = conn.query_row("SELECT COUNT(*) FROM grid_points WHERE status = 'converged'", [], |r| r.get(0)).unwrap_or(0); - let failed: i64 = conn.query_row("SELECT COUNT(*) FROM grid_points WHERE status = 'failed'", [], |r| r.get(0)).unwrap_or(0); - let cold_run_converged: i64 = conn.query_row("SELECT COUNT(*) FROM grid_points WHERE status = 'converged' AND success_method = 'cold_run'", [], |r| r.get(0)).unwrap_or(0); - let seed_step_converged: i64 = conn.query_row("SELECT COUNT(*) FROM grid_points WHERE status = 'converged' AND success_method = 'seed_step'", [], |r| r.get(0)).unwrap_or(0); + // 合并原先 7 条独立 COUNT 查询为单次扫描,用 SUM(CASE WHEN ...) 一次性聚合所有口径, + // 显著降低 status API 的数据库往返与锁竞争开销。 + let row = match &wf { + Some(name) => conn.query_row( + "SELECT + COUNT(*) AS total, + SUM(CASE WHEN status IN ('pending', 'queued') THEN 1 ELSE 0 END) AS pending, + SUM(CASE WHEN status = 'running' THEN 1 ELSE 0 END) AS running, + SUM(CASE WHEN status = 'converged' THEN 1 ELSE 0 END) AS converged, + SUM(CASE WHEN status = 'failed' THEN 1 ELSE 0 END) AS failed, + SUM(CASE WHEN status = 'converged' AND success_method = 'cold_run' THEN 1 ELSE 0 END) AS cold_run_converged, + SUM(CASE WHEN status = 'converged' AND success_method = 'seed_step' THEN 1 ELSE 0 END) AS seed_step_converged + FROM grid_points WHERE workflow_name = ?1", + params![name], + |r| { + let n = |i: usize| -> i64 { r.get::<_, Option>(i).unwrap_or(None).unwrap_or(0) }; + Ok((n(0), n(1), n(2), n(3), n(4), n(5), n(6))) + }, + ), + None => conn.query_row( + "SELECT + COUNT(*) AS total, + SUM(CASE WHEN status IN ('pending', 'queued') THEN 1 ELSE 0 END) AS pending, + SUM(CASE WHEN status = 'running' THEN 1 ELSE 0 END) AS running, + SUM(CASE WHEN status = 'converged' THEN 1 ELSE 0 END) AS converged, + SUM(CASE WHEN status = 'failed' THEN 1 ELSE 0 END) AS failed, + SUM(CASE WHEN status = 'converged' AND success_method = 'cold_run' THEN 1 ELSE 0 END) AS cold_run_converged, + SUM(CASE WHEN status = 'converged' AND success_method = 'seed_step' THEN 1 ELSE 0 END) AS seed_step_converged + FROM grid_points", + [], + |r| { + let n = |i: usize| -> i64 { r.get::<_, Option>(i).unwrap_or(None).unwrap_or(0) }; + Ok((n(0), n(1), n(2), n(3), n(4), n(5), n(6))) + }, + ), + }?; + let (total, pending, running, converged, failed, cold_run_converged, seed_step_converged) = row; Ok(serde_json::json!({ "total": total, @@ -774,6 +1567,76 @@ impl Database { }) .await? } + + pub async fn backup_database(&self, backup_dir: &str) -> Result<()> { + let pool = self.pool.clone(); + let dir = backup_dir.to_string(); + + tokio::task::spawn_blocking(move || -> Result<()> { + let path = std::path::Path::new(&dir); + if !path.exists() { + std::fs::create_dir_all(path)?; + } + + let now = std::time::SystemTime::now(); + let seven_days = std::time::Duration::from_secs(7 * 24 * 3600); + if let Ok(entries) = std::fs::read_dir(path) { + for entry in entries.flatten() { + let file_path = entry.path(); + if file_path.is_file() { + if let Some(ext) = file_path.extension() { + if ext == "db" { + if let Ok(meta) = entry.metadata() { + if let Ok(modified) = meta.modified() { + if now + .duration_since(modified) + .unwrap_or(std::time::Duration::from_secs(0)) + > seven_days + { + let _ = std::fs::remove_file(file_path); + } + } + } + } + } + } + } + } + + let timestamp = chrono::Local::now().format("%Y%m%d_%H%M%S"); + let backup_file = path.join(format!("dcts_backup_{}.db", timestamp)); + + let conn = pool + .get() + .map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?; + conn.execute( + "VACUUM INTO ?1", + params![backup_file.to_string_lossy().to_string()], + )?; + + Ok(()) + }) + .await??; + + Ok(()) + } +} + +/// 管理视图:节点信息 + 凭据状态(用于 admin 列表 / 吊销管理)。 +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +pub struct NodeCredentialView { + pub node_id: String, + pub host_name: String, + pub max_slots: i32, + pub active_slots: i32, + pub status: String, + pub cpu_usage: f32, + pub memory_usage: f32, + pub last_heartbeat: chrono::DateTime, + /// 凭据状态:none(无凭据记录) / active(有效) / revoked(已吊销) + pub token_status: String, + /// 凭据颁发时间(ISO 字符串,无凭据时为 None) + pub token_issued_at: Option, } #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] @@ -813,6 +1676,7 @@ mod tests { max_slots: 4, }; db.register_node(®_req).await.unwrap(); + db.approve_node("node-test-1").await.unwrap(); let active_nodes = db.get_active_nodes().await.unwrap(); assert_eq!(active_nodes.len(), 1); @@ -835,9 +1699,9 @@ mod tests { logn: -2.0, logo: -2.0, }; - db.upsert_grid_point(¶ms, 0).await.unwrap(); + db.upsert_grid_point(¶ms, 0, "test_wf").await.unwrap(); - let pending = db.get_pending_grid_points().await.unwrap(); + let pending = db.get_pending_grid_points("test_wf").await.unwrap(); assert_eq!(pending.len(), 1); assert_eq!(pending[0].0, params.model_name()); @@ -855,14 +1719,21 @@ mod tests { error_message: None, summary_json: "{}".to_string(), }; - db.record_task_report(&report).await.unwrap(); + db.record_task_report(&report, "test_wf").await.unwrap(); // Check grid point is marked converged - let pending_after = db.get_pending_grid_points().await.unwrap(); + let pending_after = db.get_pending_grid_points("test_wf").await.unwrap(); assert_eq!(pending_after.len(), 0); // 3. Workflow CRUD - db.upsert_workflow("test_wf", Some("Test Workflow"), "grid:\n teff: [35000]", "idle").await.unwrap(); + db.upsert_workflow( + "test_wf", + Some("Test Workflow"), + "grid:\n teff: [35000]", + "idle", + ) + .await + .unwrap(); let wf = db.get_workflow("test_wf").await.unwrap(); assert!(wf.is_some()); assert_eq!(wf.unwrap().name, "test_wf"); @@ -870,4 +1741,381 @@ mod tests { db.delete_workflow("test_wf").await.unwrap(); assert!(db.get_workflow("test_wf").await.unwrap().is_none()); } + + /// 多工作流分区隔离测试(#3 修复核心验证): + /// 1. 同一物理点写入两个工作流,互不覆盖(复合唯一约束)。 + /// 2. reset_queued_grid_points_to_pending 按 workflow 隔离:重置 wf_a 不影响 wf_b。 + /// 3. update_grid_status 按 workflow 隔离:改 wf_a 的点不影响 wf_b 同名点。 + /// 4. reset_specific_grid_points_to_pending 按 workflow 隔离。 + #[tokio::test] + async fn test_multi_workflow_grid_isolation() { + let temp_dir = tempfile::tempdir().unwrap(); + let db_path = temp_dir.path().join("iso_db.db"); + let db = Database::new(&db_path.to_string_lossy()).await.unwrap(); + + let params = GridPointParams { + teff: 35000.0, + logg: 5.5, + loghe: -1.0, + logc: -2.0, + logn: -2.0, + logo: -2.0, + }; + let name = params.model_name(); + + // 1. 两个工作流写入同一物理点 —— 应各自独立存在(复合唯一 (wf, name))。 + db.upsert_grid_point(¶ms, 0, "wf_a").await.unwrap(); + db.upsert_grid_point(¶ms, 0, "wf_b").await.unwrap(); + assert_eq!(db.get_pending_grid_points("wf_a").await.unwrap().len(), 1); + assert_eq!(db.get_pending_grid_points("wf_b").await.unwrap().len(), 1); + + // 2. 把 wf_a 的点置 queued,wf_b 保持 pending;reset wf_a 的 queued 不应波及 wf_b。 + db.update_grid_status(&name, GridPointStatus::Queued, "wf_a") + .await + .unwrap(); + assert_eq!( + db.get_grid_point_status(&name, "wf_a") + .await + .unwrap() + .unwrap() + .0, + "queued" + ); + assert_eq!( + db.get_grid_point_status(&name, "wf_b") + .await + .unwrap() + .unwrap() + .0, + "pending" + ); + let reset_cnt = db + .reset_queued_grid_points_to_pending("wf_a") + .await + .unwrap(); + assert_eq!(reset_cnt, 1); + assert_eq!( + db.get_grid_point_status(&name, "wf_a") + .await + .unwrap() + .unwrap() + .0, + "pending" + ); + // wf_b 仍是 pending(未被误改) + assert_eq!( + db.get_grid_point_status(&name, "wf_b") + .await + .unwrap() + .unwrap() + .0, + "pending" + ); + + // 3. update_grid_status 按 workflow 隔离:把 wf_a 标 failed,wf_b 不受影响。 + db.update_grid_status(&name, GridPointStatus::Failed, "wf_a") + .await + .unwrap(); + assert_eq!( + db.get_grid_point_status(&name, "wf_a") + .await + .unwrap() + .unwrap() + .0, + "failed" + ); + assert_eq!( + db.get_grid_point_status(&name, "wf_b") + .await + .unwrap() + .unwrap() + .0, + "pending" + ); + + // 4. reset_specific 按 workflow 隔离:wf_a 的 queued→running 点被重置,wf_b 同名点不动。 + db.update_grid_status(&name, GridPointStatus::Queued, "wf_a") + .await + .unwrap(); + db.update_grid_status(&name, GridPointStatus::Queued, "wf_b") + .await + .unwrap(); + let cnt = db + .reset_specific_grid_points_to_pending(std::slice::from_ref(&name), "wf_a") + .await + .unwrap(); + assert_eq!(cnt, 1, "仅 wf_a 的 queued 点被重置"); + assert_eq!( + db.get_grid_point_status(&name, "wf_a") + .await + .unwrap() + .unwrap() + .0, + "pending" + ); + assert_eq!( + db.get_grid_point_status(&name, "wf_b") + .await + .unwrap() + .unwrap() + .0, + "queued" + ); + } + + /// grid_summary_stats 按 workflow 聚合 + 全局聚合测试。 + #[tokio::test] + async fn test_grid_summary_stats_workflow_scoping() { + let temp_dir = tempfile::tempdir().unwrap(); + let db = Database::new(&temp_dir.path().join("stats_db.db").to_string_lossy()) + .await + .unwrap(); + let p = GridPointParams { + teff: 35000.0, + logg: 5.5, + loghe: -1.0, + logc: -2.0, + logn: -2.0, + logo: -2.0, + }; + db.upsert_grid_point(&p, 0, "wf_a").await.unwrap(); + db.upsert_grid_point(&p, 0, "wf_b").await.unwrap(); + + // 全局(None):合计 2 个 pending + let all = db.get_grid_summary_stats(None).await.unwrap(); + assert_eq!(all["total"], 2); + assert_eq!(all["pending"], 2); + // 单工作流:各 1 + let a = db.get_grid_summary_stats(Some("wf_a")).await.unwrap(); + assert_eq!(a["total"], 1); + assert_eq!(a["pending"], 1); + // 不存在的工作流:0 + let none = db + .get_grid_summary_stats(Some("nonexistent")) + .await + .unwrap(); + assert_eq!(none["total"], 0); + } + + /// seed_finder exact_family 索引与全量扫描结果一致性测试(#7 优化正确性)。 + /// 构造一个 exact_family 候选 + 一个 global 候选,验证索引路径仍命中正确结果。 + #[tokio::test] + async fn test_seed_index_exact_and_global_consistency() { + let temp_dir = tempfile::tempdir().unwrap(); + let db = Database::new(&temp_dir.path().join("seed_db.db").to_string_lossy()) + .await + .unwrap(); + + // exact_family 种子:与 target 同 teff/logg/loghe,仅 CNO 略有差异。 + let exact = GridPointParams { + teff: 35000.0, + logg: 5.5, + loghe: -1.0, + logc: -2.0, + logn: -2.0, + logo: -2.0, + }; + let target = GridPointParams { + teff: 35000.0, + logg: 5.5, + loghe: -1.0, + logc: -2.1, + logn: -2.0, + logo: -2.0, + }; + db.insert_seed(&exact, "/tmp/exact.7").await.unwrap(); + + // 应命中 exact_family(走索引路径) + let m = db.find_best_seed_from_db(&target).await.unwrap(); + assert!(m.is_some(), "exact_family 索引路径应命中"); + assert_eq!(m.unwrap().name, exact.model_name()); + } + + /// 工作流计算完成自动状态迁移测试(running -> completed): + /// 当工作流处于 running,且所有网格点全部到达终态(converged/failed)时, + /// list_workflows/sync_all_running_workflows_completion 应自动将状态翻转为 completed。 + #[tokio::test] + async fn test_workflow_auto_completion_status_transition() { + let temp_dir = tempfile::tempdir().unwrap(); + let db = Database::new(&temp_dir.path().join("auto_comp_db.db").to_string_lossy()) + .await + .unwrap(); + + let params = GridPointParams { + teff: 35000.0, + logg: 5.5, + loghe: -1.0, + logc: -2.0, + logn: -2.0, + logo: -2.0, + }; + let name = params.model_name(); + + // 1. 注册并启动工作流 auto_wf + db.upsert_workflow("auto_wf", Some("Auto Comp Test"), "config", "idle") + .await + .unwrap(); + db.update_workflow_status("auto_wf", "running") + .await + .unwrap(); + db.upsert_grid_point(¶ms, 0, "auto_wf").await.unwrap(); + + // 此时网格点为 pending,工作流应保持 running + let list = db.list_workflows().await.unwrap(); + assert_eq!(list[0].status, "running"); + + // 2. 网格点完成计算(converged) + db.update_grid_status(&name, GridPointStatus::Converged, "auto_wf") + .await + .unwrap(); + + // 3. 执行同步巡检,应当触发自动翻转为 completed + db.sync_all_running_workflows_completion().await.unwrap(); + let list_after = db.list_workflows().await.unwrap(); + assert_eq!( + list_after[0].status, "completed", + "所有点完成计算后,工作流状态应自动转换为 completed" + ); + + // 4. get_workflow 也返回 completed + let item = db.get_workflow("auto_wf").await.unwrap().unwrap(); + assert_eq!(item.status, "completed"); + } + + #[tokio::test] + async fn test_take_pending_node_token_atomic() { + let temp_dir = tempfile::tempdir().unwrap(); + let db_path = temp_dir.path().join("token_test.db"); + let db = Database::new(&db_path.to_string_lossy()).await.unwrap(); + + let reg = NodeRegisterRequest { + node_id: "node-atomic-test".to_string(), + host_name: "host1".to_string(), + max_slots: 2, + }; + db.register_node(®).await.unwrap(); + let token = db.approve_node("node-atomic-test").await.unwrap(); + assert!(!token.is_empty()); + + // 第一次调用:返回 token + let pending1 = db + .take_pending_node_token("node-atomic-test") + .await + .unwrap(); + assert_eq!(pending1, Some(token)); + + // 第二次调用:已被置为 NULL,返回 None + let pending2 = db + .take_pending_node_token("node-atomic-test") + .await + .unwrap(); + assert_eq!(pending2, None); + } + + #[tokio::test] + async fn test_heartbeat_node_status_check() { + let temp_dir = tempfile::tempdir().unwrap(); + let db_path = temp_dir.path().join("hb_check.db"); + let db = Database::new(&db_path.to_string_lossy()).await.unwrap(); + + let reg = NodeRegisterRequest { + node_id: "node-hb-test".to_string(), + host_name: "host1".to_string(), + max_slots: 2, + }; + db.register_node(®).await.unwrap(); + db.approve_node("node-hb-test").await.unwrap(); + + // 此时 node 状态为 online,心跳正常更新 + let hb_req = NodeHeartbeatRequest { + node_id: "node-hb-test".to_string(), + active_slots: 1, + cpu_usage: 10.0, + memory_usage: 20.0, + }; + db.heartbeat_node(&hb_req).await.unwrap(); + let nodes = db.get_active_nodes().await.unwrap(); + assert_eq!(nodes.len(), 1); + + // 人为修改节点状态为 rejected + let pool = db.pool.clone(); + tokio::task::spawn_blocking(move || { + let conn = pool.get().unwrap(); + conn.execute( + "UPDATE nodes SET status = 'rejected' WHERE node_id = 'node-hb-test'", + [], + ) + .unwrap(); + }) + .await + .unwrap(); + + // 再次发心跳:不应重置 status 为 online + db.heartbeat_node(&hb_req).await.unwrap(); + let nodes2 = db.get_active_nodes().await.unwrap(); + assert_eq!( + nodes2.len(), + 0, + "status 为 rejected 的节点心跳时不应更新为 online" + ); + } + + #[tokio::test] + async fn test_delete_workflow_cascade_cleanup() { + let temp_dir = tempfile::tempdir().unwrap(); + let db_path = temp_dir.path().join("del_wf.db"); + let db = Database::new(&db_path.to_string_lossy()).await.unwrap(); + + let wf_name = "wf_to_delete"; + db.upsert_workflow(wf_name, Some("Test"), "config", "idle") + .await + .unwrap(); + + let params = GridPointParams { + teff: 35000.0, + logg: 5.5, + loghe: -1.0, + logc: -2.0, + logn: -2.0, + logo: -2.0, + }; + db.upsert_grid_point(¶ms, 0, wf_name).await.unwrap(); + + let spec = common::models::TaskSpec { + task_id: uuid::Uuid::new_v4(), + point_name: params.model_name(), + params: params.clone(), + task_type: common::models::TaskType::ColdRun, + seed_point_name: None, + timeout_sec: 600, + workflow_name: Some(wf_name.to_string()), + }; + db.insert_task(&spec).await.unwrap(); + + // 确认插入成功 + assert!(db.get_workflow(wf_name).await.unwrap().is_some()); + assert_eq!(db.get_pending_grid_points(wf_name).await.unwrap().len(), 1); + + // 删除工作流 + db.delete_workflow(wf_name).await.unwrap(); + + // 验证 workflows, grid_points, tasks 被级联清理 + assert!(db.get_workflow(wf_name).await.unwrap().is_none()); + assert_eq!(db.get_pending_grid_points(wf_name).await.unwrap().len(), 0); + + let pool = db.pool.clone(); + let wf_owned = wf_name.to_string(); + let task_cnt: i64 = tokio::task::spawn_blocking(move || { + let conn = pool.get().unwrap(); + conn.query_row( + "SELECT COUNT(*) FROM tasks WHERE workflow_name = ?1", + params![wf_owned], + |r| r.get(0), + ) + .unwrap() + }) + .await + .unwrap(); + assert_eq!(task_cnt, 0, "关联 tasks 记录应被清理"); + } } diff --git a/crates/server/src/lib.rs b/crates/server/src/lib.rs index bcfac53..8e7ce70 100644 --- a/crates/server/src/lib.rs +++ b/crates/server/src/lib.rs @@ -1,4 +1,4 @@ pub mod api; +pub mod cors; pub mod db; pub mod scheduler; - diff --git a/crates/server/src/main.rs b/crates/server/src/main.rs index 1c4c54a..2dd8334 100644 --- a/crates/server/src/main.rs +++ b/crates/server/src/main.rs @@ -3,8 +3,9 @@ use server::api::{self, AppState}; use server::db::Database; use server::scheduler::GridScheduler; - use axum::{ + extract::DefaultBodyLimit, + http::HeaderValue, routing::{get, post}, Router, }; @@ -16,12 +17,15 @@ use std::net::SocketAddr; use std::path::{Path, PathBuf}; use std::sync::Arc; use tokio::time::{sleep, Duration}; -use tower_http::cors::CorsLayer; use tower_http::services::{ServeDir, ServeFile}; use tracing::info; #[derive(Parser, Debug)] -#[command(name = "server", version = "0.1.0", about = "Distributed Computing TLUSTY/SYNSPEC (DCTS) Server")] +#[command( + name = "server", + version = "0.1.0", + about = "Distributed Computing TLUSTY/SYNSPEC (DCTS) Server" +)] struct CliArgs { /// Optional path to workflow configuration YAML file to auto-register on startup #[arg(short = 'w', long = "workflow")] @@ -61,12 +65,15 @@ async fn main() -> Result<()> { let default_wf_path = Path::new(&server_cfg.grid_config); if default_wf_path.is_file() { if let Ok(yaml_content) = std::fs::read_to_string(default_wf_path) { - if let Err(e) = db.upsert_workflow( - "sdB_cno", - Some("sdB CNO 6D Stellar Atmosphere Grid"), - &yaml_content, - "idle", - ).await { + if let Err(e) = db + .upsert_workflow( + "sdB_cno", + Some("sdB CNO 6D Stellar Atmosphere Grid"), + &yaml_content, + "idle", + ) + .await + { tracing::warn!("预注册默认工作流失败: {}", e); } else { info!("已在数据库中成功预注册默认工作流 'sdB_cno'"); @@ -74,15 +81,38 @@ async fn main() -> Result<()> { } } + // 弱口令凭据安全警告检测 + let is_weak_token = |t: Option<&str>| -> bool { + match t { + Some(s) => { + s.len() < 12 || s == "fmqi123" || s == "admin" || s == "123456" || s == "secret" + } + None => false, + } + }; + if is_weak_token(server_cfg.auth_token.as_deref()) + || is_weak_token(server_cfg.admin_token.as_deref()) + { + tracing::warn!("⚠️ 检测到系统当前正在使用弱口令凭据或默认 Token!建议生产环境在 .env 中配置使用 openssl rand -hex 32 生成的高强度 Token!"); + } + + let rate_limiter = api::rate_limit::RateLimiter::new(5, std::time::Duration::from_secs(300)); + let state = AppState { db, queue: queue.clone(), scheduler: scheduler.clone(), - results_dir: server_cfg.results_dir, + results_dir: server_cfg.results_dir.clone(), + rate_limiter, auth_token: server_cfg.auth_token.clone(), + admin_token: server_cfg.admin_token.clone(), + auth_disabled: server_cfg.auth_disabled, + admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new( + std::collections::HashMap::new(), + )), }; - // Background loop for stale task requeueing, offline node detection, and scheduler checking + // Background maintenance & scheduling with Exponential Backoff let bg_db = state.db.clone(); let bg_queue = queue.clone(); let bg_scheduler = scheduler.clone(); @@ -90,58 +120,226 @@ async fn main() -> Result<()> { let node_stale_sec = server_cfg.node_stale_sec; tokio::spawn(async move { + let mut fail_count: u32 = 0; + let mut first_run = true; loop { - sleep(Duration::from_secs(30)).await; - if let Ok(requeued_points) = bg_queue.requeue_stale_tasks(stale_sec).await { - if !requeued_points.is_empty() { - info!("重新将 {} 个超时/掉线任务放回待计算队列", requeued_points.len()); - let _ = bg_db.reset_specific_grid_points_to_pending(&requeued_points).await; - } + if first_run { + first_run = false; + } else { + let base_delay = 30u64; + let current_delay = if fail_count == 0 { + base_delay + } else { + (base_delay * (1u64 << fail_count.min(4))).min(300) + }; + sleep(Duration::from_secs(current_delay)).await; } - if let Ok(offline) = bg_db.mark_stale_nodes_offline(node_stale_sec).await { - if offline > 0 { - info!("已标记 {} 个心跳超时的计算节点为离线状态", offline); + + let bg_db_clone = bg_db.clone(); + let bg_queue_clone = bg_queue.clone(); + let bg_scheduler_clone = bg_scheduler.clone(); + + let join_handle = tokio::spawn(async move { + let mut has_error = false; + match bg_queue_clone.requeue_stale_tasks(stale_sec).await { + Ok(requeued) => { + if !requeued.is_empty() { + info!("重新将 {} 个超时/掉线任务放回待计算队列", requeued.len()); + let mut by_wf: std::collections::HashMap> = + std::collections::HashMap::new(); + for (point, wf) in &requeued { + by_wf + .entry(wf.clone().unwrap_or_default()) + .or_default() + .push(point.clone()); + } + for (wf, points) in by_wf { + let _ = bg_db_clone + .reset_specific_grid_points_to_pending(&points, &wf) + .await; + } + } + } + Err(e) => { + tracing::warn!("重投超时任务失败: {}", e); + has_error = true; + } + } + + match bg_db_clone.mark_stale_nodes_offline(node_stale_sec).await { + Ok(offline) => { + if offline > 0 { + info!("已标记 {} 个心跳超时的计算节点为离线状态", offline); + } + } + Err(e) => { + tracing::warn!("标记超时节点离线失败: {}", e); + has_error = true; + } + } + + if let Err(e) = bg_scheduler_clone.schedule_pending_tasks().await { + tracing::warn!("后台定时性任务调度检测失败: {}", e); + has_error = true; + } + + if let Err(e) = bg_db_clone.sync_all_running_workflows_completion().await { + tracing::warn!("后台同步已完成工作流状态失败: {}", e); + has_error = true; + } + + has_error + }); + + match join_handle.await { + Ok(has_error) => { + if has_error { + fail_count = fail_count.saturating_add(1); + } else { + fail_count = 0; + } + } + Err(e) => { + tracing::error!("后台维护任务内部发生 Panic: {:?}", e); + fail_count = fail_count.saturating_add(1); } - } - if let Err(e) = bg_scheduler.schedule_pending_tasks().await { - tracing::warn!("后台定时性任务调度检测失败: {}", e); } } }); + // 每天自动触发一次数据库备份。 + // 备份目录跟随 server_cfg.backup_dir(DCTS_BACKUP_DIR,默认 data/backups), + // 与 DB_PATH 解耦,避免 DB 卷与备份卷不一致时备份落到未持久化层。 + // 首次延迟 1 小时,避免频繁重启(如调试阶段)短时间堆积备份文件;backup_database + // 自身还带有 7 天保留期清理兜底。 + let backup_db = state.db.clone(); + let backup_dir = server_cfg.backup_dir.clone(); + tokio::spawn(async move { + sleep(Duration::from_secs(3600)).await; + loop { + if let Err(e) = backup_db.backup_database(&backup_dir).await { + tracing::warn!("自动备份数据库失败: {}", e); + } + sleep(Duration::from_secs(24 * 3600)).await; + } + }); + + // 大体积上传端点单独拎出,套用更宽松的 body limit(256MB,覆盖收敛种子 .7 文件量级) + // 并限制并发数:每个 report 请求最多 256MB 驻留内存,无并发上限时 N 个请求可耗尽内存。 + // 限流后超出并发数的请求排队等待(而非直接拒绝),保证正常业务不被误伤。 + // 其余 API 用 10MB 默认上限,防止大文件内存耗尽 DoS。 + const REPORT_BODY_LIMIT: usize = 256 * 1024 * 1024; + const DEFAULT_BODY_LIMIT: usize = 10 * 1024 * 1024; + const REPORT_MAX_CONCURRENCY: usize = 4; + + let report_router = Router::new() + .route("/task/report", post(api::task::report_task)) + .layer(DefaultBodyLimit::max(REPORT_BODY_LIMIT)) + .layer(tower::ServiceBuilder::new().concurrency_limit(REPORT_MAX_CONCURRENCY)); + + // 节点注册接口独立 IP 限流保护(每分钟最多 10 次申请,无论成败都计数,防恶意频繁注册) + // 使用 new_count_all:此 limiter 专挂 /node/register,对注册路径的所有响应计入窗口。 + // 通用 API 限流器(见下方 auth_enabled 分支)用 new 构造(count_all=false),不会因 + // 成功注册把 IP 锁出整个 /api/*,避免跨端点连锁限流。 + let register_limiter = + api::rate_limit::RateLimiter::new_count_all(10, std::time::Duration::from_secs(60)); + let register_rate_limit_layer = axum::middleware::from_fn_with_state( + register_limiter, + api::rate_limit::rate_limit_middleware, + ); + let api_router = Router::new() + // Auth API + .route("/login", post(api::auth::login)) + .route("/auth/check", get(api::auth::check_auth)) // Core Node & Task API - .route("/node/register", post(api::node::register_node)) + .route( + "/node/register", + post(api::node::register_node).layer(register_rate_limit_layer), + ) + .route("/node/check_status", post(api::node::check_node_status)) .route("/node/heartbeat", post(api::node::heartbeat_node)) .route("/task/claim", post(api::task::claim_task)) - .route("/task/report", post(api::task::report_task)) .route("/seed/:name", get(api::seed::download_seed)) .route("/status", get(api::status::get_status)) // Static Data API - .route("/data/file/*filename", get(api::data::download_single_data_file)) + .route( + "/data/file/*filename", + get(api::data::download_single_data_file), + ) .route("/data/linelist", get(api::data::download_linelist)) // Workflow Management CRUD API - .route("/workflows", get(api::workflow::list_workflows).post(api::workflow::save_workflow)) - .route("/workflows/:name", get(api::workflow::get_workflow).put(api::workflow::save_workflow).delete(api::workflow::delete_workflow)) - .route("/workflows/:name/start", post(api::workflow::start_workflow)) - .route("/workflows/:name/stop", post(api::workflow::stop_workflow)); + .route( + "/workflows", + get(api::workflow::list_workflows).post(api::workflow::save_workflow), + ) + .route( + "/workflows/:name", + get(api::workflow::get_workflow) + .put(api::workflow::save_workflow) + .delete(api::workflow::delete_workflow), + ) + .route( + "/workflows/:name/start", + post(api::workflow::start_workflow), + ) + .route("/workflows/:name/stop", post(api::workflow::stop_workflow)) + // Admin Management API(节点凭据查看/审批/吊销/重发,均要求 Admin 角色) + .route("/admin/nodes", get(api::admin::list_nodes)) + .route( + "/admin/nodes/:node_id/approve", + post(api::admin::approve_node), + ) + .route( + "/admin/nodes/:node_id/reject", + post(api::admin::reject_node), + ) + .route( + "/admin/nodes/:node_id/revoke", + post(api::admin::revoke_node), + ) + .route( + "/admin/nodes/:node_id/reissue", + post(api::admin::reissue_node), + ) + // 合并大体积上报路由(继承各自的 body limit) + .merge(report_router) + .layer(DefaultBodyLimit::max(DEFAULT_BODY_LIMIT)); - let api_router = if state.auth_token.is_some() { - info!("已为 DCTS 服务端 API 路由启用 Bearer Token / X-API-Key 访问控制鉴权"); + // 鉴权启用条件:未应急关闭,且配置了 admin 凭据。 + let auth_enabled = !state.auth_disabled && state.admin_token.is_some(); + + let api_router = if auth_enabled { + info!("已启用 API 身份鉴权保护(Admin 端点需 admin token 验证;Node 节点免 Token 提交申请,经 Dashboard 管理员审批授权下发)"); + // 鉴权失败限流(防 token 在线暴力):外层先判 IP 限流,内层再做鉴权。 + // 限流状态为 20 次/分钟(按 IP),超阈值返回 429。 + let limiter = api::rate_limit::RateLimiter::new(20, std::time::Duration::from_secs(60)); + let rate_limit_layer = + axum::middleware::from_fn_with_state(limiter, api::rate_limit::rate_limit_middleware); let auth_layer = axum::middleware::from_fn_with_state(state.clone(), api::auth_middleware); - api_router.layer(auth_layer) + api_router.layer(auth_layer).layer(rate_limit_layer) } else { - tracing::warn!("⚠️ 警告:未检测到 DCTS_AUTH_TOKEN 环境变量,服务端目前运行在【内网无鉴权模式】!所有 REST API 接口均为公开可访问状态。"); + tracing::warn!( + "⚠️ 警告:未配置 DCTS_ADMIN_TOKEN / DCTS_ENROLLMENT_TOKEN(且未启用 DCTS_AUTH_DISABLE),\ + 服务端运行在【无鉴权模式】!公网部署务必配置凭据。" + ); api_router }; // Host Dashboard SPA static files from dashboard/dist if directory exists or fallback to index.html - let serve_dir = ServeDir::new("dashboard/dist") - .fallback(ServeFile::new("dashboard/dist/index.html")); + let serve_dir = + ServeDir::new("dashboard/dist").fallback(ServeFile::new("dashboard/dist/index.html")); + + // 安全响应头(CSP / nosniff / DENY / Referrer-Policy)。 + let security_headers = axum::middleware::from_fn(security_headers_middleware); let app = Router::new() + // 独立健康检查端点:不走鉴权、不走 CORS/body 限制,专供 docker healthcheck 与外部监控探测。 + // 开启鉴权后 /api/status 会返回 401,导致容器被判定不健康而反复重启,故单独提供 /healthz。 + .route("/healthz", get(api::status::healthz)) .nest("/api", api_router) - .layer(CorsLayer::permissive()) + .layer(server::cors::build_cors_layer()) + .layer(security_headers) .fallback_service(serve_dir) .with_state(state); @@ -149,14 +347,53 @@ async fn main() -> Result<()> { info!("DCTS 服务端已在 http://{} 启动监听", addr); let listener = tokio::net::TcpListener::bind(addr).await?; - axum::serve(listener, app) - .with_graceful_shutdown(async { - let _ = tokio::signal::ctrl_c().await; - info!("收到 Ctrl+C 终止信号,DCTS 服务端准备优雅关闭..."); - }) - .await?; + // into_make_service_with_connect_info:让限流中间件能从连接拿到客户端 IP(反代场景则用 X-Forwarded-For) + axum::serve( + listener, + app.into_make_service_with_connect_info::(), + ) + .with_graceful_shutdown(async { + let _ = tokio::signal::ctrl_c().await; + info!("收到 Ctrl+C 终止信号,DCTS 服务端准备优雅关闭..."); + }) + .await?; info!("DCTS 服务端已安全关闭。"); Ok(()) } +/// 注入安全响应头的中间件函数。 +async fn security_headers_middleware( + req: axum::http::Request, + next: axum::middleware::Next, +) -> axum::response::Response { + let mut resp = next.run(req).await; + + let headers = resp.headers_mut(); + // CSP:default-src 'self';放行 Google Fonts(index.html 引用);允许 data: 图片。 + // 已移除 'unsafe-eval':dashboard 构建产物不使用 eval/new Function(已核实),保留它会 + // 显著削弱 CSP 的脚本注入防护。'unsafe-inline' 暂留(静态 SPA 内联脚本/handler 需要), + // 彻底方案需前端改造为外链 + per-request nonce 注入,见 docs TODO。 + headers + .entry(axum::http::header::CONTENT_SECURITY_POLICY) + .or_insert_with(|| { + HeaderValue::from_static( + "default-src 'self'; script-src 'self' 'unsafe-inline'; \ + style-src 'self' 'unsafe-inline' https://fonts.googleapis.com; \ + font-src 'self' data: https://fonts.gstatic.com; \ + connect-src 'self'; img-src 'self' data: blob:; \ + frame-ancestors 'none'", + ) + }); + headers + .entry(axum::http::header::X_CONTENT_TYPE_OPTIONS) + .or_insert_with(|| HeaderValue::from_static("nosniff")); + headers + .entry(axum::http::header::X_FRAME_OPTIONS) + .or_insert_with(|| HeaderValue::from_static("DENY")); + headers + .entry(axum::http::HeaderName::from_static("referrer-policy")) + .or_insert_with(|| HeaderValue::from_static("strict-origin-when-cross-origin")); + + resp +} diff --git a/crates/server/src/scheduler.rs b/crates/server/src/scheduler.rs index ce80c40..c535499 100644 --- a/crates/server/src/scheduler.rs +++ b/crates/server/src/scheduler.rs @@ -23,13 +23,32 @@ impl GridScheduler { } } - /// Expands grid points from config and registers them into the database - pub async fn initialize_grid(&self, cfg: &GridConfig) -> Result<()> { - if let Err(e) = self.queue.clear_queue().await { - tracing::warn!("初始化网格时清理闲置排队记录发生警告: {}", e); + /// Expands grid points from config and registers them into the database. + /// + /// 多工作流分区(#3 修复): + /// - 仅清理**本工作流**的排队任务(clear_queue_by_workflow),不再 clear_queue() 全局清空, + /// 避免启动工作流 B 时误删工作流 A 的在队任务。 + /// - 仅重置**本工作流**的 queued 点为 pending(reset_queued_grid_points_to_pending 带 wf), + /// 避免误伤其他工作流。 + /// - upsert 带 workflow_name,使同一物理点可属于多个工作流。 + pub async fn initialize_grid(&self, cfg: &GridConfig, workflow_name: &str) -> Result<()> { + if let Err(e) = self.queue.clear_queue_by_workflow(workflow_name).await { + tracing::warn!( + "初始化工作流 {} 网格时清理该流闲置排队记录发生警告: {}", + workflow_name, + e + ); } - if let Err(e) = self.db.reset_queued_grid_points_to_pending().await { - tracing::warn!("重置网格状态到 pending 处理过程遇到异常: {}", e); + if let Err(e) = self + .db + .reset_queued_grid_points_to_pending(workflow_name) + .await + { + tracing::warn!( + "重置工作流 {} 网格状态到 pending 处理过程遇到异常: {}", + workflow_name, + e + ); } let mut points = Vec::new(); @@ -59,9 +78,21 @@ impl GridScheduler { a.cno_sum() .partial_cmp(&b.cno_sum()) .unwrap_or(std::cmp::Ordering::Equal) - .then_with(|| a.teff.partial_cmp(&b.teff).unwrap_or(std::cmp::Ordering::Equal)) - .then_with(|| b.logg.partial_cmp(&a.logg).unwrap_or(std::cmp::Ordering::Equal)) - .then_with(|| a.loghe.partial_cmp(&b.loghe).unwrap_or(std::cmp::Ordering::Equal)) + .then_with(|| { + a.teff + .partial_cmp(&b.teff) + .unwrap_or(std::cmp::Ordering::Equal) + }) + .then_with(|| { + b.logg + .partial_cmp(&a.logg) + .unwrap_or(std::cmp::Ordering::Equal) + }) + .then_with(|| { + a.loghe + .partial_cmp(&b.loghe) + .unwrap_or(std::cmp::Ordering::Equal) + }) }); // Group into Waves by cno_sum @@ -79,46 +110,100 @@ impl GridScheduler { current_cno = Some(cno); } - self.db.upsert_grid_point(pt, wave_idx).await?; + // upsert 是幂等的 ON CONFLICT DO NOTHING:若 initialize_grid 中途失败, + // 重新 start 该工作流会自然补齐(#4 半初始化回退由幂等性消解)。 + self.db + .upsert_grid_point(pt, wave_idx, workflow_name) + .await?; } - info!("已在数据库中成功初始化并记录 {} 个恒星大气网格点", points.len()); + info!( + "已在数据库中成功初始化并记录工作流 {} 的 {} 个恒星大气网格点", + workflow_name, + points.len() + ); Ok(()) } - async fn get_active_timeout_sec(&self) -> u64 { - if let Ok(yamls) = self.db.get_running_workflow_config_yamls().await { - for yaml in yamls { - if let Ok(cfg) = serde_yaml::from_str::(&yaml) { - return cfg.timeout_sec; - } + /// 读取指定工作流的 timeout_sec(按工作流分区:多工作流各有自己的超时配置)。 + async fn get_workflow_timeout_sec(&self, workflow_name: &str) -> u64 { + if let Ok(Some(wf)) = self.db.get_workflow(workflow_name).await { + if let Ok(cfg) = serde_yaml::from_str::(&wf.config_yaml) { + return cfg.timeout_sec; } } 7200 } - /// Enqueues pending grid points into MQ with active seed detection and batching + /// 读取指定工作流的 seed_step_fallback 配置。 + async fn get_workflow_seed_step_fallback(&self, workflow_name: &str) -> bool { + if let Ok(Some(wf)) = self.db.get_workflow(workflow_name).await { + if let Ok(cfg) = serde_yaml::from_str::(&wf.config_yaml) { + return cfg.seed_step_fallback; + } + } + true + } + + /// Enqueues pending grid points into MQ with active seed detection and batching. + /// + /// 多工作流分区(#3 修复):对**每个** running/initializing 工作流分别派发任务, + /// 替代原来「全局只一个 running workflow」的 LIMIT 1 假设。各工作流独立 batch、 + /// 独立 seed 匹配(seeds 仍是全局共享的物理资源池)。 pub async fn schedule_pending_tasks(&self) -> Result { - if !self.db.has_running_workflow().await? { + let workflows = self.db.get_running_workflow_names().await?; + if workflows.is_empty() { return Ok(0); } - let timeout_sec = self.get_active_timeout_sec().await; let batch_limit: usize = std::env::var("DCTS_BATCH_LIMIT") .ok() .and_then(|v| v.parse().ok()) .unwrap_or(100); - // SQL 层直接附加 LIMIT = batch_limit 筛选,完全免除数万点位无谓内存反序列化和空耗对象释放开销 - let pending = self.db.get_pending_grid_points_limit(batch_limit).await?; + let mut total_dispatched = 0; + + for wf in &workflows { + let dispatched = self + .schedule_pending_tasks_for_workflow(wf, batch_limit) + .await?; + total_dispatched += dispatched; + } + + if total_dispatched > 0 { + info!( + "已成功将 {} 个待计算网格点推进任务队列(跨 {} 个工作流)", + total_dispatched, + workflows.len() + ); + } + + Ok(total_dispatched) + } + + /// 为单个工作流派发 pending 点。 + async fn schedule_pending_tasks_for_workflow( + &self, + workflow_name: &str, + batch_limit: usize, + ) -> Result { + let timeout_sec = self.get_workflow_timeout_sec(workflow_name).await; + + // SQL 层直接附加 LIMIT = batch_limit + workflow_name 筛选,完全免除数万点位无谓内存反序列化 + let pending = self + .db + .get_pending_grid_points_limit(batch_limit, workflow_name) + .await?; let mut dispatched = 0; for (name, params, _wave) in pending { - - // Check if any seed is available in DB for active SeedStep scheduling + // Check if any seed is available in DB for active SeedStep scheduling(seeds 全局共享) let (task_type, seed_name) = match self.db.find_best_seed_from_db(¶ms).await { Ok(Some(seed_match)) => { - info!("网格点 {} 匹配到数据库近邻种子 {} (距离: {:.2}),安排 SeedStep 热启动调度", name, seed_match.name, seed_match.distance); + info!( + "工作流 {} 网格点 {} 匹配到数据库近邻种子 {} (距离: {:.2}),安排 SeedStep 热启动调度", + workflow_name, name, seed_match.name, seed_match.distance + ); (TaskType::SeedStep, Some(seed_match.name)) } _ => (TaskType::ColdRun, None), @@ -131,48 +216,96 @@ impl GridScheduler { task_type, seed_point_name: seed_name, timeout_sec, + workflow_name: Some(workflow_name.to_string()), }; self.db.insert_task(&task_spec).await?; // 采用先标记 DB 状态为 Queued 后发 MQ 的时序,防止推入 MQ 后数据库修改异常导向下一轮误重投 - self.db.update_grid_status(&name, common::models::GridPointStatus::Queued).await?; + self.db + .update_grid_status( + &name, + common::models::GridPointStatus::Queued, + workflow_name, + ) + .await?; match self.queue.push_task(&task_spec).await { Ok(_) => { dispatched += 1; } Err(e) => { - tracing::warn!("将任务 {} 推入 MQ 队列失败,回滚网格点状态: {}", name, e); - let _ = self.db.update_grid_status(&name, common::models::GridPointStatus::Pending).await; + tracing::warn!( + "将任务 {} 推入 MQ 队列失败,执行严格状态回滚以避免脏数据: {}", + name, + e + ); + if let Err(db_e) = self + .db + .update_grid_status( + &name, + common::models::GridPointStatus::Pending, + workflow_name, + ) + .await + { + tracing::error!( + "关键性回滚异常:任务 {} 无法重置回 Pending: {}", + name, + db_e + ); + } let _ = self.queue.remove_task(&task_spec.task_id.to_string()).await; } } } - if dispatched > 0 { - info!("已成功将 {} 个待计算网格点推进任务队列", dispatched); - } - Ok(dispatched) } - /// Triggers seed_step fallback for a failed point if a seed is available - pub async fn trigger_seed_step_fallback(&self, params: &GridPointParams) -> Result { - if !self.db.has_running_workflow().await? { + /// Triggers seed_step fallback for a failed point if a seed is available. + /// + /// 多工作流分区(#3 修复):传入 `workflow_name` 明确该失败点所属工作流, + /// 用该工作流自身的 timeout / seed_step_fallback 配置,并把 TaskSpec.workflow_name + /// 绑定到该工作流。 + /// + /// 语义(种子回退仅一次): + /// - 仅当该工作流配置 `seed_step_fallback: true` 时才考虑回退; + /// - 仅当该点**尚未**派发过任何 seed_step 任务时才回退一次; + /// - 找不到合适近邻种子则不回退,由调用方保持 failed 终态。 + pub async fn trigger_seed_step_fallback( + &self, + params: &GridPointParams, + workflow_name: &str, + ) -> Result { + // 该工作流须仍处于 running 态才回退(避免 stop 后继续派发) + let still_running = self + .db + .get_running_workflow_names() + .await? + .iter() + .any(|w| w == workflow_name); + if !still_running { + return Ok(false); + } + + if !self.get_workflow_seed_step_fallback(workflow_name).await { return Ok(false); } let name = params.model_name(); - if let Ok(Some((status, attempt_count))) = self.db.get_grid_point_status(&name).await { - if status == "failed" || attempt_count >= 3 { - info!("网格点 {} 已达到最大重试次数 ({}) 或处于 failed 状态,跳过种子热启动回退", name, attempt_count); - return Ok(false); - } + // 种子回退仅一次:该点在该工作流中已经派发过 seed_step 任务就不再触发新的回退 + if self.db.has_seed_step_attempt(&name, workflow_name).await? { + info!( + "网格点 {} 已使用过一次种子热启动回退,不再重复回退,保持 failed 终态", + name + ); + return Ok(false); } + // seeds 全局共享:跨工作流复用已收敛的邻近种子 let seed_match_opt = self.db.find_best_seed_from_db(params).await.ok().flatten(); if let Some(seed_match) = seed_match_opt { - let timeout_sec = self.get_active_timeout_sec().await; + let timeout_sec = self.get_workflow_timeout_sec(workflow_name).await; let name = params.model_name(); let task_spec = TaskSpec { task_id: Uuid::new_v4(), @@ -181,22 +314,38 @@ impl GridScheduler { task_type: TaskType::SeedStep, seed_point_name: Some(seed_match.name.clone()), timeout_sec, + workflow_name: Some(workflow_name.to_string()), }; self.db.insert_task(&task_spec).await?; - self.db.update_grid_status(&name, common::models::GridPointStatus::Queued).await?; + self.db + .update_grid_status( + &name, + common::models::GridPointStatus::Queued, + workflow_name, + ) + .await?; if let Err(e) = self.queue.push_task(&task_spec).await { - let _ = self.db.update_grid_status(&name, common::models::GridPointStatus::Pending).await; + let _ = self + .db + .update_grid_status( + &name, + common::models::GridPointStatus::Pending, + workflow_name, + ) + .await; let _ = self.queue.remove_task(&task_spec.task_id.to_string()).await; return Err(e); } - info!("触发种子步进 (seed_step):网格点 {} 将使用 6 维近邻种子 {} 热启动重试", name, seed_match.name); + info!( + "触发种子步进 (seed_step):工作流 {} 网格点 {} 将使用 6 维近邻种子 {} 热启动重试", + workflow_name, name, seed_match.name + ); Ok(true) } else { Ok(false) } } - } #[cfg(test)] @@ -212,8 +361,16 @@ mod tests { let results_dir = temp_dir.path().join("results"); let db = Database::new(&db_path.to_string_lossy()).await.unwrap(); - let queue = Arc::new(SqliteTaskQueue::new(&queue_db_path.to_string_lossy()).await.unwrap()); - let scheduler = GridScheduler::new(db.clone(), queue.clone(), results_dir.to_string_lossy().to_string()); + let queue = Arc::new( + SqliteTaskQueue::new(&queue_db_path.to_string_lossy()) + .await + .unwrap(), + ); + let scheduler = GridScheduler::new( + db.clone(), + queue.clone(), + results_dir.to_string_lossy().to_string(), + ); let cfg = GridConfig { grid: GridAxesConfig { @@ -238,17 +395,141 @@ mod tests { linelist: None, }; - scheduler.initialize_grid(&cfg).await.unwrap(); - db.upsert_workflow("test_wf", None, "", "running").await.unwrap(); + scheduler.initialize_grid(&cfg, "test_wf").await.unwrap(); + db.upsert_workflow("test_wf", None, "", "running") + .await + .unwrap(); - let pending = db.get_pending_grid_points().await.unwrap(); + let pending = db.get_pending_grid_points("test_wf").await.unwrap(); assert_eq!(pending.len(), 1); let dispatched = scheduler.schedule_pending_tasks().await.unwrap(); assert_eq!(dispatched, 1); - let popped = queue.pop_task().await.unwrap(); + let popped = queue.pop_task("test-node").await.unwrap(); assert!(popped.is_some()); } -} + /// 多工作流分区调度测试(#3 修复验证): + /// 1. wf_a 调度推入队列的任务,在初始化 wf_b 后依然存在(initialize_grid 改用 + /// clear_queue_by_workflow,不再全局 clear_queue)。 + /// 2. 两个 running 工作流的 pending 点都能被 schedule_pending_tasks 派发。 + #[tokio::test] + async fn test_multi_workflow_dispatch_isolation() { + let temp_dir = tempfile::tempdir().unwrap(); + let db = Database::new(&temp_dir.path().join("mw_db.db").to_string_lossy()) + .await + .unwrap(); + let queue = Arc::new( + SqliteTaskQueue::new(&temp_dir.path().join("mw_queue.db").to_string_lossy()) + .await + .unwrap(), + ); + let scheduler = GridScheduler::new(db.clone(), queue.clone(), "results".to_string()); + + let mk_cfg = |teff: f64| GridConfig { + grid: GridAxesConfig { + teff: vec![teff], + logg: vec![5.5], + loghe: vec![-1.0], + logc: vec![-2.0], + logn: vec![-2.0], + logo: vec![-2.0], + }, + chain: vec![], + synspec: None, + nworkers: 4, + timeout_sec: 3600, + resume: true, + seed_step_fallback: true, + results: None, + itek_fallback: vec![], + niter: Some(100), + template: None, + fort55: None, + linelist: None, + }; + + // wf_a 初始化并推入队列 + scheduler + .initialize_grid(&mk_cfg(35000.0), "wf_a") + .await + .unwrap(); + db.upsert_workflow("wf_a", None, "", "running") + .await + .unwrap(); + let d_a = scheduler.schedule_pending_tasks().await.unwrap(); + assert_eq!(d_a, 1); + // 任务已在队 + assert!(queue.pop_task("node-a").await.unwrap().is_some()); + + // 重新推一个 wf_a 任务(上一行 pop 掉了),再初始化 wf_b + db.update_grid_status( + &GridPointParams { + teff: 35000.0, + logg: 5.5, + loghe: -1.0, + logc: -2.0, + logn: -2.0, + logo: -2.0, + } + .model_name(), + common::models::GridPointStatus::Pending, + "wf_a", + ) + .await + .unwrap(); + let _ = scheduler + .schedule_pending_tasks_for_workflow("wf_a", 100) + .await + .unwrap(); + // 此时 wf_a 队列里应有一个任务 + assert_eq!( + queue + .pop_task("node-a") + .await + .unwrap() + .and_then(|t| t.workflow_name), + Some("wf_a".to_string()) + ); + + // 关键断言:把 wf_a 任务重新推回队列后,初始化 wf_b 不应清空它。 + db.update_grid_status( + &GridPointParams { + teff: 35000.0, + logg: 5.5, + loghe: -1.0, + logc: -2.0, + logn: -2.0, + logo: -2.0, + } + .model_name(), + common::models::GridPointStatus::Pending, + "wf_a", + ) + .await + .unwrap(); + let _ = scheduler + .schedule_pending_tasks_for_workflow("wf_a", 100) + .await + .unwrap(); + + // 初始化 wf_b(内部 clear_queue_by_workflow("wf_b"),不该动 wf_a 的任务) + scheduler + .initialize_grid(&mk_cfg(40000.0), "wf_b") + .await + .unwrap(); + db.upsert_workflow("wf_b", None, "", "running") + .await + .unwrap(); + + // wf_a 的任务仍在队:可被 node 弹出,且 workflow_name == wf_a + let popped_a = queue.pop_task("node-a").await.unwrap(); + assert!(popped_a.is_some(), "初始化 wf_b 不应清空 wf_a 的队列任务"); + assert_eq!(popped_a.unwrap().workflow_name, Some("wf_a".to_string())); + + // wf_b 的点也能被调度(两个 running 工作流并存) + let d_b = scheduler.schedule_pending_tasks().await.unwrap(); + assert!(d_b >= 1, "wf_b 的 pending 点应被派发"); + } +} diff --git a/crates/server/tests/api_tests.rs b/crates/server/tests/api_tests.rs index 731d68d..05f8255 100644 --- a/crates/server/tests/api_tests.rs +++ b/crates/server/tests/api_tests.rs @@ -16,29 +16,65 @@ async fn test_server_api_flow() { std::fs::create_dir_all(&results_dir).unwrap(); let db = Database::new(&db_path.to_string_lossy()).await.unwrap(); - let queue = Arc::new(SqliteTaskQueue::new(&queue_db_path.to_string_lossy()).await.unwrap()); - let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone(), results_dir.to_string_lossy().to_string())); + let queue = Arc::new( + SqliteTaskQueue::new(&queue_db_path.to_string_lossy()) + .await + .unwrap(), + ); + let scheduler = Arc::new(GridScheduler::new( + db.clone(), + queue.clone(), + results_dir.to_string_lossy().to_string(), + )); let state = AppState { db, queue, scheduler, results_dir: results_dir.to_string_lossy().to_string(), + rate_limiter: server::api::rate_limit::RateLimiter::new( + 5, + std::time::Duration::from_secs(300), + ), auth_token: None, + admin_token: None, + auth_disabled: false, + admin_sessions: Arc::new(tokio::sync::RwLock::new(std::collections::HashMap::new())), }; let app = axum::Router::new() - .route("/api/node/register", axum::routing::post(server::api::node::register_node)) - .route("/api/node/heartbeat", axum::routing::post(server::api::node::heartbeat_node)) - .route("/api/task/claim", axum::routing::post(server::api::task::claim_task)) - .route("/api/status", axum::routing::get(server::api::status::get_status)) - .route("/api/workflows", axum::routing::get(server::api::workflow::list_workflows).post(server::api::workflow::save_workflow)) + .route( + "/api/node/register", + axum::routing::post(server::api::node::register_node), + ) + .route( + "/api/node/heartbeat", + axum::routing::post(server::api::node::heartbeat_node), + ) + .route( + "/api/task/claim", + axum::routing::post(server::api::task::claim_task), + ) + .route( + "/api/status", + axum::routing::get(server::api::status::get_status), + ) + .route( + "/api/workflows", + axum::routing::get(server::api::workflow::list_workflows) + .post(server::api::workflow::save_workflow), + ) .with_state(state); // 1. Check status API let response = app .clone() - .oneshot(Request::builder().uri("/api/status").body(Body::empty()).unwrap()) + .oneshot( + Request::builder() + .uri("/api/status") + .body(Body::empty()) + .unwrap(), + ) .await .unwrap(); @@ -96,21 +132,41 @@ async fn test_auth_middleware_scope_and_running_status() { std::fs::create_dir_all(&results_dir).unwrap(); let db = Database::new(&db_path.to_string_lossy()).await.unwrap(); - let queue = Arc::new(SqliteTaskQueue::new(&queue_db_path.to_string_lossy()).await.unwrap()); - let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone(), results_dir.to_string_lossy().to_string())); + let queue = Arc::new( + SqliteTaskQueue::new(&queue_db_path.to_string_lossy()) + .await + .unwrap(), + ); + let scheduler = Arc::new(GridScheduler::new( + db.clone(), + queue.clone(), + results_dir.to_string_lossy().to_string(), + )); let state = AppState { db: db.clone(), queue: queue.clone(), scheduler, results_dir: results_dir.to_string_lossy().to_string(), + rate_limiter: server::api::rate_limit::RateLimiter::new( + 5, + std::time::Duration::from_secs(300), + ), auth_token: Some("secret_token_123".to_string()), + admin_token: Some("secret_token_123".to_string()), + auth_disabled: false, + admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new( + std::collections::HashMap::new(), + )), }; - let api_router = axum::Router::new() - .route("/status", axum::routing::get(server::api::status::get_status)); + let api_router = axum::Router::new().route( + "/status", + axum::routing::get(server::api::status::get_status), + ); - let auth_layer = axum::middleware::from_fn_with_state(state.clone(), server::api::auth_middleware); + let auth_layer = + axum::middleware::from_fn_with_state(state.clone(), server::api::auth_middleware); let api_router = api_router.layer(auth_layer); let app = axum::Router::new() @@ -120,7 +176,12 @@ async fn test_auth_middleware_scope_and_running_status() { // Unauthenticated API request -> 401 Unauthorized let res = app .clone() - .oneshot(Request::builder().uri("/api/status").body(Body::empty()).unwrap()) + .oneshot( + Request::builder() + .uri("/api/status") + .body(Body::empty()) + .unwrap(), + ) .await .unwrap(); assert_eq!(res.status(), StatusCode::UNAUTHORIZED); @@ -148,9 +209,1106 @@ async fn test_auth_middleware_scope_and_running_status() { logn: -2.0, logo: -2.0, }; - db.upsert_grid_point(¶ms, 0).await.unwrap(); - db.mark_grid_point_running(¶ms.model_name()).await.unwrap(); + db.upsert_grid_point(¶ms, 0, "test_wf").await.unwrap(); + db.mark_grid_point_running(¶ms.model_name(), "test_wf") + .await + .unwrap(); - let stats = db.get_grid_summary_stats().await.unwrap(); + let stats = db.get_grid_summary_stats(None).await.unwrap(); assert_eq!(stats["running"], 1); } + +/// L2 鉴权核心流程测试: +/// 注册节点 → 颁发专属 token → 用 token 调 heartbeat(200)→ 吊销 → 再调(401)。 +#[tokio::test] +async fn test_l2_node_token_issue_revoke_flow() { + let temp_dir = tempfile::tempdir().unwrap(); + let db_path = temp_dir.path().join("l2_db.db"); + + let db = Database::new(&db_path.to_string_lossy()).await.unwrap(); + + // 1. 注册节点 + let reg = common::models::NodeRegisterRequest { + node_id: "node-l2-test".to_string(), + host_name: "l2-host".to_string(), + max_slots: 4, + }; + db.register_node(®).await.unwrap(); + + // 2. 颁发专属 token,返回明文 + let token = db.issue_node_token("node-l2-test").await.unwrap(); + assert!(!token.is_empty()); + + // 2b. 一次性取走暂存明文(take_pending_node_token):首次取到与颁发一致的明文, + // 再次取为 None(取走即焚)。验证 #4 简化为单一 UPDATE...RETURNING 后行为一致。 + let pending = db.take_pending_node_token("node-l2-test").await.unwrap(); + assert_eq!(pending.as_deref(), Some(token.as_str())); + let pending2 = db.take_pending_node_token("node-l2-test").await.unwrap(); + assert!( + pending2.is_none(), + "取走即焚:第二次 take_pending 必须返回 None" + ); + // 不存在的 node take 也应返回 None(不报错) + assert!(db + .take_pending_node_token("node-not-exist") + .await + .unwrap() + .is_none()); + + // 3. token 可反查到 node_id(此调用会把 token 写入内存缓存) + let found = db.find_node_by_token(&token).await; + assert_eq!(found.as_deref(), Some("node-l2-test")); + + // 4. 错误 token 查不到 + assert!(db.find_node_by_token("wrong-token").await.is_none()); + + // 5. 吊销后 token 立即失效 + db.revoke_node_token("node-l2-test").await.unwrap(); + assert!(db.find_node_by_token(&token).await.is_none()); + + // 6. 重新颁发后恢复(覆盖式) + let token2 = db.issue_node_token("node-l2-test").await.unwrap(); + assert_eq!( + db.find_node_by_token(&token2).await.as_deref(), + Some("node-l2-test") + ); + // 旧 token 仍失效(已被覆盖) + assert!(db.find_node_by_token(&token).await.is_none()); +} + +/// 验证中间件对 node token 的端到端鉴权: +/// 用真实 node token 调 /node/heartbeat 应 200;吊销后再调应 401。 +#[tokio::test] +async fn test_middleware_node_role_with_token() { + let temp_dir = tempfile::tempdir().unwrap(); + let db_path = temp_dir.path().join("l2_mw_db.db"); + let queue_db_path = temp_dir.path().join("l2_mw_queue.db"); + let results_dir = temp_dir.path().join("results"); + std::fs::create_dir_all(&results_dir).unwrap(); + + let db = Database::new(&db_path.to_string_lossy()).await.unwrap(); + let queue = Arc::new( + SqliteTaskQueue::new(&queue_db_path.to_string_lossy()) + .await + .unwrap(), + ); + let scheduler = Arc::new(GridScheduler::new( + db.clone(), + queue.clone(), + results_dir.to_string_lossy().to_string(), + )); + + // 注册并颁发 token + let reg = common::models::NodeRegisterRequest { + node_id: "node-mw-test".to_string(), + host_name: "mw-host".to_string(), + max_slots: 2, + }; + db.register_node(®).await.unwrap(); + let token = db.issue_node_token("node-mw-test").await.unwrap(); + + let state = AppState { + db: db.clone(), + queue, + scheduler, + results_dir: results_dir.to_string_lossy().to_string(), + rate_limiter: server::api::rate_limit::RateLimiter::new( + 5, + std::time::Duration::from_secs(300), + ), + auth_token: Some("placeholder".to_string()), + admin_token: None, + auth_disabled: false, + admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new( + std::collections::HashMap::new(), + )), + }; + + let api_router = axum::Router::new().route( + "/node/heartbeat", + axum::routing::post(server::api::node::heartbeat_node), + ); + let auth_layer = + axum::middleware::from_fn_with_state(state.clone(), server::api::auth_middleware); + let api_router = api_router.layer(auth_layer); + let app = axum::Router::new() + .nest("/api", api_router) + .with_state(state); + + // 用有效 node token 调 heartbeat → 200 + let hb_body = serde_json::json!({ + "node_id": "node-mw-test", "active_slots": 1, "cpu_usage": 10.0, "memory_usage": 20.0 + }); + let res = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/node/heartbeat") + .header("authorization", format!("Bearer {}", token)) + .header("content-type", "application/json") + .body(Body::from(serde_json::to_vec(&hb_body).unwrap())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::OK); + + // 吊销后同样请求 → 401 + db.revoke_node_token("node-mw-test").await.unwrap(); + let res = app + .oneshot( + Request::builder() + .method("POST") + .uri("/api/node/heartbeat") + .header("authorization", format!("Bearer {}", token)) + .header("content-type", "application/json") + .body(Body::from(serde_json::to_vec(&hb_body).unwrap())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::UNAUTHORIZED); +} + +/// 管理 API(Admin)端到端测试: +/// 列出节点 → 吊销 → 重发 token → 鉴权(admin 放行,node/匿名 401)。 +#[tokio::test] +async fn test_admin_node_management_api() { + let temp_dir = tempfile::tempdir().unwrap(); + let db_path = temp_dir.path().join("admin_db.db"); + let queue_db_path = temp_dir.path().join("admin_queue.db"); + let results_dir = temp_dir.path().join("results"); + std::fs::create_dir_all(&results_dir).unwrap(); + + let db = Database::new(&db_path.to_string_lossy()).await.unwrap(); + let queue = Arc::new( + SqliteTaskQueue::new(&queue_db_path.to_string_lossy()) + .await + .unwrap(), + ); + let scheduler = Arc::new(GridScheduler::new( + db.clone(), + queue.clone(), + results_dir.to_string_lossy().to_string(), + )); + + // 预置一个已注册并颁发 token 的节点 + let reg = common::models::NodeRegisterRequest { + node_id: "node-admin-test".to_string(), + host_name: "admin-host".to_string(), + max_slots: 2, + }; + db.register_node(®).await.unwrap(); + let original_token = db.issue_node_token("node-admin-test").await.unwrap(); + + let state = AppState { + db: db.clone(), + queue, + scheduler, + results_dir: results_dir.to_string_lossy().to_string(), + rate_limiter: server::api::rate_limit::RateLimiter::new( + 5, + std::time::Duration::from_secs(300), + ), + auth_token: Some("admin-secret".to_string()), + admin_token: Some("admin-secret".to_string()), + auth_disabled: false, + admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new( + std::collections::HashMap::new(), + )), + }; + + let api_router = axum::Router::new() + .route( + "/admin/nodes", + axum::routing::get(server::api::admin::list_nodes), + ) + .route( + "/admin/nodes/:node_id/revoke", + axum::routing::post(server::api::admin::revoke_node), + ) + .route( + "/admin/nodes/:node_id/reissue", + axum::routing::post(server::api::admin::reissue_node), + ); + let auth_layer = + axum::middleware::from_fn_with_state(state.clone(), server::api::auth_middleware); + let app = axum::Router::new() + .nest("/api", api_router.layer(auth_layer)) + .with_state(state); + + // 1. 匿名访问 /admin/nodes → 401 + let res = app + .clone() + .oneshot( + Request::builder() + .uri("/api/admin/nodes") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::UNAUTHORIZED); + + // 2. admin 访问 /admin/nodes → 200,且能看到预置节点 token_status=active + let res = app + .clone() + .oneshot( + Request::builder() + .uri("/api/admin/nodes") + .header("authorization", "Bearer admin-secret") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::OK); + let body: serde_json::Value = serde_json::from_slice( + &axum::body::to_bytes(res.into_body(), usize::MAX) + .await + .unwrap(), + ) + .unwrap(); + let nodes = body["data"].as_array().expect("data 应为数组"); + let target = nodes + .iter() + .find(|n| n["node_id"] == "node-admin-test") + .expect("应包含预置节点"); + assert_eq!(target["token_status"], "active"); + + // 3. admin 吊销节点 token → 200 + let res = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/admin/nodes/node-admin-test/revoke") + .header("authorization", "Bearer admin-secret") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::OK); + // 旧 token 立即失效 + assert!(db.find_node_by_token(&original_token).await.is_none()); + + // 4. admin 重发 token → 200,返回新明文 token + let res = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/admin/nodes/node-admin-test/reissue") + .header("authorization", "Bearer admin-secret") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::OK); + let body: serde_json::Value = serde_json::from_slice( + &axum::body::to_bytes(res.into_body(), usize::MAX) + .await + .unwrap(), + ) + .unwrap(); + let new_token = body["node_token"].as_str().expect("应返回 node_token"); + assert!(!new_token.is_empty()); + assert_ne!(new_token, original_token); + // 新 token 可用 + assert_eq!( + db.find_node_by_token(new_token).await.as_deref(), + Some("node-admin-test") + ); + + // 5. node token 不能访问 admin API → 401(即便持有有效 node token) + let res = app + .oneshot( + Request::builder() + .uri("/api/admin/nodes") + .header("authorization", format!("Bearer {}", new_token)) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::UNAUTHORIZED); +} + +/// S1 身份绑定测试:节点 A 用自己 token 冒充节点 B 发心跳 → 403;一致时 → 200。 +#[tokio::test] +async fn test_node_heartbeat_node_id_binding() { + let temp_dir = tempfile::tempdir().unwrap(); + let db_path = temp_dir.path().join("bind_db.db"); + let queue_db_path = temp_dir.path().join("bind_queue.db"); + let results_dir = temp_dir.path().join("results"); + std::fs::create_dir_all(&results_dir).unwrap(); + + let db = Database::new(&db_path.to_string_lossy()).await.unwrap(); + let queue = Arc::new( + SqliteTaskQueue::new(&queue_db_path.to_string_lossy()) + .await + .unwrap(), + ); + let scheduler = Arc::new(GridScheduler::new( + db.clone(), + queue, + results_dir.to_string_lossy().to_string(), + )); + + // 预置节点 A,颁发 token + let reg = common::models::NodeRegisterRequest { + node_id: "node-A".to_string(), + host_name: "h".to_string(), + max_slots: 2, + }; + db.register_node(®).await.unwrap(); + let token_a = db.issue_node_token("node-A").await.unwrap(); + + let state = AppState { + db: db.clone(), + queue: Arc::new(SqliteTaskQueue::new(":memory:").await.unwrap()), + scheduler, + results_dir: results_dir.to_string_lossy().to_string(), + rate_limiter: server::api::rate_limit::RateLimiter::new( + 5, + std::time::Duration::from_secs(300), + ), + auth_token: Some("x".into()), + admin_token: None, + auth_disabled: false, + admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new( + std::collections::HashMap::new(), + )), + }; + let api_router = axum::Router::new().route( + "/node/heartbeat", + axum::routing::post(server::api::node::heartbeat_node), + ); + let auth_layer = + axum::middleware::from_fn_with_state(state.clone(), server::api::auth_middleware); + let app = axum::Router::new() + .nest("/api", api_router.layer(auth_layer)) + .with_state(state); + + // A 用自己 token,但 body 声称 node_id=node-B(冒充)→ 403 + let hb = + serde_json::json!({"node_id":"node-B","active_slots":1,"cpu_usage":0.0,"memory_usage":0.0}); + let res = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/node/heartbeat") + .header("authorization", format!("Bearer {}", token_a)) + .header("content-type", "application/json") + .body(Body::from(serde_json::to_vec(&hb).unwrap())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::FORBIDDEN); + + // A 用自己 token,body node_id=node-A(一致)→ 200 + let hb_ok = + serde_json::json!({"node_id":"node-A","active_slots":1,"cpu_usage":0.0,"memory_usage":0.0}); + let res = app + .oneshot( + Request::builder() + .method("POST") + .uri("/api/node/heartbeat") + .header("authorization", format!("Bearer {}", token_a)) + .header("content-type", "application/json") + .body(Body::from(serde_json::to_vec(&hb_ok).unwrap())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::OK); +} + +#[tokio::test] +async fn test_admin_login_flow() { + let temp_dir = tempfile::tempdir().unwrap(); + let db_path = temp_dir.path().join("login_db.db"); + let queue_db_path = temp_dir.path().join("login_queue.db"); + let results_dir = temp_dir.path().join("results"); + std::fs::create_dir_all(&results_dir).unwrap(); + + let db = Database::new(&db_path.to_string_lossy()).await.unwrap(); + let queue = Arc::new( + SqliteTaskQueue::new(&queue_db_path.to_string_lossy()) + .await + .unwrap(), + ); + let scheduler = Arc::new(GridScheduler::new( + db.clone(), + queue.clone(), + results_dir.to_string_lossy().to_string(), + )); + + let admin_pass = "my_short_admin_password_123"; + let state = AppState { + db, + queue, + scheduler, + results_dir: results_dir.to_string_lossy().to_string(), + rate_limiter: server::api::rate_limit::RateLimiter::new( + 5, + std::time::Duration::from_secs(300), + ), + auth_token: Some(admin_pass.to_string()), + admin_token: Some(admin_pass.to_string()), + auth_disabled: false, + admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new( + std::collections::HashMap::new(), + )), + }; + + let api_router = axum::Router::new() + .route("/login", axum::routing::post(server::api::auth::login)) + .route( + "/auth/check", + axum::routing::get(server::api::auth::check_auth), + ) + .layer(axum::middleware::from_fn_with_state( + state.clone(), + server::api::auth_middleware, + )); + + let app = axum::Router::new() + .nest("/api", api_router) + .with_state(state); + + // 1. 密码错误 ➔ 401 + let wrong_body = serde_json::json!({ "password": "wrong_password" }); + let res = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/login") + .header("content-type", "application/json") + .extension(axum::extract::ConnectInfo(std::net::SocketAddr::from(( + [127, 0, 0, 1], + 12345, + )))) + .body(Body::from(serde_json::to_vec(&wrong_body).unwrap())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::UNAUTHORIZED); + + // 2. 正确密码 ➔ 200 + token + let right_body = serde_json::json!({ "password": admin_pass }); + let res = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/login") + .header("content-type", "application/json") + .extension(axum::extract::ConnectInfo(std::net::SocketAddr::from(( + [127, 0, 0, 1], + 12345, + )))) + .body(Body::from(serde_json::to_vec(&right_body).unwrap())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::OK); + let bytes = axum::body::to_bytes(res.into_body(), 1024 * 1024) + .await + .unwrap(); + let json: serde_json::Value = serde_json::from_slice(&bytes).unwrap(); + assert_eq!(json["success"], true); + let token = json["token"].as_str().unwrap(); + assert!(token.len() == 64); + + // 3. 携带拿到的 Token 访问受保护的 /api/auth/check ➔ 200 + let res = app + .oneshot( + Request::builder() + .method("GET") + .uri("/api/auth/check") + .header("authorization", format!("Bearer {}", token)) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::OK); +} + +#[tokio::test] +async fn test_node_approval_workflow() { + let temp_dir = tempfile::tempdir().unwrap(); + let db_path = temp_dir.path().join("appr_db.db"); + let queue_db_path = temp_dir.path().join("appr_queue.db"); + let results_dir = temp_dir.path().join("results"); + std::fs::create_dir_all(&results_dir).unwrap(); + + let db = Database::new(&db_path.to_string_lossy()).await.unwrap(); + let queue = Arc::new( + SqliteTaskQueue::new(&queue_db_path.to_string_lossy()) + .await + .unwrap(), + ); + let scheduler = Arc::new(GridScheduler::new( + db.clone(), + queue.clone(), + results_dir.to_string_lossy().to_string(), + )); + + let admin_pass = "admin_approval_secret"; + let state = AppState { + db, + queue, + scheduler, + results_dir: results_dir.to_string_lossy().to_string(), + rate_limiter: server::api::rate_limit::RateLimiter::new( + 5, + std::time::Duration::from_secs(300), + ), + auth_token: Some(admin_pass.to_string()), + admin_token: Some(admin_pass.to_string()), + auth_disabled: false, + admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new( + std::collections::HashMap::new(), + )), + }; + + let api_router = axum::Router::new() + .route( + "/node/register", + axum::routing::post(server::api::node::register_node), + ) + .route( + "/node/check_status", + axum::routing::post(server::api::node::check_node_status), + ) + .route( + "/node/heartbeat", + axum::routing::post(server::api::node::heartbeat_node), + ) + .route( + "/admin/nodes/:node_id/approve", + axum::routing::post(server::api::admin::approve_node), + ) + .layer(axum::middleware::from_fn_with_state( + state.clone(), + server::api::auth_middleware, + )); + + let app = axum::Router::new() + .nest("/api", api_router) + .with_state(state); + + // 1. 新节点免凭据申请注册 ➔ 200 + status: pending_approval + let reg_body = serde_json::json!({ + "node_id": "node-pending-01", + "host_name": "worker-host", + "max_slots": 4 + }); + let res = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/node/register") + .header("content-type", "application/json") + .body(Body::from(serde_json::to_vec(®_body).unwrap())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::OK); + let bytes = axum::body::to_bytes(res.into_body(), 1024 * 1024) + .await + .unwrap(); + let json: serde_json::Value = serde_json::from_slice(&bytes).unwrap(); + assert_eq!(json["status"], "pending_approval"); + + // 2. Node 端轮询查状态 ➔ status: pending_approval + let check_body = serde_json::json!({ "node_id": "node-pending-01" }); + let res = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/node/check_status") + .header("content-type", "application/json") + .body(Body::from(serde_json::to_vec(&check_body).unwrap())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::OK); + let bytes = axum::body::to_bytes(res.into_body(), 1024 * 1024) + .await + .unwrap(); + let json: serde_json::Value = serde_json::from_slice(&bytes).unwrap(); + assert_eq!(json["status"], "pending_approval"); + + // 3. 管理员在 Dashboard 点击同意 ➔ 200 + let res = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/admin/nodes/node-pending-01/approve") + .header("authorization", format!("Bearer {}", admin_pass)) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::OK); + let bytes = axum::body::to_bytes(res.into_body(), 1024 * 1024) + .await + .unwrap(); + println!("Step 3 body: {}", String::from_utf8_lossy(&bytes)); + + // 4. Node 端再次轮询查状态 ➔ status: approved + 获取 node_token + let res = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/node/check_status") + .header("content-type", "application/json") + .body(Body::from(serde_json::to_vec(&check_body).unwrap())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::OK); + let bytes = axum::body::to_bytes(res.into_body(), 1024 * 1024) + .await + .unwrap(); + let json: serde_json::Value = serde_json::from_slice(&bytes).unwrap(); + println!("Step 4 status returned: {:?}", json); + + assert_eq!(json["status"], "approved"); + let node_token = json["node_token"].as_str().unwrap(); + + // 5. Node 携带拿到到的专属 Token 发送心跳 ➔ 200 + let hb_body = serde_json::json!({ + "node_id": "node-pending-01", + "active_slots": 1, + "cpu_usage": 10.0, + "memory_usage": 20.0 + }); + let res = app + .oneshot( + Request::builder() + .method("POST") + .uri("/api/node/heartbeat") + .header("authorization", format!("Bearer {}", node_token)) + .header("content-type", "application/json") + .body(Body::from(serde_json::to_vec(&hb_body).unwrap())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::OK); +} + +#[tokio::test] +async fn test_admin_sessions_capacity_limit() { + let temp_dir = tempfile::tempdir().unwrap(); + let db_path = temp_dir.path().join("sess_db.db"); + let queue_db_path = temp_dir.path().join("sess_queue.db"); + let results_dir = temp_dir.path().join("results"); + std::fs::create_dir_all(&results_dir).unwrap(); + + let db = Database::new(&db_path.to_string_lossy()).await.unwrap(); + let queue = Arc::new( + SqliteTaskQueue::new(&queue_db_path.to_string_lossy()) + .await + .unwrap(), + ); + let scheduler = Arc::new(GridScheduler::new( + db.clone(), + queue.clone(), + results_dir.to_string_lossy().to_string(), + )); + + let admin_pass = "admin_capacity_secret"; + let state = AppState { + db, + queue, + scheduler, + results_dir: results_dir.to_string_lossy().to_string(), + rate_limiter: server::api::rate_limit::RateLimiter::new( + 1000, + std::time::Duration::from_secs(300), + ), + auth_token: Some(admin_pass.to_string()), + admin_token: Some(admin_pass.to_string()), + auth_disabled: false, + admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new( + std::collections::HashMap::new(), + )), + }; + + let api_router = + axum::Router::new().route("/login", axum::routing::post(server::api::auth::login)); + let app = axum::Router::new() + .nest("/api", api_router) + .with_state(state.clone()); + + // 连续登录 105 次,超过 100 容量上限 + for _ in 0..105 { + let right_body = serde_json::json!({ "password": admin_pass }); + let res = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/login") + .header("content-type", "application/json") + .extension(axum::extract::ConnectInfo(std::net::SocketAddr::from(( + [127, 0, 0, 1], + 12345, + )))) + .body(Body::from(serde_json::to_vec(&right_body).unwrap())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::OK); + } + + // 验证 sessions 集合保存的数量不超过 MAX_ADMIN_SESSIONS (100) + let sessions = state.admin_sessions.read().await; + assert!(sessions.len() <= server::api::MAX_ADMIN_SESSIONS); + assert_eq!(sessions.len(), 100); +} + +#[tokio::test] +async fn test_node_register_rate_limit() { + let temp_dir = tempfile::tempdir().unwrap(); + let db_path = temp_dir.path().join("reg_limit_db.db"); + let queue_db_path = temp_dir.path().join("reg_limit_queue.db"); + let results_dir = temp_dir.path().join("results"); + std::fs::create_dir_all(&results_dir).unwrap(); + + let db = Database::new(&db_path.to_string_lossy()).await.unwrap(); + let queue = Arc::new( + SqliteTaskQueue::new(&queue_db_path.to_string_lossy()) + .await + .unwrap(), + ); + let scheduler = Arc::new(GridScheduler::new( + db.clone(), + queue.clone(), + results_dir.to_string_lossy().to_string(), + )); + + let state = AppState { + db, + queue, + scheduler, + results_dir: results_dir.to_string_lossy().to_string(), + rate_limiter: server::api::rate_limit::RateLimiter::new( + 5, + std::time::Duration::from_secs(300), + ), + auth_token: None, + admin_token: None, + auth_disabled: false, + admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new( + std::collections::HashMap::new(), + )), + }; + + // 注册端点专用限流器(count_all=true):对成功请求也计数,模拟生产 main.rs 的注册节流配置。 + let register_limiter = + server::api::rate_limit::RateLimiter::new_count_all(5, std::time::Duration::from_secs(60)); + let register_rate_limit_layer = axum::middleware::from_fn_with_state( + register_limiter, + server::api::rate_limit::rate_limit_middleware, + ); + + let app = axum::Router::new() + .route( + "/api/node/register", + axum::routing::post(server::api::node::register_node).layer(register_rate_limit_layer), + ) + .with_state(state); + + let reg_body = serde_json::json!({ + "node_id": "test-limit-node", + "host_name": "limit-host", + "max_slots": 4 + }); + + // 5 次以内的注册尝试 ➔ 200 + for _ in 0..5 { + let res = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/node/register") + .header("content-type", "application/json") + .extension(axum::extract::ConnectInfo(std::net::SocketAddr::from(( + [127, 0, 0, 1], + 12345, + )))) + .body(Body::from(serde_json::to_vec(®_body).unwrap())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::OK); + } + + // 第 6 次触发限流 ➔ 429 TOO_MANY_REQUESTS + let res = app + .oneshot( + Request::builder() + .method("POST") + .uri("/api/node/register") + .header("content-type", "application/json") + .extension(axum::extract::ConnectInfo(std::net::SocketAddr::from(( + [127, 0, 0, 1], + 12345, + )))) + .body(Body::from(serde_json::to_vec(®_body).unwrap())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::TOO_MANY_REQUESTS); +} + +/// 回归保护:通用 API 限流器(count_all=false)不应因 /node/register 的成功响应计数。 +/// 历史缺陷:此前中间件对 is_register 无条件计数,导致通用限流器复用时成功注册会把 IP +/// 锁出整个 /api/*(跨端点连锁)。此测试断言多次成功注册后通用限流器仍放行。 +#[tokio::test] +async fn test_general_limiter_does_not_count_successful_register() { + let temp_dir = tempfile::tempdir().unwrap(); + let db_path = temp_dir.path().join("gen_limit_db.db"); + let queue_db_path = temp_dir.path().join("gen_limit_queue.db"); + let results_dir = temp_dir.path().join("results"); + std::fs::create_dir_all(&results_dir).unwrap(); + + let db = Database::new(&db_path.to_string_lossy()).await.unwrap(); + let queue = Arc::new( + SqliteTaskQueue::new(&queue_db_path.to_string_lossy()) + .await + .unwrap(), + ); + let scheduler = Arc::new(GridScheduler::new( + db.clone(), + queue.clone(), + results_dir.to_string_lossy().to_string(), + )); + + let state = AppState { + db, + queue, + scheduler, + results_dir: results_dir.to_string_lossy().to_string(), + rate_limiter: server::api::rate_limit::RateLimiter::new( + 5, + std::time::Duration::from_secs(300), + ), + auth_token: None, + admin_token: None, + auth_disabled: false, + admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new( + std::collections::HashMap::new(), + )), + }; + + // 通用限流器(new = count_all=false),阈值仅 3,远低于下面的请求数。 + // 若仍对成功注册计数,第 4 次就会 429。 + let general_limiter = + server::api::rate_limit::RateLimiter::new(3, std::time::Duration::from_secs(60)); + let rate_limit_layer = axum::middleware::from_fn_with_state( + general_limiter, + server::api::rate_limit::rate_limit_middleware, + ); + + let app = axum::Router::new() + .route( + "/api/node/register", + axum::routing::post(server::api::node::register_node), + ) + .layer(rate_limit_layer) + .with_state(state); + + let reg_body = serde_json::json!({ + "node_id": "gen-limit-node", + "host_name": "gen-host", + "max_slots": 4 + }); + + // 连续 6 次成功注册(远超阈值 3):通用限流器不应计数成功响应,全部应为 200。 + for i in 0..6 { + // 每次用不同 node_id 避免重复注册逻辑干扰 + let body = serde_json::json!({ + "node_id": format!("gen-limit-node-{}", i), + "host_name": "gen-host", + "max_slots": 4 + }); + let res = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/node/register") + .header("content-type", "application/json") + .extension(axum::extract::ConnectInfo(std::net::SocketAddr::from(( + [127, 0, 0, 1], + 12345, + )))) + .body(Body::from(serde_json::to_vec(&body).unwrap())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!( + res.status(), + StatusCode::OK, + "第 {} 次注册应成功,通用限流器不应计数成功响应", + i + 1 + ); + } + + // 静默 reg_body 未使用的警告(保留以对齐其它测试结构) + let _ = ®_body; +} + +#[tokio::test] +async fn test_cors_same_origin_and_local_policy() { + let temp_dir = tempfile::tempdir().unwrap(); + let db_path = temp_dir.path().join("cors_db.db"); + let queue_db_path = temp_dir.path().join("cors_queue.db"); + let results_dir = temp_dir.path().join("results"); + std::fs::create_dir_all(&results_dir).unwrap(); + + let db = Database::new(&db_path.to_string_lossy()).await.unwrap(); + let queue = Arc::new( + SqliteTaskQueue::new(&queue_db_path.to_string_lossy()) + .await + .unwrap(), + ); + let scheduler = Arc::new(GridScheduler::new( + db.clone(), + queue.clone(), + results_dir.to_string_lossy().to_string(), + )); + + let state = AppState { + db, + queue, + scheduler, + results_dir: results_dir.to_string_lossy().to_string(), + rate_limiter: server::api::rate_limit::RateLimiter::new( + 100, + std::time::Duration::from_secs(300), + ), + auth_token: None, + admin_token: None, + auth_disabled: false, + admin_sessions: Arc::new(tokio::sync::RwLock::new(std::collections::HashMap::new())), + }; + + let app = axum::Router::new() + .route( + "/api/status", + axum::routing::get(server::api::status::get_status), + ) + .layer(server::cors::build_cors_layer()) + .with_state(state); + + // 1. 本地 localhost 请求 -> 允许 CORS + let res = app + .clone() + .oneshot( + Request::builder() + .uri("/api/status") + .header("origin", "http://localhost:5173") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::OK); + assert_eq!( + res.headers() + .get("access-control-allow-origin") + .unwrap() + .to_str() + .unwrap(), + "http://localhost:5173" + ); + + // 2. 本地 127.0.0.1 请求 -> 允许 CORS + let res = app + .clone() + .oneshot( + Request::builder() + .uri("/api/status") + .header("origin", "http://127.0.0.1:3000") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::OK); + assert_eq!( + res.headers() + .get("access-control-allow-origin") + .unwrap() + .to_str() + .unwrap(), + "http://127.0.0.1:3000" + ); + + // 3. 同源请求 (Origin 匹配 Host 标头) -> 允许 CORS + let res = app + .clone() + .oneshot( + Request::builder() + .uri("/api/status") + .header("host", "192.168.1.100:8090") + .header("origin", "http://192.168.1.100:8090") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::OK); + assert_eq!( + res.headers() + .get("access-control-allow-origin") + .unwrap() + .to_str() + .unwrap(), + "http://192.168.1.100:8090" + ); + + // 4. 外部非法跨源请求 -> 拒绝 CORS + let res = app + .oneshot( + Request::builder() + .uri("/api/status") + .header("host", "192.168.1.100:8090") + .header("origin", "https://attacker.example.com") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(res.status(), StatusCode::OK); + assert!(res.headers().get("access-control-allow-origin").is_none()); +} diff --git a/dashboard/index.html b/dashboard/index.html index 5978fe2..e555708 100644 --- a/dashboard/index.html +++ b/dashboard/index.html @@ -11,12 +11,37 @@
+ + +
- - + + + + + + + + +
@@ -44,6 +69,14 @@ 刷新 +
@@ -87,7 +120,7 @@
- 已收敛网格模型 + 已计算网格模型
@@ -97,7 +130,7 @@
0 - 全网网格总数: 0 + 网格总数: 0
@@ -121,40 +154,44 @@
- +
-

计算节点集群 (Worker Nodes)

- 0 个节点 +

计算节点集群与凭据管理

+
+ 0 个节点 + +
- - - - - - - + + + + + + + + - +
节点 ID主机名CPU 槽位CPU 使用率内存使用率心跳时间状态节点 ID主机名CPU 槽位CPU/内存节点状态Token 凭据心跳时间管理操作
正在连接服务端获取计算节点...正在连接服务端获取计算节点...
- -
+ +
-

恒星大气网格工作流 (Workflows)

+

恒星大气网格工作流

@@ -162,40 +199,20 @@
- - -
-
-

大气与谱线数据资源 (Data Resources)

-
-
-

- 如需本地运行计算节点或检验谱线列表,可直接通过服务端 API 获取静态配分函数及原子文件: -

- -
-
-