- Delete MLFeatureExtractor (1,294 lines) and SimpleDQNAdapter (235 lines) - Delete 830 lines of inline tests for deleted types - Remove legacy_feature_extractor field from SharedMLStrategy - Replace new() and new_with_production_extractor() with new(extractor, models, threshold) - Single constructor accepts injected models via Vec<Box<dyn MLModelAdapter>> - Update all callers: backtesting_service, 2 integration tests, 2 trading_service tests - Fix doc comments referencing MLFeatureExtractor - Fix feature count test: real extractor produces 51 features, not 225 Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
320 lines
9.2 KiB
Rust
320 lines
9.2 KiB
Rust
//! 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 anyhow::Result;
|
|
use chrono::{DateTime, Utc};
|
|
use common::ml_strategy::{
|
|
MLModelAdapter, MLPrediction, ProductionFeatureExtractor225, SharedMLStrategy,
|
|
};
|
|
use std::sync::Arc;
|
|
|
|
/// Mock adapter that produces deterministic predictions from features
|
|
struct MockAdapter {
|
|
id: String,
|
|
}
|
|
|
|
impl MockAdapter {
|
|
fn new(id: &str) -> Self {
|
|
Self {
|
|
id: id.to_string(),
|
|
}
|
|
}
|
|
}
|
|
|
|
impl MLModelAdapter for MockAdapter {
|
|
fn predict(&self, features: &[f64]) -> Result<MLPrediction> {
|
|
let sum: f64 = features.iter().sum::<f64>() / features.len().max(1) as f64;
|
|
let prediction_value = 1.0 / (1.0 + (-sum).exp());
|
|
let confidence = 0.5 + (prediction_value - 0.5).abs() * 0.8;
|
|
Ok(MLPrediction {
|
|
model_id: self.id.clone(),
|
|
prediction_value,
|
|
confidence,
|
|
features: features.to_vec(),
|
|
timestamp: Utc::now(),
|
|
inference_latency_us: 10,
|
|
})
|
|
}
|
|
|
|
fn model_id(&self) -> &str {
|
|
&self.id
|
|
}
|
|
|
|
fn validate_prediction(&mut self, _prediction: &MLPrediction, _actual_outcome: bool) {}
|
|
}
|
|
|
|
/// Mock extractor returning 225 features
|
|
struct MockExtractor;
|
|
|
|
impl ProductionFeatureExtractor225 for MockExtractor {
|
|
fn update(&mut self, _price: f64, _volume: f64, _timestamp: DateTime<Utc>) -> Result<()> {
|
|
Ok(())
|
|
}
|
|
fn extract_features(&mut self) -> Result<Vec<f64>> {
|
|
Ok(vec![0.1; 225])
|
|
}
|
|
}
|
|
|
|
fn make_strategy(threshold: f64) -> SharedMLStrategy {
|
|
SharedMLStrategy::new(
|
|
Box::new(MockExtractor),
|
|
vec![Box::new(MockAdapter::new("mock_v1"))],
|
|
threshold,
|
|
)
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_single_strategy_both_services() {
|
|
let strategy = Arc::new(make_strategy(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");
|
|
|
|
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");
|
|
|
|
if !predictions.is_empty() {
|
|
backtesting_strategy.calculate_ensemble_vote(&predictions);
|
|
}
|
|
|
|
predictions.len()
|
|
});
|
|
|
|
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"
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_concurrent_access_from_multiple_services() {
|
|
let strategy = Arc::new(make_strategy(0.0));
|
|
|
|
let mut handles = Vec::new();
|
|
|
|
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);
|
|
}
|
|
|
|
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 strategy = make_strategy(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();
|
|
|
|
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 strategy = Arc::new(make_strategy(0.0));
|
|
|
|
// 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");
|
|
|
|
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");
|
|
|
|
strategy.validate_predictions(&predictions, -0.02).await;
|
|
}
|
|
|
|
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_threshold_strategy = make_strategy(0.95);
|
|
let low_threshold_strategy = make_strategy(0.1);
|
|
|
|
let high_predictions = high_threshold_strategy
|
|
.get_ensemble_prediction(100.0, 1000.0, Utc::now())
|
|
.await
|
|
.expect("Should get 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 strategy = Arc::new(make_strategy(0.0));
|
|
|
|
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(10)).await;
|
|
|
|
let predictions2 = strategy
|
|
.get_ensemble_prediction(100.0, 1000.0, Utc::now())
|
|
.await
|
|
.expect("Should get predictions");
|
|
|
|
assert_eq!(
|
|
predictions1.len(),
|
|
predictions2.len(),
|
|
"Should have consistent number of predictions"
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_empty_prediction_handling() {
|
|
let strategy = make_strategy(0.99);
|
|
|
|
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 strategy = make_strategy(0.0);
|
|
|
|
let prediction = MLPrediction {
|
|
model_id: "test_model".to_string(),
|
|
prediction_value: 0.7,
|
|
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);
|
|
}
|