Files
DCTS/crates/server/tests/api_tests.rs
T
fmq cd370d88e7 feat(all): 任务引擎双阶段解耦、僵尸涡旋修复、动态 CPU 配额与前端详情页重构
将 TLUSTY/SYNSPEC 拆为各自独立的 enabled/policy/strategies 阶段,
以策略链自动弹栈取代单级 seed_step 布尔回退;定向修复 2026-08-02
僵尸任务涡旋事故;新增节点并发配额热调;前端详情页从 1412 行巨型
视图拆为薄控制器 + detail 子模块,并补齐工具层与单测。

引擎与调度(task_engine_decoupling_design.md)
- models.rs: 新增 StagePolicy / EngineStageConfig / TaskSpec 阶段字段、
  normalize_compat() 校正旧版在途消息策略链、failed_stage 归因
- scheduler.rs: resolve_dispatchable_chain 派发门控、
  trigger_strategy_fallback 按 failed_stage 精确弹栈;启动期
  force_recompute/skip_converged(默认)/skip_failed 三策略
- db.rs: tasks 表 +7 列持久化阶段配置;终态守卫
  (mark_grid_point_running 仅 pending/queued→running;
  record_task_report 拒绝迟到失败翻黑 converged);策略弹栈快照

僵尸涡旋修复(runbook-20260802-zombie-vortex-fix.md)
- 全链路跨库活性交叉校验:派发/claim/孤儿回收/回退统一查 MQ 队列活性,
  活则放行、死则清僵尸,结构性消除"每点重复派发"
- stop/重启卫生:清队列同步 delete_tasks_by_ids,杜绝遗留 pending 行
- report_task: 幂等吸收 + 409 区分迟到冗余结果,仅 state_changed 时回退
- MQ: NULL workflow_name 回填 __legacy__、requeue 后迟到上报被 403 竞态修复

动态 CPU 配额(dynamic_cpu_slots_design.md)
- admin.rs: POST /admin/nodes/:id/quota(Option<Option<i32>> 区分
  缺字段/显式 null);nodes 表 +admin_max_slots
- worker.rs: effective_max_slots = min(admin, physical),心跳下发原子生效

科学产物保全(tlusty_result_artifacts.md)
- runner.rs: SYNSPEC 启动前快照 fort.12/fort.14 → .bfac/.emflux 防覆盖
- 半失败点(大气收敛+光谱失败)改判 Failed 并写入 note;仅 SYNSPEC
  场景不再恒判失败;撤销归档 LRU 200 上限改为永久保留
- executor.rs: 透传 synspec_params 数值参数(此前固定 None)

前端(dashboard/)
- workflowDetail.js 1412→328 行,拆出 views/detail/{ctx,overview,
  pointsTable,parSets,pointPanel}.js,AbortController 治理监听/请求生命周期
- 删除 wfActions.js,新增 wfEnginePanel.js(双阶段三维配置编辑面板)
- 新增 utils/{errors,format,icons,polling,yamlStage}.js 纯函数模块
- 路由级动态 import 代码分割;节点配额三点菜单 + Modal 管理
- 首次引入 node:test 单测(format/polling/yamlStage/psCache,644 行)
- 系统性补齐 a11y:skip-link、ARIA、Tab 键盘漫游、toast 关闭、退出动画

文档与工具
- 新增 6 篇设计/调研:引擎解耦、动态配额、涡旋 runbook、
  光谱正确性分析、收敛判断、产物归档
- PIPELINE/design/api/database 等协同重写为分布式 C/S 架构口径
- scripts/fetch_results.sh 跨节点产物备份;import_results 按 cno 升序导入
- workflows/sdB_cno.yaml: 新增 tlusty/synspec_stage 配置块,修正 wstart 笔误
2026-08-04 23:40:52 +08:00

3419 lines
119 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
use axum::{
body::Body,
http::{Request, StatusCode},
};
use mq::sqlite_queue::SqliteTaskQueue;
use server::{api::AppState, db::Database, scheduler::GridScheduler};
use std::sync::Arc;
use tower::ServiceExt; // for oneshot
#[tokio::test]
async fn test_server_api_flow() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("api_db.db");
let queue_db_path = temp_dir.path().join("api_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
let state = AppState {
db,
queue,
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
5,
std::time::Duration::from_secs(300),
),
admin_token: None,
auth_disabled: false,
admin_sessions: Arc::new(tokio::sync::RwLock::new(std::collections::HashMap::new())),
};
let app = axum::Router::new()
.route(
"/api/node/register",
axum::routing::post(server::api::node::register_node),
)
.route(
"/api/node/heartbeat",
axum::routing::post(server::api::node::heartbeat_node),
)
.route(
"/api/task/claim",
axum::routing::post(server::api::task::claim_task),
)
.route(
"/api/status",
axum::routing::get(server::api::status::get_status),
)
.route(
"/api/workflows",
axum::routing::get(server::api::workflow::list_workflows)
.post(server::api::workflow::save_workflow),
)
.with_state(state);
// 1. Check status API
let response = app
.clone()
.oneshot(
Request::builder()
.uri("/api/status")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
// 2. Register node API
let reg_body = serde_json::json!({
"node_id": "test-node-api",
"max_slots": 8
});
let response = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/node/register")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&reg_body).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
// 3. Save Workflow API
let wf_body = serde_json::json!({
"name": "test_api_wf",
"description": "Test Workflow Description",
"config_yaml": "grid:\n teff: [35000]\n logg: [5.5]\n loghe: [-1]\n logc: [-2]\n logn: [-2]\n logo: [-2]"
});
let response = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/workflows")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&wf_body).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_auth_middleware_scope_and_running_status() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("auth_db.db");
let queue_db_path = temp_dir.path().join("auth_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
let state = AppState {
db: db.clone(),
queue: queue.clone(),
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
5,
std::time::Duration::from_secs(300),
),
admin_token: Some("secret_token_123".to_string()),
auth_disabled: false,
admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new(
std::collections::HashMap::new(),
)),
};
let api_router = axum::Router::new().route(
"/status",
axum::routing::get(server::api::status::get_status),
);
let auth_layer =
axum::middleware::from_fn_with_state(state.clone(), server::api::auth_middleware);
let api_router = api_router.layer(auth_layer);
let app = axum::Router::new()
.nest("/api", api_router)
.with_state(state);
// Unauthenticated API request -> 401 Unauthorized
let res = app
.clone()
.oneshot(
Request::builder()
.uri("/api/status")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::UNAUTHORIZED);
// Authenticated API request -> 200 OK
let res = app
.clone()
.oneshot(
Request::builder()
.uri("/api/status")
.header("authorization", "Bearer secret_token_123")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
// Test mark_grid_point_running
let params = common::models::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(),
};
db.upsert_grid_point(&params, 0, "test_wf").await.unwrap();
db.mark_grid_point_running(&params.model_name(), "test_wf")
.await
.unwrap();
let stats = db.get_grid_summary_stats(None).await.unwrap();
assert_eq!(stats["running"], 1);
}
/// L2 鉴权核心流程测试:
/// 注册节点 → 颁发专属 token → 反查校验 → 重发后旧 token 失效(覆盖式,无需独立吊销)。
#[tokio::test]
async fn test_l2_node_token_issue_reissue_flow() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("l2_db.db");
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
// 1. 注册节点
let reg = common::models::NodeRegisterRequest {
node_id: "node-l2-test".to_string(),
max_slots: 4,
};
let (_, _, secret) = db.register_node(&reg).await.unwrap();
// 2. 颁发专属 token,返回明文
let token = db.issue_node_token("node-l2-test").await.unwrap();
assert!(!token.is_empty());
// 2b. 一次性取走暂存明文(take_pending_node_token):首次取到与颁发一致的明文,
// 再次取为 None(取走即焚)。验证 #4 简化为单一 UPDATE...RETURNING 后行为一致。
// 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", secret.as_deref())
.await
.unwrap();
assert!(
pending2.is_none(),
"取走即焚:第二次 take_pending 必须返回 None"
);
// 不存在的 node take 也应返回 None(不报错)
assert!(db
.take_pending_node_token("node-not-exist", None)
.await
.unwrap()
.is_none());
// 3. token 可反查到 node_id(此调用会把 token 写入内存缓存)
let found = db.find_node_by_token(&token).await;
assert_eq!(found.as_deref(), Some("node-l2-test"));
// 4. 错误 token 查不到
assert!(db.find_node_by_token("wrong-token").await.is_none());
// 5. 重发(reissue)后新 token 生效;旧 token 因 hash 被覆盖而立即失效。
// 这正是「重发取代吊销」的核心证据:无需独立 revoked 标记,覆盖 hash 即让旧 token 失效。
let token2 = db.issue_node_token("node-l2-test").await.unwrap();
assert_eq!(
db.find_node_by_token(&token2).await.as_deref(),
Some("node-l2-test")
);
// 旧 token 仍失效(已被覆盖)
assert!(db.find_node_by_token(&token).await.is_none());
}
/// 验证中间件对 node token 的端到端鉴权:
/// 用真实 node token 调 /node/heartbeat 应 200;重发(覆盖旧 token)后再调应 401。
#[tokio::test]
async fn test_middleware_node_role_with_token() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("l2_mw_db.db");
let queue_db_path = temp_dir.path().join("l2_mw_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
// 注册并颁发 token
let reg = common::models::NodeRegisterRequest {
node_id: "node-mw-test".to_string(),
max_slots: 2,
};
db.register_node(&reg).await.unwrap();
let token = db.issue_node_token("node-mw-test").await.unwrap();
let state = AppState {
db: db.clone(),
queue,
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
5,
std::time::Duration::from_secs(300),
),
admin_token: None,
auth_disabled: false,
admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new(
std::collections::HashMap::new(),
)),
};
let api_router = axum::Router::new().route(
"/node/heartbeat",
axum::routing::post(server::api::node::heartbeat_node),
);
let auth_layer =
axum::middleware::from_fn_with_state(state.clone(), server::api::auth_middleware);
let api_router = api_router.layer(auth_layer);
let app = axum::Router::new()
.nest("/api", api_router)
.with_state(state);
// 用有效 node token 调 heartbeat → 200
let hb_body = serde_json::json!({
"node_id": "node-mw-test", "active_slots": 1, "cpu_usage": 10.0, "memory_usage": 20.0
});
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/node/heartbeat")
.header("authorization", format!("Bearer {}", token))
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&hb_body).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
// 重发(reissue)后旧 token 经中间件 → 401reissue 覆盖 token_hash
// 旧 token 明文 hash 不再匹配,find_node_by_token 返回 None → 鉴权失败。
// 这是「重发取代吊销」的端到端证据(经真实中间件而非仅 DB 层)。
db.issue_node_token("node-mw-test").await.unwrap();
let res = app
.oneshot(
Request::builder()
.method("POST")
.uri("/api/node/heartbeat")
.header("authorization", format!("Bearer {}", token))
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&hb_body).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::UNAUTHORIZED);
}
/// 管理 APIAdmin)端到端测试:
/// 列出节点 → 重发 token → 鉴权(admin 放行,node/匿名 401)。
#[tokio::test]
async fn test_admin_node_management_api() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("admin_db.db");
let queue_db_path = temp_dir.path().join("admin_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
// 预置一个已注册并颁发 token 的节点
let reg = common::models::NodeRegisterRequest {
node_id: "node-admin-test".to_string(),
max_slots: 2,
};
db.register_node(&reg).await.unwrap();
let original_token = db.issue_node_token("node-admin-test").await.unwrap();
let state = AppState {
db: db.clone(),
queue,
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
5,
std::time::Duration::from_secs(300),
),
admin_token: Some("admin-secret".to_string()),
auth_disabled: false,
admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new(
std::collections::HashMap::new(),
)),
};
let api_router = axum::Router::new()
.route(
"/admin/nodes",
axum::routing::get(server::api::admin::list_nodes),
)
.route(
"/admin/nodes/:node_id/reissue",
axum::routing::post(server::api::admin::reissue_node),
);
let auth_layer =
axum::middleware::from_fn_with_state(state.clone(), server::api::auth_middleware);
let app = axum::Router::new()
.nest("/api", api_router.layer(auth_layer))
.with_state(state);
// 1. 匿名访问 /admin/nodes → 401
let res = app
.clone()
.oneshot(
Request::builder()
.uri("/api/admin/nodes")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::UNAUTHORIZED);
// 2. admin 访问 /admin/nodes → 200,且能看到预置节点 token_status=active
let res = app
.clone()
.oneshot(
Request::builder()
.uri("/api/admin/nodes")
.header("authorization", "Bearer admin-secret")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let body: serde_json::Value = serde_json::from_slice(
&axum::body::to_bytes(res.into_body(), usize::MAX)
.await
.unwrap(),
)
.unwrap();
let nodes = body["data"].as_array().expect("data 应为数组");
let target = nodes
.iter()
.find(|n| n["node_id"] == "node-admin-test")
.expect("应包含预置节点");
assert_eq!(target["token_status"], "active");
// 3. admin 重发 token → 200,返回新明文 token
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/admin/nodes/node-admin-test/reissue")
.header("authorization", "Bearer admin-secret")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let body: serde_json::Value = serde_json::from_slice(
&axum::body::to_bytes(res.into_body(), usize::MAX)
.await
.unwrap(),
)
.unwrap();
let new_token = body["node_token"].as_str().expect("应返回 node_token");
assert!(!new_token.is_empty());
assert_ne!(new_token, original_token);
// 新 token 可用
assert_eq!(
db.find_node_by_token(new_token).await.as_deref(),
Some("node-admin-test")
);
// 重发使旧 token 立即失效(hash 被覆盖)——这正是「重发取代吊销」的核心语义
assert!(
db.find_node_by_token(&original_token).await.is_none(),
"重发后旧 token 必须立即失效"
);
// 4. node token 不能访问 admin API → 401(即便持有有效 node token
let res = app
.oneshot(
Request::builder()
.uri("/api/admin/nodes")
.header("authorization", format!("Bearer {}", new_token))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::UNAUTHORIZED);
}
/// S1 身份绑定测试:节点 A 用自己 token 冒充节点 B 发心跳 → 403;一致时 → 200。
#[tokio::test]
async fn test_node_heartbeat_node_id_binding() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("bind_db.db");
let queue_db_path = temp_dir.path().join("bind_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue));
// 预置节点 A,颁发 token
let reg = common::models::NodeRegisterRequest {
node_id: "node-A".to_string(),
max_slots: 2,
};
db.register_node(&reg).await.unwrap();
let token_a = db.issue_node_token("node-A").await.unwrap();
let state = AppState {
db: db.clone(),
queue: Arc::new(SqliteTaskQueue::new(":memory:").await.unwrap()),
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
5,
std::time::Duration::from_secs(300),
),
admin_token: None,
auth_disabled: false,
admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new(
std::collections::HashMap::new(),
)),
};
let api_router = axum::Router::new().route(
"/node/heartbeat",
axum::routing::post(server::api::node::heartbeat_node),
);
let auth_layer =
axum::middleware::from_fn_with_state(state.clone(), server::api::auth_middleware);
let app = axum::Router::new()
.nest("/api", api_router.layer(auth_layer))
.with_state(state);
// A 用自己 token,但 body 声称 node_id=node-B(冒充)→ 403
let hb =
serde_json::json!({"node_id":"node-B","active_slots":1,"cpu_usage":0.0,"memory_usage":0.0});
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/node/heartbeat")
.header("authorization", format!("Bearer {}", token_a))
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&hb).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::FORBIDDEN);
// A 用自己 tokenbody node_id=node-A(一致)→ 200
let hb_ok =
serde_json::json!({"node_id":"node-A","active_slots":1,"cpu_usage":0.0,"memory_usage":0.0});
let res = app
.oneshot(
Request::builder()
.method("POST")
.uri("/api/node/heartbeat")
.header("authorization", format!("Bearer {}", token_a))
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&hb_ok).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_admin_login_flow() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("login_db.db");
let queue_db_path = temp_dir.path().join("login_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
let admin_pass = "my_short_admin_password_123";
let state = AppState {
db,
queue,
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
5,
std::time::Duration::from_secs(300),
),
admin_token: Some(admin_pass.to_string()),
auth_disabled: false,
admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new(
std::collections::HashMap::new(),
)),
};
let api_router = axum::Router::new()
.route("/login", axum::routing::post(server::api::auth::login))
.route(
"/auth/check",
axum::routing::get(server::api::auth::check_auth),
)
.layer(axum::middleware::from_fn_with_state(
state.clone(),
server::api::auth_middleware,
));
let app = axum::Router::new()
.nest("/api", api_router)
.with_state(state);
// 1. 密码错误 ➔ 401
let wrong_body = serde_json::json!({ "password": "wrong_password" });
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/login")
.header("content-type", "application/json")
.extension(axum::extract::ConnectInfo(std::net::SocketAddr::from((
[127, 0, 0, 1],
12345,
))))
.body(Body::from(serde_json::to_vec(&wrong_body).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::UNAUTHORIZED);
// 2. 正确密码 ➔ 200 + token
let right_body = serde_json::json!({ "password": admin_pass });
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/login")
.header("content-type", "application/json")
.extension(axum::extract::ConnectInfo(std::net::SocketAddr::from((
[127, 0, 0, 1],
12345,
))))
.body(Body::from(serde_json::to_vec(&right_body).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let bytes = axum::body::to_bytes(res.into_body(), 1024 * 1024)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
assert_eq!(json["success"], true);
let token = json["token"].as_str().unwrap();
assert!(token.len() == 64);
// 3. 携带拿到的 Token 访问受保护的 /api/auth/check ➔ 200
let res = app
.oneshot(
Request::builder()
.method("GET")
.uri("/api/auth/check")
.header("authorization", format!("Bearer {}", token))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_node_approval_workflow() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("appr_db.db");
let queue_db_path = temp_dir.path().join("appr_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
let admin_pass = "admin_approval_secret";
let state = AppState {
db,
queue,
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
5,
std::time::Duration::from_secs(300),
),
admin_token: Some(admin_pass.to_string()),
auth_disabled: false,
admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new(
std::collections::HashMap::new(),
)),
};
let api_router = axum::Router::new()
.route(
"/node/register",
axum::routing::post(server::api::node::register_node),
)
.route(
"/node/check_status",
axum::routing::post(server::api::node::check_node_status),
)
.route(
"/node/heartbeat",
axum::routing::post(server::api::node::heartbeat_node),
)
.route(
"/admin/nodes/:node_id/approve",
axum::routing::post(server::api::admin::approve_node),
)
.layer(axum::middleware::from_fn_with_state(
state.clone(),
server::api::auth_middleware,
));
let app = axum::Router::new()
.nest("/api", api_router)
.with_state(state);
// 1. 新节点免凭据申请注册 ➔ 200 + status: pending_approval
let reg_body = serde_json::json!({
"node_id": "node-pending-01",
"max_slots": 4
});
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/node/register")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&reg_body).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let bytes = axum::body::to_bytes(res.into_body(), 1024 * 1024)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
assert_eq!(json["status"], "pending_approval");
// 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",
"registration_secret": reg_secret,
});
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/node/check_status")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&check_body).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let bytes = axum::body::to_bytes(res.into_body(), 1024 * 1024)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
assert_eq!(json["status"], "pending_approval");
// 3. 管理员在 Dashboard 点击同意 ➔ 200
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/admin/nodes/node-pending-01/approve")
.header("authorization", format!("Bearer {}", admin_pass))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let bytes = axum::body::to_bytes(res.into_body(), 1024 * 1024)
.await
.unwrap();
println!("Step 3 body: {}", String::from_utf8_lossy(&bytes));
// 4. Node 端再次轮询查状态 ➔ status: approved + 获取 node_token
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/node/check_status")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&check_body).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let bytes = axum::body::to_bytes(res.into_body(), 1024 * 1024)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
println!("Step 4 status returned: {:?}", json);
assert_eq!(json["status"], "approved");
let node_token = json["node_token"].as_str().unwrap();
// 5. Node 携带拿到到的专属 Token 发送心跳 ➔ 200
let hb_body = serde_json::json!({
"node_id": "node-pending-01",
"active_slots": 1,
"cpu_usage": 10.0,
"memory_usage": 20.0
});
let res = app
.oneshot(
Request::builder()
.method("POST")
.uri("/api/node/heartbeat")
.header("authorization", format!("Bearer {}", node_token))
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&hb_body).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_admin_sessions_capacity_limit() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("sess_db.db");
let queue_db_path = temp_dir.path().join("sess_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
let admin_pass = "admin_capacity_secret";
let state = AppState {
db,
queue,
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
1000,
std::time::Duration::from_secs(300),
),
admin_token: Some(admin_pass.to_string()),
auth_disabled: false,
admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new(
std::collections::HashMap::new(),
)),
};
let api_router =
axum::Router::new().route("/login", axum::routing::post(server::api::auth::login));
let app = axum::Router::new()
.nest("/api", api_router)
.with_state(state.clone());
// 连续登录 105 次,超过 100 容量上限
for _ in 0..105 {
let right_body = serde_json::json!({ "password": admin_pass });
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/login")
.header("content-type", "application/json")
.extension(axum::extract::ConnectInfo(std::net::SocketAddr::from((
[127, 0, 0, 1],
12345,
))))
.body(Body::from(serde_json::to_vec(&right_body).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
}
// 验证 sessions 集合保存的数量不超过 MAX_ADMIN_SESSIONS (100)
let sessions = state.admin_sessions.read().await;
assert!(sessions.len() <= server::api::MAX_ADMIN_SESSIONS);
assert_eq!(sessions.len(), 100);
}
#[tokio::test]
async fn test_node_register_rate_limit() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("reg_limit_db.db");
let queue_db_path = temp_dir.path().join("reg_limit_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
let state = AppState {
db,
queue,
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
5,
std::time::Duration::from_secs(300),
),
admin_token: None,
auth_disabled: false,
admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new(
std::collections::HashMap::new(),
)),
};
// 注册端点专用限流器(count_all=true):对成功请求也计数,模拟生产 main.rs 的注册节流配置。
let register_limiter =
server::api::rate_limit::RateLimiter::new_count_all(5, std::time::Duration::from_secs(60));
let register_rate_limit_layer = axum::middleware::from_fn_with_state(
register_limiter,
server::api::rate_limit::rate_limit_middleware,
);
let app = axum::Router::new()
.route(
"/api/node/register",
axum::routing::post(server::api::node::register_node).layer(register_rate_limit_layer),
)
.with_state(state);
let reg_body = serde_json::json!({
"node_id": "test-limit-node",
"max_slots": 4
});
// 5 次以内的注册尝试 ➔ 200
for _ in 0..5 {
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/node/register")
.header("content-type", "application/json")
.extension(axum::extract::ConnectInfo(std::net::SocketAddr::from((
[127, 0, 0, 1],
12345,
))))
.body(Body::from(serde_json::to_vec(&reg_body).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
}
// 第 6 次触发限流 ➔ 429 TOO_MANY_REQUESTS
let res = app
.oneshot(
Request::builder()
.method("POST")
.uri("/api/node/register")
.header("content-type", "application/json")
.extension(axum::extract::ConnectInfo(std::net::SocketAddr::from((
[127, 0, 0, 1],
12345,
))))
.body(Body::from(serde_json::to_vec(&reg_body).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::TOO_MANY_REQUESTS);
}
/// 回归保护:通用 API 限流器(count_all=false)不应因 /node/register 的成功响应计数。
/// 历史缺陷:此前中间件对 is_register 无条件计数,导致通用限流器复用时成功注册会把 IP
/// 锁出整个 /api/*(跨端点连锁)。此测试断言多次成功注册后通用限流器仍放行。
#[tokio::test]
async fn test_general_limiter_does_not_count_successful_register() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("gen_limit_db.db");
let queue_db_path = temp_dir.path().join("gen_limit_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
let state = AppState {
db,
queue,
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
5,
std::time::Duration::from_secs(300),
),
admin_token: None,
auth_disabled: false,
admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new(
std::collections::HashMap::new(),
)),
};
// 通用限流器(new = count_all=false),阈值仅 3,远低于下面的请求数。
// 若仍对成功注册计数,第 4 次就会 429。
let general_limiter =
server::api::rate_limit::RateLimiter::new(3, std::time::Duration::from_secs(60));
let rate_limit_layer = axum::middleware::from_fn_with_state(
general_limiter,
server::api::rate_limit::rate_limit_middleware,
);
let app = axum::Router::new()
.route(
"/api/node/register",
axum::routing::post(server::api::node::register_node),
)
.layer(rate_limit_layer)
.with_state(state);
let reg_body = serde_json::json!({
"node_id": "gen-limit-node",
"max_slots": 4
});
// 连续 6 次成功注册(远超阈值 3):通用限流器不应计数成功响应,全部应为 200。
for i in 0..6 {
// 每次用不同 node_id 避免重复注册逻辑干扰
let body = serde_json::json!({
"node_id": format!("gen-limit-node-{}", i),
"max_slots": 4
});
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/node/register")
.header("content-type", "application/json")
.extension(axum::extract::ConnectInfo(std::net::SocketAddr::from((
[127, 0, 0, 1],
12345,
))))
.body(Body::from(serde_json::to_vec(&body).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(
res.status(),
StatusCode::OK,
"第 {} 次注册应成功,通用限流器不应计数成功响应",
i + 1
);
}
// 静默 reg_body 未使用的警告(保留以对齐其它测试结构)
let _ = &reg_body;
}
#[tokio::test]
async fn test_cors_same_origin_and_local_policy() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("cors_db.db");
let queue_db_path = temp_dir.path().join("cors_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
let state = AppState {
db,
queue,
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
100,
std::time::Duration::from_secs(300),
),
admin_token: None,
auth_disabled: false,
admin_sessions: Arc::new(tokio::sync::RwLock::new(std::collections::HashMap::new())),
};
let app = axum::Router::new()
.route(
"/api/status",
axum::routing::get(server::api::status::get_status),
)
.layer(server::cors::build_cors_layer())
.with_state(state);
// 1. 本地 localhost 请求 -> 允许 CORS
let res = app
.clone()
.oneshot(
Request::builder()
.uri("/api/status")
.header("origin", "http://localhost:5173")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
assert_eq!(
res.headers()
.get("access-control-allow-origin")
.unwrap()
.to_str()
.unwrap(),
"http://localhost:5173"
);
// 2. 本地 127.0.0.1 请求 -> 允许 CORS
let res = app
.clone()
.oneshot(
Request::builder()
.uri("/api/status")
.header("origin", "http://127.0.0.1:3000")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
assert_eq!(
res.headers()
.get("access-control-allow-origin")
.unwrap()
.to_str()
.unwrap(),
"http://127.0.0.1:3000"
);
// 3. 同源请求 (Origin 匹配 Host 标头) -> 允许 CORS
let res = app
.clone()
.oneshot(
Request::builder()
.uri("/api/status")
.header("host", "192.168.1.100:8090")
.header("origin", "http://192.168.1.100:8090")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
assert_eq!(
res.headers()
.get("access-control-allow-origin")
.unwrap()
.to_str()
.unwrap(),
"http://192.168.1.100:8090"
);
// 4. 外部非法跨源请求 -> 拒绝 CORS
let res = app
.oneshot(
Request::builder()
.uri("/api/status")
.header("host", "192.168.1.100:8090")
.header("origin", "https://attacker.example.com")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
assert!(res.headers().get("access-control-allow-origin").is_none());
}
/// 构造一个合法的旧版 conv.json (ModelSummary) 文本,用于 import_seed 测试。
/// 关键:name 用源精度真名(`g5.0`),验证服务端逐字符保真落库。
fn make_legacy_conv_json(name: &str, converged: bool) -> String {
let summary = serde_json::json!({
"name": name,
"params": {"teff": 20000.0, "logg": 5.0, "loghe": -2.0, "logc": -4.0, "logn": -4.0, "logo": -4.0},
"stages": [],
"converged": converged,
"final_max_relc": 0.000321,
"final_chmax": null,
"seed": null,
"atmosphere_has_nan": false,
"synspec_rc": 0,
"synspec_error": null,
"synspec_sec": 1.5,
"elapsed_sec": 42.0,
"note": null,
});
serde_json::to_string(&summary).unwrap()
}
/// 构造**旧版 Python run_one.py 形态**的 conv.jsonstage 含 itek_attempts/final 嵌套,
/// 顶层含 elapsed_sec/synspec_*)——验证 import_seed 对真实历史数据的解析兼容性。
fn make_python_legacy_conv_json(name: &str) -> String {
serde_json::json!({
"name": name,
"params": {"teff": 20000.0, "logg": 5.0, "loghe": -2.0, "logc": -4.0, "logn": -4.0, "logo": -4.0},
"stages": [
{"label": "lte", "chmax": null, "lte": "T",
"itek_attempts": [{"itek": null, "rc": 0, "converged": true, "max_relc": 0.0}],
"converged": true,
"final": {"itek": null, "rc": 0, "converged": true, "max_relc": 0.0},
"best_max_relc": 0.0, "elapsed_sec": 2.1},
{"label": "nl", "chmax": null, "lte": "F",
"itek_attempts": [{"itek": null, "rc": 0, "converged": true, "max_relc": 0.000321,
"worst_depth": 1, "last_iter": 17, "n_depths": 50}],
"converged": true,
"final": {"itek": null, "rc": 0, "converged": true, "max_relc": 0.000321,
"worst_depth": 1, "last_iter": 17, "n_depths": 50},
"best_max_relc": 0.000321, "elapsed_sec": 640.0}
],
"converged": true,
"final_max_relc": 0.000321,
"final_chmax": null,
"seed": null,
"atmosphere_has_nan": false,
"synspec_rc": 0,
"synspec_sec": 3.1,
"elapsed_sec": 715.0,
})
.to_string()
}
/// 构造一个 multipart/form-data body,含 report(JSON 文本) + seed_file(二进制)。
fn make_import_multipart(
boundary: &str,
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());
body.extend_from_slice(b"Content-Disposition: form-data; name=\"report\"\r\n");
body.extend_from_slice(b"Content-Type: application/json\r\n\r\n");
body.extend_from_slice(report_json.as_bytes());
body.extend_from_slice(b"\r\n");
body.extend_from_slice(format!("--{}\r\n", boundary).as_bytes());
body.extend_from_slice(
format!(
"Content-Disposition: form-data; name=\"seed_file\"; filename=\"{}\"\r\n",
seed_name
)
.as_bytes(),
);
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
}
#[tokio::test]
async fn test_import_seed_admin_endpoint() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("import_db.db");
let queue_db_path = temp_dir.path().join("import_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
let state = AppState {
db: db.clone(),
queue,
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
5,
std::time::Duration::from_secs(300),
),
admin_token: Some("admin-secret".to_string()),
auth_disabled: false,
admin_sessions: Arc::new(tokio::sync::RwLock::new(std::collections::HashMap::new())),
};
let api_router = axum::Router::new().route(
"/admin/import_seed",
axum::routing::post(server::api::task::import_seed),
);
let auth_layer =
axum::middleware::from_fn_with_state(state.clone(), server::api::auth_middleware);
let app = axum::Router::new()
.nest("/api", api_router.layer(auth_layer))
.with_state(state);
// 1. 无 admin token → 401
let conv = make_legacy_conv_json("t20000_g5.0_he-2_c-4_n-4_o-4", true);
let body_bytes = make_import_multipart(
"boundary1",
&conv,
b"FAKE_ATMOS_7",
"t20000_g5.0_he-2_c-4_n-4_o-4.7",
"cold_run",
);
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/admin/import_seed?workflow=wf_import")
.header("content-type", "multipart/form-data; boundary=boundary1")
.body(Body::from(body_bytes))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::UNAUTHORIZED);
// 2. admin token + 收敛点 → 200conv.json + .7 落地,grid_points=converged
let body_bytes = make_import_multipart(
"boundary2",
&conv,
b"FAKE_ATMOS_7",
"t20000_g5.0_he-2_c-4_n-4_o-4.7",
"cold_run",
);
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/admin/import_seed?workflow=wf_import")
.header("authorization", "Bearer admin-secret")
.header("content-type", "multipart/form-data; boundary=boundary2")
.body(Body::from(body_bytes))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
// 关键断言:磁盘目录名用旧 conv.json 的源精度真名(g5.0),conv.json 与 .7 均落地。
let point_dir = seeds_dir.join("t20000_g5.0_he-2_c-4_n-4_o-4");
assert!(point_dir.join("conv.json").is_file(), "conv.json 应落地");
assert!(
point_dir.join("t20000_g5.0_he-2_c-4_n-4_o-4.7").is_file(),
".7 种子文件应落地"
);
// grid_points 应被幂等 upsert 且标记为 converged。
let gp = db
.get_grid_point_status("t20000_g5.0_he-2_c-4_n-4_o-4", "wf_import")
.await
.unwrap()
.expect("grid_points 应存在");
assert_eq!(gp.0, "converged", "导入的收敛点应为 converged 状态");
// 3. 幂等:重复导入同名点不应报错,状态仍 converged。
let body_bytes = make_import_multipart(
"boundary3",
&conv,
b"FAKE_ATMOS_7_AGAIN",
"t20000_g5.0_he-2_c-4_n-4_o-4.7",
"cold_run",
);
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/admin/import_seed?workflow=wf_import")
.header("authorization", "Bearer admin-secret")
.header("content-type", "multipart/form-data; boundary=boundary3")
.body(Body::from(body_bytes))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let gp = db
.get_grid_point_status("t20000_g5.0_he-2_c-4_n-4_o-4", "wf_import")
.await
.unwrap()
.unwrap();
assert_eq!(gp.0, "converged");
// 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", "cold_run");
let res = app
.oneshot(
Request::builder()
.method("POST")
.uri("/api/admin/import_seed?workflow=wf_import")
.header("authorization", "Bearer admin-secret")
.header("content-type", "multipart/form-data; boundary=boundary4")
.body(Body::from(body_bytes))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let fail_dir = seeds_dir.join("t20000_g5.0_he-2_c-4_n-4_o-4_fail");
assert!(
fail_dir.join("conv.json").is_file(),
"未收敛点仍应记录 conv.json"
);
assert!(
!fail_dir
.join("t20000_g5.0_he-2_c-4_n-4_o-4_fail.7")
.exists(),
"未收敛点不应写 .7 种子"
);
}
/// 旧版 Python run_one.py conv.json 的完整导入链路(import_results 工具的服务端侧):
/// 嵌套 stagesitek_attempts/final)应被正常解析,收敛标记落地,且旧版的
/// elapsed_sec 经 P3 迁移链路进入 grid_points.last_elapsed_sec(迁移完整性)。
#[tokio::test]
async fn test_import_seed_python_legacy_conv_json() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("import_legacy_db.db");
let queue_db_path = temp_dir.path().join("import_legacy_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
let state = AppState {
db: db.clone(),
queue,
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
5,
std::time::Duration::from_secs(300),
),
admin_token: Some("admin-secret".to_string()),
auth_disabled: false,
admin_sessions: Arc::new(tokio::sync::RwLock::new(std::collections::HashMap::new())),
};
let api_router = axum::Router::new().route(
"/admin/import_seed",
axum::routing::post(server::api::task::import_seed),
);
let auth_layer =
axum::middleware::from_fn_with_state(state.clone(), server::api::auth_middleware);
let app = axum::Router::new()
.nest("/api", api_router.layer(auth_layer))
.with_state(state);
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"),
"cold_run",
);
let res = app
.oneshot(
Request::builder()
.method("POST")
.uri("/api/admin/import_seed?workflow=wf_legacy")
.header("authorization", "Bearer admin-secret")
.header("content-type", "multipart/form-data; boundary=boundaryL")
.body(Body::from(body_bytes))
.unwrap(),
)
.await
.unwrap();
assert_eq!(
res.status(),
StatusCode::OK,
"旧版嵌套 stages 的 conv.json 应被接受"
);
// grid_pointsconverged + cold_run 手段(旧版 conv.json 无 seed_nc 阶段)+ 旧版 elapsed_sec 已落库
let row = db
.get_workflow_point_row("wf_legacy", name)
.await
.unwrap()
.expect("grid_points 应存在");
assert_eq!(row.status, "converged");
assert_eq!(row.success_method.as_deref(), Some("cold_run"));
assert_eq!(
row.last_elapsed_sec,
Some(715.0),
"旧版 conv.json 的单点耗时应完整迁移到 last_elapsed_sec"
);
// conv.json 原文落盘(嵌套诊断数据 itek_attempts/final 不丢失,详情页解析走兼容路径)
assert!(seeds_dir.join(name).join("conv.json").is_file());
assert!(seeds_dir.join(name).join(format!("{name}.7")).is_file());
}
#[tokio::test]
async fn test_node_disable_enable_flow() {
use common::models::{GridAxisValue, GridPointParams, TaskSpec, TaskType};
use uuid::Uuid;
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("disable_db.db");
let queue_db_path = temp_dir.path().join("disable_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
// 1. 注册并审批 node-disable-test(→ online),颁发 token
let reg = common::models::NodeRegisterRequest {
node_id: "node-disable-test".to_string(),
max_slots: 2,
};
db.register_node(&reg).await.unwrap();
let token = db.approve_node("node-disable-test").await.unwrap();
// 预置一个 pending_approval 节点(用于后续 409 断言)
let reg_pending = common::models::NodeRegisterRequest {
node_id: "node-pending".to_string(),
max_slots: 2,
};
db.register_node(&reg_pending).await.unwrap();
let state = AppState {
db: db.clone(),
queue: queue.clone(),
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
5,
std::time::Duration::from_secs(300),
),
admin_token: Some("admin-secret".to_string()),
auth_disabled: false,
admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new(
std::collections::HashMap::new(),
)),
};
let api_router = axum::Router::new()
.route(
"/task/claim",
axum::routing::post(server::api::task::claim_task),
)
.route(
"/node/heartbeat",
axum::routing::post(server::api::node::heartbeat_node),
)
.route(
"/admin/nodes/:node_id/disable",
axum::routing::post(server::api::admin::disable_node),
)
.route(
"/admin/nodes/:node_id/enable",
axum::routing::post(server::api::admin::enable_node),
);
let auth_layer =
axum::middleware::from_fn_with_state(state.clone(), server::api::auth_middleware);
let app = axum::Router::new()
.nest("/api", api_router.layer(auth_layer))
.with_state(state);
// 2. 基线:向队列 push 一个任务,claim 应返回 {"status":"ok"}
let task = TaskSpec {
task_id: Uuid::new_v4(),
point_name: "p_disable_test".to_string(),
params: GridPointParams {
teff: GridAxisValue::from_value(20000.0),
logg: GridAxisValue::from_value(5.0),
loghe: GridAxisValue::from_value(-2.0),
logc: GridAxisValue::from_value(-2.0),
logn: GridAxisValue::from_value(-2.0),
logo: GridAxisValue::from_value(-2.0),
},
task_type: TaskType::ColdRun,
seed_point_name: None,
timeout_sec: 60,
workflow_name: Some("wf_test".to_string()),
wave: 0,
..Default::default()
};
queue.push_task(&task).await.unwrap();
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/task/claim")
.header("authorization", format!("Bearer {}", token))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let body: serde_json::Value = serde_json::from_slice(
&axum::body::to_bytes(res.into_body(), usize::MAX)
.await
.unwrap(),
)
.unwrap();
assert_eq!(body["status"], "ok", "online 节点应能领用任务");
// 3. admin 停用节点 → 200;状态切为 disabled
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/admin/nodes/node-disable-test/disable")
.header("authorization", "Bearer admin-secret")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
assert!(
db.is_node_disabled("node-disable-test").await.unwrap(),
"停用后 is_node_disabled 应为 true"
);
// 3a. 被停用节点从「在线节点」列表中移除(/api/status 在线计数依据)
let active = db.get_active_nodes().await.unwrap();
assert!(
!active.iter().any(|n| n.node_id == "node-disable-test"),
"被停用节点不应出现在在线列表"
);
// 4. 被停用节点再次 claim → {"status":"disabled"}(不再分发,HTTP 仍 200
// 再 push 一个任务确认它真的领不到。
let task2 = TaskSpec {
task_id: Uuid::new_v4(),
..task.clone()
};
queue.push_task(&task2).await.unwrap();
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/task/claim")
.header("authorization", format!("Bearer {}", token))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let body: serde_json::Value = serde_json::from_slice(
&axum::body::to_bytes(res.into_body(), usize::MAX)
.await
.unwrap(),
)
.unwrap();
assert_eq!(
body["status"], "disabled",
"被停用节点 claim 必须返回 disabled(不可分发任务)"
);
// 队列里应仍有 pending 任务(未被领用):注册第二个 online 节点领用,
// 应当拿到刚才 push 的 task2,证明停用节点没有消费它。
let reg2 = common::models::NodeRegisterRequest {
node_id: "node-helper".to_string(),
max_slots: 2,
};
db.register_node(&reg2).await.unwrap();
let token2 = db.approve_node("node-helper").await.unwrap();
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/task/claim")
.header("authorization", format!("Bearer {}", token2))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
let body: serde_json::Value = serde_json::from_slice(
&axum::body::to_bytes(res.into_body(), usize::MAX)
.await
.unwrap(),
)
.unwrap();
assert_eq!(
body["status"], "ok",
"另一在线节点应能领用到停用节点未消费的任务"
);
// 5. 被停用节点心跳:状态保持 disabled(不被复活),但 cpu/mem/心跳已刷新
let hb = serde_json::json!({"node_id":"node-disable-test","active_slots":0,"cpu_usage":77.7,"memory_usage":88.8});
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/node/heartbeat")
.header("authorization", format!("Bearer {}", token))
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&hb).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let nodes = db.list_nodes_with_credentials().await.unwrap();
let me = nodes
.iter()
.find(|n| n.node_id == "node-disable-test")
.unwrap();
assert_eq!(me.status, "disabled", "心跳不得把 disabled 复活为 online");
assert!(
(me.cpu_usage - 77.7).abs() < 0.01,
"心跳应刷新 cpu_usage(管理员仍可见节点存活)"
);
// 6. admin 启用节点 → 200;状态切为 offline(待心跳自然翻为 online
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/admin/nodes/node-disable-test/enable")
.header("authorization", "Bearer admin-secret")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
// 7. 启用后心跳 → status 翻为 online,重新出现在在线列表,可恢复领用
let hb = serde_json::json!({"node_id":"node-disable-test","active_slots":0,"cpu_usage":10.0,"memory_usage":20.0});
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/node/heartbeat")
.header("authorization", format!("Bearer {}", token))
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&hb).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
assert!(
db.get_active_nodes()
.await
.unwrap()
.iter()
.any(|n| n.node_id == "node-disable-test"),
"启用并心跳后应重新出现在线"
);
// 8. 状态幂等保护:对 online 节点再次 enable → 409
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/admin/nodes/node-disable-test/enable")
.header("authorization", "Bearer admin-secret")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(
res.status(),
StatusCode::CONFLICT,
"对非 disabled 节点 enable 应返回 409"
);
// 9. 对 pending_approval 节点 disable → 409(审批流程独立)
let res = app
.oneshot(
Request::builder()
.method("POST")
.uri("/api/admin/nodes/node-pending/disable")
.header("authorization", "Bearer admin-secret")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(
res.status(),
StatusCode::CONFLICT,
"待审批节点不可停用(仅 online/offline 可停用)"
);
}
#[tokio::test]
async fn test_admin_set_node_quota_and_heartbeat_sync() {
// 验证动态 CPU 槽位配额下发链路(见 docs/dynamic_cpu_slots_design.md):
// admin POST /quota 设 admin_max_slots → 落库 nodes.admin_max_slots →
// list_nodes_with_credentials 返回配额 → 心跳响应体透传 admin_max_slots。
// 覆盖:设值 / 清除(null) / 负数 400 / 不存在节点 404。
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("quota_db.db");
let queue_db_path = temp_dir.path().join("quota_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
// 注册并审批 node-quota-test,颁发 token
let reg = common::models::NodeRegisterRequest {
node_id: "node-quota-test".to_string(),
max_slots: 4,
};
db.register_node(&reg).await.unwrap();
let token = db.approve_node("node-quota-test").await.unwrap();
let state = AppState {
db: db.clone(),
queue: queue.clone(),
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
5,
std::time::Duration::from_secs(300),
),
admin_token: Some("admin-secret".to_string()),
auth_disabled: false,
admin_sessions: std::sync::Arc::new(tokio::sync::RwLock::new(
std::collections::HashMap::new(),
)),
};
let api_router = axum::Router::new()
.route(
"/node/heartbeat",
axum::routing::post(server::api::node::heartbeat_node),
)
.route(
"/admin/nodes/:node_id/quota",
axum::routing::post(server::api::admin::set_node_quota),
);
let auth_layer =
axum::middleware::from_fn_with_state(state.clone(), server::api::auth_middleware);
let app = axum::Router::new()
.nest("/api", api_router.layer(auth_layer))
.with_state(state);
// 1. 基线心跳:admin_max_slots 应为 null(无配额限制)
let hb = serde_json::json!({"node_id":"node-quota-test","active_slots":0,"cpu_usage":10.0,"memory_usage":20.0});
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/node/heartbeat")
.header("authorization", format!("Bearer {}", token))
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&hb).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let body: serde_json::Value = serde_json::from_slice(
&axum::body::to_bytes(res.into_body(), usize::MAX)
.await
.unwrap(),
)
.unwrap();
assert_eq!(body["status"], "ok");
assert!(
body["admin_max_slots"].is_null(),
"未设配额时心跳响应 admin_max_slots 应为 null"
);
// 2. admin 设配额 = 2 → 200
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/admin/nodes/node-quota-test/quota")
.header("authorization", "Bearer admin-secret")
.header("content-type", "application/json")
.body(Body::from(
serde_json::to_vec(&serde_json::json!({"admin_max_slots": 2})).unwrap(),
))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
// 2a. 配额落库(list_nodes_with_credentials
let nodes = db.list_nodes_with_credentials().await.unwrap();
let me = nodes
.iter()
.find(|n| n.node_id == "node-quota-test")
.unwrap();
assert_eq!(me.admin_max_slots, Some(2), "配额应已落库为 Some(2)");
// 3. 心跳响应透传 admin_max_slots = 2
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/node/heartbeat")
.header("authorization", format!("Bearer {}", token))
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&hb).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let body: serde_json::Value = serde_json::from_slice(
&axum::body::to_bytes(res.into_body(), usize::MAX)
.await
.unwrap(),
)
.unwrap();
assert_eq!(body["admin_max_slots"], 2, "心跳响应应透传配额 2");
// 3a. 请求体缺 admin_max_slots 字段(畸形输入)→ 400,绝不能静默当作「清除配额」
// (审查修复 M1Option<i32> 缺字段默认 None 会把畸形请求误判为清除限制)。
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/admin/nodes/node-quota-test/quota")
.header("authorization", "Bearer admin-secret")
.header("content-type", "application/json")
.body(Body::from(
serde_json::to_vec(&serde_json::json!({})).unwrap(),
))
.unwrap(),
)
.await
.unwrap();
assert_eq!(
res.status(),
StatusCode::BAD_REQUEST,
"缺字段应 400 而非静默清除配额"
);
// 畸形请求无副作用:既有配额仍为 Some(2),未被静默清除
let nodes = db.list_nodes_with_credentials().await.unwrap();
let me = nodes
.iter()
.find(|n| n.node_id == "node-quota-test")
.unwrap();
assert_eq!(me.admin_max_slots, Some(2), "畸形请求不应改动既有配额");
// 4. 设 null 清除配额 → 心跳响应 admin_max_slots = null
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/admin/nodes/node-quota-test/quota")
.header("authorization", "Bearer admin-secret")
.header("content-type", "application/json")
.body(Body::from(
serde_json::to_vec(&serde_json::json!({"admin_max_slots": null})).unwrap(),
))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let nodes = db.list_nodes_with_credentials().await.unwrap();
let me = nodes
.iter()
.find(|n| n.node_id == "node-quota-test")
.unwrap();
assert_eq!(me.admin_max_slots, None, "清除后配额应为 None");
// 5. 负数 → 400
let res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/admin/nodes/node-quota-test/quota")
.header("authorization", "Bearer admin-secret")
.header("content-type", "application/json")
.body(Body::from(
serde_json::to_vec(&serde_json::json!({"admin_max_slots": -1})).unwrap(),
))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::BAD_REQUEST, "负数配额应 400");
// 6. 不存在节点 → 404
let res = app
.oneshot(
Request::builder()
.method("POST")
.uri("/api/admin/nodes/ghost-node/quota")
.header("authorization", "Bearer admin-secret")
.header("content-type", "application/json")
.body(Body::from(
serde_json::to_vec(&serde_json::json!({"admin_max_slots": 1})).unwrap(),
))
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::NOT_FOUND, "不存在节点应 404");
}
// ===== 工作流执行观测端点测试(stats / points / point detail =====
/// 测试专用:走真实写路径派发并回报一个网格点任务。
///
/// `insert_task` → `record_task_report` 会回填 `grid_points.success_method`、
/// `attempt_count` 与 `tasks.completed_at`,与生产链路一致(不绕过任何状态机逻辑)。
async fn dispatch_and_report(
db: &Database,
wf: &str,
p: &common::models::GridPointParams,
task_type: common::models::TaskType,
seed: Option<String>,
converged: bool,
) {
let task_id = uuid::Uuid::new_v4();
let spec = common::models::TaskSpec {
task_id,
point_name: p.model_name(),
params: p.clone(),
task_type,
seed_point_name: seed,
timeout_sec: 7200,
workflow_name: Some(wf.to_string()),
wave: 0,
..Default::default()
};
db.insert_task(&spec).await.unwrap();
let report = common::models::TaskReport {
task_id,
point_name: p.model_name(),
params: Some(p.clone()),
node_id: "test-node".to_string(),
status: if converged {
common::models::TaskStatus::Completed
} else {
common::models::TaskStatus::Failed
},
converged,
max_relc: if converged { Some(0.0005) } else { Some(9.5e5) },
atmosphere_has_nan: false,
elapsed_sec: 120.0,
error_message: if converged {
None
} else {
Some("nl stage diverged".to_string())
},
summary_json: "{}".to_string(),
failed_stage: None,
};
db.record_task_report(&report, wf).await.unwrap();
}
fn wf_stats_test_params(teff: f64, logc: f64) -> common::models::GridPointParams {
common::models::GridPointParams {
teff: teff.into(),
logg: 5.0.into(),
loghe: (-2.0).into(),
logc: logc.into(),
logn: (-4.0).into(),
logo: (-4.0).into(),
}
}
/// GET /api/workflows/:name/stats:鉴权、pending/queued 拆分、收敛手段归因、
/// wave 分布、近似 ETA、空工作流与未知工作流。
#[tokio::test]
async fn test_wf_stats_endpoint() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("wfstats_db.db");
let queue_db_path = temp_dir.path().join("wfstats_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
let state = AppState {
db: db.clone(),
queue,
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
5,
std::time::Duration::from_secs(300),
),
admin_token: Some("admin-secret".to_string()),
auth_disabled: false,
admin_sessions: Arc::new(tokio::sync::RwLock::new(std::collections::HashMap::new())),
};
let api_router = axum::Router::new().route(
"/workflows/:name/stats",
axum::routing::get(server::api::workflow::get_workflow_stats),
);
let auth_layer =
axum::middleware::from_fn_with_state(state.clone(), server::api::auth_middleware);
let app = axum::Router::new()
.nest("/api", api_router.layer(auth_layer))
.with_state(state);
// ---- 数据准备:wf_stats 7 个点,覆盖全部状态与三种收敛手段 ----
let yaml = "grid:\n teff: [20000]\n logg: [5.0]\n loghe: [-2]\n logc: [-4]\n logn: [-4]\n logo: [-4]";
db.upsert_workflow("wf_stats", Some("统计测试"), yaml, "running")
.await
.unwrap();
db.upsert_workflow("wf_empty", None, yaml, "idle")
.await
.unwrap();
let p_pending = wf_stats_test_params(20000.0, -4.0); // wave 0, pending
let p_queued = wf_stats_test_params(25000.0, -4.0); // wave 0, queued
let p_running = wf_stats_test_params(30000.0, -4.0); // wave 0, running
let p_cold = wf_stats_test_params(35000.0, -3.0); // wave 1, cold_run 收敛
let p_seed = wf_stats_test_params(40000.0, -3.0); // wave 1, seed_step 收敛
let p_failed = wf_stats_test_params(45000.0, -2.0); // wave 2, failed
let p_imported = wf_stats_test_params(50000.0, -2.0); // wave 2, imported 收敛
db.upsert_grid_point(&p_pending, 0, "wf_stats")
.await
.unwrap();
db.upsert_grid_point(&p_queued, 0, "wf_stats")
.await
.unwrap();
db.upsert_grid_point(&p_running, 0, "wf_stats")
.await
.unwrap();
db.upsert_grid_point(&p_cold, 1, "wf_stats").await.unwrap();
db.upsert_grid_point(&p_seed, 1, "wf_stats").await.unwrap();
db.upsert_grid_point(&p_failed, 2, "wf_stats")
.await
.unwrap();
db.upsert_grid_point(&p_imported, 2, "wf_stats")
.await
.unwrap();
db.update_grid_status(
&p_queued.model_name(),
common::models::GridPointStatus::Queued,
"wf_stats",
)
.await
.unwrap();
db.mark_grid_point_running(&p_running.model_name(), "wf_stats")
.await
.unwrap();
dispatch_and_report(
&db,
"wf_stats",
&p_cold,
common::models::TaskType::ColdRun,
None,
true,
)
.await;
dispatch_and_report(
&db,
"wf_stats",
&p_seed,
common::models::TaskType::SeedStep,
Some(p_cold.model_name()),
true,
)
.await;
dispatch_and_report(
&db,
"wf_stats",
&p_failed,
common::models::TaskType::ColdRun,
None,
false,
)
.await;
db.mark_grid_point_imported(&p_imported.model_name(), "wf_stats", None, "cold_run")
.await
.unwrap();
// 回拨任务创建时间:测试内 insert/report 同秒完成,墙钟差为 0 会被 ETA 估算
// 过滤(avg 必须 > 0);造 120s 的真实感样本,使 avg_point_sec/eta_sec 非空。
{
let conn = rusqlite::Connection::open(&db_path).unwrap();
conn.execute(
"UPDATE tasks SET created_at = datetime('now', '-120 seconds') WHERE workflow_name = 'wf_stats'",
[],
)
.unwrap();
}
// ---- 1. 无 token → 401 ----
let res = app
.clone()
.oneshot(
Request::builder()
.uri("/api/workflows/wf_stats/stats")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::UNAUTHORIZED);
// ---- 2. admin token → 200,各口径分别计数 ----
let res = app
.clone()
.oneshot(
Request::builder()
.uri("/api/workflows/wf_stats/stats")
.header("authorization", "Bearer admin-secret")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let body = axum::body::to_bytes(res.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(json["success"], true);
let data = &json["data"];
assert_eq!(data["name"], "wf_stats");
assert_eq!(data["status"], "running");
assert_eq!(data["total"], 7);
assert_eq!(data["pending"], 1, "pending 与 queued 必须分开计数");
assert_eq!(data["queued"], 1);
assert_eq!(data["running"], 1);
assert_eq!(data["converged"], 3);
assert_eq!(data["failed"], 1);
assert_eq!(data["cold_run_converged"], 2);
assert_eq!(data["seed_step_converged"], 1);
// wave 分布:3 个波次,wave0 共 3 点全未收敛
let waves = data["waves"].as_array().unwrap();
assert_eq!(waves.len(), 3);
assert_eq!(waves[0]["wave"], 0);
assert_eq!(waves[0]["total"], 3);
assert_eq!(waves[0]["converged"], 0);
assert_eq!(waves[1]["converged"], 2);
assert_eq!(waves[2]["failed"], 1);
// 有已完成任务且有剩余点 → ETA 可估算。
// P3 后 avg 取精确 elapsed_sec(夹具每次回报 120s),无在线节点按串行兜底:
// eta = 120 × 剩余 3 点 ÷ 1 槽位 = 360s。
let avg = data["avg_point_sec"].as_f64().expect("应给出平均耗时");
assert!((avg - 120.0).abs() < 1.0, "avg 应为精确 elapsed_sec 均值");
let eta = data["eta_sec"].as_f64().expect("应给出 ETA");
assert!((eta - 360.0).abs() < 5.0, "并发感知 ETA(串行兜底)");
// ---- 3. 空工作流:total=0、waves 空、ETA 为 null ----
let res = app
.clone()
.oneshot(
Request::builder()
.uri("/api/workflows/wf_empty/stats")
.header("authorization", "Bearer admin-secret")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let body = axum::body::to_bytes(res.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
let data = &json["data"];
assert_eq!(data["total"], 0);
assert!(data["waves"].as_array().unwrap().is_empty());
assert!(data["avg_point_sec"].is_null());
assert!(data["eta_sec"].is_null());
// ---- 4. 未知工作流 → 404 ----
let res = app
.oneshot(
Request::builder()
.uri("/api/workflows/nonexistent/stats")
.header("authorization", "Bearer admin-secret")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::NOT_FOUND);
}
/// 观测端点(points / point detail)共享夹具:6 个点覆盖全部状态与三种收敛手段,
/// 其中 `rescued` 有 2 次尝试(冷启动失败 → 种子步进救回),用于验证"最近尝试"取值。
///
/// 时间戳回拨:除 rescued 的 seed_step 尝试外全部 -60s,使"最近尝试"与时间排序确定
/// (测试内回报同秒落库,datetime('now') 秒级精度下必须人造先后)。
struct ObsNames {
pending: String,
cold: String,
seed: String,
failed: String,
rescued: String,
imported: String,
}
async fn seed_obs_fixture(db: &Database, db_path: &std::path::Path, wf: &str) -> ObsNames {
let yaml = "grid:\n teff: [20000]\n logg: [5.0]\n loghe: [-2]\n logc: [-4]\n logn: [-4]\n logo: [-4]";
db.upsert_workflow(wf, Some("观测测试"), yaml, "running")
.await
.unwrap();
let p_pending = wf_stats_test_params(20000.0, -4.0);
let p_cold = wf_stats_test_params(25000.0, -4.0);
let p_seed = wf_stats_test_params(30000.0, -4.0);
let p_failed = wf_stats_test_params(35000.0, -4.0);
let p_rescued = wf_stats_test_params(40000.0, -4.0);
let p_imported = wf_stats_test_params(50000.0, -4.0);
for p in [&p_pending, &p_cold, &p_seed, &p_failed, &p_rescued] {
db.upsert_grid_point(p, 0, wf).await.unwrap();
}
db.upsert_grid_point(&p_imported, 1, wf).await.unwrap();
dispatch_and_report(
db,
wf,
&p_cold,
common::models::TaskType::ColdRun,
None,
true,
)
.await;
dispatch_and_report(
db,
wf,
&p_seed,
common::models::TaskType::SeedStep,
Some(p_cold.model_name()),
true,
)
.await;
dispatch_and_report(
db,
wf,
&p_failed,
common::models::TaskType::ColdRun,
None,
false,
)
.await;
// rescued:先冷启动失败,再种子步进救回(2 次尝试,最终 converged/seed_step
dispatch_and_report(
db,
wf,
&p_rescued,
common::models::TaskType::ColdRun,
None,
false,
)
.await;
dispatch_and_report(
db,
wf,
&p_rescued,
common::models::TaskType::SeedStep,
Some(p_cold.model_name()),
true,
)
.await;
db.mark_grid_point_imported(&p_imported.model_name(), wf, None, "cold_run")
.await
.unwrap();
let conn = rusqlite::Connection::open(db_path).unwrap();
conn.execute(
"UPDATE tasks SET created_at = datetime('now','-60 seconds'),
completed_at = datetime('now','-60 seconds')
WHERE workflow_name = ?1 AND NOT (point_name = ?2 AND task_type = 'seed_step')",
rusqlite::params![wf, p_rescued.model_name()],
)
.unwrap();
drop(conn);
ObsNames {
pending: p_pending.model_name(),
cold: p_cold.model_name(),
seed: p_seed.model_name(),
failed: p_failed.model_name(),
rescued: p_rescued.model_name(),
imported: p_imported.model_name(),
}
}
/// GET /api/workflows/:name/points:过滤 / 分页 / 排序 / 白名单 400 / 最近尝试 JOIN。
#[tokio::test]
async fn test_wf_points_endpoint() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("wfpts_db.db");
let queue_db_path = temp_dir.path().join("wfpts_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
let state = AppState {
db: db.clone(),
queue,
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
5,
std::time::Duration::from_secs(300),
),
admin_token: Some("admin-secret".to_string()),
auth_disabled: false,
admin_sessions: Arc::new(tokio::sync::RwLock::new(std::collections::HashMap::new())),
};
let api_router = axum::Router::new().route(
"/workflows/:name/points",
axum::routing::get(server::api::workflow::get_workflow_points),
);
let auth_layer =
axum::middleware::from_fn_with_state(state.clone(), server::api::auth_middleware);
let app = axum::Router::new()
.nest("/api", api_router.layer(auth_layer))
.with_state(state);
let n = seed_obs_fixture(&db, &db_path, "wf_pts").await;
let yaml = "grid:\n teff: [20000]\n logg: [5.0]\n loghe: [-2]\n logc: [-4]\n logn: [-4]\n logo: [-4]";
db.upsert_workflow("wf_pts_empty", None, yaml, "idle")
.await
.unwrap();
/// 带 admin token 请求 points 端点,返回 (状态码, data JSON)。
async fn get_points(app: &axum::Router, uri: &str) -> (StatusCode, serde_json::Value) {
let res = app
.clone()
.oneshot(
Request::builder()
.uri(uri)
.header("authorization", "Bearer admin-secret")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
let status = res.status();
let body = axum::body::to_bytes(res.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
(status, json["data"].clone())
}
fn find_point<'a>(data: &'a serde_json::Value, name: &str) -> &'a serde_json::Value {
data["points"]
.as_array()
.unwrap()
.iter()
.find(|p| p["name"] == name)
.expect("点应存在于返回页")
}
// ---- 1. 无 token → 401 ----
let res = app
.clone()
.oneshot(
Request::builder()
.uri("/api/workflows/wf_pts/points")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::UNAUTHORIZED);
// ---- 2. 默认查询:total=6,按 wave→cno_sum→teff 排序 ----
let (st, data) = get_points(&app, "/api/workflows/wf_pts/points").await;
assert_eq!(st, StatusCode::OK);
assert_eq!(data["total"], 6);
let pts = data["points"].as_array().unwrap();
assert_eq!(pts.len(), 6);
assert_eq!(
pts.first().unwrap()["name"],
n.pending,
"最低 teff 应排首位"
);
assert_eq!(pts.last().unwrap()["name"], n.imported, "wave=1 应排末位");
// ---- 3. 状态过滤:failed 仅 1 个 ----
let (st, data) = get_points(&app, "/api/workflows/wf_pts/points?status=failed").await;
assert_eq!(st, StatusCode::OK);
assert_eq!(data["total"], 1);
assert_eq!(data["points"][0]["name"], n.failed);
// ---- 4. 手段过滤:seed_step 收敛 2 个(seed + rescued ----
let (st, data) = get_points(&app, "/api/workflows/wf_pts/points?method=seed_step").await;
assert_eq!(st, StatusCode::OK);
assert_eq!(data["total"], 2);
let names: Vec<&str> = data["points"]
.as_array()
.unwrap()
.iter()
.map(|p| p["name"].as_str().unwrap())
.collect();
assert!(names.contains(&n.seed.as_str()));
assert!(names.contains(&n.rescued.as_str()));
// ---- 5. 点名搜索 ----
let (st, data) = get_points(&app, "/api/workflows/wf_pts/points?q=t35000").await;
assert_eq!(st, StatusCode::OK);
assert_eq!(data["total"], 1);
assert_eq!(data["points"][0]["name"], n.failed);
// ---- 6. wave 过滤 ----
let (st, data) = get_points(&app, "/api/workflows/wf_pts/points?wave=1").await;
assert_eq!(st, StatusCode::OK);
assert_eq!(data["total"], 1);
assert_eq!(data["points"][0]["name"], n.imported);
// ---- 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);
assert_eq!(page0["points"].as_array().unwrap().len(), 2);
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(
&app,
"/api/workflows/wf_pts/points?sort=last_completed_at&order=desc",
)
.await;
assert_eq!(st, StatusCode::OK);
assert_eq!(data["points"][0]["name"], n.rescued);
// ---- 9. 非法枚举 → 400 ----
let (st, _) = get_points(&app, "/api/workflows/wf_pts/points?status=bogus").await;
assert_eq!(st, StatusCode::BAD_REQUEST);
let (st, _) = get_points(&app, "/api/workflows/wf_pts/points?sort=bogus").await;
assert_eq!(st, StatusCode::BAD_REQUEST);
// ---- 10. 最近尝试 JOINrescued 双尝试取最新(seed_step 救回) ----
let (_, data) = get_points(&app, "/api/workflows/wf_pts/points").await;
let rescued = find_point(&data, &n.rescued);
assert_eq!(rescued["status"], "converged");
assert_eq!(rescued["success_method"], "seed_step");
assert_eq!(rescued["attempt_count"], 2, "两次尝试都应计数");
assert_eq!(
rescued["last_task_type"], "seed_step",
"最近尝试应为种子步进"
);
assert_eq!(rescued["seed_point_name"], n.cold, "种子来源应为 cold 点");
assert_eq!(rescued["last_max_relc"], 0.0005);
assert_eq!(rescued["last_elapsed_sec"], 120.0, "真实墙钟耗时应落库");
// pending 点无任何尝试 → last_* 全 null
let pending = find_point(&data, &n.pending);
assert!(pending["last_task_type"].is_null());
assert!(pending["last_completed_at"].is_null());
assert!(pending["last_elapsed_sec"].is_null());
assert_eq!(pending["attempt_count"], 0);
// ---- 11. 空工作流与未知工作流 ----
let (st, data) = get_points(&app, "/api/workflows/wf_pts_empty/points").await;
assert_eq!(st, StatusCode::OK);
assert_eq!(data["total"], 0);
assert!(data["points"].as_array().unwrap().is_empty());
let (st, _) = get_points(&app, "/api/workflows/nonexistent/points").await;
assert_eq!(st, StatusCode::NOT_FOUND);
}
/// GET /api/workflows/:name/points/:point:尝试历史 + conv.json 解析、
/// conv 缺失降级 null、路径穿越 400、未知点 404。
#[tokio::test]
async fn test_point_detail_endpoint() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("wfpd_db.db");
let queue_db_path = temp_dir.path().join("wfpd_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
let state = AppState {
db: db.clone(),
queue,
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
5,
std::time::Duration::from_secs(300),
),
admin_token: Some("admin-secret".to_string()),
auth_disabled: false,
admin_sessions: Arc::new(tokio::sync::RwLock::new(std::collections::HashMap::new())),
};
let api_router = axum::Router::new().route(
"/workflows/:name/points/:point",
axum::routing::get(server::api::workflow::get_workflow_point_detail),
);
let auth_layer =
axum::middleware::from_fn_with_state(state.clone(), server::api::auth_middleware);
let app = axum::Router::new()
.nest("/api", api_router.layer(auth_layer))
.with_state(state);
let n = seed_obs_fixture(&db, &db_path, "wf_pd").await;
// 给 cold 点写一份合法 conv.jsonrescued 刻意不写,验证 null 降级)
let cold_dir = seeds_dir.join(&n.cold);
std::fs::create_dir_all(&cold_dir).unwrap();
std::fs::write(
cold_dir.join("conv.json"),
make_legacy_conv_json(&n.cold, true),
)
.unwrap();
async fn get_detail(app: &axum::Router, uri: &str) -> (StatusCode, serde_json::Value) {
let res = app
.clone()
.oneshot(
Request::builder()
.uri(uri)
.header("authorization", "Bearer admin-secret")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
let status = res.status();
let body = axum::body::to_bytes(res.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
(status, json["data"].clone())
}
// ---- 1. 无 token → 401 ----
let uri = format!("/api/workflows/wf_pd/points/{}", n.cold);
let res = app
.clone()
.oneshot(Request::builder().uri(&uri).body(Body::empty()).unwrap())
.await
.unwrap();
assert_eq!(res.status(), StatusCode::UNAUTHORIZED);
// ---- 2. cold 点:点行 + 1 次尝试 + conv 解析成功 ----
let (st, data) = get_detail(&app, &uri).await;
assert_eq!(st, StatusCode::OK);
assert_eq!(data["point"]["name"], n.cold);
assert_eq!(data["point"]["status"], "converged");
assert_eq!(data["point"]["success_method"], "cold_run");
let attempts = data["attempts"].as_array().unwrap();
assert_eq!(attempts.len(), 1);
assert_eq!(attempts[0]["task_type"], "cold_run");
assert_eq!(attempts[0]["status"], "completed");
assert_eq!(data["conv"]["converged"], true, "conv.json 应被解析");
assert_eq!(data["conv"]["final_max_relc"], 0.000321);
// ---- 3. rescued 点:2 次尝试按时间升序(冷启失败 → 种子步进救回),conv 为 null ----
let uri = format!("/api/workflows/wf_pd/points/{}", n.rescued);
let (st, data) = get_detail(&app, &uri).await;
assert_eq!(st, StatusCode::OK);
assert_eq!(data["point"]["attempt_count"], 2);
let attempts = data["attempts"].as_array().unwrap();
assert_eq!(attempts.len(), 2);
assert_eq!(attempts[0]["task_type"], "cold_run", "首次应为冷启动");
assert_eq!(attempts[0]["status"], "failed");
assert_eq!(attempts[0]["elapsed_sec"], 120.0, "每次尝试耗时应落库");
assert_eq!(attempts[1]["task_type"], "seed_step", "第二次应为种子步进");
assert_eq!(attempts[1]["status"], "completed");
assert_eq!(
attempts[1]["seed_point_name"], n.cold,
"种子来源应为 cold 点"
);
assert!(data["conv"].is_null(), "无 conv.json 应返回 null 而非报错");
// ---- 4. 路径穿越 / 非法字符 → 400 ----
for bad in ["..%2Fevil", "a%2Fb", ".hidden", ".."] {
let uri = format!("/api/workflows/wf_pd/points/{}", bad);
let (st, _) = get_detail(&app, &uri).await;
assert_eq!(st, StatusCode::BAD_REQUEST, "点名 {} 应被拒绝", bad);
}
// ---- 5. 未知点 → 404 ----
let (st, _) = get_detail(
&app,
"/api/workflows/wf_pd/points/t99999_g5.0_he-2_c-4_n-4_o-4",
)
.await;
assert_eq!(st, StatusCode::NOT_FOUND);
}
/// GET /api/workflows:每工作流内联网格统计——有点的给 stats 对象,无点的为 null。
#[tokio::test]
async fn test_list_workflows_inline_stats() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("wflist_db.db");
let queue_db_path = temp_dir.path().join("wflist_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
let state = AppState {
db: db.clone(),
queue,
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
5,
std::time::Duration::from_secs(300),
),
admin_token: Some("admin-secret".to_string()),
auth_disabled: false,
admin_sessions: Arc::new(tokio::sync::RwLock::new(std::collections::HashMap::new())),
};
let api_router = axum::Router::new().route(
"/workflows",
axum::routing::get(server::api::workflow::list_workflows),
);
let auth_layer =
axum::middleware::from_fn_with_state(state.clone(), server::api::auth_middleware);
let app = axum::Router::new()
.nest("/api", api_router.layer(auth_layer))
.with_state(state);
// wf_list_a 有 2 点(1 cold_run 收敛 + 1 pending);wf_list_b 无点
let yaml = "grid:\n teff: [20000]\n logg: [5.0]\n loghe: [-2]\n logc: [-4]\n logn: [-4]\n logo: [-4]";
db.upsert_workflow("wf_list_a", None, yaml, "running")
.await
.unwrap();
db.upsert_workflow("wf_list_b", None, yaml, "idle")
.await
.unwrap();
let p1 = wf_stats_test_params(20000.0, -4.0);
let p2 = wf_stats_test_params(25000.0, -4.0);
db.upsert_grid_point(&p1, 0, "wf_list_a").await.unwrap();
db.upsert_grid_point(&p2, 0, "wf_list_a").await.unwrap();
dispatch_and_report(
&db,
"wf_list_a",
&p1,
common::models::TaskType::ColdRun,
None,
true,
)
.await;
// 无 token → 401
let res = app
.clone()
.oneshot(
Request::builder()
.uri("/api/workflows")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::UNAUTHORIZED);
// 带 token → 200,stats 内联(两工作流一次返回,无 N+1)
let res = app
.clone()
.oneshot(
Request::builder()
.uri("/api/workflows")
.header("authorization", "Bearer admin-secret")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let body = axum::body::to_bytes(res.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
let list = json["data"].as_array().unwrap();
let wf_a = list.iter().find(|w| w["name"] == "wf_list_a").unwrap();
let wf_b = list.iter().find(|w| w["name"] == "wf_list_b").unwrap();
assert_eq!(wf_a["stats"]["total"], 2);
assert_eq!(wf_a["stats"]["converged"], 1);
assert_eq!(wf_a["stats"]["cold_run_converged"], 1);
assert_eq!(wf_a["stats"]["failed"], 0);
assert_eq!(wf_a["stats"]["running"], 0);
assert!(wf_b["stats"].is_null(), "无网格点的工作流 stats 应为 null");
}
/// GET /api/workflows/:name/progress:快照去重写入、经验速率、停滞时长、保留期清理、鉴权。
#[tokio::test]
async fn test_wf_progress_endpoint() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("wfprog_db.db");
let queue_db_path = temp_dir.path().join("wfprog_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
let state = AppState {
db: db.clone(),
queue,
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
5,
std::time::Duration::from_secs(300),
),
admin_token: Some("admin-secret".to_string()),
auth_disabled: false,
admin_sessions: Arc::new(tokio::sync::RwLock::new(std::collections::HashMap::new())),
};
let api_router = axum::Router::new().route(
"/workflows/:name/progress",
axum::routing::get(server::api::workflow::get_workflow_progress),
);
let auth_layer =
axum::middleware::from_fn_with_state(state.clone(), server::api::auth_middleware);
let app = axum::Router::new()
.nest("/api", api_router.layer(auth_layer))
.with_state(state);
// ---- 快照写入与去重 ----
let yaml = "grid:\n teff: [20000]\n logg: [5.0]\n loghe: [-2]\n logc: [-4]\n logn: [-4]\n logo: [-4]";
db.upsert_workflow("wf_prog", None, yaml, "running")
.await
.unwrap();
let p1 = wf_stats_test_params(20000.0, -4.0);
let p2 = wf_stats_test_params(25000.0, -4.0);
db.upsert_grid_point(&p1, 0, "wf_prog").await.unwrap();
db.upsert_grid_point(&p2, 0, "wf_prog").await.unwrap();
assert!(
db.record_progress_snapshot("wf_prog").await.unwrap(),
"首次记录应写入"
);
dispatch_and_report(
&db,
"wf_prog",
&p1,
common::models::TaskType::ColdRun,
None,
true,
)
.await;
assert!(
db.record_progress_snapshot("wf_prog").await.unwrap(),
"计数变化应写入"
);
assert!(
!db.record_progress_snapshot("wf_prog").await.unwrap(),
"计数未变应去重"
);
// 首条快照回拨 2 小时 → 经验速率 ≈ 0.5 点/小时(1 点收敛 / 2h
{
let conn = rusqlite::Connection::open(&db_path).unwrap();
conn.execute(
"UPDATE workflow_progress_snapshots SET ts = datetime('now', '-2 hours') \
WHERE workflow_name = 'wf_prog' AND converged = 0",
[],
)
.unwrap();
}
// ---- 无 token → 401 ----
let res = app
.clone()
.oneshot(
Request::builder()
.uri("/api/workflows/wf_prog/progress")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::UNAUTHORIZED);
// ---- 带 token → 200:序列 2 条、速率 ≈0.5、停滞 ≈0(末条即最新增长) ----
let res = app
.clone()
.oneshot(
Request::builder()
.uri("/api/workflows/wf_prog/progress")
.header("authorization", "Bearer admin-secret")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let body = axum::body::to_bytes(res.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
let data = &json["data"];
assert_eq!(data["hours"], 24);
assert_eq!(data["series"].as_array().unwrap().len(), 2);
let rate = data["rate_per_hour"].as_f64().expect("应有经验速率");
assert!((rate - 0.5).abs() < 0.05, "速率应 ≈ 0.5 点/小时");
let stalled = data["stalled_minutes"].as_f64().expect("应有停滞时长");
assert!(stalled < 1.0, "终态刚增长,停滞应 ≈ 0");
// ---- 保留期清理:回拨一条 -10 天 → purge(7) 后仅剩 1 条 ----
{
let conn = rusqlite::Connection::open(&db_path).unwrap();
conn.execute(
"UPDATE workflow_progress_snapshots SET ts = datetime('now', '-10 days') \
WHERE workflow_name = 'wf_prog' AND converged = 0",
[],
)
.unwrap();
}
db.purge_progress_snapshots(7).await.unwrap();
let series = db.get_progress_series("wf_prog", 168).await.unwrap();
assert_eq!(series.len(), 1, "超 7 天的快照应被清理");
// ---- 未知工作流 → 404 ----
let res = app
.oneshot(
Request::builder()
.uri("/api/workflows/nonexistent/progress")
.header("authorization", "Bearer admin-secret")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(res.status(), StatusCode::NOT_FOUND);
}
/// 构造仅含 report 字段的 multipart body(模拟节点上报,不带 seed_file)。
fn make_report_only_multipart(boundary: &str, report_json: &str) -> Vec<u8> {
let mut body = Vec::new();
body.extend_from_slice(format!("--{}\r\n", boundary).as_bytes());
body.extend_from_slice(b"Content-Disposition: form-data; name=\"report\"\r\n");
body.extend_from_slice(b"Content-Type: application/json\r\n\r\n");
body.extend_from_slice(report_json.as_bytes());
body.extend_from_slice(b"\r\n");
body.extend_from_slice(format!("--{}--\r\n", boundary).as_bytes());
body
}
/// 端到端回归(2026-08-02 涡旋事故):已收敛点收到迟到的重复失败报告时,
/// 走完整 HTTP 链路(claim → report)后状态保持 converged,且不触发种子回退
/// state_changed=false 跳过 fallback;纵有可用种子也不产生 seed_step 任务)。
#[tokio::test]
async fn test_duplicate_failure_report_cannot_flip_converged() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("flip_db.db");
let queue_db_path = temp_dir.path().join("flip_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
// 预置节点与 token。
let reg = common::models::NodeRegisterRequest {
node_id: "node-flip".to_string(),
max_slots: 2,
};
db.register_node(&reg).await.unwrap();
let token = db.issue_node_token("node-flip").await.unwrap();
let state = AppState {
db: db.clone(),
queue: queue.clone(),
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
100,
std::time::Duration::from_secs(300),
),
admin_token: None,
auth_disabled: false,
admin_sessions: Arc::new(tokio::sync::RwLock::new(std::collections::HashMap::new())),
};
let api_router = axum::Router::new()
.route(
"/task/claim",
axum::routing::post(server::api::task::claim_task),
)
.route(
"/task/report",
axum::routing::post(server::api::task::report_task),
);
let auth_layer =
axum::middleware::from_fn_with_state(state.clone(), server::api::auth_middleware);
let app = axum::Router::new()
.nest("/api", api_router.layer(auth_layer))
.with_state(state);
// 工作流 running + 网格点 + 可用种子(若回退被错误触发,必然能匹配到种子并派发任务,
// 使"无 seed_step 任务"断言成为强证据)。
db.upsert_workflow("wf_flip", None, "", "running")
.await
.unwrap();
let params = common::models::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 name = params.model_name();
db.upsert_grid_point(&params, 0, "wf_flip").await.unwrap();
db.insert_seed_named(&name, &params, "/tmp/flip_seed.7")
.await
.unwrap();
// ---- 任务 A:正常收敛 ----
let task_a = common::models::TaskSpec {
task_id: uuid::Uuid::new_v4(),
point_name: name.clone(),
params: params.clone(),
task_type: common::models::TaskType::ColdRun,
seed_point_name: None,
timeout_sec: 7200,
workflow_name: Some("wf_flip".to_string()),
wave: 0,
..Default::default()
};
db.insert_task(&task_a).await.unwrap();
queue.push_task(&task_a).await.unwrap();
let claim_res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/task/claim")
.header("authorization", format!("Bearer {}", token))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(claim_res.status(), StatusCode::OK);
let report_a = common::models::TaskReport {
task_id: task_a.task_id,
point_name: name.clone(),
params: Some(params.clone()),
node_id: "node-flip".to_string(),
status: common::models::TaskStatus::Completed,
converged: true,
max_relc: Some(0.0005),
atmosphere_has_nan: false,
elapsed_sec: 120.0,
error_message: None,
summary_json: "{}".to_string(),
failed_stage: None,
};
let body_a =
make_report_only_multipart("flipbound1", &serde_json::to_string(&report_a).unwrap());
let report_res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/task/report")
.header("authorization", format!("Bearer {}", token))
.header("content-type", "multipart/form-data; boundary=flipbound1")
.body(Body::from(body_a))
.unwrap(),
)
.await
.unwrap();
assert_eq!(report_res.status(), StatusCode::OK);
assert_eq!(
db.get_grid_point_status(&name, "wf_flip")
.await
.unwrap()
.unwrap()
.0,
"converged"
);
// ---- 任务 B:迟到的重复失败报告(涡旋残留任务的典型行为)----
let task_b = common::models::TaskSpec {
task_id: uuid::Uuid::new_v4(),
point_name: name.clone(),
params: params.clone(),
task_type: common::models::TaskType::ColdRun,
seed_point_name: None,
timeout_sec: 7200,
workflow_name: Some("wf_flip".to_string()),
wave: 0,
..Default::default()
};
db.insert_task(&task_b).await.unwrap();
queue.push_task(&task_b).await.unwrap();
let claim_b = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/task/claim")
.header("authorization", format!("Bearer {}", token))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(claim_b.status(), StatusCode::OK);
let report_b = common::models::TaskReport {
task_id: task_b.task_id,
point_name: name.clone(),
params: Some(params.clone()),
node_id: "node-flip".to_string(),
status: common::models::TaskStatus::Failed,
converged: false,
max_relc: Some(9.5e5),
atmosphere_has_nan: false,
elapsed_sec: 130.0,
error_message: Some("nl stage diverged".to_string()),
summary_json: "{}".to_string(),
failed_stage: None,
};
let body_b =
make_report_only_multipart("flipbound2", &serde_json::to_string(&report_b).unwrap());
let report_b_res = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/task/report")
.header("authorization", format!("Bearer {}", token))
.header("content-type", "multipart/form-data; boundary=flipbound2")
.body(Body::from(body_b))
.unwrap(),
)
.await
.unwrap();
assert_eq!(
report_b_res.status(),
StatusCode::OK,
"上报本身应成功(吸收)"
);
// 核心断言:converged 不被翻黑,且未触发任何种子回退任务。
assert_eq!(
db.get_grid_point_status(&name, "wf_flip")
.await
.unwrap()
.unwrap()
.0,
"converged",
"迟到失败报告不得翻黑 converged 点"
);
assert!(
db.has_pending_tasks_for_point(&name, "wf_flip", Some("seed_step"))
.await
.unwrap()
.is_empty(),
"不得因迟到失败报告派发 seed_step"
);
assert!(
!db.has_seed_step_attempt(&name, "wf_flip").await.unwrap(),
"全程不应产生任何 seed_step 行"
);
}
/// save_workflow 阶段配置合法性校验(修复审查 #4/#5):
/// - 双阶段全关 → 400(无任何计算可执行);
/// - 启用阶段空策略链 → 400(调度无顺位可派,任务必失败);
/// - 合法配置(synspec-only 场景 B / 默认双开)→ 200。
#[tokio::test]
async fn test_save_workflow_validates_stage_configs() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("wfval_db.db");
let queue_db_path = temp_dir.path().join("wfval_queue.db");
let seeds_dir = temp_dir.path().join("results");
std::fs::create_dir_all(&seeds_dir).unwrap();
let db = Database::new(&db_path.to_string_lossy()).await.unwrap();
let queue = Arc::new(
SqliteTaskQueue::new(&queue_db_path.to_string_lossy())
.await
.unwrap(),
);
let scheduler = Arc::new(GridScheduler::new(db.clone(), queue.clone()));
let state = AppState {
db: db.clone(),
queue,
scheduler,
seeds_dir: seeds_dir.to_string_lossy().to_string(),
rate_limiter: server::api::rate_limit::RateLimiter::new(
5,
std::time::Duration::from_secs(300),
),
admin_token: Some("admin-secret".to_string()),
auth_disabled: false,
admin_sessions: Arc::new(tokio::sync::RwLock::new(std::collections::HashMap::new())),
};
let app = axum::Router::new()
.route(
"/api/workflows",
axum::routing::post(server::api::workflow::save_workflow),
)
.with_state(state);
let base_yaml = "grid:\n teff: [35000]\n logg: [5.5]\n loghe: [-1]\n logc: [-2]\n logn: [-2]\n logo: [-2]";
let post_save = |name: &str, yaml: &str| {
let app = app.clone();
let body = serde_json::json!({
"name": name,
"description": null,
"config_yaml": yaml,
});
async move {
app.oneshot(
Request::builder()
.method("POST")
.uri("/api/workflows")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&body).unwrap()))
.unwrap(),
)
.await
.unwrap()
}
};
// 1. 双阶段全关 → 400
let both_off = format!(
"{}\ntlusty:\n enabled: false\n policy: skip_converged\n strategies: [cold_run, seed_step]\n\
synspec_stage:\n enabled: false\n policy: skip_converged\n strategies: [standard]\n",
base_yaml
);
let res = post_save("wf_val_a", &both_off).await;
assert_eq!(res.status(), StatusCode::BAD_REQUEST, "双阶段全关应 400");
// 2. 启用阶段空策略链 → 400
let empty_chain = format!(
"{}\ntlusty:\n enabled: true\n policy: skip_converged\n strategies: []\n",
base_yaml
);
let res = post_save("wf_val_b", &empty_chain).await;
assert_eq!(
res.status(),
StatusCode::BAD_REQUEST,
"启用阶段空策略链应 400"
);
// 3. 合法:TLUSTY 关 + SYNSPEC 启(设计 §2.2 场景 B:仅更新光谱)→ 200
let syn_only = format!(
"{}\ntlusty:\n enabled: false\n policy: skip_converged\n strategies: [cold_run, seed_step]\n\
synspec_stage:\n enabled: true\n policy: force_recompute\n strategies: [standard]\n",
base_yaml
);
let res = post_save("wf_val_c", &syn_only).await;
assert_eq!(res.status(), StatusCode::OK, "synspec-only 合法配置应 200");
// 4. 合法:无阶段块(默认双开)→ 200
let res = post_save("wf_val_d", base_yaml).await;
assert_eq!(res.status(), StatusCode::OK, "默认配置应 200");
}