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(®_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(¶ms, 0, "test_wf").await.unwrap(); db.mark_grid_point_running(¶ms.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, }; db.register_node(®).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 后行为一致。 let pending = db.take_pending_node_token("node-l2-test").await.unwrap(); assert_eq!(pending.as_deref(), Some(token.as_str())); let pending2 = db.take_pending_node_token("node-l2-test").await.unwrap(); assert!( pending2.is_none(), "取走即焚:第二次 take_pending 必须返回 None" ); // 不存在的 node take 也应返回 None(不报错) assert!(db .take_pending_node_token("node-not-exist") .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(®).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 经中间件 → 401:reissue 覆盖 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); } /// 管理 API(Admin)端到端测试: /// 列出节点 → 重发 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(®).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(®).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 用自己 token,body 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(®_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"); // 2. Node 端轮询查状态 ➔ status: pending_approval let check_body = serde_json::json!({ "node_id": "node-pending-01" }); 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(®_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(®_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 _ = ®_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.json(stage 含 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, ) -> Vec { 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"); 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", ); 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 + 收敛点 → 200,conv.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", ); 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", ); 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"); 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 工具的服务端侧): /// 嵌套 stages(itek_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")); 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_points:converged + imported 手段 + 旧版 elapsed_sec 已落库 let row = db .get_workflow_point_row("wf_legacy", name) .await .unwrap() .expect("grid_points 应存在"); assert_eq!(row.status, "converged"); assert_eq!(row.success_method.as_deref(), Some("imported")); assert_eq!( row.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(®).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(®_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, }; 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(®2).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 可停用)" ); } // ===== 工作流执行观测端点测试(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, 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, }; 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(), }; 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) .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"], 1); assert_eq!(data["seed_step_converged"], 1); assert_eq!(data["imported_converged"], 1); // wave 分布:3 个波次,wave0 共 3 点全未收敛 let waves = data["waves"].as_array().unwrap(); assert_eq!(waves.len(), 3); 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) .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 钳位 ≤500 ---- 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); // ---- 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. 最近尝试 JOIN:rescued 双尝试取最新(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.json(rescued 刻意不写,验证 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); }