From abcbf84509446ff0f49dfaa73d72292ba690a080 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sat, 21 Feb 2026 18:24:52 +0100 Subject: [PATCH] feat(trading_service): wire MLEngine with real EnsembleCoordinator Replace empty MLEngine struct with actual EnsembleCoordinator from ml crate. Registers DQN, PPO, TFT, Mamba2 models with configurable weights loaded from config repository. Co-Authored-By: Claude Opus 4.6 --- services/trading_service/src/state.rs | 121 +++++++++++++++++++++++++- 1 file changed, 117 insertions(+), 4 deletions(-) diff --git a/services/trading_service/src/state.rs b/services/trading_service/src/state.rs index d638428ee..78033e514 100644 --- a/services/trading_service/src/state.rs +++ b/services/trading_service/src/state.rs @@ -757,10 +757,30 @@ impl PositionStateValidator { } } -/// ML engine and model registry placeholder -#[derive(Debug, Default)] +/// ML engine backed by the real `ml::ensemble::EnsembleCoordinator`. +/// +/// The coordinator is lazily initialized during +/// [`initialize_with_config_repository`](MLEngine::initialize_with_config_repository) +/// because it needs async model registration. Before that call, +/// [`coordinator()`](MLEngine::coordinator) returns `None`. pub struct MLEngine { - // ML models and predictions + /// The real ensemble coordinator from the `ml` crate. + /// `None` until [`initialize_with_config_repository`] completes. + coordinator: Option, +} + +impl std::fmt::Debug for MLEngine { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("MLEngine") + .field("coordinator", &self.coordinator) + .finish() + } +} + +impl Default for MLEngine { + fn default() -> Self { + Self { coordinator: None } + } } impl MLEngine { @@ -768,11 +788,15 @@ impl MLEngine { Self::default() } + /// Initialize the ML engine by creating an `EnsembleCoordinator` and + /// registering the default production models (DQN, PPO, TFT, MAMBA-2). + /// + /// Model weights and the inference timeout are loaded from the config + /// repository. If a config key is missing the default value is used. pub async fn initialize_with_config_repository( &mut self, config_repository: &Arc, ) -> TradingServiceResult<()> { - // Load and initialize ML models from config repository (no direct database access) // Load ML inference timeout from config repository if let Ok(Some(inference_timeout)) = config_repository .get_config_u64("MachineLearning", "inference_timeout_ms") @@ -784,8 +808,97 @@ impl MLEngine { ); } + // Create the real ensemble coordinator + let coordinator = ml::ensemble::EnsembleCoordinator::new(); + + // Load per-model weights from config, falling back to equal weights + let default_weight = 0.25_f64; + + let dqn_weight = config_repository + .get_config_f64("MachineLearning", "dqn_weight") + .await + .ok() + .flatten() + .unwrap_or(default_weight); + + let ppo_weight = config_repository + .get_config_f64("MachineLearning", "ppo_weight") + .await + .ok() + .flatten() + .unwrap_or(default_weight); + + let tft_weight = config_repository + .get_config_f64("MachineLearning", "tft_weight") + .await + .ok() + .flatten() + .unwrap_or(default_weight); + + let mamba2_weight = config_repository + .get_config_f64("MachineLearning", "mamba2_weight") + .await + .ok() + .flatten() + .unwrap_or(default_weight); + + // Register the four production models + if let Err(e) = coordinator + .register_model("DQN".to_string(), dqn_weight) + .await + { + tracing::warn!("Failed to register DQN model: {}", e); + } + if let Err(e) = coordinator + .register_model("PPO".to_string(), ppo_weight) + .await + { + tracing::warn!("Failed to register PPO model: {}", e); + } + if let Err(e) = coordinator + .register_model("TFT".to_string(), tft_weight) + .await + { + tracing::warn!("Failed to register TFT model: {}", e); + } + if let Err(e) = coordinator + .register_model("MAMBA-2".to_string(), mamba2_weight) + .await + { + tracing::warn!("Failed to register MAMBA-2 model: {}", e); + } + + let model_count = coordinator.model_count().await; + tracing::info!( + "ML engine EnsembleCoordinator initialized with {} models \ + (DQN={:.2}, PPO={:.2}, TFT={:.2}, MAMBA-2={:.2})", + model_count, + dqn_weight, + ppo_weight, + tft_weight, + mamba2_weight, + ); + + self.coordinator = Some(coordinator); + Ok(()) } + + /// Get a shared reference to the underlying `EnsembleCoordinator`. + /// + /// Returns `None` if the engine has not been initialized yet. + pub fn coordinator(&self) -> Option<&ml::ensemble::EnsembleCoordinator> { + self.coordinator.as_ref() + } + + /// Get a mutable reference to the underlying `EnsembleCoordinator`. + /// + /// Useful for adding inference adapters via + /// [`EnsembleCoordinator::add_adapter`]. + /// Returns `None` if the engine has not been initialized yet. + pub fn coordinator_mut(&mut self) -> Option<&mut ml::ensemble::EnsembleCoordinator> { + self.coordinator.as_mut() + } } /// Market data manager with multiple providers