//! Model Factory for Testing //! //! This module provides factory functions for creating model instances //! primarily for testing purposes. use crate::{Features, MLModel, MLResult, ModelMetadata, ModelPrediction, ModelType}; use std::sync::Arc; /// Simple `DQN` wrapper for testing #[derive(Debug)] pub struct DQNWrapper { model_id: String, } impl DQNWrapper { /// Create a new `DQN` wrapper pub fn new(model_id: String) -> Self { Self { model_id } } } #[async_trait::async_trait] impl MLModel for DQNWrapper { fn name(&self) -> &str { &self.model_id } fn model_type(&self) -> ModelType { ModelType::DQN } async fn predict(&self, _features: &Features) -> MLResult { // Simple stub implementation for testing Ok(ModelPrediction::new( self.model_id.clone(), 0.5, // prediction value 0.8, // confidence )) } fn get_confidence(&self) -> f64 { 0.8 } fn get_metadata(&self) -> ModelMetadata { ModelMetadata::new( ModelType::DQN, "1.0.0".to_string(), 10, // features_used 128.0, // memory_usage_mb ) } } /// Create a `DQN` wrapper for testing pub fn create_dqn_wrapper() -> MLResult> { Ok(Arc::new(DQNWrapper::new("test_dqn".to_string()))) } /// Create a `DQN` wrapper with specific model ID pub fn create_dqn_wrapper_with_id(model_id: String) -> MLResult> { Ok(Arc::new(DQNWrapper::new(model_id))) } /// Simple PPO wrapper for testing #[derive(Debug)] pub struct PPOWrapper { model_id: String, } impl PPOWrapper { /// Create a new PPO wrapper pub fn new(model_id: String) -> Self { Self { model_id } } } #[async_trait::async_trait] impl MLModel for PPOWrapper { fn name(&self) -> &str { &self.model_id } fn model_type(&self) -> ModelType { ModelType::PPO } async fn predict(&self, _features: &Features) -> MLResult { // Simple stub implementation for testing Ok(ModelPrediction::new( self.model_id.clone(), 0.6, // prediction value 0.85, // confidence )) } fn get_confidence(&self) -> f64 { 0.85 } fn get_metadata(&self) -> ModelMetadata { ModelMetadata::new( ModelType::PPO, "1.0.0".to_string(), 15, // features_used 145.0, // memory_usage_mb ) } } /// Create a PPO wrapper for testing pub fn create_ppo_wrapper() -> MLResult> { Ok(Arc::new(PPOWrapper::new("test_ppo".to_string()))) } /// Create a PPO wrapper with specific model ID pub fn create_ppo_wrapper_with_id(model_id: String) -> MLResult> { Ok(Arc::new(PPOWrapper::new(model_id))) } /// Simple TFT wrapper for testing #[derive(Debug)] pub struct TFTWrapper { model_id: String, } impl TFTWrapper { /// Create a new TFT wrapper pub fn new(model_id: String) -> Self { Self { model_id } } } #[async_trait::async_trait] impl MLModel for TFTWrapper { fn name(&self) -> &str { &self.model_id } fn model_type(&self) -> ModelType { ModelType::TFT } async fn predict(&self, _features: &Features) -> MLResult { // Simple stub implementation for testing Ok(ModelPrediction::new( self.model_id.clone(), 0.55, // prediction value 0.82, // confidence )) } fn get_confidence(&self) -> f64 { 0.82 } fn get_metadata(&self) -> ModelMetadata { ModelMetadata::new( ModelType::TFT, "1.0.0".to_string(), 20, // features_used 125.0, // memory_usage_mb ) } } /// Create a TFT wrapper for testing pub fn create_tft_wrapper() -> MLResult> { Ok(Arc::new(TFTWrapper::new("test_tft".to_string()))) } /// Create a TFT wrapper with specific model ID pub fn create_tft_wrapper_with_id(model_id: String) -> MLResult> { Ok(Arc::new(TFTWrapper::new(model_id))) } /// Simple MAMBA wrapper for testing #[derive(Debug)] pub struct MambaWrapper { model_id: String, } impl MambaWrapper { /// Create a new MAMBA wrapper pub fn new(model_id: String) -> Self { Self { model_id } } } #[async_trait::async_trait] impl MLModel for MambaWrapper { fn name(&self) -> &str { &self.model_id } fn model_type(&self) -> ModelType { ModelType::MAMBA } async fn predict(&self, _features: &Features) -> MLResult { // Simple stub implementation for testing Ok(ModelPrediction::new( self.model_id.clone(), 0.58, // prediction value 0.87, // confidence )) } fn get_confidence(&self) -> f64 { 0.87 } fn get_metadata(&self) -> ModelMetadata { ModelMetadata::new( ModelType::MAMBA, "1.0.0".to_string(), 25, // features_used 164.0, // memory_usage_mb ) } } /// Create a MAMBA wrapper for testing pub fn create_mamba_wrapper() -> MLResult> { Ok(Arc::new(MambaWrapper::new("test_mamba".to_string()))) } /// Create a MAMBA wrapper with specific model ID pub fn create_mamba_wrapper_with_id(model_id: String) -> MLResult> { Ok(Arc::new(MambaWrapper::new(model_id))) } #[cfg(test)] mod tests { use super::*; #[tokio::test] async fn test_create_dqn_wrapper() { let model = create_dqn_wrapper().unwrap(); assert_eq!(model.name(), "test_dqn"); assert_eq!(model.model_type(), ModelType::DQN); assert!(model.is_ready()); } #[tokio::test] async fn test_dqn_wrapper_prediction() { let model = create_dqn_wrapper().unwrap(); let features = Features::new( vec![1.0, 2.0, 3.0], vec!["f1".to_string(), "f2".to_string(), "f3".to_string()], ); let prediction = model.predict(&features).await.unwrap(); assert_eq!(prediction.value, 0.5); assert_eq!(prediction.confidence, 0.8); } #[tokio::test] async fn test_create_ppo_wrapper() { let model = create_ppo_wrapper().unwrap(); assert_eq!(model.name(), "test_ppo"); assert_eq!(model.model_type(), ModelType::PPO); assert!(model.is_ready()); } #[tokio::test] async fn test_ppo_wrapper_prediction() { let model = create_ppo_wrapper().unwrap(); let features = Features::new( vec![1.0, 2.0, 3.0], vec!["f1".to_string(), "f2".to_string(), "f3".to_string()], ); let prediction = model.predict(&features).await.unwrap(); assert_eq!(prediction.value, 0.6); assert_eq!(prediction.confidence, 0.85); } #[tokio::test] async fn test_create_tft_wrapper() { let model = create_tft_wrapper().unwrap(); assert_eq!(model.name(), "test_tft"); assert_eq!(model.model_type(), ModelType::TFT); assert!(model.is_ready()); } #[tokio::test] async fn test_tft_wrapper_prediction() { let model = create_tft_wrapper().unwrap(); let features = Features::new( vec![1.0, 2.0, 3.0], vec!["f1".to_string(), "f2".to_string(), "f3".to_string()], ); let prediction = model.predict(&features).await.unwrap(); assert_eq!(prediction.value, 0.55); assert_eq!(prediction.confidence, 0.82); } #[tokio::test] async fn test_create_mamba_wrapper() { let model = create_mamba_wrapper().unwrap(); assert_eq!(model.name(), "test_mamba"); assert_eq!(model.model_type(), ModelType::MAMBA); assert!(model.is_ready()); } #[tokio::test] async fn test_mamba_wrapper_prediction() { let model = create_mamba_wrapper().unwrap(); let features = Features::new( vec![1.0, 2.0, 3.0], vec!["f1".to_string(), "f2".to_string(), "f3".to_string()], ); let prediction = model.predict(&features).await.unwrap(); assert_eq!(prediction.value, 0.58); assert_eq!(prediction.confidence, 0.87); } #[tokio::test] async fn test_all_wrappers_with_custom_ids() { let dqn = create_dqn_wrapper_with_id("custom_dqn".to_string()).unwrap(); let ppo = create_ppo_wrapper_with_id("custom_ppo".to_string()).unwrap(); let tft = create_tft_wrapper_with_id("custom_tft".to_string()).unwrap(); let mamba = create_mamba_wrapper_with_id("custom_mamba".to_string()).unwrap(); assert_eq!(dqn.name(), "custom_dqn"); assert_eq!(ppo.name(), "custom_ppo"); assert_eq!(tft.name(), "custom_tft"); assert_eq!(mamba.name(), "custom_mamba"); // All models should be ready assert!(dqn.is_ready()); assert!(ppo.is_ready()); assert!(tft.is_ready()); assert!(mamba.is_ready()); } }