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
+19 -2
View File
@@ -328,11 +328,28 @@ impl Default for ServerConfig {
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 秒缓冲,避免两边的超时检测同时触发冲突
// 默认设置为 21600 秒6 小时),是计算任务默认超时(7200 秒)的 3 倍缓冲。
// 历史值 7800s 仅比 7200s 高 600s,在科学计算网格的长尾难收敛点上极易出现:
// 一个耗时接近 timeout 的任务,在上报链路抖动/排队时被后台 requeue_stale_tasks
// 重置为 pending → 原节点完成上报时 verify_task_claim 因 status≠claimed 拒绝 →
// 昂贵的计算结果被静默丢弃、重投重算。3 倍缓冲可基本消除该竞态。
// 若显式配置了 DCTS_STALE_SEC,但小于 timeout 默认 7200 的 1.5 倍,给出告警。
let stale_sec = std::env::var("DCTS_STALE_SEC")
.ok()
.and_then(|v| v.parse::<u64>().ok())
.unwrap_or(7800);
.unwrap_or(21600);
const DEFAULT_TIMEOUT_SEC: u64 = 7200;
if stale_sec < DEFAULT_TIMEOUT_SEC * 3 / 2 {
tracing::warn!(
"DCTS_STALE_SEC={} 过小(小于默认 timeout {} 的 1.5 倍 {}):\
耗时接近 timeout 的任务在上报前可能被重投,导致计算结果被丢弃重算。\
建议 ≥ {}timeout 的 3 倍)。",
stale_sec,
DEFAULT_TIMEOUT_SEC,
DEFAULT_TIMEOUT_SEC * 3 / 2,
DEFAULT_TIMEOUT_SEC * 3
);
}
let node_stale_sec = std::env::var("DCTS_NODE_STALE_SEC")
.ok()
.and_then(|v| v.parse::<u64>().ok())
+174 -8
View File
@@ -7,6 +7,44 @@ use std::sync::OnceLock;
static FORT9_RE: OnceLock<Regex> = OnceLock::new();
static NAN_RE: OnceLock<Regex> = OnceLock::new();
/// 匹配 Fortran 无-E 科学记数法的尾数+指数部分(归一化用,见 parse_fortran_float)。
static NO_E_EXP_RE: OnceLock<Regex> = OnceLock::new();
/// 解析 fort.9 / fort.7 中的数值字符串为 f64,兼容 Fortran 的**无-E 科学记数法**。
///
/// # 背景
/// Fortran 在数值指数 ≥ 100 时会输出省略 `E` 的科学记数法(字段宽度不够把 E 挤掉),
/// 例如把 `-1.35E+118` 输出成 `-1.35+118`、`9.99E+108` 输出成 `9.99+108`。这出现在
/// tlusty 迭代发散时的 fort.9 中(布居数变化达 ±100~±200 量级)。
///
/// 历史问题:直接 `s.parse::<f64>()` 对 `-1.35+118` 返回 ErrRust 的 f64::from_str
/// 不认无-E 记数法),旧代码 `Err(_) => continue` 会**静默跳过该行**。若被跳过的恰好是
/// 本次迭代最差的 depth 行,会低估 max_relc,把发散误判为收敛(科学正确性 bug)。
///
/// # 规则
/// 1. 先尝试标准 `parse`(覆盖所有正常情况:整数、小数、带 E/e 的科学记数法)。
/// 2. 失败则归一化:形如 `[前导符号?]<尾数>[+/-]<指数>`(无 E)的,在尾部符号前插 E
/// 再 parse。如 `-2.37+192` → `-2.37E+192`。
/// 3. 指数导致 f64 上溢时返回 ±Infinity(对收敛判断是正确的:发散值 → max_relc=Inf →
/// is_finite()=false → converged=false,标记为未收敛而非崩溃)。
fn parse_fortran_float(s: &str) -> Option<f64> {
let s = s.trim();
// 1. 标准 parse
if let Ok(v) = s.parse::<f64>() {
return Some(v);
}
// 2. 归一化无-E 记数法:[前导符号?]<尾数>(含小数点或多位数字)[+/-]<指数>
let re = NO_E_EXP_RE.get_or_init(|| {
Regex::new(r"^([+-]?[\d.]+)([+-]\d+)$").unwrap()
});
if let Some(caps) = re.captures(s) {
let normalized = format!("{}E{}", &caps[1], &caps[2]);
if let Ok(v) = normalized.parse::<f64>() {
return Some(v); // 含 Inf/-Inf(溢出)
}
}
None
}
#[derive(Debug, Clone)]
struct Fort9Row {
@@ -16,6 +54,24 @@ struct Fort9Row {
/// Parses `fort.9` and evaluates convergence against `chmax`
pub fn check_fort9(path: &Path, chmax: f64) -> ConvCheckResult {
// M10:非法 chmax 防御。chmax ≤ 0(误配为负或 0)会使 `max_relc < chmax` 永远为假,
// 整个工作流的所有 stage 都永不收敛却无任何提示,表现为"全部失败"难以定位。
// 在此显式拒绝并返回错误,让调用方/日志能立即看出是配置问题而非物理发散。
if !chmax.is_finite() || chmax <= 0.0 {
return ConvCheckResult {
converged: false,
max_relc: f64::INFINITY,
worst_depth: -1,
last_iter: None,
n_depths: 0,
chmax,
error: Some(format!(
"非法 chmax={}(须为正有限数):请检查工作流 YAML 中该 stage 的 chmax 配置",
chmax
)),
};
}
let file = match File::open(path) {
Ok(f) => f,
Err(e) => {
@@ -52,9 +108,9 @@ pub fn check_fort9(path: &Path, chmax: f64) -> ConvCheckResult {
Ok(v) => v,
Err(_) => continue,
};
let maximum: f64 = match caps[7].parse() {
Ok(v) => v,
Err(_) => continue,
let maximum: f64 = match parse_fortran_float(&caps[7]) {
Some(v) => v,
None => continue,
};
if cur_iter != Some(iter) {
@@ -120,7 +176,16 @@ pub fn check_fort9(path: &Path, chmax: f64) -> ConvCheckResult {
}
}
/// Checks if an atmosphere file (.7) contains NaN lines (>10% NaN lines = invalid) using exact word boundary
/// Checks if an atmosphere file (.7) contains invalid numeric lines (NaN / Inf / Fortran overflow `***`).
///
/// 检测内容:
/// - NaNFortran 写出 `NaN`/`nan`/`NAN`,数值发散的典型产物)。
/// - Inf / InfinityFortran 除零/溢出)。
/// - Fortran 字段宽度溢出标记 `***`(如 `********`):Tlusty 数值溢出发散时常以星号填满
/// 字段而非写 NaN。历史上只检 `\bnan\b`,全溢出发散的大气会被判"无 NaN"→converged
/// 产出物理上完全错误的大气。
///
/// 超过 10% 的行命中任一标记即判定无效。
///
/// 文件缺失时返回 `false`(语义:不存在 NaN 内容)。这与“含 NaN 导致无效”是不同语义;
/// 调用方需先自行确认文件存在性,不应将“缺失”与“含 NaN”混为一谈。
@@ -131,13 +196,15 @@ pub fn atmosphere_has_nan(path: &Path) -> bool {
};
let reader = BufReader::new(file);
let mut total_lines = 0;
let mut nan_lines = 0;
let nan_re = NAN_RE.get_or_init(|| Regex::new(r"(?i)\bnan\b").unwrap());
let mut bad_lines = 0;
let nan_re = NAN_RE.get_or_init(|| {
Regex::new(r"(?i)(\bnan\b|\binf(?:inity)?\b|\*{3,})").unwrap()
});
for line in reader.lines().map_while(Result::ok) {
total_lines += 1;
if nan_re.is_match(&line) {
nan_lines += 1;
bad_lines += 1;
}
}
@@ -145,7 +212,7 @@ pub fn atmosphere_has_nan(path: &Path) -> bool {
return true;
}
(nan_lines as f64) > (total_lines as f64 * 0.1)
(bad_lines as f64) > (total_lines as f64 * 0.1)
}
#[cfg(test)]
@@ -171,5 +238,104 @@ mod tests {
// Missing file returns false (absence != contains NaN)
let missing_path = dir.path().join("missing.7");
assert!(!atmosphere_has_nan(&missing_path));
// M11: Fortran 字段溢出标记 ***(数值发散时的常见形态,历史上漏检)
let overflow_file_path = dir.path().join("overflow.7");
std::fs::write(&overflow_file_path, "******** 2 3\n******** 5 6\n7 8 9\n").unwrap();
assert!(atmosphere_has_nan(&overflow_file_path));
// M11: Inf / Infinity 也应被识别为无效数值
let inf_file_path = dir.path().join("inf.7");
std::fs::write(&inf_file_path, "Inf 2 3\nInfinity 5 6\n7 8 9\n").unwrap();
assert!(atmosphere_has_nan(&inf_file_path));
}
/// 验证 parse_fortran_float 对各种数值格式(含 Fortran 无-E 记数法)的解析。
#[test]
fn test_parse_fortran_float() {
// 标准 parse 能覆盖的
assert_eq!(parse_fortran_float("100"), Some(100.0));
assert_eq!(parse_fortran_float("-0.001"), Some(-0.001));
assert_eq!(parse_fortran_float("3.14"), Some(3.14));
assert_eq!(parse_fortran_float("-5.42E+72"), Some(-5.42e72));
assert_eq!(parse_fortran_float("1.5e-99"), Some(1.5e-99));
assert_eq!(parse_fortran_float("0"), Some(0.0));
// Fortran 无-E 记数法(指数≥100 时 E 被挤掉)
assert_eq!(parse_fortran_float("-2.37+192"), Some(-2.37e192));
assert_eq!(parse_fortran_float("-1.35+118"), Some(-1.35e118));
assert_eq!(parse_fortran_float("9.99+108"), Some(9.99e108));
// 负指数的无-E 记数法
assert_eq!(parse_fortran_float("1.5-99"), Some(1.5e-99));
// 指数导致 f64 上溢 → Inf(发散值,应被识别为无效→未收敛)
assert_eq!(parse_fortran_float("1.5+400"), Some(f64::INFINITY));
assert_eq!(parse_fortran_float("-9.99+500"), Some(f64::NEG_INFINITY));
// Rust 的 f64::from_str 接受 "NaN"/"inf",返回 NaN/Inf(非 None)。
// 这些在 check_fort9 中会被 is_finite() 判为无效 → converged=false,行为正确。
assert!(parse_fortran_float("NaN").map(|v| v.is_nan()).unwrap_or(false));
assert_eq!(parse_fortran_float("inf"), Some(f64::INFINITY));
// 无法解析的垃圾 → None(调用方 continue 跳过)
assert_eq!(parse_fortran_float("abc"), None);
assert_eq!(parse_fortran_float("--1.0"), None);
assert_eq!(parse_fortran_float("1.2.3"), None);
}
/// 回归测试:fort.9 含 Fortran 无-E 记数法的发散行不得被静默跳过。
///
/// 历史 bug:旧代码 `caps[7].parse::<f64>()` 对 `-1.35+118` 返回 Err → `continue`
/// 跳过该行。若被跳过的是最差 depth 行,会低估 max_relc,把发散误判为收敛。
/// 修复后:无-E 记数法被正确归一化为 -1.35E+118,发散行参与 max_relc 计算 →
/// converged=false(而非崩溃或误判收敛)。
#[test]
fn test_fort9_with_no_e_notation_diverged() {
let dir = tempfile::tempdir().unwrap();
let fort9 = dir.path().join("diverged.9");
// 构造一个发散的 fort.9:某次迭代的 maximum 列含无-E 记数法(极大值)。
// 格式:iter depth fr0 chg tcorr pop max_relc ...max_relc 是 caps[7]
// 这里 caps[7] 放 -1.35+118(无-E= -1.35E+118,发散)。
std::fs::write(
&fort9,
// 行格式参照真实 fort.9:两列整数后跟若干数值列 + 末尾两列整数
// 第7个数值字段(caps[7]) 是 maximum
" 10 1 0 1.0E-3 1.0E-3 1.0E-3 -1.35+118 5 10\n\
10 2 0 1.0E-3 1.0E-3 1.0E-3 2.0E-4 5 10\n",
)
.unwrap();
let res = check_fort9(&fort9, 0.001);
// 发散行(-1.35E+118)未被静默跳过,max_relc 应为 1.35e118(极大,远超 chmax
assert!(
!res.converged,
"含发散行(无-E 记数法)的迭代不得被误判为收敛"
);
assert!(
res.max_relc > 1.0e100,
"发散行应参与 max_relc 计算,实际 max_relc={}",
res.max_relc
);
assert!(
res.error.is_none(),
"归一化后的发散值不应产生解析错误,实际 error={:?}",
res.error
);
}
/// 验证正常收敛的 fort.9(无无-E 记数法)不受影响。
#[test]
fn test_fort9_normal_converged_unchanged() {
let dir = tempfile::tempdir().unwrap();
let fort9 = dir.path().join("converged.9");
// max_relc(caps[7]) 远小于 chmax → 收敛
std::fs::write(
&fort9,
" 10 1 0 1.0E-3 1.0E-3 1.0E-3 1.0E-5 5 10\n",
)
.unwrap();
let res = check_fort9(&fort9, 0.001);
assert!(res.converged, "正常收敛行应判为收敛");
assert!((res.max_relc - 1.0e-5).abs() < 1e-15);
}
}
+18 -5
View File
@@ -151,17 +151,30 @@ fn write_if_changed(target_path: &Path, content: &[u8], executable: bool) -> Res
};
if should_write {
let mut file = File::create(target_path)?;
file.write_all(content)?;
file.flush()?;
// 原子写:先写 .tmp 再 rename,避免半写后崩溃留下截断/损坏的可执行二进制。
// 历史上直接 File::create 覆盖目标,若 write_all 中途进程被杀/磁盘满,会留下
// 截断的 tlusty_static,且权限可能已设 0o755(可执行但损坏),node 尝试运行时
// 产生难以定位的 Fortran 崩溃。tmp + rename 保证目标要么是完整旧版、要么是完整新版。
let tmp_path = target_path.with_extension("tmp.write");
{
let mut file = File::create(&tmp_path)
.with_context(|| format!("创建临时文件失败: {}", tmp_path.display()))?;
file.write_all(content)?;
file.flush()?;
// drop 前 sync_all 确保数据落盘,降低断电丢数据的概率。
let _ = file.sync_all();
}
#[cfg(unix)]
if executable {
use std::os::unix::fs::PermissionsExt;
let mut perms = fs::metadata(target_path)?.permissions();
let mut perms = fs::metadata(&tmp_path)?.permissions();
perms.set_mode(0o755);
fs::set_permissions(target_path, perms)?;
fs::set_permissions(&tmp_path, perms)?;
}
fs::rename(&tmp_path, target_path)
.with_context(|| format!("原子重命名 {} -> {} 失败", tmp_path.display(), target_path.display()))?;
}
Ok(())
+65 -3
View File
@@ -239,12 +239,23 @@ pub fn make_input5(
for ion in &ions {
let ilast = if ion.nlevs == 1 { 1 } else { 0 };
let ilvl = if ion.nlevs == 1 { 0 } else { ilvlin };
// ions 行字段右对齐到固定列,与真实 tlusty 输入文件
// (tests/tlusty/hhe/fort.5) 逐字节一致:
// iat 结束于 col4(|iat|=4), iz col10(+6), nlevs col16(+6),
// ilast col23(+7), ilvl col30(+7), nonstd col37(+7)。
// 实测:此宽列宽与窄列宽对 tlusty 输出(fort.7/9 等)完全相同(list-directed I/O
// 列宽无关),但对齐真实文件便于与历史参考 diff、符合 tlusty 官方输入惯例。
// 历史上曾用 `" {} {:2} {:5}..."`(窄列宽,与 Python 旧实现一致但偏离真实 fort.5)。
ions_block.push_str(&format!(
" {:2} {:2} {:5} {:5} {:5} 0 '{}' '{}'\n",
ion.iat, ion.iz, ion.nlevs, ilast, ilvl, ion.typion, ion.filei
"{:>4}{:>6}{:>6}{:>7}{:>7}{:>7} '{}' '{}'\n",
ion.iat, ion.iz, ion.nlevs, ilast, ilvl, 0, ion.typion, ion.filei
));
}
ions_block.push_str(" 0 0 0 -1 0 0 ' ' ' '\n");
// 终止行:(iat,iz,nlevs,ilast,ilvl,nonstd)=(0,0,0,-1,0,0)typion/filei 全空。
ions_block.push_str(&format!(
"{:>4}{:>6}{:>6}{:>7}{:>7}{:>7} '{}' '{}'\n",
0, 0, 0, -1, 0, 0, " ", " "
));
format!(
"{:.1} {:.1} ! TEFF, GRAV\n \
@@ -285,4 +296,55 @@ mod tests {
assert!(input5.contains("data/n1.dat"));
assert!(input5.contains("data/o1_23+10lev.dat"));
}
/// ions 行必须与真实 tlusty 输入文件 (tests/tlusty/hhe/fort.5) 的宽列宽逐字节一致。
/// 字段右对齐到固定列:iat→col4, iz→col10, nlevs→col16, ilast→col23, ilvl→col30, nonstd→col37。
/// 一旦此处失败,说明 ions 格式被改动且偏离了真实 tlusty 输入惯例(虽不影响 list-directed
/// 解析结果,但破坏与历史参考文件的可 diff 性)。
#[test]
fn test_ions_block_byte_perfect_with_real_fort5() {
// H+He 仅(metals="")对应真实 hhe fort.5 的 ions 集。
let params = GridPointParams {
teff: 35000.0.into(),
logg: 4.0.into(),
loghe: (-2.0).into(),
logc: (-1.0).into(),
logn: (-1.0).into(),
logo: (-1.0).into(),
};
let input5 = make_input5(&params, "T", "T", "", 100);
let lines: Vec<&str> = input5.lines().collect();
// 真实 fort.5 的 ions 数据行(数值部分 + typion/filei)。
// 注:filei 真实用 './data/...'Rust/Python 生成用 'data/...'DCTS 靠 data 软链解析),
// 这是既有独立差异,本测试只校验数值列宽 + typion。
let expected_ions_num: &[&str] = &[
" 1 0 9 0 100 0",
" 1 1 1 1 0 0",
" 2 0 14 0 100 0",
" 2 1 14 0 100 0",
" 2 2 1 1 0 0",
" 0 0 0 -1 0 0",
];
// 从生成结果中提取 ions 数据行的数值部分(引号前)。
let gen_ions_num: Vec<String> = lines
.iter()
.filter(|l| {
// ions 数据行:含引号且首 token 是整数
l.contains('\'')
&& l
.split_whitespace()
.next()
.map(|t| t.parse::<i32>().is_ok())
.unwrap_or(false)
})
.map(|l| l.split('\'').next().unwrap().trim_end().to_string())
.collect();
assert_eq!(
gen_ions_num.as_slice(),
expected_ions_num,
"ions 行数值部分必须与真实 fort.5 逐字节一致(宽列宽)"
);
}
}
+36 -2
View File
@@ -152,12 +152,25 @@ impl<'de> Deserialize<'de> for GridAxisValue {
}
// 主路径:YAML/JSON 标量原文(`5.0`、`20000`、`-2`、`"-4"`)。
// 文本必须可解析为 f64,否则返回反序列化错误——而非静默退化为 NaN。
// 旧实现对不可解析文本(如 YAML 里写成 `abc`)静默产生 NaN,NaN 随后流入
// cno_sum 求和、调度排序(partial_cmp(...).unwrap_or(Equal) 视 NaN 为相等)、
// 种子距离计算,导致难以定位的诡异行为。
fn visit_str<E: serde::de::Error>(self, v: &str) -> Result<Self::Value, E> {
Ok(GridAxisValue::from_text(v))
match v.parse::<f64>() {
Ok(value) => Ok(GridAxisValue {
value,
text: v.into(),
}),
Err(_) => Err(serde::de::Error::custom(format!(
"网格轴值不是合法数值: {:?}",
v
))),
}
}
fn visit_string<E: serde::de::Error>(self, v: String) -> Result<Self::Value, E> {
Ok(GridAxisValue::from_text(&v))
self.visit_str(&v)
}
// 兜底路径:纯数值来源(无原文)。先尝试解析原 deserializer 文本不可得,
@@ -614,6 +627,27 @@ mod tests {
assert!((sum - 6.0).abs() < 1e-9);
}
/// 回归:反序列化路径必须拒绝不可解析为 f64 的文本,而非静默退化为 NaN。
///
/// 旧实现 `from_text` 对非法文本用 `unwrap_or(f64::NAN)`NaN 随后流入 cno_sum
/// 求和、调度排序(partial_cmp().unwrap_or(Equal) 视 NaN 为相等)、种子距离计算,
/// 造成难以定位的诡异行为。现在字符串反序列化路径显式校验并返回错误。
#[test]
fn test_grid_axis_value_rejects_non_numeric_text() {
// 合法数值文本仍正常解析并保留源精度
let ok: Result<GridAxisValue, _> = serde_json::from_str("\"5.0\"");
assert!(ok.is_ok());
assert!((ok.unwrap().value() - 5.0).abs() < 1e-9);
// 非法文本应被拒绝(不再静默产生 NaN)
let bad: Result<GridAxisValue, _> = serde_json::from_str("\"abc\"");
assert!(bad.is_err(), "非法文本应触发反序列化错误而非产生 NaN");
// 确认不会再产生 NaN 值
let nan_text: Result<GridAxisValue, _> = serde_json::from_str("\"NaN-text\"");
assert!(nan_text.is_err());
}
#[test]
fn test_grid_point_status_display_and_conversion() {
assert_eq!(GridPointStatus::Pending.to_string(), "pending");
+111 -40
View File
@@ -98,20 +98,81 @@ pub fn default_seed_chain() -> Vec<StageConfig> {
]
}
/// 运行子进程,带超时与优雅退出(shutdown)感知。
///
/// 三种终止路径:
/// 1. 子进程正常结束 → 返回 ExitStatus。
/// 2. 超时(timeout_sec)→ SIGKILL 子进程 + 二级 30s 等待 reap,超时则放弃 Childkill_on_drop 兜底)。
/// 3. shutdown 信号(节点收到 SIGTERM/SIGINT)→ 立即 SIGKILL 子进程并快速返回 Err
/// 让上层尽快退出(在途任务的结果会丢失,由服务端 stale 重投兜底)。
///
/// 历史 bug:超时 kill 后 `child.wait().await` 无二级超时,Fortran 进程若卡死
/// OpenMP hang / ptrace)会使 wait 永久阻塞,超时机制名存实亡、slot 永久泄漏。
async fn run_child_async_with_timeout(
mut child: tokio::process::Child,
timeout_sec: u64,
shutdown: Option<std::sync::Arc<std::sync::atomic::AtomicBool>>,
) -> Result<std::process::ExitStatus> {
match tokio::time::timeout(tokio::time::Duration::from_secs(timeout_sec), child.wait()).await {
Ok(res) => Ok(res?),
Err(_) => {
let timeout_fut = tokio::time::timeout(
tokio::time::Duration::from_secs(timeout_sec),
child.wait(),
);
// 若提供了 shutdown 标志,则与超时/正常结束三路 select;否则只等超时/正常结束。
let outcome: Result<std::process::ExitStatus, ShutdownOrTimeout> = if let Some(flag) = shutdown {
let shutdown_watcher = async move {
// 轮询 shutdown 标志(10ms 粒度足够灵敏,开销可忽略)。
loop {
if flag.load(std::sync::atomic::Ordering::Acquire) {
return;
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
};
tokio::select! {
biased; // 优先响应 shutdown
_ = shutdown_watcher => Err(ShutdownOrTimeout::Shutdown),
r = timeout_fut => match r {
Ok(res) => Ok(res?),
Err(_) => Err(ShutdownOrTimeout::Timeout),
},
}
} else {
match timeout_fut.await {
Ok(res) => Ok(res?),
Err(_) => Err(ShutdownOrTimeout::Timeout),
}
};
match outcome {
Ok(status) => Ok(status),
Err(ShutdownOrTimeout::Shutdown) => {
let _ = child.start_kill();
let _ = child.wait().await;
let _ = tokio::time::timeout(
std::time::Duration::from_secs(30),
child.wait(),
)
.await;
anyhow::bail!("节点收到退出信号,子进程已被终止");
}
Err(ShutdownOrTimeout::Timeout) => {
let _ = child.start_kill();
let _ = tokio::time::timeout(
std::time::Duration::from_secs(30),
child.wait(),
)
.await;
anyhow::bail!("进程计算超时 (上限: {} 秒)", timeout_sec);
}
}
}
#[derive(Debug)]
enum ShutdownOrTimeout {
Shutdown,
Timeout,
}
pub struct ExecutionRunner<'a> {
pub runtime: &'a RuntimePaths,
pub work_dir: PathBuf,
@@ -139,6 +200,7 @@ impl<'a> ExecutionRunner<'a> {
seed_atmos,
synspec_cfg,
7200,
None,
)
.await
}
@@ -153,6 +215,7 @@ impl<'a> ExecutionRunner<'a> {
seed_atmos: Option<&Path>,
synspec_cfg: Option<&SynspecConfig>,
timeout_sec: u64,
shutdown: Option<std::sync::Arc<std::sync::atomic::AtomicBool>>,
) -> Result<ModelSummary> {
// `name` 取自权威的 TaskSpec.point_nameDB 的 grid_points.name 列,源精度正确),
// 而非 params.model_name()。原因:服务端把 GridPointParams 存成 6 个 REAL 数值列,
@@ -160,14 +223,10 @@ impl<'a> ExecutionRunner<'a> {
// 产出错误名(g5 而非 g5.0)。point_name 走独立 TEXT 列,精度全程保留。
// 下游(沙盒子目录、各阶段快照、conv.json.name、归档目录)全部用此 name,
// 故只需在此处用权威 name 即可让整条链精度正确。
let derived = params.model_name();
if derived != name {
warn!(
"网格点权威名 {} 与 params 重推名 {} 不一致(DB REAL 列回读丢精度所致),\
采用权威 point_name",
name, derived
);
}
//
// 历史:此处曾把 params.model_name() 与 name 对比并 warn 不一致。但该不一致是
// DB REAL 列回读丢精度的已知现象(runner 端无法修复,根治需改 DB schema 存原文),
// 且 runner 已全程采用权威 name,对比结果不参与任何决策——故移除这段噪音 warn。
let model_dir = self.work_dir.join(name);
tokio::fs::create_dir_all(&model_dir).await?;
@@ -276,7 +335,7 @@ impl<'a> ExecutionRunner<'a> {
.kill_on_drop(true)
.spawn()?;
let status_res = run_child_async_with_timeout(child, timeout_sec).await;
let status_res = run_child_async_with_timeout(child, timeout_sec, shutdown.clone()).await;
let rc = match status_res {
Ok(st) => st.code().unwrap_or(-1),
Err(e) => {
@@ -393,38 +452,49 @@ impl<'a> ExecutionRunner<'a> {
if final_7.is_file() {
let syn_t0 = Instant::now();
let _ = tokio::fs::copy(&final_7, model_dir.join("fort.8")).await;
let _ = tokio::fs::remove_file(model_dir.join("fort.7")).await;
// H10synspec 输入文件(fort.8 大气 / fort.55 控制卡)写入失败不可静默吞掉。
// 历史上用 `let _ =` 忽略错误,磁盘满/inode 耗尽时 synspec 会读到旧/缺失的
// fort.8 产出垃圾光谱,却仍生成 .spec 并被归档为"成功"。现改为写入失败即记
// synspec_err 并跳过 synspec 阶段,避免产出物理上错误的谱。
if let Err(e) = tokio::fs::copy(&final_7, model_dir.join("fort.8")).await {
warn!("synspec 输入 fort.8 (大气) 复制失败,跳过 synspec: {}", e);
synspec_err = Some(format!("fort.8 copy failed: {}", e));
} else {
let _ = tokio::fs::remove_file(model_dir.join("fort.7")).await;
// Fort.55 parameter generation or symlink
let fort55_path = model_dir.join("fort.55");
let fort19_path = model_dir.join("fort.19");
// Fort.55 parameter generation or symlink
let fort55_path = model_dir.join("fort.55");
let fort19_path = model_dir.join("fort.19");
let _ = tokio::fs::remove_file(&fort55_path).await;
let _ = tokio::fs::remove_file(&fort19_path).await;
let _ = tokio::fs::remove_file(&fort55_path).await;
let _ = tokio::fs::remove_file(&fort19_path).await;
let default_cfg = SynspecConfig {
wstart: 1400.0,
wend: 1410.0,
imode: 0,
idrv: 50,
ifreq: 1,
rel_cutoff: 0.0001,
abs_cutoff: 0.01,
};
let fort55_text = generate_fort55_content(synspec_cfg.unwrap_or(&default_cfg));
let _ = tokio::fs::write(&fort55_path, &fort55_text).await;
#[cfg(unix)]
{
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 default_cfg = SynspecConfig {
wstart: 1400.0,
wend: 1410.0,
imode: 0,
idrv: 50,
ifreq: 1,
rel_cutoff: 0.0001,
abs_cutoff: 0.01,
};
let fort55_text = generate_fort55_content(synspec_cfg.unwrap_or(&default_cfg));
if let Err(e) = tokio::fs::write(&fort55_path, &fort55_text).await {
warn!("synspec 输入 fort.55 (控制卡) 写入失败,跳过 synspec: {}", e);
synspec_err = Some(format!("fort.55 write failed: {}", e));
} else {
#[cfg(unix)]
{
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() {
if synspec_err.is_none() && 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?
@@ -440,7 +510,8 @@ impl<'a> ExecutionRunner<'a> {
.spawn()?;
let synspec_timeout_sec = 600_u64.min(timeout_sec);
let status_res = run_child_async_with_timeout(child, synspec_timeout_sec).await;
let status_res =
run_child_async_with_timeout(child, synspec_timeout_sec, shutdown.clone()).await;
let rc = match status_res {
Ok(st) => st.code().unwrap_or(-1),
Err(e) => {
+95 -6
View File
@@ -10,29 +10,118 @@ pub struct SeedMatch {
pub const MAX_GLOBAL_SEED_DISTANCE: f64 = 3.0;
/// CNO 有向距离:富金属方向(目标比种子富)重罚,贫金属方向(目标比种子贫)轻罚。
///
/// **数据标定依据**——对历史 1191 个真实 seed_step(种子,目标)配对的成败统计:
/// - 贫金属方向(种子更富、目标往贫走,delta=目标−种子 < 0):成功率 **4254%**
/// - 富金属方向(目标更富、delta > 0):成功率仅 **311%**
/// (每个 loghe 分层该规律独立成立,he=−4 时贫方向 54% vs 富方向 3%,差 18 倍)
///
/// **物理解释**:从高金属丰度的收敛解出发**减少**金属(贫方向)是稳定微扰;
/// 反过来从贫金属种子**增加**金属(富方向),新增的紫外谱线辐射驱动会破坏已建立的
/// 辐射平衡,导致大气发散出现 NaN。这是 NLTE 辐射流体计算的已知特性。
///
/// **回测**:对 871 个失败点的 exact_family 种子池回测,旧版绝对值距离有 55% 选了糟糕的
/// 富方向种子;改用此非对称距离后,81% 选贫方向,净改善 314 个点(旧选富→新选贫)。
fn directed_cno_distance(cand: &GridPointParams, target: &GridPointParams) -> f64 {
const RICH_PENALTY: f64 = 4.0; // 目标比种子富 → 该方向微扰不稳定,重罚
const POOR_PENALTY: f64 = 1.0; // 目标比种子贫 → 该方向微扰稳定,轻罚
// delta = target cand:正 = 目标更富(坏方向),负 = 目标更贫(好方向)
let penalize = |delta: f64| if delta > 0.0 { delta * RICH_PENALTY } else { -delta * POOR_PENALTY };
penalize(target.logc.value() - cand.logc.value())
+ penalize(target.logn.value() - cand.logn.value())
+ penalize(target.logo.value() - cand.logo.value())
}
pub fn calculate_seed_distance(cand: &GridPointParams, target: &GridPointParams) -> (bool, f64) {
let d_teff = (cand.teff.value() - target.teff.value()).abs();
let d_logg = (cand.logg.value() - target.logg.value()).abs();
let d_loghe = (cand.loghe.value() - target.loghe.value()).abs();
let d_cno = (cand.logc.value() - target.logc.value()).abs()
+ (cand.logn.value() - target.logn.value()).abs()
+ (cand.logo.value() - target.logo.value()).abs();
// 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)
// 同物理族内仅 CNO 不同:用有向距离优先匹配贫金属方向的种子(见 directed_cno_distance)。
(true, directed_cno_distance(cand, target))
} else {
// 距离公式物理意义与标定阐释:
// 在恒星非局部热力学平衡(NLTE)辐射流体力学与光谱大气计算中,不同物理自由度对于迭代收敛过程的基本影响层级截然相反:
// 1. Teff (有效温度) 通常达数千至数十万 K,主导连续谱黑体势函数与强激发电离步阶,故除以 5000.0 归一化为基底主控距离量;
// 2. logg (表面重力加速度) 对静力学与辐射光致压差梯度的平衡破坏力极烈,压强差稍高会触发极大激波不平衡,因此乘上 2.0 予以最高维权惩罚;
// 3. loghe (氦丰度) 对自由电子密度与热库贡献次于 H-He 电离梯度,乘 0.5 作为次要控制项;
// 4. CNO 金属元素虽然影响紫外谱线辐射驱动但整体状态基本可作次优微扰微增系数看待,乘 0.1;
// 通过上述尺度正态映射可挑选得到高收敛继承性的初态迭代种子模型
// 4. CNO 金属元素影响紫外谱线辐射驱动,其方向性同样关键(富方向不稳定,见 directed_cno_distance),
// 乘 0.1 归一化后纳入全局距离
let d_cno = directed_cno_distance(cand, target);
let global_d = (d_teff / 5000.0) + (d_logg * 2.0) + (d_loghe * 0.5) + (d_cno * 0.1);
(false, global_d)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::models::GridAxisValue;
fn params(teff: f64, logg: f64, loghe: f64, logc: f64, logn: f64, logo: f64) -> GridPointParams {
GridPointParams {
teff: GridAxisValue::from_value(teff),
logg: GridAxisValue::from_value(logg),
loghe: GridAxisValue::from_value(loghe),
logc: GridAxisValue::from_value(logc),
logn: GridAxisValue::from_value(logn),
logo: GridAxisValue::from_value(logo),
}
}
/// exact_family 内,同 CNO 差幅度下,贫方向距离应远小于富方向距离(4:1)。
#[test]
fn test_directed_cno_distance_favors_poor_metal_direction() {
let seed = params(40000.0, 6.0, -2.0, -2.0, -2.0, -2.0); // 种子 CNO=-6
// 贫方向:目标 CNO=-7(种子更富,目标往贫走),|Δ|=1
let poor_target = params(40000.0, 6.0, -2.0, -3.0, -2.0, -2.0);
// 富方向:目标 CNO=-5(目标更富),|Δ|=1,同一分量、同幅度
let rich_target = params(40000.0, 6.0, -2.0, -1.0, -2.0, -2.0);
let d_poor = directed_cno_distance(&seed, &poor_target);
let d_rich = directed_cno_distance(&seed, &rich_target);
assert!(d_poor < d_rich, "贫方向应更近");
assert!(
(d_rich / d_poor - 4.0).abs() < 1e-9,
"富/贫方向同幅度距离比应为 RICH/POOR=4.0,实际 {} / {} = {}",
d_rich,
d_poor,
d_rich / d_poor
);
}
/// exact_family 选种应在多个候选中优先选贫方向种子,即便其 CNO 绝对差更大。
#[test]
fn test_exact_family_prefers_poor_direction_seed() {
let target = params(40000.0, 6.0, -2.0, -2.0, -2.0, -2.0); // 目标 CNO=-6
// 候选A:富方向种子(目标比种子富),CNO 绝对差=1
let rich_seed = params(40000.0, 6.0, -2.0, -3.0, -2.0, -2.0); // CNO=-7, 目标更富
// 候选B:贫方向种子(目标比种子贫),CNO 绝对差=2(更大)
let poor_seed = params(40000.0, 6.0, -2.0, -1.0, -1.0, -2.0); // CNO=-4, 目标更贫
let (_, d_rich) = calculate_seed_distance(&rich_seed, &target);
let (_, d_poor) = calculate_seed_distance(&poor_seed, &target);
// 贫种子虽 CNO 绝对差更大(2 vs 1),但因方向有利,距离应更小
assert!(
d_poor < d_rich,
"贫方向种子距离 {} 应小于富方向 {}(即便 CNO 绝对差更大)",
d_poor,
d_rich
);
}
/// 跨 teff(非 exact)时仍标记 global,且方向性体现在 global_d 里。
#[test]
fn test_global_branch_marks_non_exact_and_keeps_direction() {
let target = params(40000.0, 6.0, -2.0, -2.0, -2.0, -2.0);
// 跨 teff 20000K(超过 exact 容忍 5000K)→ global
let cand = params(60000.0, 6.0, -2.0, -2.0, -2.0, -2.0);
let (is_exact, d) = calculate_seed_distance(&cand, &target);
assert!(!is_exact, "跨 teff 20000K 应为 global 分支");
assert!(d > 0.0);
}
}
+101 -4
View File
@@ -179,7 +179,30 @@ impl SqliteTaskQueue {
Err(e) => return Err(e.into()),
};
let task: TaskSpec = serde_json::from_str(&payload)?;
// 毒消息防护:payload 解析失败时,若直接返回 Err 会让事务回滚,
// 该行保持 pending,下次 pop_task 因 ORDER BY 确定性排序会再次
// 选中同一行 → 无限循环,阻塞全部任务分发(毒消息死循环)。
// 修复:把不可解析的行标记为 'dead_letter' 并提交事务,使其脱离
// pending 候选集,让出队能继续前进。dead_letter 行不再被
// pop_task/requeue_stale_tasks 触及,由运维侧事后审计清理。
let task: TaskSpec = match serde_json::from_str(&payload) {
Ok(t) => t,
Err(e) => {
tracing::error!(
"任务队列 payload 解析失败,标记为死信以避免毒消息死循环: \
task_id={} 错误: {}",
task_id,
e
);
tx.execute(
"UPDATE task_queue SET status = 'dead_letter' WHERE task_id = ?1",
params![task_id],
)?;
tx.commit()?;
// 死信已落库,循环继续尝试 pop 下一行 pending 任务。
continue;
}
};
// 记录任务归属:claim 时写入领用方 node_id,供 report 阶段校验,
// 杜绝「节点 A 领用、节点 B 上报」的跨节点伪造结果投毒。
@@ -322,9 +345,15 @@ impl SqliteTaskQueue {
Ok(())
}
/// 仅清理指定工作流的排队任务。
/// 仅清理指定工作流的**未被领取**的排队任务。
///
/// 用于 stop_workflow 按工作流隔离清理,避免在多工作流场景下误清其他工作流的任务。
/// 用于 stop_workflow / initialize_grid 按工作流隔离清理,避免在多工作流场景下
/// 误清其他工作流的任务。
///
/// 只删 `pending`(未被领取)行,保留 `claimed` 行:节点已领取正在执行的任务
/// 必须保留领用凭证,否则节点上报时 verify_task_claim 找不到记录 → 403 →
/// 计算结果丢失、网格点永久卡在 running(#6 修复)。
/// `claimed` 行会在上报成功后由 remove_task 自然清理。
pub async fn clear_queue_by_workflow(&self, workflow_name: &str) -> Result<()> {
let pool = self.pool.clone();
let wf_owned = workflow_name.to_string();
@@ -333,7 +362,7 @@ impl SqliteTaskQueue {
.get()
.map_err(|e| anyhow::anyhow!("Queue DB pool error: {}", e))?;
conn.execute(
"DELETE FROM task_queue WHERE workflow_name = ?1",
"DELETE FROM task_queue WHERE workflow_name = ?1 AND status = 'pending'",
params![wf_owned],
)?;
Ok(())
@@ -517,6 +546,74 @@ mod tests {
}
/// 出队同工作流内按 wave ASC:低难度 wave 优先,即使它 created_at 更晚。
/// 毒消息防护:payload 损坏的 pending 行不应无限阻塞出队。
///
/// 旧实现中 pop_task 遇到不可解析的 payload 会回滚事务,该行保持 pending,
/// 因 ORDER BY wave, created_at 确定性排序,下次仍选中同一行 → 死循环阻塞整个队列。
/// 现实现把不可解析行标记为 'dead_letter' 并提交,使其脱离 pending 候选集,
/// 后续 pending 任务能正常出队。
#[tokio::test]
async fn test_pop_task_skips_poison_message() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("poison.db");
let queue = SqliteTaskQueue::new(&db_path.to_string_lossy())
.await
.unwrap();
// 直接写一行 payload 损坏的 pending 任务(绕过 push_task 的合法序列化)
{
let pool = queue.pool.clone();
tokio::task::spawn_blocking(move || -> Result<()> {
let conn = pool.get().unwrap();
conn.execute(
"INSERT INTO task_queue (task_id, payload, status, created_at, wave)
VALUES ('poison-1', 'this is not valid json', 'pending', datetime('now'), 0)",
[],
)?;
Ok(())
})
.await
.unwrap()
.unwrap();
}
// 再推一个合法任务(wave=0,但 created_at 晚于 poison,故 pop 会先撞上 poison 行)
queue
.push_task(&mk_task("wf_ok", 0, "ok_point"))
.await
.unwrap();
// 第一次 pop:应跳过 poison(标记为 dead_letter)并返回合法任务
let popped = queue.pop_task("n1").await.unwrap();
assert!(
popped.is_some(),
"毒消息不应阻塞合法任务出队"
);
let task = popped.unwrap();
assert_eq!(task.point_name, "ok_point");
// 确认 poison 行已被标记为 dead_letter,不再处于 pending
{
let pool = queue.pool.clone();
let status: String = tokio::task::spawn_blocking(move || -> Result<String> {
let conn = pool.get().unwrap();
let s: String = conn.query_row(
"SELECT status FROM task_queue WHERE task_id = 'poison-1'",
[],
|r| r.get(0),
)?;
Ok(s)
})
.await
.unwrap()
.unwrap();
assert_eq!(status, "dead_letter", "毒消息应被标记为 dead_letter");
}
// 队列已空
assert!(queue.pop_task("n1").await.unwrap().is_none());
}
#[tokio::test]
async fn test_pop_wave_priority_within_workflow() {
let temp_dir = tempfile::tempdir().unwrap();
+84 -42
View File
@@ -5,7 +5,7 @@ use common::models::{ModelSummary, TaskSpec, TaskType};
use common::runner::ExecutionRunner;
use reqwest::Client;
use std::path::{Path, PathBuf};
use tracing::{info, warn};
use tracing::{debug, info, warn};
pub async fn execute_task(
client: &Client,
@@ -13,6 +13,7 @@ pub async fn execute_task(
runtime: &RuntimePaths,
work_dir: &Path,
task: &TaskSpec,
shutdown: Option<std::sync::Arc<std::sync::atomic::AtomicBool>>,
) -> Result<(ModelSummary, Option<Vec<u8>>)> {
info!(
"开始执行计算任务 {} (网格点: {})",
@@ -48,46 +49,80 @@ pub async fn execute_task(
let mut seed_atmos_path: Option<PathBuf> = None;
// 2. If seed_step, download seed .7 file from server using atomic file rename
if task.task_type == TaskType::SeedStep {
if let Some(ref seed_name) = task.seed_point_name {
let seed_url = format!("{}/api/seed/{}", server_url, seed_name);
info!("正在从服务端下载种子大气文件: {}", seed_url);
match client.get(&seed_url).send().await {
Ok(resp) if resp.status().is_success() => {
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?;
// 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?;
seed_atmos_path = Some(final_seed_path);
}
}
Ok(resp) => {
warn!("下载种子文件失败: HTTP {}", resp.status());
}
Err(e) => {
warn!("下载种子文件失败: {}", e);
}
}
}
}
// 3. Isolated task sandbox directory per slot to prevent multi-slot race collisions
// 提前创建 per-slot 隔离沙盒目录:种子下载后需复制一份私有副本进沙盒(见下方),
// 故沙盒必须先于种子下载就绪。
let slot_work_dir = work_dir.join(format!("task_{}", task.task_id));
tokio::fs::create_dir_all(&slot_work_dir).await?;
// 2. If seed_step, download seed .7 file from server using atomic file rename.
// 调度入口已改为「冷启动优先」,SeedStep 仅作为冷启动失败后的救援任务出现,服务端
// 派发前已经 find_best_seed_from_db 确认种子存在于 DB。因此下载失败(缺种子名/HTTP
// 错误/网络异常/响应体读取失败)按硬错误处理,直接失败该任务(fail-fast),不再无种
// 子继续运行:default_seed_chain 首阶段 seed_nc 的 ltgray="F" 依赖 fort.8,无种子时
// Tlusty 在无初始大气下运行必然崩溃。任务失败后由服务端走既有上报路径,
// has_seed_step_attempt 阻止重复回退,网格点保持 failed 终态。
if task.task_type == TaskType::SeedStep {
let seed_name = task.seed_point_name.as_deref().ok_or_else(|| {
anyhow::anyhow!(
"SeedStep 任务 {} 缺少 seed_point_name,无法热启动",
task.point_name
)
})?;
let seed_url = format!("{}/api/seed/{}", server_url, seed_name);
info!("正在从服务端下载种子大气文件: {}", seed_url);
let resp = client.get(&seed_url).send().await.map_err(|e| {
anyhow::anyhow!(
"SeedStep 任务 {} 下载种子文件 {} 失败: {}",
task.point_name,
seed_url,
e
)
})?;
if !resp.status().is_success() {
anyhow::bail!(
"SeedStep 任务 {} 下载种子文件 {} 失败: HTTP {}",
task.point_name,
seed_url,
resp.status()
);
}
let bytes = resp.bytes().await.map_err(|e| {
anyhow::anyhow!(
"SeedStep 任务 {} 读取种子文件 {} 响应体失败: {}",
task.point_name,
seed_url,
e
)
})?;
let temp_seed_dir = work_dir.join(".seed_cache");
tokio::fs::create_dir_all(&temp_seed_dir).await?;
// 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?;
// 关键:复制一份种子到本任务沙盒私有副本,让 seed_atmos_path
// 指向私有副本而非共享缓存。此后 runner 的 current_seed 全程
// 只引用沙盒内文件,与 .seed_cache 完全解耦——这样 LRU 清理
// (含并发竞争)即便删掉该缓存文件,也不会破坏正在使用该种子
// 的 in-flight 任务。.seed_cache 退化为纯粹的下载去重缓存。
let private_seed = slot_work_dir.join("seed_atmos.seed.7");
tokio::fs::copy(&final_seed_path, &private_seed).await?;
seed_atmos_path = Some(private_seed);
}
// 3. slot_work_dir 已在种子下载前提前创建,种子私有副本亦已落盘于沙盒内。)
let runner = ExecutionRunner::new(runtime, slot_work_dir.clone());
let summary = runner
.run_model_with_timeout(
@@ -96,10 +131,13 @@ pub async fn execute_task(
// 而非 task.params.model_name()(后者经 DB REAL 列回读已丢精度 "5.0"→"5")。
&task.point_name,
task.task_type.clone(),
// custom_chain 恒为 None:执行链由 task_type 决定(ColdRun→default_cold_chain、
// SeedStep→default_seed_chain),节点端不再做「缺种子回退冷启动链」的降级。
None,
seed_atmos_path.as_deref(),
None,
task.timeout_sec,
shutdown,
)
.await?;
@@ -368,10 +406,14 @@ pub async fn cleanup_seed_cache(seed_dir: &Path) {
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());
match tokio::fs::remove_file(&path).await {
Ok(()) => info!("LRU 清理种子缓存文件: {}", path.display()),
// 文件已被并发删除(多 slot 同时清理同一批最旧文件):清理目标已达成,
// 视为成功,不再误报 warn。其他真实 IO 错误才需要告警。
Err(e) if e.kind() == std::io::ErrorKind::NotFound => {
debug!("种子缓存文件已被并发删除: {}", path.display())
}
Err(e) => warn!("清理种子缓存文件 {} 失败: {}", path.display(), e),
}
}
}
+26 -2
View File
@@ -33,16 +33,31 @@ async fn main() -> Result<()> {
None => {
info!("本地未发现 node token,准备向服务端提交注册申请并等待管理员审批...");
let public_client = build_client_with_token(None);
let issued = NodeWorker::register_and_fetch_token(
// 读取上次注册持久化的 registration_secret(若存在)。
// 场景:节点曾注册获批、本地 .node_token 因 reissue 失效被删除后重启。
// 重新注册时服务端对"已存在节点"不再下发新 secret,但本地旧 secret 仍有效,
// 复用它才能在 check_status 取回 reissue 产生的新 token(否则取不回,H8)。
let secret_path = runtime_dir.join(".node_secret");
let existing_secret = read_node_token(&secret_path);
let (issued, reg_secret) = NodeWorker::register_and_fetch_token(
&public_client,
&node_cfg.server_url,
&node_cfg.node_id,
existing_secret.clone(),
)
.await
.context("向服务端提交申请或获取专属 token 失败")?;
write_node_token(&token_path, &issued)?;
info!("已持久化获批的专属 node token 到 {}", token_path.display());
// 持久化 registration_secretH8):服务端对新节点会下发一次性 secret,
// 对已存在节点不下发(复用本地旧 secret)。无论新旧,落盘以便下次重启复用。
let secret_to_persist = reg_secret.or(existing_secret);
if let Some(secret) = &secret_to_persist {
if let Err(e) = write_node_token(&secret_path, secret) {
warn!("持久化 registration_secret 失败(不影响本次启动): {}", e);
}
}
issued
}
};
@@ -91,8 +106,17 @@ fn write_node_token(path: &Path, token: &str) -> Result<()> {
}
/// 构造带 Authorization: Bearer 头的 reqwest client。
///
/// 设置连接与请求超时,避免心跳/领用/种子下载/上报在网络层挂起时无限等待。
/// 旧实现无任何超时:一旦连接 stall,心跳永远不发,节点被误判离线;
/// 上报挂起则活动 slot 永不归还。
fn build_client_with_token(token: Option<&str>) -> Client {
let mut builder = Client::builder();
let mut builder = Client::builder()
// 连接建立阶段(TCP 握手 + TLS)超时,防止对端不可达时长时间挂起。
.connect_timeout(std::time::Duration::from_secs(15))
// 单个请求整体超时。种子下载(.7 大气文件)可能较大,故给到 10 分钟;
// 心跳/领用/上报这类小请求会远早于此完成。
.timeout(std::time::Duration::from_secs(600));
if let Some(t) = token {
let mut headers = reqwest::header::HeaderMap::new();
if let Ok(val) = reqwest::header::HeaderValue::from_str(&format!("Bearer {}", t)) {
+129 -14
View File
@@ -12,6 +12,26 @@ use std::sync::Arc;
use tokio::time::{sleep, Duration};
use tracing::{info, warn};
/// 活动 slot 计数的 RAII 守卫:仅负责 drop 时 -1。
///
/// **自增的时机**:领用主循环在 `tokio::spawn` 之前、紧跟容量检查之后同步调用
/// `active_slots.fetch_add(1, ...)`,与「检查是否还有空闲容量」紧邻成原子操作,
/// 杜绝「检查通过 → spawn 排队 → 再回循环检查时计数尚未自增」的竞态窗口
/// (曾导致满载时瞬时 `active_slots` 冲到 `max_slots + 1`,前端显示「预设 + 1」)。
///
/// **归还的可靠性**:自增后立即构造本守卫并移入 spawned future,即便 future 内
/// 任意代码 panicexecute_task / save_result_artifacts / 归档清理),Drop 也会执行,
/// 确保计数始终归还,不会像「future 末尾手动 fetch_sub」那样泄漏到死锁。
struct SlotGuard {
counter: Arc<AtomicI32>,
}
impl Drop for SlotGuard {
fn drop(&mut self) {
self.counter.fetch_sub(1, Ordering::AcqRel);
}
}
/// 领用请求的归一化结果。
///
/// 区分「被管理员停用」与「暂无任务」:前者节点保持存活、空闲待命(拉长轮询),
@@ -50,11 +70,19 @@ impl NodeWorker {
/// 仅注册并领取专属 token(供 main.rs 在本地无 token 时调用)。
/// 支持免凭据申请注册并轮询等待管理员在 Web Dashboard 上点击同意。
///
/// `existing_secret`:本地持久化的旧 registration_secret(若有)。重新注册(节点已存在)
/// 时服务端不下发新 secret,此时复用旧 secret 才能在 check_status 取回 reissue 产生的新
/// tokenH8:防止已存在节点 reissue 后取不回 token)。
///
/// 返回 (node_token, registration_secret)registration_secret 是注册时服务端下发的一次性
/// 凭据(新节点)或复用的旧 secret(已存在节点),取走待发 token 时须回传。
pub async fn register_and_fetch_token(
client: &Client,
server_url: &str,
node_id: &str,
) -> Result<String> {
existing_secret: Option<String>,
) -> Result<(String, Option<String>)> {
info!(
"正在向服务端 {} 提交计算节点 {} 的注册申请...",
server_url, node_id
@@ -77,10 +105,16 @@ impl NodeWorker {
let json: Value = resp.json().await?;
let status = json.get("status").and_then(|v| v.as_str()).unwrap_or("");
let registration_secret = json
.get("registration_secret")
.and_then(|v| v.as_str())
.map(|s| s.to_string())
// 已存在节点重新注册:服务端不下发新 secret,复用本地旧 secretH8)。
.or_else(|| existing_secret.clone());
if status == "approved" {
if let Some(t) = json.get("node_token").and_then(|v| v.as_str()) {
return Ok(t.to_string());
return Ok((t.to_string(), registration_secret));
}
}
@@ -93,7 +127,10 @@ impl NodeWorker {
loop {
sleep(Duration::from_secs(5)).await;
let check_req = serde_json::json!({ "node_id": node_id });
let check_req = serde_json::json!({
"node_id": node_id,
"registration_secret": registration_secret,
});
let resp = match client
.post(format!("{}/api/node/check_status", server_url))
.json(&check_req)
@@ -123,7 +160,7 @@ impl NodeWorker {
"🎉 节点 {} 已成功获取管理员授权!专属访问 Token 接收完成。",
node_id
);
return Ok(token.to_string());
return Ok((token.to_string(), registration_secret));
}
} else if check_status == "rejected" {
anyhow::bail!("节点 {} 的注册申请已被管理员拒绝或清理", node_id);
@@ -253,20 +290,80 @@ impl NodeWorker {
tokio::fs::create_dir_all(&work_dir).await?;
let result_dir = PathBuf::from(&self.config.result_dir);
// 启动时清理上次崩溃/强杀遗留的 task_* 沙盒目录。
// 背景:节点非正常退出(SIGKILL / 断电)时,正在跑的任务沙盒不会清理,
// 每次重启都留下一份完整的 TLUSTY 工作目录(含 data 软链 + fort 文件,
// 单任务可达数十~上百 MB)。.seed_cache 和 result_dir 都有 LRU 治理,
// 唯独 work_dir 无任何清理,长期累积会写满磁盘导致新任务失败。
// 此时节点刚启动、无任何活动任务,删除 task_* 子目录是安全的。
if let Ok(mut rd) = tokio::fs::read_dir(&work_dir).await {
while let Ok(Some(entry)) = rd.next_entry().await {
let name = entry.file_name();
let name_str = name.to_string_lossy();
if name_str.starts_with("task_") {
let p = entry.path();
match tokio::fs::remove_dir_all(&p).await {
Ok(_) => warn!("启动清理上次残留沙盒: {}", p.display()),
Err(e) => warn!("启动清理残留沙盒 {} 失败: {}", p.display(), e),
}
}
}
}
let shutting_down = Arc::new(std::sync::atomic::AtomicBool::new(false));
let shutdown_signal = shutting_down.clone();
// 信号处理:同时监听 SIGINTCtrl+C)与 SIGTERM。
// 生产编排(Docker stop / k8s pod 终止 / systemd stop / kill <pid>)发送的是
// **SIGTERM** 而非 SIGINT。历史上只监听 ctrl_c() 会导致收到 SIGTERM 时信号处理器
// 完全不触发 → shutting_down 永远 false → 主循环继续领任务 → grace period 后被
// SIGKILL 强杀 → 正在跑的长任务结果直接丢失。
// 此 shutdown 标志同时传给 runner,让在途 Fortran 子进程在收到信号后被立即 kill,
// 避免等满 30s 优雅期仍超时强退(C5)。在途任务结果会丢失,由服务端 stale 重投兜底。
tokio::spawn(async move {
if tokio::signal::ctrl_c().await.is_ok() {
info!("收到 Ctrl+C 终止信号,停止领用新任务,准备优雅退出 (再次按 Ctrl+C 可强制立即退出)...");
shutdown_signal.store(true, Ordering::Release);
// 二次 Ctrl+C 强行立即退出
if tokio::signal::ctrl_c().await.is_ok() {
warn!("再次收到 Ctrl+C 终止信号,强行立即中断退出!");
use tokio::signal::unix::{signal, SignalKind};
let mut sigterm = match signal(SignalKind::terminate()) {
Ok(s) => s,
Err(e) => {
warn!("无法注册 SIGTERM 监听: {}(将仅响应 SIGINT", e);
// 退化为只监听 SIGINT:保留"首次信号优雅退出 + 二次信号强退"完整能力。
if tokio::signal::ctrl_c().await.is_ok() {
info!("收到 Ctrl+C 终止信号,停止领用新任务,准备优雅退出 (再次按 Ctrl+C 可强制立即退出)...");
shutdown_signal.store(true, Ordering::Release);
if tokio::signal::ctrl_c().await.is_ok() {
warn!("再次收到 Ctrl+C,强行立即中断退出!");
std::process::exit(130);
}
}
return;
}
};
let mut sigint = match signal(SignalKind::interrupt()) {
Ok(s) => s,
Err(_) => {
warn!("无法注册 SIGINT 监听,仅响应 SIGTERM");
sigterm.recv().await;
info!("收到 SIGTERM 终止信号,停止领用新任务,准备优雅退出 (再次发送 SIGTERM 可强制立即退出)...");
shutdown_signal.store(true, Ordering::Release);
sigterm.recv().await;
warn!("再次收到 SIGTERM,强行立即中断退出!");
std::process::exit(130);
}
};
// 第一次信号:触发优雅退出(停止领用 + kill 在途 child
tokio::select! {
_ = sigterm.recv() => info!("收到 SIGTERM 终止信号,停止领用新任务,准备优雅退出 (再次发送 SIGTERM/SIGINT 可强制立即退出)..."),
_ = sigint.recv() => info!("收到 Ctrl+C (SIGINT) 终止信号,停止领用新任务,准备优雅退出 (再次按 Ctrl+C 可强制立即退出)..."),
}
shutdown_signal.store(true, Ordering::Release);
// 第二次信号:强行立即退出
tokio::select! {
_ = sigterm.recv() => warn!("再次收到 SIGTERM,强行立即中断退出!"),
_ = sigint.recv() => warn!("再次收到 Ctrl+C,强行立即中断退出!"),
}
std::process::exit(130);
});
let mut was_disconnected = false;
@@ -285,7 +382,6 @@ impl NodeWorker {
info!("与服务端恢复网络连接,已自动重新上线并开始领用计算任务!");
was_disconnected = false;
}
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();
@@ -293,11 +389,22 @@ impl NodeWorker {
let work_dir = work_dir.clone();
let result_dir = result_dir.clone();
let slots_counter = self.active_slots.clone();
let shutdown = shutting_down.clone();
// 在 spawn 前同步自增,与容量检查紧邻成原子操作:避免
// 「检查通过 → spawn 排队 → 回循环再检查时计数尚未自增」的竞态
// 导致瞬时 active_slots 冲到 max_slots + 1。
slots_counter.fetch_add(1, Ordering::AcqRel);
// RAII 守卫仅负责 drop 时 -1(含 panic 展栈),杜绝 spawned future
// panic 导致的活动 slot 永久泄漏。
let slot_guard = SlotGuard { counter: slots_counter.clone() };
tokio::spawn(async move {
// 把守卫移入 future,确保任务结束(含 panic)时归还 slot。
let _slot_guard = slot_guard;
let slot_work_dir = work_dir.join(format!("task_{}", task.task_id));
let res =
execute_task(&client, &server_url, &runtime, &work_dir, &task)
execute_task(&client, &server_url, &runtime, &work_dir, &task, Some(shutdown.clone()))
.await
.map_err(|e| e.to_string());
@@ -353,7 +460,7 @@ impl NodeWorker {
);
}
}
slots_counter.fetch_sub(1, Ordering::AcqRel);
// _slot_guard 在此作用域结束时 drop,归还活动 slot 计数。
});
}
Ok(ClaimOutcome::Empty) => {
@@ -432,7 +539,15 @@ impl NodeWorker {
std::process::exit(1);
}
if status.is_server_error() {
// 服务端 5xx(DB 故障 / 内部异常)≠「暂无任务」。归一化为 Empty 会让节点静默
// 空转且日志看不出"服务端故障",队列里有任务却不出。走 Err 触发 was_disconnected
// 退避路径,日志里能看到"请求领用任务出错",便于运维定位。
anyhow::bail!("服务端临时故障 (HTTP {})", status);
}
if !status.is_success() {
// 其余 4xx(如 429 限流)按 Empty 处理,下次轮询重试。
return Ok(ClaimOutcome::Empty);
}
+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();