diff --git a/web-gateway/src/routes/backtesting.rs b/web-gateway/src/routes/backtesting.rs index c4ff2fc4a..2d2dc68a0 100644 --- a/web-gateway/src/routes/backtesting.rs +++ b/web-gateway/src/routes/backtesting.rs @@ -13,8 +13,8 @@ use crate::state::AppState; pub fn router() -> Router { Router::new() .route("/start", post(start_backtest)) - .route("/{id}/status", get(get_status)) - .route("/{id}/results", get(get_results)) + .route("/:id/status", get(get_status)) + .route("/:id/results", get(get_results)) } #[derive(serde::Deserialize)] @@ -95,3 +95,174 @@ async fn get_results( serde_json::to_value(response.into_inner()).map_err(|e| AppError::Internal(e.into()))?, )) } + +#[cfg(test)] +#[allow(clippy::unwrap_used)] +mod tests { + use axum::body::Body; + use axum::http::{Request, StatusCode}; + use axum::middleware; + use axum::Router; + use jsonwebtoken::{encode, EncodingKey, Header}; + use std::sync::Arc; + use tower::ServiceExt; + + use super::*; + use crate::auth::claims::Claims; + use crate::auth::middleware::auth_middleware; + use crate::config::GatewayConfig; + + const SECRET: &str = "backtesting-test-secret"; + + fn test_state() -> AppState { + let config = GatewayConfig { + jwt_secret: SECRET.to_string(), + trading_service_url: String::new(), + ..GatewayConfig::default() + }; + let (ws_broadcast, _) = tokio::sync::broadcast::channel(16); + AppState { + config: Arc::new(config), + trading_channel: None, + backtesting_channel: None, + ml_training_channel: None, + ws_broadcast, + } + } + + fn test_app(state: AppState) -> Router { + Router::new() + .nest("/backtesting", router()) + .layer(middleware::from_fn_with_state( + state.config.clone(), + auth_middleware, + )) + .with_state(state) + } + + fn make_token() -> String { + let claims = Claims { + sub: "backtester-1".into(), + exp: u64::MAX, + iat: 1_700_000_000, + jti: "bt1".into(), + roles: vec!["analyst".into()], + permissions: vec!["write:backtesting".into()], + }; + let key = EncodingKey::from_secret(SECRET.as_bytes()); + encode(&Header::default(), &claims, &key).unwrap() + } + + fn valid_start_body() -> serde_json::Value { + serde_json::json!({ + "strategy_name": "momentum", + "symbols": ["ES.FUT"], + "start_date_unix_nanos": 1_700_000_000_000_000_000_i64, + "end_date_unix_nanos": 1_700_086_400_000_000_000_i64, + "initial_capital": 100_000.0, + "parameters": {}, + "save_results": false, + "description": "test run", + }) + } + + #[tokio::test] + async fn test_start_backtest_without_auth_returns_401() { + let state = test_state(); + let app = test_app(state); + let req = Request::builder() + .method("POST") + .uri("/backtesting/start") + .header("Content-Type", "application/json") + .body(Body::from(serde_json::to_string(&valid_start_body()).unwrap())) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn test_start_backtest_no_service_returns_500() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + let req = Request::builder() + .method("POST") + .uri("/backtesting/start") + .header("Authorization", format!("Bearer {token}")) + .header("Content-Type", "application/json") + .body(Body::from(serde_json::to_string(&valid_start_body()).unwrap())) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR); + } + + #[tokio::test] + async fn test_start_backtest_invalid_body_returns_422() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + // Missing many required fields + let body = serde_json::json!({"strategy_name": "momentum"}); + let req = Request::builder() + .method("POST") + .uri("/backtesting/start") + .header("Authorization", format!("Bearer {token}")) + .header("Content-Type", "application/json") + .body(Body::from(serde_json::to_string(&body).unwrap())) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNPROCESSABLE_ENTITY); + } + + #[tokio::test] + async fn test_get_status_without_auth_returns_401() { + let state = test_state(); + let app = test_app(state); + let req = Request::builder() + .uri("/backtesting/bt-123/status") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn test_get_status_no_service_returns_500() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + let req = Request::builder() + .uri("/backtesting/bt-123/status") + .header("Authorization", format!("Bearer {token}")) + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR); + } + + #[tokio::test] + async fn test_get_results_without_auth_returns_401() { + let state = test_state(); + let app = test_app(state); + let req = Request::builder() + .uri("/backtesting/bt-123/results") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn test_get_results_no_service_returns_500() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + let req = Request::builder() + .uri("/backtesting/bt-123/results") + .header("Authorization", format!("Bearer {token}")) + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR); + } +} diff --git a/web-gateway/src/routes/config.rs b/web-gateway/src/routes/config.rs index 9dffda4e5..9ef648974 100644 --- a/web-gateway/src/routes/config.rs +++ b/web-gateway/src/routes/config.rs @@ -59,3 +59,137 @@ async fn update_config( serde_json::to_value(response.into_inner()).map_err(|e| AppError::Internal(e.into()))?, )) } + +#[cfg(test)] +#[allow(clippy::unwrap_used)] +mod tests { + use axum::body::Body; + use axum::http::{Request, StatusCode}; + use axum::middleware; + use axum::Router; + use jsonwebtoken::{encode, EncodingKey, Header}; + use std::sync::Arc; + use tower::ServiceExt; + + use super::*; + use crate::auth::claims::Claims; + use crate::auth::middleware::auth_middleware; + use crate::config::GatewayConfig; + + const SECRET: &str = "config-test-secret"; + + fn test_state() -> AppState { + let config = GatewayConfig { + jwt_secret: SECRET.to_string(), + trading_service_url: String::new(), + ..GatewayConfig::default() + }; + let (ws_broadcast, _) = tokio::sync::broadcast::channel(16); + AppState { + config: Arc::new(config), + trading_channel: None, + backtesting_channel: None, + ml_training_channel: None, + ws_broadcast, + } + } + + fn test_app(state: AppState) -> Router { + Router::new() + .nest("/config", router()) + .layer(middleware::from_fn_with_state( + state.config.clone(), + auth_middleware, + )) + .with_state(state) + } + + fn make_token() -> String { + let claims = Claims { + sub: "admin-1".into(), + exp: u64::MAX, + iat: 1_700_000_000, + jti: "cfg1".into(), + roles: vec!["admin".into()], + permissions: vec!["read:config".into(), "write:config".into()], + }; + let key = EncodingKey::from_secret(SECRET.as_bytes()); + encode(&Header::default(), &claims, &key).unwrap() + } + + #[tokio::test] + async fn test_get_config_without_auth_returns_401() { + let state = test_state(); + let app = test_app(state); + let req = Request::builder() + .uri("/config") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn test_get_config_no_service_returns_500() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + let req = Request::builder() + .uri("/config") + .header("Authorization", format!("Bearer {token}")) + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR); + } + + #[tokio::test] + async fn test_update_config_without_auth_returns_401() { + let state = test_state(); + let app = test_app(state); + let body = serde_json::json!({"parameters": {}, "persist": false}); + let req = Request::builder() + .method("PUT") + .uri("/config") + .header("Content-Type", "application/json") + .body(Body::from(serde_json::to_string(&body).unwrap())) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn test_update_config_no_service_returns_500() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + let body = serde_json::json!({"parameters": {"key": "value"}, "persist": true}); + let req = Request::builder() + .method("PUT") + .uri("/config") + .header("Authorization", format!("Bearer {token}")) + .header("Content-Type", "application/json") + .body(Body::from(serde_json::to_string(&body).unwrap())) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR); + } + + #[tokio::test] + async fn test_update_config_invalid_body_returns_422() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + // Missing required persist field; parameters is also missing + let body = serde_json::json!({"something_else": 42}); + let req = Request::builder() + .method("PUT") + .uri("/config") + .header("Authorization", format!("Bearer {token}")) + .header("Content-Type", "application/json") + .body(Body::from(serde_json::to_string(&body).unwrap())) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNPROCESSABLE_ENTITY); + } +} diff --git a/web-gateway/src/routes/ml.rs b/web-gateway/src/routes/ml.rs index d0606c162..4959ddae7 100644 --- a/web-gateway/src/routes/ml.rs +++ b/web-gateway/src/routes/ml.rs @@ -93,3 +93,169 @@ async fn get_performance( serde_json::to_value(response.into_inner()).map_err(|e| AppError::Internal(e.into()))?, )) } + +#[cfg(test)] +#[allow(clippy::unwrap_used)] +mod tests { + use axum::body::Body; + use axum::http::{Request, StatusCode}; + use axum::middleware; + use axum::Router; + use jsonwebtoken::{encode, EncodingKey, Header}; + use std::sync::Arc; + use tower::ServiceExt; + + use super::*; + use crate::auth::claims::Claims; + use crate::auth::middleware::auth_middleware; + use crate::config::GatewayConfig; + + const SECRET: &str = "ml-test-secret"; + + fn test_state() -> AppState { + let config = GatewayConfig { + jwt_secret: SECRET.to_string(), + trading_service_url: String::new(), + ..GatewayConfig::default() + }; + let (ws_broadcast, _) = tokio::sync::broadcast::channel(16); + AppState { + config: Arc::new(config), + trading_channel: None, + backtesting_channel: None, + ml_training_channel: None, + ws_broadcast, + } + } + + fn test_app(state: AppState) -> Router { + Router::new() + .nest("/ml", router()) + .layer(middleware::from_fn_with_state( + state.config.clone(), + auth_middleware, + )) + .with_state(state) + } + + fn make_token() -> String { + let claims = Claims { + sub: "ml-user-1".into(), + exp: u64::MAX, + iat: 1_700_000_000, + jti: "ml1".into(), + roles: vec!["trader".into()], + permissions: vec!["read:ml".into()], + }; + let key = EncodingKey::from_secret(SECRET.as_bytes()); + encode(&Header::default(), &claims, &key).unwrap() + } + + #[tokio::test] + async fn test_predictions_without_auth_returns_401() { + let state = test_state(); + let app = test_app(state); + let req = Request::builder() + .uri("/ml/predictions?symbol=ES.FUT") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn test_predictions_no_service_returns_500() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + let req = Request::builder() + .uri("/ml/predictions?symbol=ES.FUT") + .header("Authorization", format!("Bearer {token}")) + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR); + } + + #[tokio::test] + async fn test_submit_ml_order_without_auth_returns_401() { + let state = test_state(); + let app = test_app(state); + let body = serde_json::json!({ + "symbol": "ES.FUT", + "account_id": "acc-1", + }); + let req = Request::builder() + .method("POST") + .uri("/ml/orders") + .header("Content-Type", "application/json") + .body(Body::from(serde_json::to_string(&body).unwrap())) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn test_submit_ml_order_no_service_returns_500() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + let body = serde_json::json!({ + "symbol": "ES.FUT", + "account_id": "acc-1", + }); + let req = Request::builder() + .method("POST") + .uri("/ml/orders") + .header("Authorization", format!("Bearer {token}")) + .header("Content-Type", "application/json") + .body(Body::from(serde_json::to_string(&body).unwrap())) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR); + } + + #[tokio::test] + async fn test_submit_ml_order_invalid_body_returns_422() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + // Missing required account_id field + let body = serde_json::json!({"symbol": "ES.FUT"}); + let req = Request::builder() + .method("POST") + .uri("/ml/orders") + .header("Authorization", format!("Bearer {token}")) + .header("Content-Type", "application/json") + .body(Body::from(serde_json::to_string(&body).unwrap())) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNPROCESSABLE_ENTITY); + } + + #[tokio::test] + async fn test_performance_without_auth_returns_401() { + let state = test_state(); + let app = test_app(state); + let req = Request::builder() + .uri("/ml/performance") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn test_performance_no_service_returns_500() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + let req = Request::builder() + .uri("/ml/performance") + .header("Authorization", format!("Bearer {token}")) + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR); + } +} diff --git a/web-gateway/src/routes/performance.rs b/web-gateway/src/routes/performance.rs index 170fff009..11ae54238 100644 --- a/web-gateway/src/routes/performance.rs +++ b/web-gateway/src/routes/performance.rs @@ -81,3 +81,139 @@ async fn get_throughput( serde_json::to_value(response.into_inner()).map_err(|e| AppError::Internal(e.into()))?, )) } + +#[cfg(test)] +#[allow(clippy::unwrap_used)] +mod tests { + use axum::body::Body; + use axum::http::{Request, StatusCode}; + use axum::middleware; + use axum::Router; + use jsonwebtoken::{encode, EncodingKey, Header}; + use std::sync::Arc; + use tower::ServiceExt; + + use super::*; + use crate::auth::claims::Claims; + use crate::auth::middleware::auth_middleware; + use crate::config::GatewayConfig; + + const SECRET: &str = "performance-test-secret"; + + fn test_state() -> AppState { + let config = GatewayConfig { + jwt_secret: SECRET.to_string(), + trading_service_url: String::new(), + ..GatewayConfig::default() + }; + let (ws_broadcast, _) = tokio::sync::broadcast::channel(16); + AppState { + config: Arc::new(config), + trading_channel: None, + backtesting_channel: None, + ml_training_channel: None, + ws_broadcast, + } + } + + fn test_app(state: AppState) -> Router { + Router::new() + .nest("/performance", router()) + .layer(middleware::from_fn_with_state( + state.config.clone(), + auth_middleware, + )) + .with_state(state) + } + + fn make_token() -> String { + let claims = Claims { + sub: "perf-user-1".into(), + exp: u64::MAX, + iat: 1_700_000_000, + jti: "pf1".into(), + roles: vec!["analyst".into()], + permissions: vec!["read:performance".into()], + }; + let key = EncodingKey::from_secret(SECRET.as_bytes()); + encode(&Header::default(), &claims, &key).unwrap() + } + + #[tokio::test] + async fn test_metrics_without_auth_returns_401() { + let state = test_state(); + let app = test_app(state); + let req = Request::builder() + .uri("/performance/metrics") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn test_metrics_no_service_returns_500() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + let req = Request::builder() + .uri("/performance/metrics") + .header("Authorization", format!("Bearer {token}")) + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR); + } + + #[tokio::test] + async fn test_latency_without_auth_returns_401() { + let state = test_state(); + let app = test_app(state); + let req = Request::builder() + .uri("/performance/latency") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn test_latency_no_service_returns_500() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + let req = Request::builder() + .uri("/performance/latency") + .header("Authorization", format!("Bearer {token}")) + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR); + } + + #[tokio::test] + async fn test_throughput_without_auth_returns_401() { + let state = test_state(); + let app = test_app(state); + let req = Request::builder() + .uri("/performance/throughput") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn test_throughput_no_service_returns_500() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + let req = Request::builder() + .uri("/performance/throughput") + .header("Authorization", format!("Bearer {token}")) + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR); + } +} diff --git a/web-gateway/src/routes/trading.rs b/web-gateway/src/routes/trading.rs index f93fcd9f0..0e591320e 100644 --- a/web-gateway/src/routes/trading.rs +++ b/web-gateway/src/routes/trading.rs @@ -16,9 +16,9 @@ use crate::state::AppState; pub fn router() -> Router { Router::new() .route("/positions", get(get_positions)) - .route("/orders/{id}", get(get_order_status)) + .route("/orders/:id", get(get_order_status)) .route("/orders", post(submit_order)) - .route("/orders/{id}", delete(cancel_order)) + .route("/orders/:id", delete(cancel_order)) .route("/account", get(get_account)) } diff --git a/web-gateway/src/routes/training.rs b/web-gateway/src/routes/training.rs index c2775f631..4c27f219e 100644 --- a/web-gateway/src/routes/training.rs +++ b/web-gateway/src/routes/training.rs @@ -16,7 +16,7 @@ pub fn router() -> Router { Router::new() .route("/jobs", get(list_jobs)) .route("/jobs", post(start_job)) - .route("/jobs/{id}", delete(cancel_job)) + .route("/jobs/:id", delete(cancel_job)) } async fn list_jobs( @@ -97,3 +97,165 @@ async fn cancel_job( serde_json::to_value(response.into_inner()).map_err(|e| AppError::Internal(e.into()))?, )) } + +#[cfg(test)] +#[allow(clippy::unwrap_used)] +mod tests { + use axum::body::Body; + use axum::http::{Request, StatusCode}; + use axum::middleware; + use axum::Router; + use jsonwebtoken::{encode, EncodingKey, Header}; + use std::sync::Arc; + use tower::ServiceExt; + + use super::*; + use crate::auth::claims::Claims; + use crate::auth::middleware::auth_middleware; + use crate::config::GatewayConfig; + + const SECRET: &str = "training-test-secret"; + + fn test_state() -> AppState { + let config = GatewayConfig { + jwt_secret: SECRET.to_string(), + trading_service_url: String::new(), + ..GatewayConfig::default() + }; + let (ws_broadcast, _) = tokio::sync::broadcast::channel(16); + AppState { + config: Arc::new(config), + trading_channel: None, + backtesting_channel: None, + ml_training_channel: None, + ws_broadcast, + } + } + + fn test_app(state: AppState) -> Router { + Router::new() + .nest("/training", router()) + .layer(middleware::from_fn_with_state( + state.config.clone(), + auth_middleware, + )) + .with_state(state) + } + + fn make_token() -> String { + let claims = Claims { + sub: "trainer-1".into(), + exp: u64::MAX, + iat: 1_700_000_000, + jti: "tr1".into(), + roles: vec!["trainer".into()], + permissions: vec!["write:training".into()], + }; + let key = EncodingKey::from_secret(SECRET.as_bytes()); + encode(&Header::default(), &claims, &key).unwrap() + } + + #[tokio::test] + async fn test_list_jobs_without_auth_returns_401() { + let state = test_state(); + let app = test_app(state); + let req = Request::builder() + .uri("/training/jobs") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn test_list_jobs_no_service_returns_500() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + let req = Request::builder() + .uri("/training/jobs") + .header("Authorization", format!("Bearer {token}")) + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR); + } + + #[tokio::test] + async fn test_start_job_without_auth_returns_401() { + let state = test_state(); + let app = test_app(state); + let body = serde_json::json!({"model_type": "dqn"}); + let req = Request::builder() + .method("POST") + .uri("/training/jobs") + .header("Content-Type", "application/json") + .body(Body::from(serde_json::to_string(&body).unwrap())) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn test_start_job_no_service_returns_500() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + let body = serde_json::json!({"model_type": "dqn"}); + let req = Request::builder() + .method("POST") + .uri("/training/jobs") + .header("Authorization", format!("Bearer {token}")) + .header("Content-Type", "application/json") + .body(Body::from(serde_json::to_string(&body).unwrap())) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR); + } + + #[tokio::test] + async fn test_start_job_invalid_body_returns_422() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + // Missing required model_type field + let body = serde_json::json!({"use_gpu": true}); + let req = Request::builder() + .method("POST") + .uri("/training/jobs") + .header("Authorization", format!("Bearer {token}")) + .header("Content-Type", "application/json") + .body(Body::from(serde_json::to_string(&body).unwrap())) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNPROCESSABLE_ENTITY); + } + + #[tokio::test] + async fn test_cancel_job_without_auth_returns_401() { + let state = test_state(); + let app = test_app(state); + let req = Request::builder() + .method("DELETE") + .uri("/training/jobs/job-abc") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn test_cancel_job_no_service_returns_500() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + let req = Request::builder() + .method("DELETE") + .uri("/training/jobs/job-abc") + .header("Authorization", format!("Bearer {token}")) + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR); + } +} diff --git a/web-gateway/src/routes/tune.rs b/web-gateway/src/routes/tune.rs index bc9d6203a..21b55de06 100644 --- a/web-gateway/src/routes/tune.rs +++ b/web-gateway/src/routes/tune.rs @@ -15,9 +15,9 @@ use crate::state::AppState; pub fn router() -> Router { Router::new() .route("/start", post(start_tune)) - .route("/{id}/status", get(get_status)) - .route("/{id}/stop", post(stop_tune)) - .route("/{id}/best", get(get_best)) + .route("/:id/status", get(get_status)) + .route("/:id/stop", post(stop_tune)) + .route("/:id/best", get(get_best)) } #[derive(serde::Deserialize)] @@ -128,3 +128,137 @@ async fn get_best( "total_trials": status.total_trials, }))) } + +#[cfg(test)] +#[allow(clippy::unwrap_used)] +mod tests { + use axum::body::Body; + use axum::http::{Request, StatusCode}; + use axum::middleware; + use axum::Router; + use jsonwebtoken::{encode, EncodingKey, Header}; + use std::sync::Arc; + use tower::ServiceExt; + + use super::*; + use crate::auth::claims::Claims; + use crate::auth::middleware::auth_middleware; + use crate::config::GatewayConfig; + + const SECRET: &str = "tune-test-secret"; + + fn test_state() -> AppState { + let config = GatewayConfig { + jwt_secret: SECRET.to_string(), + trading_service_url: String::new(), + ..GatewayConfig::default() + }; + let (ws_broadcast, _) = tokio::sync::broadcast::channel(16); + AppState { + config: Arc::new(config), + trading_channel: None, + backtesting_channel: None, + ml_training_channel: None, + ws_broadcast, + } + } + + fn test_app(state: AppState) -> Router { + Router::new() + .nest("/tune", router()) + .layer(middleware::from_fn_with_state( + state.config.clone(), + auth_middleware, + )) + .with_state(state) + } + + fn make_token() -> String { + let claims = Claims { + sub: "tuner-1".into(), + exp: u64::MAX, + iat: 1_700_000_000, + jti: "tn1".into(), + roles: vec!["trainer".into()], + permissions: vec!["write:tuning".into()], + }; + let key = EncodingKey::from_secret(SECRET.as_bytes()); + encode(&Header::default(), &claims, &key).unwrap() + } + + #[tokio::test] + async fn test_start_tune_without_auth_returns_401() { + let state = test_state(); + let app = test_app(state); + let body = serde_json::json!({"model_type": "dqn"}); + let req = Request::builder() + .method("POST") + .uri("/tune/start") + .header("Content-Type", "application/json") + .body(Body::from(serde_json::to_string(&body).unwrap())) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn test_start_tune_no_service_returns_500() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + let body = serde_json::json!({"model_type": "dqn"}); + let req = Request::builder() + .method("POST") + .uri("/tune/start") + .header("Authorization", format!("Bearer {token}")) + .header("Content-Type", "application/json") + .body(Body::from(serde_json::to_string(&body).unwrap())) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR); + } + + #[tokio::test] + async fn test_start_tune_invalid_body_returns_422() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + // Missing required model_type field + let body = serde_json::json!({"num_trials": 10}); + let req = Request::builder() + .method("POST") + .uri("/tune/start") + .header("Authorization", format!("Bearer {token}")) + .header("Content-Type", "application/json") + .body(Body::from(serde_json::to_string(&body).unwrap())) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNPROCESSABLE_ENTITY); + } + + #[tokio::test] + async fn test_get_status_without_auth_returns_401() { + let state = test_state(); + let app = test_app(state); + let req = Request::builder() + .uri("/tune/tune-job-1/status") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn test_get_status_no_service_returns_500() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + let req = Request::builder() + .uri("/tune/tune-job-1/status") + .header("Authorization", format!("Bearer {token}")) + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR); + } +}