diff --git a/web-dashboard/src/pages/ConfigDashboard.tsx b/web-dashboard/src/pages/ConfigDashboard.tsx index 83b203a12..aa5412fdd 100644 --- a/web-dashboard/src/pages/ConfigDashboard.tsx +++ b/web-dashboard/src/pages/ConfigDashboard.tsx @@ -16,6 +16,13 @@ export function ConfigDashboard() { return (
+ {config.isError && ( +
+ Failed to load configuration: {config.error.message} + +
+ )} + {/* Tabs */}
{(['trading', 'risk', 'ml', 'system'] as const).map((tab) => ( diff --git a/web-gateway/src/routes/backtesting.rs b/web-gateway/src/routes/backtesting.rs index 2d2dc68a0..72bae93bb 100644 --- a/web-gateway/src/routes/backtesting.rs +++ b/web-gateway/src/routes/backtesting.rs @@ -34,6 +34,8 @@ async fn start_backtest( Extension(_claims): Extension, Json(body): Json, ) -> Result, AppError> { + validate_backtest(&body)?; + let channel = state .backtesting_channel .as_ref() @@ -56,6 +58,38 @@ async fn start_backtest( )) } +/// Validate backtest request fields before forwarding to gRPC. +fn validate_backtest(body: &StartBacktestBody) -> Result<(), AppError> { + if body.strategy_name.trim().is_empty() { + return Err(AppError::BadRequest( + "Strategy name is required".to_owned(), + )); + } + if body.symbols.is_empty() { + return Err(AppError::BadRequest( + "At least one symbol is required".to_owned(), + )); + } + for sym in &body.symbols { + if !sym.chars().all(|c| c.is_ascii_alphanumeric() || c == '.') { + return Err(AppError::BadRequest(format!( + "Invalid symbol: {sym}" + ))); + } + } + if body.start_date_unix_nanos >= body.end_date_unix_nanos { + return Err(AppError::BadRequest( + "Start date must be before end date".to_owned(), + )); + } + if !body.initial_capital.is_finite() || body.initial_capital <= 0.0 { + return Err(AppError::BadRequest( + "Initial capital must be a positive number".to_owned(), + )); + } + Ok(()) +} + async fn get_status( State(state): State, Extension(_claims): Extension, @@ -265,4 +299,94 @@ mod tests { let resp = app.oneshot(req).await.unwrap(); assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR); } + + #[test] + fn test_validate_backtest_valid() { + let body = StartBacktestBody { + strategy_name: "momentum".to_string(), + symbols: vec!["ES.FUT".to_string()], + start_date_unix_nanos: 1_000, + end_date_unix_nanos: 2_000, + initial_capital: 100_000.0, + parameters: Default::default(), + save_results: false, + description: String::new(), + }; + assert!(validate_backtest(&body).is_ok()); + } + + #[test] + fn test_validate_backtest_empty_strategy() { + let body = StartBacktestBody { + strategy_name: String::new(), + symbols: vec!["ES.FUT".to_string()], + start_date_unix_nanos: 1_000, + end_date_unix_nanos: 2_000, + initial_capital: 100_000.0, + parameters: Default::default(), + save_results: false, + description: String::new(), + }; + assert!(validate_backtest(&body).is_err()); + } + + #[test] + fn test_validate_backtest_no_symbols() { + let body = StartBacktestBody { + strategy_name: "momentum".to_string(), + symbols: vec![], + start_date_unix_nanos: 1_000, + end_date_unix_nanos: 2_000, + initial_capital: 100_000.0, + parameters: Default::default(), + save_results: false, + description: String::new(), + }; + assert!(validate_backtest(&body).is_err()); + } + + #[test] + fn test_validate_backtest_bad_date_range() { + let body = StartBacktestBody { + strategy_name: "momentum".to_string(), + symbols: vec!["ES.FUT".to_string()], + start_date_unix_nanos: 2_000, + end_date_unix_nanos: 1_000, + initial_capital: 100_000.0, + parameters: Default::default(), + save_results: false, + description: String::new(), + }; + assert!(validate_backtest(&body).is_err()); + } + + #[test] + fn test_validate_backtest_negative_capital() { + let body = StartBacktestBody { + strategy_name: "momentum".to_string(), + symbols: vec!["ES.FUT".to_string()], + start_date_unix_nanos: 1_000, + end_date_unix_nanos: 2_000, + initial_capital: -1.0, + parameters: Default::default(), + save_results: false, + description: String::new(), + }; + assert!(validate_backtest(&body).is_err()); + } + + #[test] + fn test_validate_backtest_bad_symbol() { + let body = StartBacktestBody { + strategy_name: "momentum".to_string(), + symbols: vec!["ES;DROP".to_string()], + start_date_unix_nanos: 1_000, + end_date_unix_nanos: 2_000, + initial_capital: 100_000.0, + parameters: Default::default(), + save_results: false, + description: String::new(), + }; + assert!(validate_backtest(&body).is_err()); + } } diff --git a/web-gateway/src/routes/training.rs b/web-gateway/src/routes/training.rs index 4c27f219e..8c93c94dc 100644 --- a/web-gateway/src/routes/training.rs +++ b/web-gateway/src/routes/training.rs @@ -52,11 +52,25 @@ struct StartJobBody { use_gpu: bool, } +/// Known model types that the ML training service supports. +const VALID_MODEL_TYPES: &[&str] = &["dqn", "ppo", "tft", "mamba2"]; + async fn start_job( State(state): State, Extension(_claims): Extension, Json(body): Json, ) -> Result, AppError> { + if body.model_type.is_empty() { + return Err(AppError::BadRequest("Model type is required".to_owned())); + } + if !VALID_MODEL_TYPES.contains(&body.model_type.to_lowercase().as_str()) { + return Err(AppError::BadRequest(format!( + "Unknown model type '{}'. Valid types: {}", + body.model_type, + VALID_MODEL_TYPES.join(", ") + ))); + } + let channel = state .ml_training_channel .as_ref() @@ -258,4 +272,21 @@ mod tests { let resp = app.oneshot(req).await.unwrap(); assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR); } + + #[tokio::test] + async fn test_start_job_unknown_model_type_returns_400() { + let state = test_state(); + let app = test_app(state); + let token = make_token(); + let body = serde_json::json!({"model_type": "unknown_model"}); + 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::BAD_REQUEST); + } }