feat(all): 科学计算正确性修复、调度竞态消除、节点优雅退出、安全加固与收敛分析重构

科学计算正确性:
- 修复 Fortran 无-E 科学记数法(指数≥100 时 E 被挤掉,如 -1.35+118)导致
  发散行被静默跳过、误判收敛的 bug;扩展大气无效检测覆盖 Inf 与 *** 溢出标记
- 种子匹配改为 CNO 有向距离(富金属方向重罚 4×、贫金属方向轻罚 1×),
  基于 1191 个真实种子配对回测标定,回测净改善 314 个点
- GridAxisValue 反序列化拒绝非法文本(不再静默 NaN);chmax≤0 显式报错
- ions 行宽列宽对齐真实 fort.5 格式

调度与队列竞态:
- 原子选点(IMMEDIATE 事务 SELECT+UPDATE)消除并发调度重复派发 (#5)
- 调度互斥锁 + 冷启动优先策略(SeedStep 仅作失败后救援,不再正常路径热启动)
- 毒消息 dead_letter 标记防出队死循环;clear_queue 保留 claimed 行 (#6)
- 孤儿 running 点回收兜底;stale_sec 默认 7800→21600s(3 倍超时缓冲)

节点生命周期:
- SIGTERM+SIGINT 双信号监听(修复 Docker stop 发 SIGTERM 不触发优雅退出)
- SlotGuard RAII 防活动 slot 泄漏;子进程超时增加二级 30s wait 防 Fortran hang
- SeedStep 种子下载 fail-fast + 沙盒私有副本解耦 LRU 清理竞争
- reqwest Client 增加连接/请求超时;启动清理残留 task_* 沙盒

安全加固:
- 节点注册 registration_secret 二次凭据 (H8),恒定时间比对防时序旁路
- token 缓存 generation 机制消除 reissue 后旧 token TOCTOU 复活窗口
- 新增 /api/auth/logout 服务端 session 即时撤销;fail-closed 鉴权启动策略
- 前端 token 迁移 sessionStorage;YAML 高亮改 DOM API 消除 XSS 注入面
- 备份文件权限收紧 0600;点表动态值全面 escapeHtml

前端 Dashboard:
- 收敛性分析从热力图重构为 Parallel Sets 平行集合图(6 维+状态轴,手写 SVG 零依赖)
- 进度曲线横轴改为真实时间(服务端 now 锚定,停滞期诚实留白);轮询指数退避
- 移除 imported 收敛途径分类,导入点按实际途径 cold_run/seed_step 归类
- 初始化时 /api/auth/check 校验 token;401 toast 提示替代静默 reload

服务端恢复与工具链:
- 启动恢复 initializing 态工作流;默认工作流 INSERT-only 不覆盖 API 编辑
- body limit 分层(10MB 不再截断 256MB report);multipart 显式错误处理
- 嵌入二进制原子写(tmp+rename)防半写损坏
- import_results 判定收敛途径透传 success_method;conv.json 格式对齐本项目
- push_import_results.sh 退出码修复 + .bat UTF-8 BOM + scp 上传
- Docker USE_MIRRORS 默认关闭;移除无用 assets 挂载;删除 hosts.ini 入库
This commit is contained in:
fmq
2026-08-01 17:01:40 +08:00
parent 1bfa240cb0
commit c8fd24b120
45 changed files with 2821 additions and 831 deletions
+52 -1
View File
@@ -3,11 +3,12 @@
//! 提供基于短密码的身份认证服务:
//! - POST /api/login:校验管理员密码,成功后返回 Admin Token,并记录 IP 错误次数防止暴力破解。
//! - GET /api/auth/check:由 auth_middleware 保护,供前端初始化时检测当前保存的 Token 是否有效。
//! - POST /api/auth/logout:撤销当前 session token(服务端立即失效),供前端登出调用。
use super::{ct_eq_str, AppState};
use axum::{
extract::{ConnectInfo, State},
http::StatusCode,
http::{HeaderMap, StatusCode},
response::IntoResponse,
Json,
};
@@ -119,3 +120,53 @@ pub async fn check_auth() -> impl IntoResponse {
})),
)
}
/// POST /api/auth/logout — 撤销当前 session token。
///
/// 由 auth_middlewareRole::Admin)校验通过后到达,从请求头取出 token(复用与中间件
/// 一致的 `extract_token_from_headers`,同时支持 Authorization: Bearer 与 X-API-Key、
/// 拒绝空值)并从 `admin_sessions` 中移除,使该 token 在服务端立即失效(而非等 24h 过期)。
/// 这样即便 token 已被窃取,登出操作也能立即阻断重放。
pub async fn logout(
State(state): State<AppState>,
headers: HeaderMap,
) -> impl IntoResponse {
// 与 auth_middleware 口径一致地提取 token(支持 X-API-Key、拒绝空值)。
let token = crate::api::extract_token_from_headers(&headers);
let removed = if let Some(t) = token {
let mut sessions = state.admin_sessions.write().await;
// 与中间件一致:遍历全部 session key 做恒定时间比对(不提前 break),消除 key
// 存在性/位置的时序旁路。命中后记录 key、遍历完成后再 remove。容量受
// MAX_ADMIN_SESSIONS 约束,遍历开销可接受。
let mut target: Option<String> = None;
for k in sessions.keys() {
if ct_eq_str(&t, k) {
target = Some(k.clone());
// 不 break:继续遍历以保持恒定时间
}
}
if let Some(k) = target {
sessions.remove(&k);
true
} else {
false
}
} else {
false
};
if removed {
info!("Admin session 已登出撤销(服务端立即失效)");
} else {
warn!("登出请求未匹配到有效 session(可能为 admin_token 主凭据或已失效)");
}
(
StatusCode::OK,
Json(serde_json::json!({
"success": true,
"message": "已登出"
})),
)
}
+18 -11
View File
@@ -85,6 +85,10 @@ fn required_role(path: &str, method: &axum::http::Method) -> Option<Role> {
if path == "/auth/check" && method == Method::GET {
return Some(Role::Admin);
}
// 登出(撤销当前 session-> Admin
if path == "/auth/logout" && method == Method::POST {
return Some(Role::Admin);
}
// 写操作 → Admin
if path == "/workflows" && (method == Method::POST || method == Method::GET) {
return Some(Role::Admin);
@@ -96,14 +100,12 @@ fn required_role(path: &str, method: &axum::http::Method) -> Option<Role> {
if path == "/status" && method == Method::GET {
return Some(Role::Admin);
}
// 管理 API(节点凭据查看/审批/重发/停用/启用)→ Admin
// 管理 API(节点凭据查看/审批/重发/停用/启用、历史种子导入等)→ Admin
// 注:所有 /admin/* 均需 Admin 鉴权(含 /admin/import_seed),统一在此判定即可,
// 无需为单个子路径重复列举(避免出现被前缀匹配遮蔽的不可达分支)。
if path.starts_with("/admin/") {
return Some(Role::Admin);
}
// 历史种子导入(run_grid.py 旧产物回灌)→ Admin
if path == "/admin/import_seed" && method == Method::POST {
return Some(Role::Admin);
}
// Node 运行态 → Node
if path == "/node/heartbeat" && method == Method::POST {
return Some(Role::Node);
@@ -151,13 +153,13 @@ pub(crate) fn is_valid_node_id(id: &str) -> bool {
.all(|c| c.is_ascii_alphanumeric() || c == '.' || c == '_' || c == '-')
}
/// 从请求头提取凭据原文(支持 `Authorization: Bearer <t>` 与 `X-API-Key: <t>`)。
/// 从 HeaderMap 提取凭据原文(支持 `Authorization: Bearer <t>` 与 `X-API-Key: <t>`)。
///
/// 安全:非 `Bearer ` 前缀的 Authorization 一律视为无 token(不再回退为裸头值比较),
/// 避免 `Authorization: Basic ...` 之类的上游代理头被误送入 token 比对。
fn extract_token(req: &Request<axum::body::Body>) -> Option<String> {
if let Some(auth) = req
.headers()
/// 避免 `Authorization: Basic ...` 之类的上游代理头被误送入 token 比对。空值也不作为 token。
/// 公开供 logout handler 等需要从头部取 token 的场景复用,保证口径与 auth_middleware 一致。
pub(crate) fn extract_token_from_headers(headers: &axum::http::HeaderMap) -> Option<String> {
if let Some(auth) = headers
.get(header::AUTHORIZATION)
.and_then(|v| v.to_str().ok())
{
@@ -168,7 +170,7 @@ fn extract_token(req: &Request<axum::body::Body>) -> Option<String> {
}
// 非 Bearer 前缀或空值:不作为 token
}
if let Some(key) = req.headers().get("x-api-key").and_then(|v| v.to_str().ok()) {
if let Some(key) = headers.get("x-api-key").and_then(|v| v.to_str().ok()) {
if !key.is_empty() {
return Some(key.to_string());
}
@@ -176,6 +178,11 @@ fn extract_token(req: &Request<axum::body::Body>) -> Option<String> {
None
}
/// 从请求头提取凭据原文(中间件路径)。
fn extract_token(req: &Request<axum::body::Body>) -> Option<String> {
extract_token_from_headers(req.headers())
}
/// Axum 鉴权中间件(L2)。
///
/// 流程:
+49 -9
View File
@@ -39,7 +39,7 @@ pub async fn register_node(
// 申请注册新节点(免凭据提交申请,进入 pending_approval 状态)
match state.db.register_node(&req).await {
Ok(true) => {
Ok((true, _, registration_secret)) => {
info!(
"接收到新节点 {} 的注册申请,已加入待审批 (pending_approval) 队列",
req.node_id
@@ -50,17 +50,41 @@ pub async fn register_node(
"status": "pending_approval",
"message": "节点注册申请已成功提交!请在管理 Dashboard 控制台上点击【同意接入】授权该节点",
"node_token": null,
// H8:下发一次性 registration_secret,节点须在 /node/check_status 取 token 时回传,
// 防止仅知道 node_id 的攻击者抢先取走待发 token。
"registration_secret": registration_secret,
})),
))
}
Ok(false) => {
// 节点已处于待审批或已存在列表
Ok((false, existing_status, _)) => {
// 节点已存在:按其真实状态如实响应,避免误导运维。
// 旧实现一律回 "pending_approval",导致已 online 的节点重新注册时被告知"等待审批"。
let (status, message) = match existing_status.as_deref() {
Some("online") => (
"approved",
"节点已授权(online),配置已更新。如需新 token 请联系管理员重发".to_string(),
),
Some("disabled") => (
"disabled",
"节点已被管理员停用,配置已更新。请联系管理员重新启用".to_string(),
),
// pending_approval 或其他未知态:仍处于待审批
_ => (
"pending_approval",
"节点注册申请等待管理员审批中".to_string(),
),
};
info!(
"已存在节点 {} 重新注册(状态: {:?}),配置已更新",
req.node_id, existing_status
);
Ok((
StatusCode::OK,
Json(json!({
"status": "pending_approval",
"message": "节点注册申请等待管理员审批中",
"status": status,
"message": message,
"node_token": null,
"registration_secret": null,
})),
))
}
@@ -71,6 +95,9 @@ pub async fn register_node(
#[derive(serde::Deserialize)]
pub struct CheckNodeStatusRequest {
pub node_id: String,
/// H8:节点注册时下发的一次性凭据,取走待发 token 前须校验。
/// 旧版节点未持有此凭据时不传,服务端对无 registration_secret 记录的旧节点保持兼容。
pub registration_secret: Option<String>,
}
/// POST /api/node/check_status — Node 端轮询检查审批结果。
@@ -84,8 +111,12 @@ pub async fn check_node_status(
));
}
// 尝试拉取取走即焚的暂存明文 Token
match state.db.take_pending_node_token(&req.node_id).await {
// 尝试拉取取走即焚的暂存明文 Token(内部校验 registration_secret
match state
.db
.take_pending_node_token(&req.node_id, req.registration_secret.as_deref())
.await
{
Ok(Some(raw_token)) => {
info!(
"节点 {} 的注册申请已被管理员审批同意,下发专属 Token",
@@ -100,8 +131,9 @@ pub async fn check_node_status(
})),
))
}
Ok(None) | Err(_) => {
// 查节点表状态
Ok(None) => {
// 可能原因:尚未审批 / registration_secret 不匹配 / token 已被取走。
// 查节点表状态以区分"待审批"与"未通过",避免暴露 secret 校验失败的具体原因。
match state.db.get_node_exists(&req.node_id).await {
Ok(true) => Ok((
StatusCode::OK,
@@ -121,6 +153,14 @@ pub async fn check_node_status(
)),
}
}
Err(_) => Ok((
StatusCode::OK,
Json(json!({
"status": "rejected",
"message": "节点注册申请未通过或已被移除",
"node_token": null,
})),
)),
}
}
+1 -1
View File
@@ -25,7 +25,7 @@ pub async fn get_status(
.await
.unwrap_or(serde_json::json!({
"total": 0, "pending": 0, "queued": 0, "running": 0, "converged": 0, "failed": 0,
"cold_run_converged": 0, "seed_step_converged": 0, "imported_converged": 0
"cold_run_converged": 0, "seed_step_converged": 0
}));
Ok(Json(json!({
+112 -22
View File
@@ -64,21 +64,54 @@ pub async fn report_task(
) -> Result<impl IntoResponse, crate::api::AppError> {
let mut report_json: Option<TaskReport> = None;
let mut seed_file_data: Option<Vec<u8>> = None;
let mut multipart_error = false;
while let Ok(Some(field)) = multipart.next_field().await {
let field_name = field.name().unwrap_or("").to_string();
if field_name == "report" {
if let Ok(bytes) = field.bytes().await {
if let Ok(report) = serde_json::from_slice::<TaskReport>(&bytes) {
report_json = Some(report);
// 遍历全部 multipart 字段。旧实现 `while let Ok(Some(field))` 在首个字段读取错误时
// 静默停止迭代,可能丢失后续的 seed_file/report 字段,产生半截请求被当成完整请求处理。
// 现在显式记录字段错误并在出现错误时拒绝该请求(multipart 流一旦出错无法继续可靠读取)。
loop {
match multipart.next_field().await {
Ok(Some(field)) => {
let field_name = field.name().unwrap_or("").to_string();
if field_name == "report" {
match field.bytes().await {
Ok(bytes) => {
if let Ok(report) = serde_json::from_slice::<TaskReport>(&bytes) {
report_json = Some(report);
} else {
warn!("report 字段 JSON 解析失败,忽略该字段");
}
}
Err(e) => {
warn!("读取 report 字段失败: {}", e);
multipart_error = true;
}
}
} else if field_name == "seed_file" {
match field.bytes().await {
Ok(bytes) => {
seed_file_data = Some(bytes.to_vec());
}
Err(e) => {
warn!("读取 seed_file 字段失败: {}", e);
multipart_error = true;
}
}
}
}
} else if field_name == "seed_file" {
if let Ok(bytes) = field.bytes().await {
seed_file_data = Some(bytes.to_vec());
Ok(None) => break,
Err(e) => {
warn!("解析 multipart 字段时出错: {}", e);
multipart_error = true;
break;
}
}
}
if multipart_error {
return Err(crate::api::AppError::BadRequest(
"multipart 请求体解析不完整(字段读取失败)".to_string(),
));
}
let mut report = match report_json {
Some(r) => r,
@@ -192,9 +225,15 @@ pub async fn report_task(
name,
seed_path.display()
);
// 用权威 point_name(即 name = report.point_name,源精度正确)入种子库,
// 而非从 params 重推 model_name()。原因:此处的 params 经 node 上报回传,
// 其 logg 等轴在服务端 DB REAL 列回读时已丢精度(5.0→"5"),重推会得到
// 降级名(g5 而非 g5.0),导致 seeds.point_name 与磁盘文件名错配,
// 后续 download_seed 拼 seeds_dir/<degraded>/<degraded>.7 → 404,种子复用失效。
// 文件已用权威 name 写入(上方 seed_path),DB 必须同名对齐。
let _ = state
.db
.insert_seed(&params, &seed_path.to_string_lossy())
.insert_seed_named(&name, &params, &seed_path.to_string_lossy())
.await;
}
}
@@ -263,19 +302,62 @@ pub async fn import_seed(
) -> Result<impl IntoResponse, crate::api::AppError> {
let mut summary_json: Option<String> = None;
let mut seed_file_data: Option<Vec<u8>> = None;
// 收敛途径(cold_run/seed_step):由 import_results 工具依据旧 conv.json 的 stages 是否
// 含 seed_nc 判定后透传。缺失或非法时兜底 cold_run(容错旧版工具 / 防注入)。
let mut success_method: Option<String> = None;
let mut multipart_error = false;
while let Ok(Some(field)) = multipart.next_field().await {
let field_name = field.name().unwrap_or("").to_string();
if field_name == "report" {
if let Ok(bytes) = field.bytes().await {
summary_json = Some(String::from_utf8_lossy(&bytes).to_string());
// 显式遍历全部字段,记录读取错误。旧实现 `while let Ok(...)` 在首字段出错时静默停止,
// 可能丢失后续 seed_file/report 字段导致半截请求被处理。
loop {
match multipart.next_field().await {
Ok(Some(field)) => {
let field_name = field.name().unwrap_or("").to_string();
if field_name == "report" {
match field.bytes().await {
Ok(bytes) => {
summary_json = Some(String::from_utf8_lossy(&bytes).to_string());
}
Err(e) => {
warn!("历史种子导入:读取 report 字段失败: {}", e);
multipart_error = true;
}
}
} else if field_name == "seed_file" {
match field.bytes().await {
Ok(bytes) => {
seed_file_data = Some(bytes.to_vec());
}
Err(e) => {
warn!("历史种子导入:读取 seed_file 字段失败: {}", e);
multipart_error = true;
}
}
} else if field_name == "success_method" {
match field.text().await {
Ok(text) => {
success_method = Some(text);
}
Err(e) => {
warn!("历史种子导入:读取 success_method 字段失败: {}", e);
multipart_error = true;
}
}
}
}
} else if field_name == "seed_file" {
if let Ok(bytes) = field.bytes().await {
seed_file_data = Some(bytes.to_vec());
Ok(None) => break,
Err(e) => {
warn!("历史种子导入:解析 multipart 字段时出错: {}", e);
multipart_error = true;
break;
}
}
}
if multipart_error {
return Err(crate::api::AppError::BadRequest(
"multipart 请求体解析不完整(字段读取失败)".to_string(),
));
}
let summary_json = match summary_json {
Some(s) => s,
@@ -367,12 +449,18 @@ pub async fn import_seed(
}
}
// 4. 更新 grid_points 状态:收敛→converged(success_method='imported');否则维持 pending
// 4. 更新 grid_points 状态:收敛→converged(success_method=工具判定的途径);否则维持 pending
// 让正常调度处理(导入未收敛点无意义,但记录其尝试)。
// 途径缺失或非法时兜底 cold_run(容错旧版工具 / 防注入),由 db 层再次白名单校验。
if converged {
// 归一化:仅接受 cold_run / seed_step,其余(含 None)一律兜底 cold_run。
let method = match success_method.as_deref() {
Some("seed_step") => "seed_step",
_ => "cold_run",
};
if let Err(e) = state
.db
.mark_grid_point_imported(&name, &workflow_name, Some(summary.elapsed_sec))
.mark_grid_point_imported(&name, &workflow_name, Some(summary.elapsed_sec), method)
.await
{
warn!("历史种子导入:标记 {} 为 converged 失败: {}", name, e);
@@ -380,8 +468,10 @@ pub async fn import_seed(
}
info!(
"历史种子导入完成:网格点 {} (workflow={}, converged={}, max_relc={:?})",
name, workflow_name, converged, max_relc
"历史种子导入完成:网格点 {} (workflow={}, converged={}, success_method={}, max_relc={:?})",
name, workflow_name, converged,
success_method.as_deref().unwrap_or("(default cold_run)"),
max_relc
);
Ok((
+15 -3
View File
@@ -303,6 +303,7 @@ pub struct ProgressQuery {
/// 工作流进度时间序列 + 经验速率 + 停滞时长(详情页概览进度曲线数据源)。
///
/// - `series`:窗口内的计数快照(超 300 点自动降采样,首末点保留);
/// - `now`:服务端当前 UTC 时刻(同 ts 格式),供前端把曲线右缘锚定为“现在”、横轴按真实时间铺开;
/// - `rate_per_hour`:窗口首末 converged 增量 ÷ 时长(快照 <2 条或时长 ≤0 为 null);
/// - `stalled_minutes`:终态数(converged+failed)最后一次增长到窗口末端的分钟数
/// (用于"进度停滞"预警;快照 <2 条为 null)。
@@ -381,6 +382,13 @@ pub async fn get_workflow_progress(
None
};
// 服务端当前 UTC 时刻(与快照 ts 同格式)。前端用它锚定曲线右缘 = “现在”,
// 使横轴是真实时间线:停滞期(无快照)会诚实显示为空白,而非被索引均分抹平。
let now = chrono::Utc::now()
.naive_utc()
.format("%Y-%m-%d %H:%M:%S")
.to_string();
Ok((
StatusCode::OK,
Json(ApiResponse {
@@ -388,6 +396,7 @@ pub async fn get_workflow_progress(
message: "成功获取进度时间序列".to_string(),
data: Some(serde_json::json!({
"hours": hours,
"now": now,
"series": series,
"rate_per_hour": rate_per_hour,
"stalled_minutes": stalled_minutes,
@@ -411,7 +420,8 @@ pub struct PointsQuery {
}
/// 工作流逐点列表:点参数 + 状态 + 收敛手段 + 最近尝试(max_relc/种子来源/节点/错误)。
/// 支持状态/手段/波次过滤、点名搜索、白名单排序与分页limit ≤ 500)。详情页点表数据源。
/// 支持状态/手段/波次过滤、点名搜索、白名单排序与分页limit 缺省时返回全量(联合分析
/// 需完整数据,截断会让结论失真);点表明细分页传显式 limit(≤500)。详情页数据源。
pub async fn get_workflow_points(
State(state): State<AppState>,
AxumPath(name): AxumPath<String>,
@@ -446,7 +456,7 @@ pub async fn get_workflow_points(
}
}
if let Some(m) = &pq.method {
if !matches!(m.as_str(), "cold_run" | "seed_step" | "imported") {
if !matches!(m.as_str(), "cold_run" | "seed_step") {
return Err(crate::api::AppError::BadRequest(format!(
"非法的 method 参数: {}",
m
@@ -488,7 +498,9 @@ pub async fn get_workflow_points(
wave: pq.wave,
q: pq.q.clone(),
order_by,
limit: pq.limit.unwrap_or(100).clamp(1, 500),
// limit 缺省 → NoneSQL 不拼 LIMIT,返回全量)。联合分析依赖完整数据,
// 故不设上限;点表明细分页传显式 limit(≤500)走分页。
limit: pq.limit.map(|l| l.max(1)),
offset: pq.offset.unwrap_or(0).max(0),
};
+509 -115
View File
@@ -19,6 +19,23 @@ fn hash_token(token: &str) -> String {
hex::encode(hasher.finalize())
}
/// 恒定时间比对两个非空字符串(先 SHA-256 摘要再比较等长摘要,消除长度时序旁路)。
/// 用于 registration_secret 校验,避免通过比对耗时探得 secret 前缀。
fn ct_eq_option(a: &str, b: &str) -> bool {
use subtle::ConstantTimeEq;
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()
}
/// 多工作流分区迁移:把旧版 grid_points 表(仅 name UNIQUE,无 workflow_name 列)
/// 重建为带 workflow_name 列、(workflow_name, name) 复合唯一的新结构。
///
@@ -131,10 +148,15 @@ pub struct SeedCacheItem {
/// 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
/// 把这三个轴量化到桶。由于 exact 要求「双侧」严格小于阈值(target 和候选两侧都可能在
/// 量化边界两侧),对每个轴都做 **floor / floor+1 双桶** 写入与查询,确保跨越量化边界的
/// 真实 exact 候选必被覆盖:
/// - teff 按 5000K 量化为整数(floor),查 floor 与 floor+1 两桶覆盖 [floor*5000, (floor+2)*5000)
/// - logg/loghe 按 0.01 精度量化(×100 后 floor),查 floor 与 floor+1。
///
/// 旧实现仅对 teff 双写,logg/loghe 用 round 单桶,导致 d_logg<0.01 但 *100 round 落在相邻
/// 整数的两点(如 5.004→500 vs 5.005→501)被分到不同桶、永不相遇——exact 候选被静默丢失,
/// 且全局回退扫描又显式跳过 exact 候选,无法补救。改为三个轴一致地双写双查后消除该边界 Bug。
///
/// 命中桶后仍在桶内做精确 distance 计算取最优,故量化只用于缩小候选集,不影响正确性。
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
@@ -145,24 +167,25 @@ struct SeedBucketKey {
}
impl SeedBucketKey {
fn from_params(params: &GridPointParams) -> [Self; 2] {
// 返回该参数应落入的两个桶(floor 与 floor+1),供插入时双写、查询时双查
// teff 量化:floor(teff/5000)。如 teff=35000 → 7teff=37499 → 7teff=37500 → 8。
/// 返回该参数应落入的全部桶键(每个轴的 floor 与 floor+1 笛卡尔积,共 8 个)。
/// 插入时对每个键写入,查询时对每个键查询,确保跨量化边界的 exact 候选必命中
fn from_params(params: &GridPointParams) -> Vec<Self> {
let teff_floor = (params.teff.value() / 5000.0).floor() as i64;
let logg_q = (params.logg.value() * 100.0).round() as i64;
let loghe_q = (params.loghe.value() * 100.0).round() as i64;
[
SeedBucketKey {
teff_bucket: teff_floor,
logg_q,
loghe_q,
},
SeedBucketKey {
teff_bucket: teff_floor + 1,
logg_q,
loghe_q,
},
]
let logg_floor = (params.logg.value() * 100.0).floor() as i64;
let loghe_floor = (params.loghe.value() * 100.0).floor() as i64;
let mut keys = Vec::with_capacity(8);
for dt in [0, 1] {
for dg in [0, 1] {
for dh in [0, 1] {
keys.push(SeedBucketKey {
teff_bucket: teff_floor + dt,
logg_q: logg_floor + dg,
loghe_q: loghe_floor + dh,
});
}
}
}
keys
}
}
@@ -175,17 +198,38 @@ pub struct Database {
seed_index: std::sync::Arc<
tokio::sync::RwLock<std::collections::HashMap<SeedBucketKey, Vec<SeedCacheItem>>>,
>,
/// node token 反查缓存:token_hash → (node_id, 插入时间)。
/// node token 反查缓存:token_hash → (node_id, 插入时间, 回填时的缓存 generation)。
/// 鉴权中间件每个 Node 请求都查 find_node_by_token,此缓存把高频心跳/领用请求
/// 的 DB 查询降为内存读。TTL 由 `TOKEN_CACHE_TTL` 控制;issue(重发)时整体失效。
///
/// cache_generation 是单调递增的"失效代次":每次 invalidate_token_cache 自增。
/// find_node_by_token 在 DB 查询前记录当前 generation,回填时若 generation 已变化
/// (说明期间发生过 reissue 导致的 invalidate),则丢弃本次回填,彻底消除
/// "旧 token_hash 复活"的 TOCTOU 窗口(旧实现仅缩小窗口、未消除)。
token_cache: std::sync::Arc<
tokio::sync::RwLock<std::collections::HashMap<String, (String, std::time::Instant)>>,
tokio::sync::RwLock<TokenCache>,
>,
}
/// token 反查缓存的单条存活时长(秒)。issue(重发)会立即整体失效,TTL 仅兜底。
const TOKEN_CACHE_TTL: std::time::Duration = std::time::Duration::from_secs(60);
/// token 反查缓存的内部结构:entries 表 + 单调递增的失效代次。
struct TokenCache {
entries: std::collections::HashMap<String, (String, std::time::Instant)>,
/// 每次 invalidate_token_cache 自增;find_node_by_token 回填时据此判断是否发生过失效。
generation: u64,
}
impl TokenCache {
fn new() -> Self {
Self {
entries: std::collections::HashMap::new(),
generation: 0,
}
}
}
impl Database {
pub async fn new(db_path: &str) -> Result<Self> {
let db_path_owned = db_path.to_string();
@@ -216,7 +260,7 @@ impl Database {
std::collections::HashMap::new(),
)),
token_cache: std::sync::Arc::new(tokio::sync::RwLock::new(
std::collections::HashMap::new(),
TokenCache::new(),
)),
};
db.init_tables().await?;
@@ -237,10 +281,19 @@ impl Database {
status TEXT NOT NULL DEFAULT 'online',
cpu_usage REAL NOT NULL DEFAULT 0.0,
memory_usage REAL NOT NULL DEFAULT 0.0,
last_heartbeat DATETIME NOT NULL
last_heartbeat DATETIME NOT NULL,
registration_secret TEXT
);",
[],
)?;
// 旧库迁移:为 nodes 表补 registration_secret 列(H8check_status 取 token 需此凭据)。
let has_reg_secret = conn
.prepare("PRAGMA table_info(nodes)")?
.query_map([], |r| r.get::<_, String>(1))?
.any(|r| r.map(|n| n == "registration_secret").unwrap_or(false));
if !has_reg_secret {
let _ = conn.execute("ALTER TABLE nodes ADD COLUMN registration_secret TEXT", []);
}
conn.execute(
"CREATE TABLE IF NOT EXISTS grid_points (
@@ -426,38 +479,56 @@ impl Database {
}
// --- Node operations ---
pub async fn register_node(&self, req: &NodeRegisterRequest) -> Result<bool> {
/// 注册/刷新节点。返回 (is_new, existing_status)
/// - 新申请:`(true, None)`
/// - 已存在(含 online 等已审批态):`(false, Some(<旧状态>))`,仅更新配置保持既有状态。
///
/// 返回旧状态供 API 层区分响应:已审批(online)的节点免凭据重新注册时,
/// 不应回 "pending_approval"(误导运维以为还需审批),而应如实告知其已是已授权节点。
pub async fn register_node(
&self,
req: &NodeRegisterRequest,
) -> Result<(bool, Option<String>, Option<String>)> {
let pool = self.pool.clone();
let req_cloned = req.clone();
let is_new = tokio::task::spawn_blocking(move || -> Result<bool> {
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<String> = stmt.query_row(params![req_cloned.node_id], |r| r.get(0)).ok();
let (is_new, existing_status, registration_secret) =
tokio::task::spawn_blocking(move || -> Result<(bool, Option<String>, Option<String>)> {
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<String> =
stmt.query_row(params![req_cloned.node_id], |r| r.get(0)).ok();
match existing_status {
Some(_st) => {
// 已存在的节点:更新配置,保持既有状态
conn.execute(
"UPDATE nodes SET max_slots = ?1, last_heartbeat = datetime('now') WHERE node_id = ?2",
params![req_cloned.max_slots, req_cloned.node_id],
)?;
Ok(false)
match &existing_status {
Some(_st) => {
// 已存在的节点:更新配置,保持既有状态
conn.execute(
"UPDATE nodes SET max_slots = ?1, last_heartbeat = datetime('now') WHERE node_id = ?2",
params![req_cloned.max_slots, req_cloned.node_id],
)?;
Ok((false, existing_status, None))
}
None => {
// 新申请节点:生成一次性 registration_secretH8)并插入待审批状态。
// registration_secret 用于 /node/check_status 取走专属 token 的二次凭据,
// 防止知道 node_id(常源自主机名,可猜测)的攻击者抢先取走待发 token。
let secret = format!(
"{}{}",
uuid::Uuid::new_v4().simple(),
uuid::Uuid::new_v4().simple()
);
conn.execute(
"INSERT INTO nodes (node_id, max_slots, status, last_heartbeat, registration_secret)
VALUES (?1, ?2, 'pending_approval', datetime('now'), ?3)",
params![req_cloned.node_id, req_cloned.max_slots, secret],
)?;
Ok((true, None, Some(secret)))
}
}
None => {
// 新申请节点:插入待审批状态 (pending_approval)
conn.execute(
"INSERT INTO nodes (node_id, max_slots, status, last_heartbeat)
VALUES (?1, ?2, 'pending_approval', datetime('now'))",
params![req_cloned.node_id, req_cloned.max_slots],
)?;
Ok(true)
}
}
})
.await??;
})
.await??;
Ok(is_new)
Ok((is_new, existing_status, registration_secret))
}
/// 管理员审批同意节点接入:将节点状态切为 online 并生成专属 node_token(返回明文 token)。
@@ -634,15 +705,45 @@ impl Database {
///
/// 注:SQLite 的 `UPDATE ... RETURNING` 返回的是列的**新值**(SET 之后),故清空后
/// RETURNING 该列只会得到 NULL,无法用于读旧值;因此这里用显式 SELECT + UPDATE。
pub async fn take_pending_node_token(&self, node_id: &str) -> Result<Option<String>> {
pub async fn take_pending_node_token(
&self,
node_id: &str,
registration_secret: Option<&str>,
) -> Result<Option<String>> {
let pool = self.pool.clone();
let node_id_owned = node_id.to_string();
let secret_owned = registration_secret.map(|s| s.to_string());
let token = tokio::task::spawn_blocking(move || -> Result<Option<String>> {
let mut conn = pool
.get()
.map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?;
let tx = conn.transaction_with_behavior(rusqlite::TransactionBehavior::Immediate)?;
// H8:取走待发 token 前校验 registration_secret(节点注册时下发的一次性凭据)。
// 仅当 nodes 表记录的 registration_secret 与请求提供的一致(恒定时间比对),
// 才允许取走 token,防止仅知道 node_id(可猜测)的攻击者抢先取走。
let stored_secret: Option<String> = {
let mut secret_stmt = tx.prepare(
"SELECT registration_secret FROM nodes WHERE node_id = ?1",
)?;
secret_stmt
.query_row(params![node_id_owned], |r| r.get::<_, Option<String>>(0))
.ok()
.flatten()
};
let secret_ok = match (&stored_secret, &secret_owned) {
(Some(a), Some(b)) => ct_eq_option(a, b),
// 旧库节点(无 registration_secret)不强制要求,保持向后兼容;
// 新节点(有 secret)必须提供正确 secret。
(None, _) => true,
(Some(_), None) => false,
};
if !secret_ok {
tx.commit()?;
return Ok(None);
}
let raw_token: Option<String> = {
let mut select_stmt = tx.prepare(
"SELECT raw_token_pending FROM node_credentials
@@ -675,23 +776,32 @@ impl Database {
///
/// 失效语义:重发(issue_node_token)会用 ON CONFLICT 覆盖该 node 的 token_hash
/// 旧 token 明文 hash 不再存在于表 → 查询返回 None → 401。无需独立的 revoked 标记。
///
/// 撤销竞态修复:历史上存在 TOCTOU 窗口——线程 A 用旧 token miss 落 DB 查到 node_id
/// 后准备回填,期间线程 B(管理员 reissue)覆盖 DB 的 token_hash 并 clear() 缓存,
/// 随后线程 A 拿到写锁把旧 token_hash 回填进缓存,导致已撤销的旧 token 在 TTL(60s)
/// 内仍能鉴权。修复:回填时在同一把写锁内重新校验该 token_hash 是否仍是 DB 当前值
/// (未被 reissue 覆盖),是才回填,杜绝旧 token 复活窗口。
pub async fn find_node_by_token(&self, token: &str) -> Option<String> {
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 let Some((node_id, inserted)) = cache.entries.get(&token_hash) {
if inserted.elapsed() < TOKEN_CACHE_TTL {
return Some(node_id.clone());
}
}
}
// 2) miss 落 DB
// 2) 记录 DB 查询前的缓存 generationmiss 落 DB
// generation 用于回填时的 TOCTOU 终极防护:若 DB 查询与回填之间发生过
// invalidatereissue),generation 会变化,本次回填将被丢弃。
let gen_before = { self.token_cache.read().await.generation };
let pool = self.pool.clone();
let hash_for_db = token_hash.clone();
let node_id: Option<String> =
let db_hit: Option<String> =
tokio::task::spawn_blocking(move || -> Result<Option<String>> {
let conn = pool
.get()
@@ -711,17 +821,31 @@ impl Database {
.and_then(|r| r.ok())
.flatten();
// 3) 命中则回填缓存(None 不缓存,避免失效态被短暂缓存)
if let Some(id) = &node_id {
// 3) 命中则回填缓存;回填前校验 generation 未变化(期间无 invalidate),
// 彻底消除"旧 token_hash 复活"窗口。generation 变化则视为已撤销,不缓存、不返回。
if let Some(id) = db_hit {
let mut cache = self.token_cache.write().await;
cache.insert(token_hash, (id.clone(), std::time::Instant::now()));
if cache.generation == gen_before {
cache
.entries
.insert(token_hash, (id.clone(), std::time::Instant::now()));
Some(id)
} else {
// 期间发生过 reissue 导致的 invalidate:旧 token_hash 已不应复活。
None
}
} else {
None
}
node_id
}
/// 清空全部 token 反查缓存。在 issuetoken 轮换使旧 token 失效)时调用。
/// 清空全部 token 反查缓存并自增 generation。在 issuetoken 轮换使旧 token 失效)时调用。
/// 自增 generation 使所有在途的 find_node_by_token 回填(gen_before 已过期)被丢弃,
/// 彻底消除"DB 读取旧 hash → reissue clear → 回填旧 hash"的 TOCTOU 复活窗口。
async fn invalidate_token_cache(&self) {
self.token_cache.write().await.clear();
let mut cache = self.token_cache.write().await;
cache.entries.clear();
cache.generation = cache.generation.wrapping_add(1);
}
/// 判断指定 node_id 是否已存在于 nodes 表(重发 token 前置校验,防幽灵 node_id)。
@@ -936,6 +1060,92 @@ impl Database {
.await?
}
/// 原子选点:在 IMMEDIATE 事务内将 pending 点标记为 queued 并返回。
///
/// 解决 `get_pending_grid_points_limit`SELECT)与 `update_grid_status`UPDATE
/// 分离导致的 TOCTOU 竞态:两个并发调度调用可能 SELECT 到同一批 pending 点,
/// 各自创建任务,产生重复派发(#5 修复)。
///
/// 与 `pop_task`sqlite_queue.rs)和 `take_pending_node_token` 同口径:
/// IMMEDIATE 事务在 BEGIN 时即获取写锁,SELECT 与 UPDATE 之间不会被其它
/// 调用方插入,从而只有一个调用方能 claiming 到某批点。
pub async fn claim_pending_grid_points(
&self,
limit: usize,
workflow_name: &str,
) -> Result<Vec<(String, GridPointParams, i32)>> {
let pool = self.pool.clone();
let wf = workflow_name.to_string();
tokio::task::spawn_blocking(move || -> Result<Vec<(String, GridPointParams, i32)>> {
let mut conn = pool
.get()
.map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?;
let tx =
conn.transaction_with_behavior(rusqlite::TransactionBehavior::Immediate)?;
let limit_param = if limit == usize::MAX {
-1i64
} else {
limit as i64
};
let mut stmt = tx.prepare(
"UPDATE grid_points
SET status = 'queued'
WHERE rowid IN (
SELECT rowid FROM grid_points
WHERE status = 'pending' AND workflow_name = ?1
ORDER BY wave ASC, cno_sum ASC, teff ASC
LIMIT ?2
)
RETURNING name, teff, logg, loghe, logc, logn, logo, wave",
)?;
let rows_iter = stmt.query_map(params![wf, limit_param], |r| {
Ok((
r.get::<_, String>(0)?,
GridPointParams {
teff: GridAxisValue::from_value(r.get::<_, f64>(1)?),
logg: GridAxisValue::from_value(r.get::<_, f64>(2)?),
loghe: GridAxisValue::from_value(r.get::<_, f64>(3)?),
logc: GridAxisValue::from_value(r.get::<_, f64>(4)?),
logn: GridAxisValue::from_value(r.get::<_, f64>(5)?),
logo: GridAxisValue::from_value(r.get::<_, f64>(6)?),
},
r.get::<_, i32>(7)?,
))
})?;
let mut list = Vec::new();
for r in rows_iter {
list.push(r?);
}
drop(stmt);
tx.commit()?;
// UPDATE...RETURNING 不保证行序(子查询 ORDER BY 仅决定 LIMIT 选取),
// 在 Rust 侧按调度优先级排序,保持与原 get_pending_grid_points_limit 同序。
list.sort_by(|a, b| {
a.2.cmp(&b.2) // wave ASC
.then_with(|| {
a.1.cno_sum()
.partial_cmp(&b.1.cno_sum())
.unwrap_or(std::cmp::Ordering::Equal)
})
.then_with(|| {
a.1.teff
.partial_cmp(&b.1.teff)
.unwrap_or(std::cmp::Ordering::Equal)
})
});
Ok(list)
})
.await?
}
/// 重置指定工作流的 queued 点为 pending(系统重启/工作流启动时使用)。
/// 按 workflow 隔离,避免误伤其他工作流(多工作流分区修复点)。
pub async fn reset_queued_grid_points_to_pending(&self, workflow_name: &str) -> Result<usize> {
@@ -982,6 +1192,38 @@ impl Database {
.await?
}
/// 回收孤儿 running 网格点(#6 修复兜底)。
///
/// 场景:网格点处于 `running` 态,但其任务在 `tasks` 表中仍为 `pending`
/// (从未被上报),且创建时间已超过 stale_sec。这意味着领用凭证(queue 行)
/// 已不存在(被误删、server 崩溃丢队列等),节点无法上报,点永远卡在 running。
///
/// 与 `requeue_stale_tasks` 互补:后者处理 queue 行仍存在但 claimed 超时的情况;
/// 本方法处理 queue 行已消失、requeue 找不到的情况。
///
/// 返回被重置的点数。
pub async fn reset_orphaned_running_points(&self, stale_sec: u64) -> Result<usize> {
let pool = self.pool.clone();
tokio::task::spawn_blocking(move || -> Result<usize> {
let conn = pool.get().map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?;
let stale_offset = format!("-{} seconds", stale_sec);
let count = conn.execute(
"UPDATE grid_points SET status = 'pending'
WHERE status = 'running'
AND name IN (
SELECT point_name FROM tasks
WHERE tasks.workflow_name = grid_points.workflow_name
AND tasks.point_name = grid_points.name
AND tasks.status = 'pending'
AND tasks.created_at < datetime('now', ?1)
)",
params![stale_offset],
)?;
Ok(count)
})
.await?
}
/// 更新指定工作流内某点的状态。按 workflow 隔离,防跨工作流误改同名点。
pub async fn update_grid_status(
&self,
@@ -1014,19 +1256,28 @@ impl Database {
.await
}
/// 历史种子导入专用:把网格点标记为 converged 并记录 success_method='imported'
/// 历史种子导入专用:把网格点标记为 converged 并记录收敛途径 `success_method`
///
/// 与正常 `record_task_report` 路径的区别:导入不走 task 队列,无 task_type 可取,
/// 故 success_method 固定为 'imported' 以区分正常计算收敛与历史回灌。
/// 故由导入工具(import_results)依据旧 conv.json 的 stages 是否含 seed_nc 判定该点
/// 当初是冷启动收敛(cold_run)还是种子步进收敛(seed_step),经 multipart 字段透传至此。
/// 导入点因此融入冷启动/种子步进统计,而非独立为 imported 分类。
///
/// `elapsed_sec`:旧版 conv.json 的单点墙钟耗时(`summary.elapsed_sec`),落入
/// `last_elapsed_sec` 列使迁移点在详情页/点表保留真实耗时;无此数据传 None。
/// - `success_method`:须为 "cold_run" 或 "seed_step",非法值兜底为 "cold_run"(防注入)。
/// - `elapsed_sec`:旧版 conv.json 的单点墙钟耗时(`summary.elapsed_sec`),落入
/// `last_elapsed_sec` 列使迁移点在详情页/点表保留真实耗时;无此数据传 None。
pub async fn mark_grid_point_imported(
&self,
name: &str,
workflow_name: &str,
elapsed_sec: Option<f64>,
success_method: &str,
) -> Result<()> {
// 白名单校验:仅接受两种合法途径,非法值兜底 cold_run(避免拼接 SQL 注入风险)。
let method = match success_method {
"seed_step" => "seed_step",
_ => "cold_run",
};
let pool = self.pool.clone();
let name_owned = name.to_string();
let wf = workflow_name.to_string();
@@ -1035,9 +1286,9 @@ impl Database {
.get()
.map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?;
conn.execute(
"UPDATE grid_points SET status = 'converged', success_method = 'imported', last_elapsed_sec = ?1 \
WHERE name = ?2 AND workflow_name = ?3",
params![elapsed_sec, name_owned, wf],
"UPDATE grid_points SET status = 'converged', success_method = ?1, last_elapsed_sec = ?2 \
WHERE name = ?3 AND workflow_name = ?4",
params![method, elapsed_sec, name_owned, wf],
)?;
Ok(())
})
@@ -1116,6 +1367,23 @@ impl Database {
Ok(())
}
/// 删除 tasks 历史表中指定 task_id 的行。
///
/// 用于调度回滚:当 push_task 失败时,grid_points 已回滚、queue 已清理,
/// 但先于 push 插入的 tasks 历史行(status='pending')会遗留,污染每点尝试计数统计。
/// 此方法在回滚路径中调用以保持三者一致。
pub async fn delete_task(&self, task_id: &uuid::Uuid) -> Result<()> {
let pool = self.pool.clone();
let id = task_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 tasks WHERE task_id = ?1", params![id])?;
Ok(())
})
.await??;
Ok(())
}
pub async fn record_task_report(&self, report: &TaskReport, workflow_name: &str) -> Result<()> {
let pool = self.pool.clone();
let report_cloned = report.clone();
@@ -1675,6 +1943,32 @@ impl Database {
.await?
}
/// 获取所有处于 `initializing` 态的工作流 (name, config_yaml)。
///
/// 用于服务端启动恢复:`start_workflow` 把状态切到 `initializing` 后在后台 spawn
/// `initialize_grid`。若进程在初始化中途崩溃/重启,工作流会永久卡在 `initializing`
/// `get_running_workflow_names` 仍把它算作可调度,但无人完成网格展开)。
/// 启动时检测到这些半初始化工作流后重新跑 `initialize_grid`(幂等,ON CONFLICT DO NOTHING
/// 把状态推进到 `running`,避免半初始化网格被调度。
pub async fn get_initializing_workflows(&self) -> Result<Vec<(String, String)>> {
let pool = self.pool.clone();
tokio::task::spawn_blocking(move || -> Result<Vec<(String, String)>> {
let conn = pool
.get()
.map_err(|e| anyhow::anyhow!("DB Pool Error: {}", e))?;
let mut stmt = conn.prepare(
"SELECT name, config_yaml FROM workflows WHERE status = 'initializing'",
)?;
let rows = stmt.query_map([], |row| Ok((row.get(0)?, row.get(1)?)))?;
let mut list = Vec::new();
for r in rows {
list.push(r?);
}
Ok(list)
})
.await?
}
/// 网格汇总统计。
///
/// `workflow_filter`
@@ -1682,8 +1976,8 @@ impl Database {
/// - `Some(wf)`:仅聚合指定工作流(按工作流隔离的进度统计)。
///
/// 口径说明:`pending` 与 `queued` **分开**计数(详情页需要区分"未入队"与"排队中");
/// 旧版前端若需合并口径,自行相加(见 dashboard state.js)。`imported_converged` 统计
/// 历史导入收敛点(success_method='imported'),与 cold_run/seed_step 并列互斥
/// 旧版前端若需合并口径,自行相加(见 dashboard state.js)。导入的历史点按其实际
/// 收敛途径(cold_run/seed_step)归类,与正常计算点一并统计——不再有独立 imported 分类
pub async fn get_grid_summary_stats(
&self,
workflow_filter: Option<&str>,
@@ -1704,13 +1998,12 @@ impl Database {
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,
SUM(CASE WHEN status = 'converged' AND success_method = 'imported' THEN 1 ELSE 0 END) AS imported_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<i64>>(i).unwrap_or(None).unwrap_or(0) };
Ok((n(0), n(1), n(2), n(3), n(4), n(5), n(6), n(7), n(8)))
Ok((n(0), n(1), n(2), n(3), n(4), n(5), n(6), n(7)))
},
),
None => conn.query_row(
@@ -1722,13 +2015,12 @@ impl Database {
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,
SUM(CASE WHEN status = 'converged' AND success_method = 'imported' THEN 1 ELSE 0 END) AS imported_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<i64>>(i).unwrap_or(None).unwrap_or(0) };
Ok((n(0), n(1), n(2), n(3), n(4), n(5), n(6), n(7), n(8)))
Ok((n(0), n(1), n(2), n(3), n(4), n(5), n(6), n(7)))
},
),
}?;
@@ -1741,7 +2033,6 @@ impl Database {
failed,
cold_run_converged,
seed_step_converged,
imported_converged,
) = row;
Ok(serde_json::json!({
@@ -1753,7 +2044,6 @@ impl Database {
"failed": failed,
"cold_run_converged": cold_run_converged,
"seed_step_converged": seed_step_converged,
"imported_converged": imported_converged,
}))
})
.await?
@@ -1848,7 +2138,6 @@ impl Database {
failed,
cold_run_converged: g("cold_run_converged"),
seed_step_converged: g("seed_step_converged"),
imported_converged: g("imported_converged"),
waves,
avg_point_sec,
eta_sec,
@@ -1908,8 +2197,24 @@ impl Database {
)?;
// 数据行:LEFT JOIN 最近一次尝试(从未派发则 last_* 全 NULL
let limit_idx = binds.len() + 1;
let offset_idx = binds.len() + 2;
let mut all_binds = binds;
// limit=None 时不拼 LIMIT 子句(联合分析需全量,截断会让分析失真)。
let limit_clause = match f.limit {
Some(lim) => {
let limit_idx = all_binds.len() + 1;
all_binds.push(Box::new(lim));
format!(" LIMIT ?{}", limit_idx)
}
None => String::new(),
};
// offset 仅在有 limit 或非零时才有意义;None-limit 全量场景强制忽略 offset。
let offset_clause = if f.limit.is_some() {
let offset_idx = all_binds.len() + 1;
all_binds.push(Box::new(f.offset));
format!(" OFFSET ?{}", offset_idx)
} else {
String::new()
};
let sql = format!(
"SELECT gp.name, gp.teff, gp.logg, gp.loghe, gp.logc, gp.logn, gp.logo,
gp.cno_sum, gp.wave, gp.status, gp.success_method, gp.attempt_count,
@@ -1922,14 +2227,10 @@ impl Database {
ORDER BY t2.completed_at IS NULL, t2.completed_at DESC, t2.created_at DESC
LIMIT 1)
WHERE {}
ORDER BY {}
LIMIT ?{} OFFSET ?{}",
where_sql, f.order_by, limit_idx, offset_idx
ORDER BY {}{}{}",
where_sql, f.order_by, limit_clause, offset_clause
);
let mut stmt = conn.prepare(&sql)?;
let mut all_binds = binds;
all_binds.push(Box::new(f.limit));
all_binds.push(Box::new(f.offset));
let rows = stmt.query_map(
rusqlite::params_from_iter(all_binds.iter().map(|b| b.as_ref())),
point_row_from_query,
@@ -2133,23 +2434,31 @@ impl Database {
let now = std::time::SystemTime::now();
let seven_days = std::time::Duration::from_secs(7 * 24 * 3600);
// 保留期清理:只删除本程序写出的 `dcts_backup_*.db` 文件中超过 7 天的。
// 旧实现会删除 backup_dir 内任何超期 *.db(含运维放置的无关 .db 文件)。
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);
}
}
}
if !file_path.is_file() {
continue;
}
// 仅匹配自身产物命名 dcts_backup_<时间戳>.db,避免误删无关 .db
let matches_name = file_path
.file_name()
.and_then(|n| n.to_str())
.map(|n| n.starts_with("dcts_backup_") && n.ends_with(".db"))
.unwrap_or(false);
if !matches_name {
continue;
}
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);
}
}
}
@@ -2167,6 +2476,19 @@ impl Database {
params![backup_file.to_string_lossy().to_string()],
)?;
// 收紧备份文件权限为 0600(仅 owner 读写),与主库口径一致。
// 备份通过 VACUUM INTO 直接由 SQLite 写出,默认沿用 umask(可能 0644),
// 可被同机其他用户读取;备份含 node 凭据 hash 与(瞬时)明文待发 token。
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
if let Ok(meta) = std::fs::metadata(&backup_file) {
let mut perms = meta.permissions();
perms.set_mode(0o600);
let _ = std::fs::set_permissions(&backup_file, perms);
}
}
Ok(())
})
.await??;
@@ -2248,7 +2570,6 @@ pub struct WorkflowStats {
pub failed: i64,
pub cold_run_converged: i64,
pub seed_step_converged: i64,
pub imported_converged: i64,
pub waves: Vec<WaveStats>,
/// 近似单点平均墙钟耗时(秒,含排队等待,仅参考);无历史数据为 None。
pub avg_point_sec: Option<f64>,
@@ -2324,7 +2645,8 @@ pub struct PointFilter {
pub q: Option<String>,
/// 编译期列名 + ASC/DESC 拼成的 ORDER BY 片段(不含用户文本)。
pub order_by: String,
pub limit: i64,
/// None = 不限制(联合分析拉全量,避免截断导致分析失真)。
pub limit: Option<i64>,
pub offset: i64,
}
@@ -2563,7 +2885,7 @@ mod tests {
}
/// grid_summary_stats 按 workflow 聚合 + 全局聚合测试。
/// 覆盖 pending/queued 拆分计数与 imported_converged 口径。
/// 覆盖 pending/queued 拆分计数与导入点按途径(cold_run/seed_step)归类的口径。
#[tokio::test]
async fn test_grid_summary_stats_workflow_scoping() {
let temp_dir = tempfile::tempdir().unwrap();
@@ -2590,29 +2912,29 @@ mod tests {
db.upsert_grid_point(&p, 0, "wf_b").await.unwrap();
db.upsert_grid_point(&p2, 1, "wf_a").await.unwrap();
// wf_b 的点入队(queued),wf_a 的 p2 走历史导入收敛(imported
// wf_b 的点入队(queued),wf_a 的 p2 走历史导入收敛(标记为 seed_step 途径
db.update_grid_status(&p.model_name(), GridPointStatus::Queued, "wf_b")
.await
.unwrap();
db.mark_grid_point_imported(&p2.model_name(), "wf_a", None)
db.mark_grid_point_imported(&p2.model_name(), "wf_a", None, "seed_step")
.await
.unwrap();
// 全局(None):3 个点,pending/queued/imported 分开计数
// 全局(None):3 个点,pending/queued/converged 分开计数
// 导入点按 seed_step 途径计入 seed_step_converged(不再有独立 imported 分类)
let all = db.get_grid_summary_stats(None).await.unwrap();
assert_eq!(all["total"], 3);
assert_eq!(all["pending"], 1);
assert_eq!(all["queued"], 1);
assert_eq!(all["converged"], 1);
assert_eq!(all["imported_converged"], 1);
assert_eq!(all["cold_run_converged"], 0);
assert_eq!(all["seed_step_converged"], 0);
// 单工作流 wf_a1 pending + 1 imported
assert_eq!(all["seed_step_converged"], 1);
// 单工作流 wf_a1 pending + 1 seed_step 收敛
let a = db.get_grid_summary_stats(Some("wf_a")).await.unwrap();
assert_eq!(a["total"], 2);
assert_eq!(a["pending"], 1);
assert_eq!(a["queued"], 0);
assert_eq!(a["imported_converged"], 1);
assert_eq!(a["seed_step_converged"], 1);
// 不存在的工作流:0
let none = db
.get_grid_summary_stats(Some("nonexistent"))
@@ -2717,23 +3039,30 @@ mod tests {
node_id: "node-atomic-test".to_string(),
max_slots: 2,
};
db.register_node(&reg).await.unwrap();
let (_, _, secret) = db.register_node(&reg).await.unwrap();
let token = db.approve_node("node-atomic-test").await.unwrap();
assert!(!token.is_empty());
// 第一次调用:返回 token
// 第一次调用:提供正确 registration_secret返回 token
let pending1 = db
.take_pending_node_token("node-atomic-test")
.take_pending_node_token("node-atomic-test", secret.as_deref())
.await
.unwrap();
assert_eq!(pending1, Some(token));
// 第二次调用:已被置为 NULL,返回 None
let pending2 = db
.take_pending_node_token("node-atomic-test")
.take_pending_node_token("node-atomic-test", secret.as_deref())
.await
.unwrap();
assert_eq!(pending2, None);
// 错误的 registration_secret:不应返回 tokenH8 防护)
let pending3 = db
.take_pending_node_token("node-atomic-test", Some("wrong-secret"))
.await
.unwrap();
assert_eq!(pending3, None);
}
#[tokio::test]
@@ -2842,4 +3171,69 @@ mod tests {
.unwrap();
assert_eq!(task_cnt, 0, "关联 tasks 记录应被清理");
}
/// 原子选点测试(#5 修复验证):
/// 1. claim_pending_grid_points 返回 pending 点并原子标记为 queued。
/// 2. 第二次 claim 返回空(点已非 pending)。
/// 3. 排序正确:wave ASC, cno_sum ASC, teff ASC。
#[tokio::test]
async fn test_claim_pending_grid_points_atomic() {
let temp_dir = tempfile::tempdir().unwrap();
let db = Database::new(&temp_dir.path().join("claim_db.db").to_string_lossy())
.await
.unwrap();
let wf = "claim_wf";
// 插入 3 个不同 wave 的点
let p1 = GridPointParams {
teff: 35000.0.into(),
logg: 5.5.into(),
loghe: (-1.0).into(),
logc: (-2.0).into(),
logn: (-2.0).into(),
logo: (-2.0).into(),
};
let p2 = GridPointParams {
teff: 40000.0.into(),
logg: 5.0.into(),
loghe: (-1.0).into(),
logc: (-3.0).into(),
logn: (-3.0).into(),
logo: (-3.0).into(),
};
let p3 = GridPointParams {
teff: 30000.0.into(),
logg: 5.5.into(),
loghe: (-1.0).into(),
logc: (-1.0).into(),
logn: (-1.0).into(),
logo: (-1.0).into(),
};
db.upsert_grid_point(&p1, 1, wf).await.unwrap();
db.upsert_grid_point(&p2, 0, wf).await.unwrap();
db.upsert_grid_point(&p3, 2, wf).await.unwrap();
// 第一次 claim:应返回全部 3 个,按 wave ASC 排序(p2 wave=0, p1 wave=1, p3 wave=2
let claimed = db.claim_pending_grid_points(100, wf).await.unwrap();
assert_eq!(claimed.len(), 3);
assert_eq!(claimed[0].0, p2.model_name(), "wave=0 应排第一");
assert_eq!(claimed[1].0, p1.model_name(), "wave=1 应排第二");
assert_eq!(claimed[2].0, p3.model_name(), "wave=2 应排第三");
// 验证点已变为 queued
let pending_after = db.get_pending_grid_points(wf).await.unwrap();
assert_eq!(pending_after.len(), 0, "claim 后不应有 pending 点");
// 第二次 claim:应返回空
let claimed_again = db.claim_pending_grid_points(100, wf).await.unwrap();
assert_eq!(claimed_again.len(), 0, "已 queued 的点不应被再次 claim");
// LIMIT 测试:重置回 pending 后只 claim 2 个
db.reset_queued_grid_points_to_pending(wf).await.unwrap();
let partial = db.claim_pending_grid_points(2, wf).await.unwrap();
assert_eq!(partial.len(), 2, "LIMIT 2 应只返回 2 个点");
let remaining = db.claim_pending_grid_points(100, wf).await.unwrap();
assert_eq!(remaining.len(), 1, "剩余 1 个点");
}
}
+105 -21
View File
@@ -10,7 +10,7 @@ use axum::{
Router,
};
use clap::Parser;
use common::config::ServerConfig;
use common::config::{GridConfig, ServerConfig};
use common::logging::init_logging;
use mq::sqlite_queue::SqliteTaskQueue;
use std::net::SocketAddr;
@@ -57,22 +57,30 @@ async fn main() -> Result<()> {
let queue = Arc::new(SqliteTaskQueue::new(&server_cfg.queue_db_path).await?);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
// Auto-register sdB_cno.yaml if exists and not yet in DB
// Auto-register sdB_cno.yaml if exists and not yet in DB.
// 仅在 DB 中尚无该工作流时注册(INSERT),绝不覆盖已存在的配置——
// 旧实现用 upsert 每次启动都用文件内容覆盖 config_yaml/description/status
// 导致管理员通过 API 编辑过的配置在重启后被静默回退。
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
{
tracing::warn!("预注册默认工作流失败: {}", e);
let already_exists = matches!(db.get_workflow("sdB_cno").await, Ok(Some(_)));
if !already_exists {
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'");
}
} else {
info!("已在数据库中成功预注册默认工作流 'sdB_cno'");
info!("默认工作流 'sdB_cno' 已存在于数据库,保留现有配置(不覆盖 API 编辑)");
}
}
}
@@ -84,6 +92,51 @@ async fn main() -> Result<()> {
tracing::warn!("⚠️ 检测到系统当前正在使用弱口令凭据或默认 Token!建议生产环境在 .env 中配置使用 openssl rand -hex 32 生成的高强度 Token");
}
// 启动恢复:把卡在 `initializing` 态的工作流重新初始化。
// 背景:`start_workflow` 把状态切到 `initializing` 后在后台 spawn `initialize_grid`
// 若进程在初始化中途崩溃/重启,工作流会永久卡在 `initializing`——`get_running_workflow_names`
// 仍把它视为可调度,但网格展开未完成,导致半初始化网格被调度。
// `initialize_grid` 是幂等的(upsert ON CONFLICT DO NOTHING),重跑可补齐缺失点并把
// 状态推进到 `running`。失败则回退为 `idle` 等待人工重启(与 start_workflow 口径一致)。
match db.get_initializing_workflows().await {
Ok(stuck) if !stuck.is_empty() => {
info!(
"检测到 {} 个卡在 initializing 态的工作流(上次启动未完成即重启),开始重新初始化...",
stuck.len()
);
for (wf_name, wf_yaml) in &stuck {
match GridConfig::from_yaml_str(wf_yaml) {
Ok(cfg) => {
match scheduler.initialize_grid(&cfg, wf_name).await {
Ok(_) => {
let _ = db.update_workflow_status(wf_name, "running").await;
info!("启动恢复:工作流 {} 已完成重新初始化并切回 running", wf_name);
}
Err(e) => {
tracing::warn!(
"启动恢复:工作流 {} 重新初始化失败,回退为 idle: {}",
wf_name,
e
);
let _ = db.update_workflow_status(wf_name, "idle").await;
}
}
}
Err(e) => {
tracing::warn!(
"启动恢复:工作流 {} 的 YAML 配置解析失败,回退为 idle: {}",
wf_name,
e
);
let _ = db.update_workflow_status(wf_name, "idle").await;
}
}
}
}
Ok(_) => {}
Err(e) => tracing::warn!("启动恢复:查询 initializing 工作流失败: {}", e),
}
let rate_limiter = api::rate_limit::RateLimiter::new(5, std::time::Duration::from_secs(300));
let state = AppState {
@@ -165,6 +218,21 @@ async fn main() -> Result<()> {
}
}
// 孤儿 running 点回收(#6 修复兜底):queue 行已消失(误删/崩溃丢队列)
// 但 grid_points 仍卡在 running 的点,requeue_stale_tasks 找不到它们,
// 在此按 tasks 表的 stale pending 记录兜底重置为 pending,让调度器重新派发。
match bg_db_clone.reset_orphaned_running_points(stale_sec).await {
Ok(reset) => {
if reset > 0 {
info!("已回收 {} 个孤儿 running 网格点(领用凭证丢失,重置为 pending)", reset);
}
}
Err(e) => {
tracing::warn!("回收孤儿 running 网格点失败: {}", e);
has_error = true;
}
}
if let Err(e) = bg_scheduler_clone.schedule_pending_tasks().await {
tracing::warn!("后台定时性任务调度检测失败: {}", e);
has_error = true;
@@ -231,7 +299,6 @@ async fn main() -> Result<()> {
// 大体积上传端点单独拎出,套用更宽松的 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;
@@ -254,10 +321,16 @@ async fn main() -> Result<()> {
api::rate_limit::rate_limit_middleware,
);
let api_router = Router::new()
// 其余 API(小体积)套用 10MB 默认上限,防止大文件内存耗尽 DoS。
// 注意:body limit layer 从外到内执行、先接触原始 body 流的层先生效。
// 必须把 10MB 限制只套在"小体积子 router"上,再与 report_router 合并,
// 合并后的外层不能再套任何全局 limit —— 否则外层 10MB 会截断 report 的 256MB body 流,
// 导致收敛种子 .7(常 >10MB)上报被 413 拒绝、结果反复重算。
let small_body_router = Router::new()
// Auth API
.route("/login", post(api::auth::login))
.route("/auth/check", get(api::auth::check_auth))
.route("/auth/logout", post(api::auth::logout))
// Core Node & Task API
.route(
"/node/register",
@@ -329,12 +402,24 @@ async fn main() -> Result<()> {
"/admin/nodes/:node_id/enable",
post(api::admin::enable_node),
)
// 合并大体积上报路由(继承各自的 body limit)
.merge(report_router)
.layer(DefaultBodyLimit::max(DEFAULT_BODY_LIMIT));
// 鉴权启用条件:未应急关闭,且配置了 admin 凭据
// 合并两个子 router:各自携带自己的 body limit,互不覆盖
let api_router = small_body_router.merge(report_router);
// 鉴权策略:fail-closed。
// - 配置了 DCTS_ADMIN_TOKEN → 启用完整鉴权。
// - 显式 DCTS_AUTH_DISABLE=1 → 无鉴权(仅本地调试,需运维主动声明承担风险)。
// - 既未配置 token、又未显式 disable → **拒绝启动**。
// 避免 .env 缺失/变量名拼错/容器未注入环境变量时服务静默退化为完全无鉴权裸奔。
let auth_enabled = !state.auth_disabled && state.admin_token.is_some();
if !auth_enabled && !state.auth_disabled {
anyhow::bail!(
"拒绝启动:未配置 DCTS_ADMIN_TOKEN 且未显式设置 DCTS_AUTH_DISABLE=1。\
生产部署必须在 .env 中配置 DCTS_ADMIN_TOKEN;若确为本地调试,\
请显式设置 DCTS_AUTH_DISABLE=1 以承担无鉴权风险。"
);
}
let api_router = if auth_enabled {
info!("已启用 API 身份鉴权保护(Admin 端点需 admin token 验证;Node 节点免 Token 提交申请,经 Dashboard 管理员审批授权下发)");
@@ -346,9 +431,8 @@ async fn main() -> Result<()> {
let auth_layer = axum::middleware::from_fn_with_state(state.clone(), api::auth_middleware);
api_router.layer(auth_layer).layer(rate_limit_layer)
} else {
tracing::warn!(
"⚠️ 警告:未配置 DCTS_ADMIN_TOKEN(且未启用 DCTS_AUTH_DISABLE),\
服务端运行在【无鉴权模式】!公网部署务必配置凭据。"
info!(
"DCTS_AUTH_DISABLE=1 已生效:服务端运行在无鉴权模式(仅限本地调试,切勿用于生产)。"
);
api_router
};
+35 -28
View File
@@ -11,11 +11,18 @@ use crate::db::Database;
pub struct GridScheduler {
db: Database,
queue: Arc<SqliteTaskQueue>,
/// 调度互斥锁:防止 start_workflow 的即时调度与后台 30s 循环并发进入
/// schedule_pending_tasks,消除 TOCTOU 竞态导致的重复派发(#5 修复)。
schedule_lock: tokio::sync::Mutex<()>,
}
impl GridScheduler {
pub fn new(db: Database, queue: Arc<SqliteTaskQueue>) -> Self {
Self { db, queue }
Self {
db,
queue,
schedule_lock: tokio::sync::Mutex::new(()),
}
}
/// Expands grid points from config and registers them into the database.
@@ -140,12 +147,17 @@ impl GridScheduler {
true
}
/// Enqueues pending grid points into MQ with active seed detection and batching.
/// Enqueues pending grid points into MQ with batching(冷启动优先).
///
/// 多工作流分区(#3 修复):对**每个** running/initializing 工作流分别派发任务,
/// 替代原来「全局只一个 running workflow」的 LIMIT 1 假设。各工作流独立 batch
/// 独立 seed 匹配(seeds 仍是全局共享的物理资源池)。
/// 替代原来「全局只一个 running workflow」的 LIMIT 1 假设。各工作流独立 batch
/// 正常路径一律派发 ColdRun 冷启动,种子匹配只发生在失败后的
/// trigger_seed_step_fallbackseeds 仍是全局共享的物理资源池)。
pub async fn schedule_pending_tasks(&self) -> Result<usize> {
// 互斥锁:start_workflow 的即时调度与后台 30s 循环可能并发调用本方法,
// 两者各自 SELECT 同一批 pending 点会产生重复任务(#5 修复)。
let _guard = self.schedule_lock.lock().await;
let workflows = self.db.get_running_workflow_names().await?;
if workflows.is_empty() {
return Ok(0);
@@ -184,46 +196,35 @@ impl GridScheduler {
) -> Result<usize> {
let timeout_sec = self.get_workflow_timeout_sec(workflow_name).await;
// SQL 层直接附加 LIMIT = batch_limit + workflow_name 筛选,完全免除数万点位无谓内存反序列化
// 原子选点(#5 修复):IMMEDIATE 事务内完成 SELECT + UPDATE status='queued'
// 替代原来 get_pending_grid_points_limitSELECT+ update_grid_statusUPDATE
// 的分离操作,杜绝两个并发调度调用 SELECT 到同一批 pending 点的 TOCTOU 竞态。
let pending = self
.db
.get_pending_grid_points_limit(batch_limit, workflow_name)
.claim_pending_grid_points(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 schedulingseeds 全局共享)
let (task_type, seed_name) = match self.db.find_best_seed_from_db(&params).await {
Ok(Some(seed_match)) => {
info!(
"工作流 {} 网格点 {} 匹配到数据库近邻种子 {} (距离: {:.2}),安排 SeedStep 热启动调度",
workflow_name, name, seed_match.name, seed_match.distance
);
(TaskType::SeedStep, Some(seed_match.name))
}
_ => (TaskType::ColdRun, None),
};
// 冷启动优先:正常调度路径一律走 ColdRun 自包含冷启动(lte 阶段 ltgray="T"
// 生成 grey start,不依赖任何种子文件)。SeedStep 仅作为冷启动失败后的救援
// 任务出现,由 trigger_seed_step_fallback(见下)派发——「冷启动优先、失败再
// 种子步进」的单向语义,避免旧版「正常路径热启动优先 + 节点端缺种子回退冷启
// 动链」的双向纠缠。
let task_spec = TaskSpec {
task_id: Uuid::new_v4(),
point_name: name.clone(),
params,
task_type,
seed_point_name: seed_name,
task_type: TaskType::ColdRun,
seed_point_name: None,
timeout_sec,
workflow_name: Some(workflow_name.to_string()),
wave,
};
self.db.insert_task(&task_spec).await?;
// 采用先标记 DB 状态Queued 后发 MQ 的时序,防止推入 MQ 后数据库修改异常导向下一轮误重投
self.db
.update_grid_status(
&name,
common::models::GridPointStatus::Queued,
workflow_name,
)
.await?;
// 点已在 claim_pending_grid_points 的 IMMEDIATE 事务中原子标记queued
// 无需再单独 update_grid_status(Queued)。push 失败时回滚为 pending 即可。
match self.queue.push_task(&task_spec).await {
Ok(_) => {
dispatched += 1;
@@ -250,6 +251,9 @@ impl GridScheduler {
);
}
let _ = self.queue.remove_task(&task_spec.task_id.to_string()).await;
// 同步清理先于 push 插入的 tasks 历史行,避免遗留 pending 历史记录
// 污染每点尝试计数统计(attempt_count 依赖 tasks 表聚合)。
let _ = self.db.delete_task(&task_spec.task_id).await;
}
}
}
@@ -333,6 +337,9 @@ impl GridScheduler {
)
.await;
let _ = self.queue.remove_task(&task_spec.task_id.to_string()).await;
// 清理先于 push 插入的 tasks 历史行,与上方 schedule_pending_tasks_for_workflow
// 回滚口径一致,避免遗留 pending 历史污染尝试计数。
let _ = self.db.delete_task(&task_spec.task_id).await;
return Err(e);
}
info!(
+39 -14
View File
@@ -221,7 +221,7 @@ async fn test_l2_node_token_issue_reissue_flow() {
node_id: "node-l2-test".to_string(),
max_slots: 4,
};
db.register_node(&reg).await.unwrap();
let (_, _, secret) = db.register_node(&reg).await.unwrap();
// 2. 颁发专属 token,返回明文
let token = db.issue_node_token("node-l2-test").await.unwrap();
@@ -229,16 +229,23 @@ async fn test_l2_node_token_issue_reissue_flow() {
// 2b. 一次性取走暂存明文(take_pending_node_token):首次取到与颁发一致的明文,
// 再次取为 None(取走即焚)。验证 #4 简化为单一 UPDATE...RETURNING 后行为一致。
let pending = db.take_pending_node_token("node-l2-test").await.unwrap();
// H8:须提供注册时下发的 registration_secret。
let pending = db
.take_pending_node_token("node-l2-test", secret.as_deref())
.await
.unwrap();
assert_eq!(pending.as_deref(), Some(token.as_str()));
let pending2 = db.take_pending_node_token("node-l2-test").await.unwrap();
let pending2 = db
.take_pending_node_token("node-l2-test", secret.as_deref())
.await
.unwrap();
assert!(
pending2.is_none(),
"取走即焚:第二次 take_pending 必须返回 None"
);
// 不存在的 node take 也应返回 None(不报错)
assert!(db
.take_pending_node_token("node-not-exist")
.take_pending_node_token("node-not-exist", None)
.await
.unwrap()
.is_none());
@@ -774,9 +781,15 @@ async fn test_node_approval_workflow() {
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
assert_eq!(json["status"], "pending_approval");
// H8:注册响应下发一次性 registration_secret,后续 check_status 取 token 须回传。
let reg_secret = json["registration_secret"].as_str().map(|s| s.to_string());
assert!(reg_secret.is_some(), "注册响应应包含 registration_secret");
// 2. Node 端轮询查状态 ➔ status: pending_approval
let check_body = serde_json::json!({ "node_id": "node-pending-01" });
let check_body = serde_json::json!({
"node_id": "node-pending-01",
"registration_secret": reg_secret,
});
let res = app
.clone()
.oneshot(
@@ -1291,6 +1304,7 @@ fn make_import_multipart(
report_json: &str,
seed_bytes: &[u8],
seed_name: &str,
success_method: &str,
) -> Vec<u8> {
let mut body = Vec::new();
body.extend_from_slice(format!("--{}\r\n", boundary).as_bytes());
@@ -1309,6 +1323,12 @@ fn make_import_multipart(
body.extend_from_slice(b"Content-Type: application/octet-stream\r\n\r\n");
body.extend_from_slice(seed_bytes);
body.extend_from_slice(b"\r\n");
// 收敛途径字段(cold_run/seed_step):模拟 import_results 工具判定后透传的途径。
body.extend_from_slice(format!("--{}\r\n", boundary).as_bytes());
body.extend_from_slice(b"Content-Disposition: form-data; name=\"success_method\"\r\n");
body.extend_from_slice(b"Content-Type: text/plain\r\n\r\n");
body.extend_from_slice(success_method.as_bytes());
body.extend_from_slice(b"\r\n");
body.extend_from_slice(format!("--{}--\r\n", boundary).as_bytes());
body
}
@@ -1360,6 +1380,7 @@ async fn test_import_seed_admin_endpoint() {
&conv,
b"FAKE_ATMOS_7",
"t20000_g5.0_he-2_c-4_n-4_o-4.7",
"cold_run",
);
let res = app
.clone()
@@ -1381,6 +1402,7 @@ async fn test_import_seed_admin_endpoint() {
&conv,
b"FAKE_ATMOS_7",
"t20000_g5.0_he-2_c-4_n-4_o-4.7",
"cold_run",
);
let res = app
.clone()
@@ -1419,6 +1441,7 @@ async fn test_import_seed_admin_endpoint() {
&conv,
b"FAKE_ATMOS_7_AGAIN",
"t20000_g5.0_he-2_c-4_n-4_o-4.7",
"cold_run",
);
let res = app
.clone()
@@ -1443,7 +1466,7 @@ async fn test_import_seed_admin_endpoint() {
// 4. 未收敛点 → 200,但不写 .7、grid_points 维持 pending(未建 converged)。
let conv_fail = make_legacy_conv_json("t20000_g5.0_he-2_c-4_n-4_o-4_fail", false);
let body_bytes = make_import_multipart("boundary4", &conv_fail, b"WONT_BE_USED", "x.7");
let body_bytes = make_import_multipart("boundary4", &conv_fail, b"WONT_BE_USED", "x.7", "cold_run");
let res = app
.oneshot(
Request::builder()
@@ -1516,7 +1539,7 @@ async fn test_import_seed_python_legacy_conv_json() {
let name = "t20000_g5.0_he-2_c-4_n-4_o-4";
let conv = make_python_legacy_conv_json(name);
let body_bytes =
make_import_multipart("boundaryL", &conv, b"FAKE_ATMOS_7", &format!("{name}.7"));
make_import_multipart("boundaryL", &conv, b"FAKE_ATMOS_7", &format!("{name}.7"), "cold_run");
let res = app
.oneshot(
Request::builder()
@@ -1535,14 +1558,14 @@ async fn test_import_seed_python_legacy_conv_json() {
"旧版嵌套 stages 的 conv.json 应被接受"
);
// grid_pointsconverged + imported 手段 + 旧版 elapsed_sec 已落库
// grid_pointsconverged + cold_run 手段(旧版 conv.json 无 seed_nc 阶段)+ 旧版 elapsed_sec 已落库
let row = db
.get_workflow_point_row("wf_legacy", name)
.await
.unwrap()
.expect("grid_points 应存在");
assert_eq!(row.status, "converged");
assert_eq!(row.success_method.as_deref(), Some("imported"));
assert_eq!(row.success_method.as_deref(), Some("cold_run"));
assert_eq!(
row.last_elapsed_sec,
Some(715.0),
@@ -2036,7 +2059,7 @@ async fn test_wf_stats_endpoint() {
false,
)
.await;
db.mark_grid_point_imported(&p_imported.model_name(), "wf_stats", None)
db.mark_grid_point_imported(&p_imported.model_name(), "wf_stats", None, "cold_run")
.await
.unwrap();
@@ -2091,9 +2114,8 @@ async fn test_wf_stats_endpoint() {
assert_eq!(data["running"], 1);
assert_eq!(data["converged"], 3);
assert_eq!(data["failed"], 1);
assert_eq!(data["cold_run_converged"], 1);
assert_eq!(data["cold_run_converged"], 2);
assert_eq!(data["seed_step_converged"], 1);
assert_eq!(data["imported_converged"], 1);
// wave 分布:3 个波次,wave0 共 3 点全未收敛
let waves = data["waves"].as_array().unwrap();
assert_eq!(waves.len(), 3);
@@ -2225,7 +2247,7 @@ async fn seed_obs_fixture(db: &Database, db_path: &std::path::Path, wf: &str) ->
true,
)
.await;
db.mark_grid_point_imported(&p_imported.model_name(), wf, None)
db.mark_grid_point_imported(&p_imported.model_name(), wf, None, "cold_run")
.await
.unwrap();
@@ -2383,7 +2405,7 @@ async fn test_wf_points_endpoint() {
assert_eq!(data["total"], 1);
assert_eq!(data["points"][0]["name"], n.imported);
// ---- 7. 分页:limit=2 翻页,total 恒定;limit 钳位 ≤500 ----
// ---- 7. 分页:limit=2 翻页,total 恒定;limit 缺省=全量(无上限,联合分析用) ----
let (_, page0) = get_points(&app, "/api/workflows/wf_pts/points?limit=2&offset=0").await;
let (_, page2) = get_points(&app, "/api/workflows/wf_pts/points?limit=2&offset=4").await;
assert_eq!(page0["total"], 6);
@@ -2391,6 +2413,9 @@ async fn test_wf_points_endpoint() {
assert_eq!(page2["points"].as_array().unwrap().len(), 2);
let (_, all) = get_points(&app, "/api/workflows/wf_pts/points?limit=9999").await;
assert_eq!(all["points"].as_array().unwrap().len(), 6);
// limit 完全缺省 → 全量(None 路径,不拼 LIMIT 子句)
let (_, full) = get_points(&app, "/api/workflows/wf_pts/points").await;
assert_eq!(full["points"].as_array().unwrap().len(), 6);
// ---- 8. 时间排序:最近完成的点(rescued 的 seed_step 尝试最新)排首位 ----
let (st, data) = get_points(
@@ -53,7 +53,7 @@ async fn test_same_workflow_name_preserves_converged() {
db.upsert_grid_point_named(name, &p, 0, "sdB_cno")
.await
.unwrap();
db.mark_grid_point_imported(name, "sdB_cno", None)
db.mark_grid_point_imported(name, "sdB_cno", None, "cold_run")
.await
.unwrap();
@@ -84,7 +84,7 @@ async fn test_different_workflow_name_causes_recompute() {
db.upsert_grid_point_named(name, &p, 0, "imported")
.await
.unwrap();
db.mark_grid_point_imported(name, "imported", None)
db.mark_grid_point_imported(name, "imported", None, "cold_run")
.await
.unwrap();
@@ -122,7 +122,7 @@ async fn test_mixed_grid_import_then_init_avoids_recompute() {
db.upsert_grid_point_named("t20000_g5.0_he-2_c-4_n-4_o-4", &p_old, 0, "sdB_cno")
.await
.unwrap();
db.mark_grid_point_imported("t20000_g5.0_he-2_c-4_n-4_o-4", "sdB_cno", None)
db.mark_grid_point_imported("t20000_g5.0_he-2_c-4_n-4_o-4", "sdB_cno", None, "cold_run")
.await
.unwrap();
@@ -176,7 +176,7 @@ async fn test_precision_diff_import_then_init_preserves_converged() {
db.upsert_grid_point_named(canonical, &p, 0, "sdB_cno")
.await
.unwrap();
db.mark_grid_point_imported(canonical, "sdB_cno", None)
db.mark_grid_point_imported(canonical, "sdB_cno", None, "cold_run")
.await
.unwrap();