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 results_dir = temp_dir.path().join("results"); std::fs::create_dir_all(&results_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(), results_dir.to_string_lossy().to_string())); let state = AppState { db, queue, scheduler, results_dir: results_dir.to_string_lossy().to_string(), auth_token: None, }; 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", "host_name": "api-host", "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 results_dir = temp_dir.path().join("results"); std::fs::create_dir_all(&results_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(), results_dir.to_string_lossy().to_string())); let state = AppState { db: db.clone(), queue: queue.clone(), scheduler, results_dir: results_dir.to_string_lossy().to_string(), auth_token: Some("secret_token_123".to_string()), }; 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, logg: 5.5, loghe: -1.0, logc: -2.0, logn: -2.0, logo: -2.0, }; db.upsert_grid_point(¶ms, 0).await.unwrap(); db.mark_grid_point_running(¶ms.model_name()).await.unwrap(); let stats = db.get_grid_summary_stats().await.unwrap(); assert_eq!(stats["running"], 1); }