feat(all): 源精度命名体系、工作流可观测台、节点停用管理与白名单归档

核心变更:

  1. GridAxisValue 源精度命名
     - 新增 GridAxisValue 类型,携带 f64 数值 + YAML 源书写文本(Deref 透明兼容算术)
     - config.rs 绕过 serde_yaml 归一化,逐 token 捕获轴值原文(logg: 5.0 → g5.0)
     - runner/executor/scheduler 全链路改用 DB TEXT 列权威 point_name,
       修复 REAL 列回读丢精度导致的 model_name 错配

  2. 工作流执行可观测台
     - 新增 stats/progress/points 三组 API(进度时间序列、经验速率 ETA、
       停滞预警、逐点明细分页、收敛性热力图数据)
     - 新增 workflow_progress_snapshots 表 + tasks/grid_points 耗时列
     - runner 携带 last_iter/worst_depth/n_depths 进 conv.json
     - 前端新增 hash 路由、工作流详情页(概览/网格点/收敛分析三 Tab)、YAML 编辑器

  3. 节点停用/启用管理
     - 新增 disabled 状态 + disable/enable API;停用节点保持心跳但停止分发,
       worker 空闲待命而非退出;移除 revoke API,token 失效统一走重发覆盖;
       移除 host_name 字段

  4. 白名单结果归档
     - 新增 result_filter 模块,只归档有语义产物,丢弃 Tlusty 中间单元(~2MB/模型)
     - executor 原子写入归档 + 200 点 LRU 上限

  5. 历史数据导入
     - sync_seeds 重写为 import_results:经 /admin/import_seed 标记 converged +
       按新版命名迁移产物树

  6. 部署与目录重规划
     - data/results→seeds、data/archive→result + migrate_data_dirs.sh
     - deploy.sh 增强(SSH 复用、Profile、远程 env);Dockerfile 瘦身

  7. 文档同步更新 api/database/architecture/deployment
This commit is contained in:
fmq
2026-07-31 01:34:05 +08:00
parent b91f1e4fa5
commit 1bfa240cb0
73 changed files with 12332 additions and 1608 deletions
+74 -35
View File
@@ -1,12 +1,11 @@
//! 管理 APIAdmin 角色)。
//!
//! 提供 node 凭据的可视化与运维操作,供 Dashboard 管理界面调用:
//! - 列出所有节点及其凭据状态(在线/token 是否有效/吊销/颁发时间)
//! - 吊销指定节点的专属 token(立即失效,不影响其他节点)
//! - 列出所有节点及其凭据状态(在线/token 是否有效/颁发时间)
//! - 重新颁发指定节点的专属 token(返回新明文,旧 token 失效)
//!
//! 这些端点均要求 Admin 角色(见 mod.rs 授权矩阵),node 自身无权操作他人或自身凭据,
//! 从而保证「吊销/重发」是管理员主动行为,避免被攻陷节点篡改凭据体系。
//! 从而保证「重发」是管理员主动行为,避免被攻陷节点篡改凭据体系。
use super::{is_valid_node_id, AppState};
use axum::{
@@ -31,37 +30,6 @@ pub async fn list_nodes(
}
}
/// POST /api/admin/nodes/:node_id/revoke — 吊销指定节点的专属 token。
///
/// 吊销后该 node 的现有 token 立即失效,须重新走注册流程领取新 token。
/// 操作幂等:对无凭据记录或已吊销的节点调用不会报错。
pub async fn revoke_node(
State(state): State<AppState>,
AxumPath(node_id): AxumPath<String>,
) -> Result<impl IntoResponse, crate::api::AppError> {
// node_id 白名单校验,防止注入或异常输入(与 register_node 的 node_id 来源口径一致)
if !is_valid_node_id(&node_id) {
return Err(crate::api::AppError::BadRequest(
"非法的节点 ID 参数".to_string(),
));
}
match state.db.revoke_node_token(&node_id).await {
Ok(_) => {
info!("管理员已吊销节点 {} 的专属 token", node_id);
Ok((
StatusCode::OK,
Json(
json!({ "success": true, "message": format!("节点 '{}' 的 token 已吊销", node_id) }),
),
))
}
Err(e) => {
warn!("吊销节点 {} token 失败: {}", node_id, e);
Err(e.into())
}
}
}
/// POST /api/admin/nodes/:node_id/reissue — 重新颁发指定节点的专属 token。
///
/// 旧 token 立即失效,返回新 token 明文(仅此一次,DB 只存 hash)。
@@ -93,7 +61,12 @@ pub async fn reissue_node(
StatusCode::OK,
Json(json!({
"success": true,
"message": format!("节点 '{}' 的 token 已重新颁发,请将新 token 同步到该节点", node_id),
// 重发使旧 token 立即失效:node 进程内存里仍握着旧 token,下一次心跳/领用会 401 退出。
// 必须显式提示恢复方式,否则管理员易困惑"为什么重发后节点反而掉了"。
"message": format!(
"节点 '{}' 的 token 已重新颁发,旧 token 立即失效(该节点下次心跳/领用将 401 退出)。\n恢复方式(二选一):① 把下方新 token 写入该节点本地 .node_token 文件后重启节点;② 在节点机器删除 .node_token 后重启(会自动重注册取回新 token,限 1 天内有效)。",
node_id
),
"node_token": new_token,
})),
))
@@ -152,3 +125,69 @@ pub async fn reject_node(
Err(e) => Err(e.into()),
}
}
/// POST /api/admin/nodes/:node_id/disable — 手动停用一个在线/离线节点。
///
/// 停用后节点保持在线心跳(Dashboard 可见其存活),但 claim 不再向其分发任务,
/// worker 收到 `{"status":"disabled"}` 后会拉长轮询、空闲待命。可随时调用 `/enable` 恢复。
/// 仅 `online`/`offline` 节点可停用;对 `pending_approval`/`disabled` 调用返回 409。
pub async fn disable_node(
State(state): State<AppState>,
AxumPath(node_id): AxumPath<String>,
) -> Result<impl IntoResponse, crate::api::AppError> {
if !is_valid_node_id(&node_id) {
return Err(crate::api::AppError::BadRequest(
"非法的节点 ID 参数".to_string(),
));
}
match state.db.set_node_enabled(&node_id, false).await {
Ok(true) => {
info!("管理员已停用节点 {}(停止分发任务,保持空闲待命)", node_id);
Ok((
StatusCode::OK,
Json(
json!({ "success": true, "message": format!("节点 '{}' 已停用,不再分发任务", node_id) }),
),
))
}
Ok(false) => Err(crate::api::AppError::Conflict(format!(
"节点 '{}' 当前状态不支持停用(仅在线/离线节点可停用)",
node_id
))),
Err(e) => Err(e.into()),
}
}
/// POST /api/admin/nodes/:node_id/enable — 重新启用被手动停用的节点。
///
/// 将节点从 `disabled` 切为 `offline`,靠下一次心跳自然翻成 online 后恢复分发。
/// 仅 `disabled` 节点可启用;对其他状态调用返回 409(幂等保护)。
pub async fn enable_node(
State(state): State<AppState>,
AxumPath(node_id): AxumPath<String>,
) -> Result<impl IntoResponse, crate::api::AppError> {
if !is_valid_node_id(&node_id) {
return Err(crate::api::AppError::BadRequest(
"非法的节点 ID 参数".to_string(),
));
}
match state.db.set_node_enabled(&node_id, true).await {
Ok(true) => {
info!(
"管理员已重新启用节点 {}(下一次心跳后恢复分发任务)",
node_id
);
Ok((
StatusCode::OK,
Json(
json!({ "success": true, "message": format!("节点 '{}' 已重新启用,将在下一次心跳后恢复分发任务", node_id) }),
),
))
}
Ok(false) => Err(crate::api::AppError::Conflict(format!(
"节点 '{}' 当前状态不支持启用(仅已停用节点可启用)",
node_id
))),
Err(e) => Err(e.into()),
}
}
+17 -17
View File
@@ -29,11 +29,10 @@ pub struct AppState {
pub db: Database,
pub queue: Arc<SqliteTaskQueue>,
pub scheduler: Arc<GridScheduler>,
pub results_dir: String,
/// server 端种子库目录(原 results_dir)。下载种子时读此路径。**永不清理**。
pub seeds_dir: String,
/// 限流与密码防暴破限速器
pub rate_limiter: rate_limit::RateLimiter,
/// 兼容字段:Some 表示「已启用某种鉴权」,用于 main.rs 决定是否挂载鉴权中间件。
pub auth_token: Option<String>,
/// Admin 凭据(管理 Dashboard / workflow 写操作)。
pub admin_token: Option<String>,
/// 应急开关:跳过全部鉴权(仅本地调试)。
@@ -70,7 +69,7 @@ enum Role {
///
/// 设计依据(最小权限):
/// - Admin 写操作(workflow CRUD / start / stop / status / approve / reject)只对 admin token 开放。
/// - Node 运行态接口只认 node 专属 token(管理员在 Dashboard 审批后颁发,绑定 node_id,可吊销)。
/// - Node 运行态接口只认 node 专属 token(管理员在 Dashboard 审批后颁发,绑定 node_id,可重发轮换)。
/// - 注册端点 /node/register 和状态轮询 /node/check_status 为 Public 免凭据(提交申请 ➔ 待管理员审批)。
///
/// 注意:路径已去掉 `/api` 前缀(nest 挂载后中间件看到的 path 不含 nest 前缀)。
@@ -97,10 +96,14 @@ 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
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);
@@ -139,7 +142,7 @@ fn ct_eq_str(a: &str, b: &str) -> bool {
}
/// node_id 白名单:字母、数字、点、下划线、连字符,长度 1-128。
/// 用于 register_node / admin revoke / reissue 统一入口校验,与 Dashboard XSS 防护口径一致。
/// 用于 register_node / admin reissue / disable / enable 统一入口校验,与 Dashboard XSS 防护口径一致。
pub(crate) fn is_valid_node_id(id: &str) -> bool {
!id.is_empty()
&& id.len() <= 128
@@ -148,14 +151,6 @@ pub(crate) fn is_valid_node_id(id: &str) -> bool {
.all(|c| c.is_ascii_alphanumeric() || c == '.' || c == '_' || c == '-')
}
/// host_name 白名单:可打印 ASCII(排除控制字符),长度 1-128。
/// 防止 host_name 携带 HTML/控制字符进入管理 Dashboard 触发存储型 XSS 或污染显示。
pub(crate) fn is_valid_host_name(name: &str) -> bool {
!name.is_empty()
&& name.len() <= 128
&& name.chars().all(|c| c.is_ascii() && !c.is_ascii_control())
}
/// 从请求头提取凭据原文(支持 `Authorization: Bearer <t>` 与 `X-API-Key: <t>`)。
///
/// 安全:非 `Bearer ` 前缀的 Authorization 一律视为无 token(不再回退为裸头值比较),
@@ -240,8 +235,13 @@ pub async fn auth_middleware(
let mut sessions = state.admin_sessions.write().await;
let now = std::time::Instant::now();
sessions.retain(|_, expiry| *expiry > now);
if sessions.contains_key(&token) {
valid = true;
// 恒定时间比对:遍历全部 session key 逐个 ct_eq_str,不提前返回
// (与 admin token 的恒定时间口径一致,消除 key 存在性的时序旁路)。
// 容量受 MAX_ADMIN_SESSIONS100)约束,遍历开销可接受。
for k in sessions.keys() {
if ct_eq_str(&token, k) {
valid = true;
}
}
}
@@ -273,7 +273,7 @@ pub async fn auth_middleware(
}
None => (
StatusCode::UNAUTHORIZED,
"Unauthorized: invalid or revoked node token",
"Unauthorized: invalid or stale node token",
)
.into_response(),
}
+1 -6
View File
@@ -1,4 +1,4 @@
use super::{is_valid_host_name, is_valid_node_id, AppState, AuthenticatedNode};
use super::{is_valid_node_id, AppState, AuthenticatedNode};
use axum::{
extract::{Extension, State},
http::StatusCode,
@@ -20,11 +20,6 @@ pub async fn register_node(
"非法的节点 ID(仅允许字母、数字、点、下划线、连字符,长度 1-128)".to_string(),
));
}
if !is_valid_host_name(&req.host_name) {
return Err(crate::api::AppError::BadRequest(
"非法的主机名(仅允许可打印 ASCII,长度 1-128".to_string(),
));
}
// 已认证已拿到 Token 的节点刷新元数据配置
if let Some(Extension(ref auth)) = auth_node {
+1 -1
View File
@@ -28,7 +28,7 @@ pub async fn download_seed(
));
}
let seed_file_path = std::path::Path::new(&state.results_dir)
let seed_file_path = std::path::Path::new(&state.seeds_dir)
.join(&name)
.join(format!("{}.7", name));
+2 -1
View File
@@ -24,7 +24,8 @@ pub async fn get_status(
.get_grid_summary_stats(None)
.await
.unwrap_or(serde_json::json!({
"total": 0, "pending": 0, "running": 0, "converged": 0, "failed": 0
"total": 0, "pending": 0, "queued": 0, "running": 0, "converged": 0, "failed": 0,
"cold_run_converged": 0, "seed_step_converged": 0, "imported_converged": 0
}));
Ok(Json(json!({
+185 -9
View File
@@ -1,10 +1,11 @@
use super::{AppState, AuthenticatedNode};
use axum::{
extract::{Extension, Multipart, State},
extract::{Extension, Multipart, Query, State},
response::IntoResponse,
Json,
};
use common::models::{GridPointParams, ModelSummary, TaskReport, TaskStatus};
use serde::Deserialize;
use serde_json::json;
use std::path::Path;
use tokio::fs;
@@ -16,6 +17,23 @@ pub async fn claim_task(
State(state): State<AppState>,
Extension(auth_node): Extension<AuthenticatedNode>,
) -> Result<impl IntoResponse, crate::api::AppError> {
// 管理员手动停用拦截:被停用的节点保持在线但不再分发任务。
// 返回 HTTP 200 + {"status":"disabled"}(绝不能用 403 —— worker 见 403 会判定
// token 失效而 exit(1),停用是运维意图而非凭据失效,应让 worker 空闲待命)。
match state.db.is_node_disabled(&auth_node.node_id).await {
Ok(true) => {
return Ok((
StatusCode::OK,
Json(json!({"status": "disabled", "task": null})),
));
}
Ok(false) => {}
Err(e) => {
tracing::error!("查询节点停用状态异常: {}", e);
return Err(crate::api::AppError::Internal(e));
}
}
// 领用时记录任务归属:pop_task 写入 claimed_by_node_id
// report 阶段据此校验「上报者确为领用者」,杜绝跨节点伪造结果。
match state.queue.pop_task(&auth_node.node_id).await {
@@ -112,12 +130,7 @@ pub async fn report_task(
let name = report.point_name.clone();
if name.is_empty()
|| name.starts_with('.')
|| !name.chars().all(|c| {
c.is_ascii_alphanumeric() || c == '.' || c == '_' || c == '-' || c == '+' || c == '@'
})
{
if !super::workflow::is_valid_point_name(&name) {
warn!(
"拒绝可能包含路径穿越或特殊非常规编码号攻击的网格点名称请求: {}",
name
@@ -157,7 +170,7 @@ pub async fn report_task(
}
// 采用原子写入模式保持 conv.json 与核心二进制数据完整落地后才揭晓真实文件名
let model_dir = Path::new(&state.results_dir).join(&name);
let model_dir = Path::new(&state.seeds_dir).join(&name);
if fs::create_dir_all(&model_dir).await.is_ok() {
let conv_tmp = model_dir.join(format!("conv.json.{}.tmp", uuid::Uuid::new_v4().simple()));
let conv_path = model_dir.join("conv.json");
@@ -197,7 +210,7 @@ pub async fn report_task(
info!("网格点 {} 计算未成功完成,检查种子回退机制...", name);
if let Err(e) = state
.scheduler
.trigger_seed_step_fallback(&params, &workflow_name)
.trigger_seed_step_fallback(&params, &name, &workflow_name)
.await
{
warn!("网格点 {} 触发种子回退机制失败: {}", name, e);
@@ -218,3 +231,166 @@ fn extract_params(report: &TaskReport) -> Option<GridPointParams> {
.ok()
.map(|summary| summary.params)
}
/// `/admin/import_seed` 的查询参数。
#[derive(Debug, Deserialize)]
pub struct ImportSeedQuery {
/// 目标工作流名(导入到此工作流的 grid_points)。缺省归入 `imported` 工作流。
#[serde(default = "default_import_workflow")]
pub workflow: String,
}
fn default_import_workflow() -> String {
"imported".to_string()
}
/// 历史种子导入端点(Admin 鉴权)。
///
/// 供 `tools/import_results` 把旧版单机 `run_grid.py` 产物(`conv.json` + `.7` 大气文件)
/// 批量回灌进 DCTS。与 `/task/report` 的关键区别:
/// - **跳过任务归属校验**`verify_task_claim`):历史数据无领用语义,导入端点不经过
/// claim/report 队列,直接幂等落库。
/// - **`point_name` 取旧 `conv.json` 的 `name` 字段**Python `gen_input5.model_name`
/// 生成的源精度真名,如 `t20000_g5.0_...`),而非从数值重推——保证迁移逐字符保真。
/// - **真实 `max_relc`** 取自 `summary.final_max_relc`(旧版已记录),不硬编码。
///
/// 幂等:`upsert_grid_point` 用 `ON CONFLICT DO NOTHING``.7`/`conv.json` 原子覆盖写,
/// 可重复运行。
pub async fn import_seed(
State(state): State<AppState>,
Query(query): Query<ImportSeedQuery>,
mut multipart: Multipart,
) -> Result<impl IntoResponse, crate::api::AppError> {
let mut summary_json: Option<String> = None;
let mut seed_file_data: Option<Vec<u8>> = None;
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());
}
} else if field_name == "seed_file" {
if let Ok(bytes) = field.bytes().await {
seed_file_data = Some(bytes.to_vec());
}
}
}
let summary_json = match summary_json {
Some(s) => s,
None => {
return Err(crate::api::AppError::BadRequest(
"请求中缺少 report 字段(旧版 conv.json 内容)".to_string(),
));
}
};
// 解析旧版 conv.jsonModelSummary 结构)取 name / params / 收敛状态 / 真实 max_relc。
let summary: ModelSummary = match serde_json::from_str(&summary_json) {
Ok(s) => s,
Err(e) => {
warn!("历史种子导入:conv.json 解析失败: {}", e);
return Err(crate::api::AppError::BadRequest(
"conv.json 解析失败,非合法 ModelSummary".to_string(),
));
}
};
// point_name 优先用旧 conv.json 的 name(源精度真名);回退到 params 规范名。
let name = if !summary.name.is_empty() {
summary.name.clone()
} else {
summary.params.model_name()
};
// 名称合法性校验(防路径穿越),与 report_task 同口径。
if !super::workflow::is_valid_point_name(&name) {
warn!("历史种子导入:拒绝非法网格点名称: {}", name);
return Err(crate::api::AppError::BadRequest(
"非法的网格点名称参数".to_string(),
));
}
let workflow_name = query.workflow;
let params = summary.params.clone();
let converged = summary.converged && !summary.atmosphere_has_nan;
let max_relc = summary.final_max_relc;
// 1. 幂等写入 grid_pointsON CONFLICT DO NOTHING):无需事先 start 工作流。
// 用权威 name(旧 conv.json 的源精度真名),而非从 params 重推——导入路径的
// params 来自旧 JSON(无源文本,model_name() 会失真)。
if let Err(e) = state
.db
.upsert_grid_point_named(&name, &params, 0, &workflow_name)
.await
{
tracing::error!("历史种子导入:upsert grid_points {} 失败: {}", name, e);
return Err(crate::api::AppError::Internal(e));
}
// 2. 落地 conv.json(原子 tmp→rename)。
let model_dir = Path::new(&state.seeds_dir).join(&name);
if fs::create_dir_all(&model_dir).await.is_ok() {
let conv_tmp = model_dir.join(format!("conv.json.{}.tmp", uuid::Uuid::new_v4().simple()));
let conv_path = model_dir.join("conv.json");
if fs::write(&conv_tmp, &summary_json).await.is_ok() {
let _ = fs::rename(&conv_tmp, &conv_path).await;
}
// 3. 收敛且干净才写 .7 + 入种子库(与 report_task 同口径)。
if converged {
if let Some(bytes) = seed_file_data {
let seed_tmp =
model_dir.join(format!("{}.7.{}.tmp", name, uuid::Uuid::new_v4().simple()));
let seed_path = model_dir.join(format!("{}.7", name));
if fs::write(&seed_tmp, bytes).await.is_ok()
&& fs::rename(&seed_tmp, &seed_path).await.is_ok()
{
info!(
"历史种子导入:网格点 {} 收敛种子已落地: {} (max_relc={:?})",
name,
seed_path.display(),
max_relc
);
let _ = state
.db
.insert_seed_named(&name, &params, &seed_path.to_string_lossy())
.await;
}
} else {
warn!(
"历史种子导入:网格点 {} 声称收敛但未上传 seed_file,跳过种子写入",
name
);
}
}
}
// 4. 更新 grid_points 状态:收敛→converged(success_method='imported');否则维持 pending
// 让正常调度处理(导入未收敛点无意义,但记录其尝试)。
if converged {
if let Err(e) = state
.db
.mark_grid_point_imported(&name, &workflow_name, Some(summary.elapsed_sec))
.await
{
warn!("历史种子导入:标记 {} 为 converged 失败: {}", name, e);
}
}
info!(
"历史种子导入完成:网格点 {} (workflow={}, converged={}, max_relc={:?})",
name, workflow_name, converged, max_relc
);
Ok((
StatusCode::OK,
Json(json!({
"status": "ok",
"point_name": name,
"converged": converged,
"max_relc": max_relc,
})),
))
}
+346 -4
View File
@@ -1,6 +1,6 @@
use super::AppState;
use axum::{
extract::{Path as AxumPath, State},
extract::{Path as AxumPath, Query as AxumQuery, State},
http::StatusCode,
response::IntoResponse,
Json,
@@ -71,6 +71,19 @@ fn is_valid_workflow_name(name: &str) -> bool {
.all(|c| c.is_ascii_alphanumeric() || c == '.' || c == '_' || c == '-')
}
/// 网格点名称白名单(全端统一口径):字母/数字/`.`/`_`/`-`/`+`/`@`
/// 非空、不以 `.` 开头(拒绝 `..` 穿越)、长度 ≤ 128。
///
/// 字符集不含 `/`、`\`,从源头杜绝路径穿越——任何磁盘路径拼接前的唯一闸门。
pub(crate) fn is_valid_point_name(name: &str) -> bool {
!name.is_empty()
&& !name.starts_with('.')
&& name.len() <= 128
&& name
.chars()
.all(|c| c.is_ascii_alphanumeric() || matches!(c, '.' | '_' | '-' | '+' | '@'))
}
pub async fn save_workflow(
State(state): State<AppState>,
Json(req): Json<CreateWorkflowRequest>,
@@ -82,8 +95,8 @@ pub async fn save_workflow(
));
}
// Validate YAML config string
if let Err(e) = serde_yaml::from_str::<GridConfig>(&req.config_yaml) {
// Validate YAML config string(用源精度解析,校验 + 保留 grid 轴书写小数位)
if let Err(e) = GridConfig::from_yaml_str(&req.config_yaml) {
return Err(crate::api::AppError::BadRequest(format!(
"无效的 YAML 配置: {}",
e
@@ -186,7 +199,7 @@ pub async fn start_workflow(
Ok(true) => {}
}
let grid_cfg: GridConfig = match serde_yaml::from_str(&item.config_yaml) {
let grid_cfg: GridConfig = match GridConfig::from_yaml_str(&item.config_yaml) {
Ok(cfg) => cfg,
Err(e) => {
let _ = state.db.update_workflow_status(&name, "idle").await;
@@ -233,6 +246,335 @@ pub async fn start_workflow(
))
}
/// 单工作流执行统计:进度(pending/queued/running/converged/failed 分开计数)、
/// 收敛手段归因(冷启动/种子步进/历史导入)、难度波次分布、近似 ETA。详情页数据源。
pub async fn get_workflow_stats(
State(state): State<AppState>,
AxumPath(name): AxumPath<String>,
) -> Result<impl IntoResponse, crate::api::AppError> {
if !is_valid_workflow_name(&name) {
return Err(crate::api::AppError::BadRequest(
"工作流名称含非法字符".to_string(),
));
}
let item = match state.db.get_workflow(&name).await {
Ok(Some(item)) => item,
Ok(None) => {
return Err(crate::api::AppError::NotFound(format!(
"工作流 '{}' 未找到",
name
)))
}
Err(e) => return Err(e.into()),
};
// 在线节点总槽位:ETA 并发感知除数(无在线节点时按串行兜底,db 层 max(1))。
let total_slots: i64 = state
.db
.get_active_nodes()
.await
.unwrap_or_default()
.iter()
.map(|n| n.max_slots as i64)
.sum();
match state
.db
.get_workflow_detail_stats(&name, &item.status, total_slots)
.await
{
Ok(stats) => Ok((
StatusCode::OK,
Json(ApiResponse {
success: true,
message: "成功获取工作流统计".to_string(),
data: Some(stats),
}),
)),
Err(e) => Err(e.into()),
}
}
/// 进度时间序列查询参数。
#[derive(Debug, serde::Deserialize)]
pub struct ProgressQuery {
/// 时间窗口(小时),默认 24,钳位 1–168(7 天)。
pub hours: Option<i64>,
}
/// 工作流进度时间序列 + 经验速率 + 停滞时长(详情页概览进度曲线数据源)。
///
/// - `series`:窗口内的计数快照(超 300 点自动降采样,首末点保留);
/// - `rate_per_hour`:窗口首末 converged 增量 ÷ 时长(快照 <2 条或时长 ≤0 为 null);
/// - `stalled_minutes`:终态数(converged+failed)最后一次增长到窗口末端的分钟数
/// (用于"进度停滞"预警;快照 <2 条为 null)。
pub async fn get_workflow_progress(
State(state): State<AppState>,
AxumPath(name): AxumPath<String>,
AxumQuery(q): AxumQuery<ProgressQuery>,
) -> Result<impl IntoResponse, crate::api::AppError> {
if !is_valid_workflow_name(&name) {
return Err(crate::api::AppError::BadRequest(
"工作流名称含非法字符".to_string(),
));
}
match state.db.get_workflow(&name).await {
Ok(Some(_)) => {}
Ok(None) => {
return Err(crate::api::AppError::NotFound(format!(
"工作流 '{}' 未找到",
name
)))
}
Err(e) => return Err(e.into()),
}
let hours = q.hours.unwrap_or(24).clamp(1, 168);
let mut series = state.db.get_progress_series(&name, hours).await?;
// 降采样:过密时等间隔抽取,首末点强制保留(曲线端点不失真)。
if series.len() > 300 {
let step = (series.len() as f64 / 300.0).ceil() as usize;
let last = series.last().cloned();
series = series.into_iter().step_by(step).collect();
if let Some(l) = last {
if series.last().map(|p| p.ts.as_str()) != Some(l.ts.as_str()) {
series.push(l);
}
}
}
// SQLite datetime('now') 为 UTC 'YYYY-MM-DD HH:MM:SS'
let parse_ts = |ts: &str| chrono::NaiveDateTime::parse_from_str(ts, "%Y-%m-%d %H:%M:%S").ok();
let rate_per_hour: Option<f64> = match (series.first(), series.last()) {
(Some(first), Some(last)) if series.len() >= 2 => {
match (parse_ts(&first.ts), parse_ts(&last.ts)) {
(Some(t0), Some(t1)) => {
let dh = (t1 - t0).num_seconds() as f64 / 3600.0;
if dh > 0.0 {
Some((last.converged - first.converged) as f64 / dh)
} else {
None
}
}
_ => None,
}
}
_ => None,
};
let stalled_minutes: Option<f64> = if series.len() >= 2 {
let mut last_progress_idx = None;
for i in 1..series.len() {
let prev = series[i - 1].converged + series[i - 1].failed;
let cur = series[i].converged + series[i].failed;
if cur > prev {
last_progress_idx = Some(i);
}
}
let anchor = last_progress_idx.unwrap_or(0);
let end = series.len() - 1;
match (parse_ts(&series[anchor].ts), parse_ts(&series[end].ts)) {
(Some(t0), Some(t1)) => Some(((t1 - t0).num_seconds() as f64 / 60.0).max(0.0)),
_ => None,
}
} else {
None
};
Ok((
StatusCode::OK,
Json(ApiResponse {
success: true,
message: "成功获取进度时间序列".to_string(),
data: Some(serde_json::json!({
"hours": hours,
"series": series,
"rate_per_hour": rate_per_hour,
"stalled_minutes": stalled_minutes,
})),
}),
))
}
/// 逐点列表查询参数。枚举类参数(status/method/sort/order)一律白名单校验后
/// 才进入 db 层;limit/offset 钳位;任何用户文本都不会拼进 SQL 字符串。
#[derive(Debug, serde::Deserialize)]
pub struct PointsQuery {
pub status: Option<String>,
pub method: Option<String>,
pub wave: Option<i32>,
pub q: Option<String>,
pub sort: Option<String>,
pub order: Option<String>,
pub limit: Option<i64>,
pub offset: Option<i64>,
}
/// 工作流逐点列表:点参数 + 状态 + 收敛手段 + 最近尝试(max_relc/种子来源/节点/错误)。
/// 支持状态/手段/波次过滤、点名搜索、白名单排序与分页(limit ≤ 500)。详情页点表数据源。
pub async fn get_workflow_points(
State(state): State<AppState>,
AxumPath(name): AxumPath<String>,
AxumQuery(pq): AxumQuery<PointsQuery>,
) -> Result<impl IntoResponse, crate::api::AppError> {
if !is_valid_workflow_name(&name) {
return Err(crate::api::AppError::BadRequest(
"工作流名称含非法字符".to_string(),
));
}
// 未知工作流返回 404(与 stats 端点一致),而非空列表
match state.db.get_workflow(&name).await {
Ok(Some(_)) => {}
Ok(None) => {
return Err(crate::api::AppError::NotFound(format!(
"工作流 '{}' 未找到",
name
)))
}
Err(e) => return Err(e.into()),
}
if let Some(s) = &pq.status {
if !matches!(
s.as_str(),
"pending" | "queued" | "running" | "converged" | "failed"
) {
return Err(crate::api::AppError::BadRequest(format!(
"非法的 status 参数: {}",
s
)));
}
}
if let Some(m) = &pq.method {
if !matches!(m.as_str(), "cold_run" | "seed_step" | "imported") {
return Err(crate::api::AppError::BadRequest(format!(
"非法的 method 参数: {}",
m
)));
}
}
// 排序白名单 → 编译期列名片段;NULL 统一靠后(IS NULL 升序前置),不受 dir 影响。
let sort = pq.sort.as_deref().unwrap_or("wave");
let dir = if pq
.order
.as_deref()
.map(|o| o.eq_ignore_ascii_case("desc"))
.unwrap_or(false)
{
"DESC"
} else {
"ASC"
};
let order_by = match sort {
"wave" => format!("gp.wave {dir}, gp.cno_sum ASC, gp.teff ASC"),
"teff" => format!("gp.teff {dir}, gp.wave ASC, gp.cno_sum ASC"),
"max_relc" => format!("t.max_relc IS NULL ASC, t.max_relc {dir}, gp.wave ASC"),
"attempts" => format!("gp.attempt_count {dir}, gp.wave ASC, gp.cno_sum ASC"),
"last_completed_at" => {
format!("t.completed_at IS NULL ASC, t.completed_at {dir}, gp.wave ASC")
}
_ => {
return Err(crate::api::AppError::BadRequest(format!(
"非法的 sort 参数: {}",
sort
)))
}
};
let filter = crate::db::PointFilter {
status: pq.status.clone(),
method: pq.method.clone(),
wave: pq.wave,
q: pq.q.clone(),
order_by,
limit: pq.limit.unwrap_or(100).clamp(1, 500),
offset: pq.offset.unwrap_or(0).max(0),
};
match state.db.list_workflow_points(&name, &filter).await {
Ok((total, points)) => Ok((
StatusCode::OK,
Json(ApiResponse {
success: true,
message: "成功获取工作流网格点列表".to_string(),
data: Some(serde_json::json!({ "total": total, "points": points })),
}),
)),
Err(e) => Err(e.into()),
}
}
/// 单网格点详情:点行 + 全部尝试历史 + conv.json 逐阶段诊断。
///
/// conv.json 读自 `seeds_dir/<point>/conv.json`(单层目录)。点名经 `is_valid_point_name`
/// 白名单(无 `/`、`\`,拒前导 `.`)——即路径穿越的前置闸门;读盘后再做 canonicalize
/// 归属兜底校验(纵深防御)。缺失/读失败/解析失败一律 `conv: null`(仍 200),
/// 前端降级显示"诊断文件不可用"。
pub async fn get_workflow_point_detail(
State(state): State<AppState>,
AxumPath((name, point)): AxumPath<(String, String)>,
) -> Result<impl IntoResponse, crate::api::AppError> {
if !is_valid_workflow_name(&name) || !is_valid_point_name(&point) {
return Err(crate::api::AppError::BadRequest(
"工作流名称或网格点名称含非法字符".to_string(),
));
}
let point_row = match state.db.get_workflow_point_row(&name, &point).await {
Ok(Some(row)) => row,
Ok(None) => {
return Err(crate::api::AppError::NotFound(format!(
"网格点 '{}' 未找到(工作流 '{}'",
point, name
)))
}
Err(e) => return Err(e.into()),
};
let attempts = state.db.list_point_attempts(&name, &point).await?;
let conv_path = std::path::Path::new(&state.seeds_dir)
.join(&point)
.join("conv.json");
let conv: Option<common::models::ModelSummary> =
match tokio::fs::read_to_string(&conv_path).await {
Ok(s) => {
let confined = std::path::Path::new(&state.seeds_dir)
.canonicalize()
.ok()
.zip(conv_path.canonicalize().ok())
.map(|(root, f)| f.starts_with(root))
.unwrap_or(false);
if confined {
match serde_json::from_str(&s) {
Ok(summary) => Some(summary),
Err(e) => {
tracing::warn!("网格点 {} 的 conv.json 解析失败: {}", point, e);
None
}
}
} else {
tracing::warn!("网格点 {} 的 conv.json 路径越界,拒绝读取", point);
None
}
}
Err(_) => None,
};
Ok((
StatusCode::OK,
Json(ApiResponse {
success: true,
message: "成功获取网格点详情".to_string(),
data: Some(serde_json::json!({
"point": point_row,
"attempts": attempts,
"conv": conv,
})),
}),
))
}
pub async fn stop_workflow(
State(state): State<AppState>,
AxumPath(name): AxumPath<String>,
+926 -202
View File
File diff suppressed because it is too large Load Diff
+52 -25
View File
@@ -55,11 +55,7 @@ async fn main() -> Result<()> {
let db = Database::new(&server_cfg.db_path).await?;
let queue = Arc::new(SqliteTaskQueue::new(&server_cfg.queue_db_path).await?);
let scheduler = Arc::new(GridScheduler::new(
db.clone(),
queue.clone(),
server_cfg.results_dir.clone(),
));
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
// Auto-register sdB_cno.yaml if exists and not yet in DB
let default_wf_path = Path::new(&server_cfg.grid_config);
@@ -81,18 +77,10 @@ async fn main() -> Result<()> {
}
}
// 弱口令凭据安全警告检测
let is_weak_token = |t: Option<&str>| -> bool {
match t {
Some(s) => {
s.len() < 12 || s == "fmqi123" || s == "admin" || s == "123456" || s == "secret"
}
None => false,
}
};
if is_weak_token(server_cfg.auth_token.as_deref())
|| is_weak_token(server_cfg.admin_token.as_deref())
{
// 弱口令凭据安全警告检测:仅按强度阈值判断(短于 16 字节视为弱口令)。
// 推荐用 `openssl rand -hex 32`64 字符)生成。
let is_weak_token = |t: Option<&str>| -> bool { t.map(|s| s.len() < 16).unwrap_or(false) };
if is_weak_token(server_cfg.admin_token.as_deref()) {
tracing::warn!("⚠️ 检测到系统当前正在使用弱口令凭据或默认 Token!建议生产环境在 .env 中配置使用 openssl rand -hex 32 生成的高强度 Token");
}
@@ -102,9 +90,8 @@ async fn main() -> Result<()> {
db,
queue: queue.clone(),
scheduler: scheduler.clone(),
results_dir: server_cfg.results_dir.clone(),
seeds_dir: server_cfg.seeds_dir.clone(),
rate_limiter,
auth_token: server_cfg.auth_token.clone(),
admin_token: server_cfg.admin_token.clone(),
auth_disabled: server_cfg.auth_disabled,
admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new(
@@ -188,6 +175,23 @@ async fn main() -> Result<()> {
has_error = true;
}
// P3 进度快照:对每个运行中工作流记录计数(record_progress_snapshot
// 内部去重——计数无变化不落库);顺带清理超过 7 天的旧快照。
// 观测性写入失败不回退调度退避(不置 has_error)。
match bg_db_clone.get_running_workflow_names().await {
Ok(names) => {
for wf in names {
if let Err(e) = bg_db_clone.record_progress_snapshot(&wf).await {
tracing::warn!("记录工作流 {} 进度快照失败: {}", wf, e);
}
}
}
Err(e) => tracing::warn!("获取运行中工作流列表失败: {}", e),
}
if let Err(e) = bg_db_clone.purge_progress_snapshots(7).await {
tracing::warn!("清理过期进度快照失败: {}", e);
}
has_error
});
@@ -234,6 +238,8 @@ async fn main() -> Result<()> {
let report_router = Router::new()
.route("/task/report", post(api::task::report_task))
// 历史种子导入同样上传 .7 大气文件,并入宽松 body limit / 并发限流组。
.route("/admin/import_seed", post(api::task::import_seed))
.layer(DefaultBodyLimit::max(REPORT_BODY_LIMIT))
.layer(tower::ServiceBuilder::new().concurrency_limit(REPORT_MAX_CONCURRENCY));
@@ -284,7 +290,24 @@ async fn main() -> Result<()> {
post(api::workflow::start_workflow),
)
.route("/workflows/:name/stop", post(api::workflow::stop_workflow))
// Admin Management API(节点凭据查看/审批/吊销/重发,均要求 Admin 角色)
// 工作流执行观测 API(进度统计 / 逐点明细 / 单点诊断,均要求 Admin 角色)
.route(
"/workflows/:name/stats",
get(api::workflow::get_workflow_stats),
)
.route(
"/workflows/:name/progress",
get(api::workflow::get_workflow_progress),
)
.route(
"/workflows/:name/points",
get(api::workflow::get_workflow_points),
)
.route(
"/workflows/:name/points/:point",
get(api::workflow::get_workflow_point_detail),
)
// Admin Management API(节点凭据查看/审批/重发/停用/启用,均要求 Admin 角色)
.route("/admin/nodes", get(api::admin::list_nodes))
.route(
"/admin/nodes/:node_id/approve",
@@ -294,14 +317,18 @@ async fn main() -> Result<()> {
"/admin/nodes/:node_id/reject",
post(api::admin::reject_node),
)
.route(
"/admin/nodes/:node_id/revoke",
post(api::admin::revoke_node),
)
.route(
"/admin/nodes/:node_id/reissue",
post(api::admin::reissue_node),
)
.route(
"/admin/nodes/:node_id/disable",
post(api::admin::disable_node),
)
.route(
"/admin/nodes/:node_id/enable",
post(api::admin::enable_node),
)
// 合并大体积上报路由(继承各自的 body limit)
.merge(report_router)
.layer(DefaultBodyLimit::max(DEFAULT_BODY_LIMIT));
@@ -320,7 +347,7 @@ async fn main() -> Result<()> {
api_router.layer(auth_layer).layer(rate_limit_layer)
} else {
tracing::warn!(
"⚠️ 警告:未配置 DCTS_ADMIN_TOKEN / DCTS_ENROLLMENT_TOKEN(且未启用 DCTS_AUTH_DISABLE),\
"⚠️ 警告:未配置 DCTS_ADMIN_TOKEN(且未启用 DCTS_AUTH_DISABLE),\
服务端运行在【无鉴权模式】!公网部署务必配置凭据。"
);
api_router
+56 -61
View File
@@ -11,16 +11,11 @@ use crate::db::Database;
pub struct GridScheduler {
db: Database,
queue: Arc<SqliteTaskQueue>,
_results_dir: String,
}
impl GridScheduler {
pub fn new(db: Database, queue: Arc<SqliteTaskQueue>, results_dir: String) -> Self {
Self {
db,
queue,
_results_dir: results_dir,
}
pub fn new(db: Database, queue: Arc<SqliteTaskQueue>) -> Self {
Self { db, queue }
}
/// Expands grid points from config and registers them into the database.
@@ -52,19 +47,19 @@ impl GridScheduler {
}
let mut points = Vec::new();
for &teff in &cfg.grid.teff {
for &logg in &cfg.grid.logg {
for &loghe in &cfg.grid.loghe {
for &logc in &cfg.grid.logc {
for &logn in &cfg.grid.logn {
for &logo in &cfg.grid.logo {
for teff in &cfg.grid.teff {
for logg in &cfg.grid.logg {
for loghe in &cfg.grid.loghe {
for logc in &cfg.grid.logc {
for logn in &cfg.grid.logn {
for logo in &cfg.grid.logo {
points.push(GridPointParams {
teff,
logg,
loghe,
logc,
logn,
logo,
teff: teff.clone(),
logg: logg.clone(),
loghe: loghe.clone(),
logc: logc.clone(),
logn: logn.clone(),
logo: logo.clone(),
});
}
}
@@ -196,7 +191,7 @@ impl GridScheduler {
.await?;
let mut dispatched = 0;
for (name, params, _wave) in pending {
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)) => {
@@ -217,6 +212,7 @@ impl GridScheduler {
seed_point_name: seed_name,
timeout_sec,
workflow_name: Some(workflow_name.to_string()),
wave,
};
self.db.insert_task(&task_spec).await?;
@@ -274,6 +270,7 @@ impl GridScheduler {
pub async fn trigger_seed_step_fallback(
&self,
params: &GridPointParams,
name: &str,
workflow_name: &str,
) -> Result<bool> {
// 该工作流须仍处于 running 态才回退(避免 stop 后继续派发)
@@ -291,9 +288,13 @@ impl GridScheduler {
return Ok(false);
}
let name = params.model_name();
// name 取自权威的 report.point_name= grid_points.name 列,源精度正确),
// 而非 params.model_name()。原因:此处 params 经 node 上报回传,其 logg 等
// 轴在服务端 DB REAL 列回读时已丢精度(5.0→"5"),重推 model_name() 会得到
// 降级名(g5 而非 g5.0),导致回退任务的 point_name 与 grid_points.name 列错配,
// 状态更新静默失败。与 runner 的修复保持同一原则:用权威 name。
// 种子回退仅一次:该点在该工作流中已经派发过 seed_step 任务就不再触发新的回退
if self.db.has_seed_step_attempt(&name, workflow_name).await? {
if self.db.has_seed_step_attempt(name, workflow_name).await? {
info!(
"网格点 {} 已使用过一次种子热启动回退,不再重复回退,保持 failed 终态",
name
@@ -306,30 +307,27 @@ impl GridScheduler {
if let Some(seed_match) = seed_match_opt {
let timeout_sec = self.get_workflow_timeout_sec(workflow_name).await;
let name = params.model_name();
let task_spec = TaskSpec {
task_id: Uuid::new_v4(),
point_name: name.clone(),
point_name: name.to_string(),
params: params.clone(),
task_type: TaskType::SeedStep,
seed_point_name: Some(seed_match.name.clone()),
timeout_sec,
workflow_name: Some(workflow_name.to_string()),
// seed_step 是失败后的回退任务,wave 设 0 不抢占正常调度队列里的低难度 wave 优先级。
wave: 0,
};
self.db.insert_task(&task_spec).await?;
self.db
.update_grid_status(
&name,
common::models::GridPointStatus::Queued,
workflow_name,
)
.update_grid_status(name, common::models::GridPointStatus::Queued, workflow_name)
.await?;
if let Err(e) = self.queue.push_task(&task_spec).await {
let _ = self
.db
.update_grid_status(
&name,
name,
common::models::GridPointStatus::Pending,
workflow_name,
)
@@ -358,7 +356,6 @@ mod tests {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("sched_db.db");
let queue_db_path = temp_dir.path().join("sched_queue.db");
let results_dir = temp_dir.path().join("results");
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
@@ -366,20 +363,17 @@ mod tests {
.await
.unwrap(),
);
let scheduler = GridScheduler::new(
db.clone(),
queue.clone(),
results_dir.to_string_lossy().to_string(),
);
let scheduler = GridScheduler::new(db.clone(), queue.clone());
#[allow(deprecated)] // results 是死字段,构造时必须填 None
let cfg = GridConfig {
grid: GridAxesConfig {
teff: vec![35000.0],
logg: vec![5.5],
loghe: vec![-1.0],
logc: vec![-2.0],
logn: vec![-2.0],
logo: vec![-2.0],
teff: vec![35000.0.into()],
logg: vec![5.5.into()],
loghe: vec![(-1.0).into()],
logc: vec![(-2.0).into()],
logn: vec![(-2.0).into()],
logo: vec![(-2.0).into()],
},
chain: vec![],
synspec: None,
@@ -425,16 +419,17 @@ mod tests {
.await
.unwrap(),
);
let scheduler = GridScheduler::new(db.clone(), queue.clone(), "results".to_string());
let scheduler = GridScheduler::new(db.clone(), queue.clone());
#[allow(deprecated)] // results 是死字段,构造时必须填 None
let mk_cfg = |teff: f64| GridConfig {
grid: GridAxesConfig {
teff: vec![teff],
logg: vec![5.5],
loghe: vec![-1.0],
logc: vec![-2.0],
logn: vec![-2.0],
logo: vec![-2.0],
teff: vec![teff.into()],
logg: vec![5.5.into()],
loghe: vec![(-1.0).into()],
logc: vec![(-2.0).into()],
logn: vec![(-2.0).into()],
logo: vec![(-2.0).into()],
},
chain: vec![],
synspec: None,
@@ -466,12 +461,12 @@ mod tests {
// 重新推一个 wf_a 任务(上一行 pop 掉了),再初始化 wf_b
db.update_grid_status(
&GridPointParams {
teff: 35000.0,
logg: 5.5,
loghe: -1.0,
logc: -2.0,
logn: -2.0,
logo: -2.0,
teff: 35000.0.into(),
logg: 5.5.into(),
loghe: (-1.0).into(),
logc: (-2.0).into(),
logn: (-2.0).into(),
logo: (-2.0).into(),
}
.model_name(),
common::models::GridPointStatus::Pending,
@@ -496,12 +491,12 @@ mod tests {
// 关键断言:把 wf_a 任务重新推回队列后,初始化 wf_b 不应清空它。
db.update_grid_status(
&GridPointParams {
teff: 35000.0,
logg: 5.5,
loghe: -1.0,
logc: -2.0,
logn: -2.0,
logo: -2.0,
teff: 35000.0.into(),
logg: 5.5.into(),
loghe: (-1.0).into(),
logc: (-2.0).into(),
logn: (-2.0).into(),
logo: (-2.0).into(),
}
.model_name(),
common::models::GridPointStatus::Pending,
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,199 @@
//! 验证「历史种子导入的工作流名」与「正式工作流名」的隔离关系。
//!
//! 用户意图:import_results 把旧 Python 计算结果导入,标记为已完成,避免重算。
//! 关键问题:导入到工作流 A,之后正式启动工作流 B(同名/异名),B 能否看到 A 标记的 converged
use common::config::GridConfig;
use common::models::GridPointParams;
use mq::sqlite_queue::SqliteTaskQueue;
use server::{db::Database, scheduler::GridScheduler};
use std::sync::Arc;
fn make_params() -> GridPointParams {
use common::models::GridAxisValue;
GridPointParams {
teff: GridAxisValue::from_value(20000.0),
logg: GridAxisValue::from_value(5.0),
loghe: GridAxisValue::from_value(-2.0),
logc: GridAxisValue::from_value(-4.0),
logn: GridAxisValue::from_value(-4.0),
logo: GridAxisValue::from_value(-4.0),
}
}
/// 构造只含一个网格点(t20000_g5.0_he-2_c-4_n-4_o-4)的 config。
fn make_grid_cfg() -> GridConfig {
let yaml = "grid:\n teff: [20000]\n logg: [5.0]\n loghe: [-2]\n logc: [-4]\n logn: [-4]\n logo: [-4]\n";
GridConfig::from_yaml_str(yaml).unwrap()
}
async fn setup() -> (Database, Arc<GridScheduler>) {
let tmp = tempfile::tempdir().unwrap();
let db = Database::new(&tmp.path().join("db.db").to_string_lossy())
.await
.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&tmp.path().join("q.db").to_string_lossy())
.await
.unwrap(),
);
let sched = Arc::new(GridScheduler::new(db.clone(), queue));
(db, sched)
}
/// 场景 1(正确用法):导入到工作流 "sdB_cno",再用同名 config initialize_grid。
/// 期望:initialize_grid 的 ON CONFLICT(workflow_name, name) DO NOTHING 保留 converged 状态。
#[tokio::test]
async fn test_same_workflow_name_preserves_converged() {
let (db, sched) = setup().await;
let name = "t20000_g5.0_he-2_c-4_n-4_o-4";
let p = make_params();
// 模拟 import_seedupsert + mark_imported,工作流名 = sdB_cno
db.upsert_grid_point_named(name, &p, 0, "sdB_cno")
.await
.unwrap();
db.mark_grid_point_imported(name, "sdB_cno", None)
.await
.unwrap();
// 之后正式启动同名工作流:initialize_grid(sdB_cno)
sched
.initialize_grid(&make_grid_cfg(), "sdB_cno")
.await
.unwrap();
let st = db.get_grid_point_status(name, "sdB_cno").await.unwrap();
assert_eq!(
st.unwrap().0,
"converged",
"同名工作流:导入的 converged 应被保留,避免重算"
);
println!("✓ 场景1(同名):status=converged,旧结果被保留,不会重算");
}
/// 场景 2(错误用法):导入到工作流 "imported",之后正式启动 "sdB_cno"。
/// 期望:sdB_cno 分区下是新插入的 pending 行,看不到 imported 分区的 converged。
#[tokio::test]
async fn test_different_workflow_name_causes_recompute() {
let (db, sched) = setup().await;
let name = "t20000_g5.0_he-2_c-4_n-4_o-4";
let p = make_params();
// 模拟 import_seed:导入到 "imported" 工作流
db.upsert_grid_point_named(name, &p, 0, "imported")
.await
.unwrap();
db.mark_grid_point_imported(name, "imported", None)
.await
.unwrap();
// 之后正式启动 "sdB_cno" 工作流
sched
.initialize_grid(&make_grid_cfg(), "sdB_cno")
.await
.unwrap();
// imported 分区:converged(种子库有,但不会被 sdB_cno 调度看到)
let st_imp = db.get_grid_point_status(name, "imported").await.unwrap();
assert_eq!(st_imp.unwrap().0, "converged");
// sdB_cno 分区:pending(重新算!看不到 imported 的 converged
let st_real = db.get_grid_point_status(name, "sdB_cno").await.unwrap();
assert_eq!(
st_real.unwrap().0,
"pending",
"异名工作流:sdB_cno 看不到 imported 的 converged,会重算"
);
println!("✓ 场景2(异名):sdB_cno 分区 status=pending,会重复计算 —— 验证了工作流名必须匹配");
}
/// 场景 3(真实意图验证):旧网格有部分点已算完(导入为 converged),
/// 新网格比旧网格多了若干点。用同名工作流启动后:
/// - 旧点保持 converged(不重算)
/// - 新点是 pending(会被调度计算)
/// 这正是「同步旧结果避免重算」的核心语义。
#[tokio::test]
async fn test_mixed_grid_import_then_init_avoids_recompute() {
let (db, sched) = setup().await;
let p_old = make_params(); // t20000_g5.0_...
// 模拟 import_seed:旧网格里这个点已收敛,导入到 sdB_cno
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)
.await
.unwrap();
// 正式启动 sdB_cno,config 比旧网格多了一个新点(t25000)
let yaml = "grid:\n teff: [20000, 25000]\n logg: [5.0]\n loghe: [-2]\n logc: [-4]\n logn: [-4]\n logo: [-4]\n";
let cfg = GridConfig::from_yaml_str(yaml).unwrap();
sched.initialize_grid(&cfg, "sdB_cno").await.unwrap();
// 旧点:converged(不重算)
let st_old = db
.get_grid_point_status("t20000_g5.0_he-2_c-4_n-4_o-4", "sdB_cno")
.await
.unwrap();
assert_eq!(
st_old.unwrap().0,
"converged",
"旧点应保持 converged 不重算"
);
// 新点:pending(会被调度)
let st_new = db
.get_grid_point_status("t25000_g5.0_he-2_c-4_n-4_o-4", "sdB_cno")
.await
.unwrap();
assert_eq!(st_new.unwrap().0, "pending", "新点应为 pending 等待计算");
println!("✓ 场景3(混合网格):旧点converged保留 + 新点pending待算 —— 完全符合避免重算的意图");
}
/// 场景 4(精度差异命门):旧 conv.json name=g5(无小数),但配置 logg=5.0 → model_name()=g5.0。
/// 导入工具必须把 name 重写为 g5.0 入库,否则 initialize_grid 插入的 g5.0 行与导入的 g5 行
/// 复合唯一键不匹配,导入的 converged 被孤立、g5.0 被重算。
/// 本测试直接模拟「入库的 grid_points.name = g5.0」(即工具重写后的状态),
/// 验证 initialize_grid(sdB_cno) 后该行保持 converged(不重算)。
#[tokio::test]
async fn test_precision_diff_import_then_init_preserves_converged() {
let (db, sched) = setup().await;
// 配置权威名(logg=5.0 → g5.0,保留小数)
let canonical = "t20000_g5.0_he-2_c-4_n-4_o-4";
use common::models::GridAxisValue;
let p = GridPointParams {
teff: GridAxisValue::from_value(20000.0),
logg: GridAxisValue::from_value(5.0),
loghe: GridAxisValue::from_value(-2.0),
logc: GridAxisValue::from_value(-4.0),
logn: GridAxisValue::from_value(-4.0),
logo: GridAxisValue::from_value(-4.0),
};
// 模拟 import_results 重写 name 后入库:grid_points.name = canonical(g5.0)
db.upsert_grid_point_named(canonical, &p, 0, "sdB_cno")
.await
.unwrap();
db.mark_grid_point_imported(canonical, "sdB_cno", None)
.await
.unwrap();
// 启动同名工作流:initialize_grid 用配置 model_name()(=g5.0) 插入
let yaml = "grid:\n teff: [20000]\n logg: [5.0]\n loghe: [-2]\n logc: [-4]\n logn: [-4]\n logo: [-4]\n";
let cfg = GridConfig::from_yaml_str(yaml).unwrap();
sched.initialize_grid(&cfg, "sdB_cno").await.unwrap();
// 关键断言:name=g5.0 的行保持 convergedON CONFLICT DO NOTHING 命中)
let st = db
.get_grid_point_status(canonical, "sdB_cno")
.await
.unwrap();
assert_eq!(
st.unwrap().0,
"converged",
"精度一致(g5.0=g5.0)时导入的 converged 必须保留,不重算"
);
println!("✓ 场景4(精度差异命门):g5.0 入库 + initialize_grid → converged 保留,避免重算");
}