//! Integration tests for SharedMLStrategy //! //! Validates that ONE SINGLE SYSTEM works for both trading and backtesting services. //! NO duplication - both services use the same SharedMLStrategy instance. use chrono::Utc; use common::ml_strategy::{MLPrediction, SharedMLStrategy}; use ml::features::ProductionFeatureExtractorAdapter; use std::sync::Arc; #[tokio::test] async fn test_single_strategy_both_services() { // Create ONE SINGLE SYSTEM with production feature extractor (225 features) let extractor = Box::new(ProductionFeatureExtractorAdapter::new()); let strategy = Arc::new(SharedMLStrategy::new_with_production_extractor(extractor, 0.3)); // Simulate trading service using the strategy let trading_strategy = Arc::clone(&strategy); let trading_handle = tokio::spawn(async move { let predictions = trading_strategy .get_ensemble_prediction(100.0, 1000.0, Utc::now()) .await .expect("Trading service should get predictions"); // Calculate vote if predictions are available if !predictions.is_empty() { trading_strategy.calculate_ensemble_vote(&predictions); } predictions.len() }); // Simulate backtesting service using the SAME strategy let backtesting_strategy = Arc::clone(&strategy); let backtesting_handle = tokio::spawn(async move { let predictions = backtesting_strategy .get_ensemble_prediction(102.0, 1100.0, Utc::now()) .await .expect("Backtesting service should get predictions"); // Calculate vote if predictions are available if !predictions.is_empty() { backtesting_strategy.calculate_ensemble_vote(&predictions); } predictions.len() }); // Both services should succeed let trading_count = trading_handle.await.expect("Trading task should complete"); let backtesting_count = backtesting_handle .await .expect("Backtesting task should complete"); assert!(trading_count > 0, "Trading should generate predictions"); assert!( backtesting_count > 0, "Backtesting should generate predictions" ); // Performance tracking would be populated after validate_predictions is called // For now, just verify the strategy is functioning } #[tokio::test] async fn test_concurrent_access_from_multiple_services() { let extractor = Box::new(ProductionFeatureExtractorAdapter::new()); let strategy = Arc::new(SharedMLStrategy::new_with_production_extractor(extractor, 0.5)); let mut handles = Vec::new(); // Spawn 10 concurrent tasks (simulating trading + backtesting + monitoring services) for i in 0..10 { let strategy_clone = Arc::clone(&strategy); let handle = tokio::spawn(async move { let price = 100.0 + (i as f64); let volume = 1000.0 + (i as f64 * 10.0); strategy_clone .get_ensemble_prediction(price, volume, Utc::now()) .await .expect("Should get predictions") }); handles.push(handle); } // Wait for all tasks for handle in handles { let predictions = handle.await.expect("Task should complete"); assert!(!predictions.is_empty(), "Should have predictions"); } } #[tokio::test] async fn test_ensemble_vote_aggregation() { let extractor = Box::new(ProductionFeatureExtractorAdapter::new()); let strategy = SharedMLStrategy::new_with_production_extractor(extractor, 0.0); let predictions = vec![ MLPrediction { model_id: "dqn_v1".to_string(), prediction_value: 0.8, confidence: 0.9, features: vec![], timestamp: Utc::now(), inference_latency_us: 50, }, MLPrediction { model_id: "dqn_v2".to_string(), prediction_value: 0.6, confidence: 0.7, features: vec![], timestamp: Utc::now(), inference_latency_us: 60, }, MLPrediction { model_id: "dqn_v3".to_string(), prediction_value: 0.7, confidence: 0.8, features: vec![], timestamp: Utc::now(), inference_latency_us: 55, }, ]; let result = strategy.calculate_ensemble_vote(&predictions); assert!(result.is_some(), "Should calculate ensemble vote"); let (vote, confidence) = result.unwrap_or_default(); // Weighted average should be between 0.6 and 0.8 assert!( (0.6..=0.8).contains(&vote), "Vote should be in expected range" ); assert!( (0.7..=0.9).contains(&confidence), "Confidence should be in expected range" ); } #[tokio::test] async fn test_performance_tracking_across_services() { let extractor = Box::new(ProductionFeatureExtractorAdapter::new()); let strategy = Arc::new(SharedMLStrategy::new_with_production_extractor(extractor, 0.5)); // Trading service generates signals for _ in 0..5 { let predictions = strategy .get_ensemble_prediction(100.0, 1000.0, Utc::now()) .await .expect("Should get predictions"); // Validate positive outcome strategy.validate_predictions(&predictions, 0.05).await; } // Backtesting service generates signals for _ in 0..5 { let predictions = strategy .get_ensemble_prediction(102.0, 1100.0, Utc::now()) .await .expect("Should get predictions"); // Validate negative outcome strategy.validate_predictions(&predictions, -0.02).await; } // Check performance summary let performance = strategy.get_performance_summary().await; for (model_id, perf) in performance.iter() { assert!( perf.total_predictions > 0, "Model {} should have predictions", model_id ); assert!( perf.accuracy_percentage >= 0.0 && perf.accuracy_percentage <= 100.0, "Accuracy should be valid percentage" ); } } #[tokio::test] async fn test_confidence_threshold_filtering() { let high_extractor = Box::new(ProductionFeatureExtractorAdapter::new()); let high_threshold_strategy = SharedMLStrategy::new_with_production_extractor(high_extractor, 0.95); let low_extractor = Box::new(ProductionFeatureExtractorAdapter::new()); let low_threshold_strategy = SharedMLStrategy::new_with_production_extractor(low_extractor, 0.1); // High threshold should filter out most predictions let high_predictions = high_threshold_strategy .get_ensemble_prediction(100.0, 1000.0, Utc::now()) .await .expect("Should get predictions"); // Low threshold should keep most predictions let low_predictions = low_threshold_strategy .get_ensemble_prediction(100.0, 1000.0, Utc::now()) .await .expect("Should get predictions"); assert!( low_predictions.len() >= high_predictions.len(), "Lower threshold should have more predictions" ); } #[tokio::test] async fn test_feature_extraction_consistency() { let extractor = Box::new(ProductionFeatureExtractorAdapter::new()); let strategy = Arc::new(SharedMLStrategy::new_with_production_extractor(extractor, 0.5)); // Generate predictions at two different times with same price/volume let predictions1 = strategy .get_ensemble_prediction(100.0, 1000.0, Utc::now()) .await .expect("Should get predictions"); tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; let predictions2 = strategy .get_ensemble_prediction(100.0, 1000.0, Utc::now()) .await .expect("Should get predictions"); // Should have same number of models responding assert_eq!( predictions1.len(), predictions2.len(), "Should have consistent number of predictions" ); } #[tokio::test] async fn test_empty_prediction_handling() { let extractor = Box::new(ProductionFeatureExtractorAdapter::new()); let strategy = SharedMLStrategy::new_with_production_extractor(extractor, 0.99); // Very high threshold let predictions = vec![]; let result = strategy.calculate_ensemble_vote(&predictions); assert!(result.is_none(), "Should return None for empty predictions"); } #[tokio::test] async fn test_model_performance_accuracy_tracking() { let extractor = Box::new(ProductionFeatureExtractorAdapter::new()); let strategy = SharedMLStrategy::new_with_production_extractor(extractor, 0.0); let prediction = MLPrediction { model_id: "test_model".to_string(), prediction_value: 0.7, // Predicts positive confidence: 0.8, features: vec![], timestamp: Utc::now(), inference_latency_us: 50, }; // Test with positive outcome (correct prediction) strategy .validate_predictions(std::slice::from_ref(&prediction), 0.05) .await; let performance = strategy.get_performance_summary().await; let model_perf = performance .get("test_model") .expect("Should have test_model performance"); assert_eq!(model_perf.total_predictions, 1); assert_eq!(model_perf.correct_predictions, 1); assert_eq!(model_perf.accuracy_percentage, 100.0); // Test with negative outcome (incorrect prediction) strategy .validate_predictions(std::slice::from_ref(&prediction), -0.05) .await; let performance = strategy.get_performance_summary().await; let model_perf = performance .get("test_model") .expect("Should have test_model performance"); assert_eq!(model_perf.total_predictions, 2); assert_eq!(model_perf.correct_predictions, 1); assert_eq!(model_perf.accuracy_percentage, 50.0); }