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:
@@ -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_middleware(Role::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": "已登出"
|
||||
})),
|
||||
)
|
||||
}
|
||||
|
||||
@@ -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)。
|
||||
///
|
||||
/// 流程:
|
||||
|
||||
@@ -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,
|
||||
})),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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
@@ -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(¶ms, &seed_path.to_string_lossy())
|
||||
.insert_seed_named(&name, ¶ms, &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((
|
||||
|
||||
@@ -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 缺省 → None(SQL 不拼 LIMIT,返回全量)。联合分析依赖完整数据,
|
||||
// 故不设上限;点表明细分页传显式 limit(≤500)走分页。
|
||||
limit: pq.limit.map(|l| l.max(1)),
|
||||
offset: pq.offset.unwrap_or(0).max(0),
|
||||
};
|
||||
|
||||
|
||||
+509
-115
@@ -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 → 7;teff=37499 → 7;teff=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 列(H8:check_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_secret(H8)并插入待审批状态。
|
||||
// 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 查询前的缓存 generation,miss 落 DB。
|
||||
// generation 用于回填时的 TOCTOU 终极防护:若 DB 查询与回填之间发生过
|
||||
// invalidate(reissue),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 反查缓存。在 issue(token 轮换使旧 token 失效)时调用。
|
||||
/// 清空全部 token 反查缓存并自增 generation。在 issue(token 轮换使旧 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_a:1 pending + 1 imported
|
||||
assert_eq!(all["seed_step_converged"], 1);
|
||||
// 单工作流 wf_a:1 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(®).await.unwrap();
|
||||
let (_, _, secret) = db.register_node(®).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:不应返回 token(H8 防护)
|
||||
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
@@ -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
|
||||
};
|
||||
|
||||
@@ -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_fallback(seeds 仍是全局共享的物理资源池)。
|
||||
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_limit(SELECT)+ update_grid_status(UPDATE)
|
||||
// 的分离操作,杜绝两个并发调度调用 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 scheduling(seeds 全局共享)
|
||||
let (task_type, seed_name) = match self.db.find_best_seed_from_db(¶ms).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!(
|
||||
|
||||
@@ -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(®).await.unwrap();
|
||||
let (_, _, secret) = db.register_node(®).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_points:converged + imported 手段 + 旧版 elapsed_sec 已落库
|
||||
// grid_points:converged + 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();
|
||||
|
||||
|
||||
Reference in New Issue
Block a user