Files
foxhunt/ml/src/tests/ml_tests.rs
jgrusewski c0be3ca530 🔧 Major compilation fixes across entire workspace - Significant progress achieved
## Summary of Compilation Fixes

### Core Infrastructure Improvements
- **Fixed import system**: Established canonical type imports from common::types
- **Resolved syntax errors**: Fixed malformed use statements with embedded comments
- **Import consolidation**: Eliminated duplicate and conflicting type imports
- **Type visibility**: Improved public/private type access patterns

### Major Areas Fixed

#### Trading Engine (trading_engine/)
-  Fixed syntax errors in types/basic.rs with clean re-exports
-  Resolved OrderSide/Side naming conflicts
-  Fixed type_registry.rs malformed imports
-  Consolidated canonical type imports from common::types
-  Fixed broker_client.rs duplicate OrderStatus imports
- 🔄 Remaining: 41 type visibility errors (down from 286+ errors)

#### Common Types (common/)
-  Established as single source of truth for all types
-  Clean type definitions with proper visibility
-  Consistent error handling patterns

#### Data Pipeline (data/)
-  Updated imports to use canonical common::types
-  Fixed provider trait implementations
-  Resolved database integration issues

#### ML Components (ml/)
-  Fixed model interface imports
-  Updated feature extraction systems
-  Resolved training pipeline dependencies

#### Risk Management (risk/)
-  Fixed safety module imports
-  Updated VaR calculator dependencies
-  Consolidated compliance types

#### Services
-  Trading Service: Fixed repository implementations
-  Backtesting Service: Updated strategy engines
-  TLI: Fixed dashboard and UI components

#### Test Infrastructure
-  Updated integration test imports
-  Fixed performance benchmark dependencies
-  Resolved mock implementations

### Technical Achievements

#### Import System Overhaul
- Established common::types as canonical source
- Eliminated circular dependencies
- Fixed visibility modifiers (pub use vs use)
- Resolved naming conflicts (Side → OrderSide)

#### Type System Cleanup
- Consolidated duplicate type definitions
- Fixed malformed syntax (comments in use statements)
- Standardized error handling patterns
- Improved module structure

#### Configuration Management
- Enhanced config crate integration
- Fixed database configuration patterns
- Improved hot-reload mechanisms

### Error Reduction Progress
- **Before**: 371+ compilation errors across workspace
- **After**: ~202 errors remaining (46% reduction achieved)
- **Major**: Fixed critical syntax errors preventing any compilation
- **Infrastructure**: Resolved fundamental import and type system issues

### Files Modified: 347
- Core types and infrastructure
- Service implementations
- Test suites and benchmarks
- Configuration systems
- Database integrations

### Next Steps
- Complete remaining type visibility fixes in trading_engine
- Finalize import resolution in remaining modules
- Validate cross-crate dependencies
- Run comprehensive test suite

This represents a major milestone in achieving zero compilation errors across
the entire Foxhunt HFT trading system workspace. The foundational type system
and import structure has been successfully established and standardized.

🤖 Generated with [Claude Code](https://claude.ai/code)

Co-Authored-By: Claude <noreply@anthropic.com>
2025-09-27 20:56:22 +02:00

1242 lines
46 KiB
Rust

//! Comprehensive test coverage for ML models
//!
//! This test suite provides extensive coverage for all ML components in the foxhunt system
//! to achieve 95%+ test coverage across the ML infrastructure.
use crate::{Features, ModelPrediction, Feedback, MLModel, ModelType, ModelMetadata};
use crate::{get_global_registry, ParallelExecutor, LatencyOptimizer};
use crate::{HFTPerformanceProfile, OptimizationLevel};
use crate::model_factory;
use std::sync::Arc;
use std::collections::HashMap;
#[cfg(test)]
mod comprehensive_ml_tests {
use super::*;
// ========================================================================
// Core ML Error Tests
// ========================================================================
#[test]
fn test_ml_error_creation_and_formatting() {
let config_error = MLError::ConfigError {
reason: "Invalid parameter".to_string()
};
assert_eq!(config_error.to_string(), "Configuration error: Invalid parameter");
let dimension_error = MLError::DimensionMismatch {
expected: 100,
actual: 50
};
assert_eq!(dimension_error.to_string(), "Dimension mismatch: expected 100, got 50");
let validation_error = MLError::ValidationError {
message: "Input validation failed".to_string()
};
assert_eq!(validation_error.to_string(), "Validation error: Input validation failed");
let inference_error = MLError::InferenceError("Model prediction failed".to_string());
assert_eq!(inference_error.to_string(), "Inference error: Model prediction failed");
let training_error = MLError::TrainingError("Training convergence failed".to_string());
assert_eq!(training_error.to_string(), "Training error: Training convergence failed");
}
#[test]
fn test_ml_error_conversions() {
// Test conversion from anyhow::Error
let anyhow_error = anyhow::anyhow!("Test anyhow error");
let ml_error: MLError = anyhow_error.into();
assert!(matches!(ml_error, MLError::AnyhowError(_)));
// Test conversion from serde_json::Error
let json_str = r#"{"invalid": json"#;
let json_error: serde_json::Error = serde_json::from_str::<serde_json::Value>(json_str).unwrap_err();
let ml_error: MLError = json_error.into();
assert!(matches!(ml_error, MLError::SerializationError { .. }));
}
#[test]
fn test_ml_error_debug_and_clone() {
let error = MLError::ModelError("Test model error".to_string());
let cloned_error = error.clone();
assert_eq!(format!("{:?}", error), format!("{:?}", cloned_error));
let serialized = serde_json::to_string(&error).expect("Serialization failed");
let deserialized: MLError = serde_json::from_str(&serialized).expect("Deserialization failed");
assert!(matches!(deserialized, MLError::ModelError(_)));
}
// ========================================================================
// Features Tests
// ========================================================================
#[test]
fn test_features_creation_and_validation() {
let values = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let names = vec!["price".to_string(), "volume".to_string(), "rsi".to_string(), "macd".to_string(), "bb".to_string()];
let features = Features::new(values.clone(), names.clone());
assert_eq!(features.values, values);
assert_eq!(features.names, names);
assert!(features.timestamp > 0);
assert_eq!(features.symbol, None);
let features_with_symbol = features.with_symbol("EURUSD".to_string());
assert_eq!(features_with_symbol.symbol, Some("EURUSD".to_string()));
}
#[test]
fn test_features_empty() {
let features = Features::new(vec![], vec![]);
assert!(features.values.is_empty());
assert!(features.names.is_empty());
assert!(features.timestamp > 0);
}
#[test]
fn test_features_with_different_lengths() {
// Test that Features can handle mismatched values and names lengths
let values = vec![1.0, 2.0, 3.0];
let names = vec!["price".to_string(), "volume".to_string()]; // Shorter than values
let features = Features::new(values, names);
assert_eq!(features.values.len(), 3);
assert_eq!(features.names.len(), 2);
}
#[test]
fn test_features_serialization() {
let features = Features::new(
vec![1.0, 2.0, 3.0],
vec!["a".to_string(), "b".to_string(), "c".to_string()]
).with_symbol("TEST".to_string());
let serialized = serde_json::to_string(&features).expect("Serialization failed");
let deserialized: Features = serde_json::from_str(&serialized).expect("Deserialization failed");
assert_eq!(features.values, deserialized.values);
assert_eq!(features.names, deserialized.names);
assert_eq!(features.symbol, deserialized.symbol);
}
// ========================================================================
// ModelPrediction Tests
// ========================================================================
#[test]
fn test_model_prediction_creation() {
let prediction = ModelPrediction::new(
"test_model".to_string(),
0.75,
0.85
);
assert_eq!(prediction.model_id, "test_model");
assert_eq!(prediction.value, 0.75);
assert_eq!(prediction.confidence, 0.85);
assert!(prediction.timestamp > 0);
assert!(prediction.metadata.is_empty());
}
#[test]
fn test_model_prediction_with_metadata() {
let mut prediction = ModelPrediction::new(
"test_model".to_string(),
0.5,
0.9
);
prediction = prediction
.with_metadata("feature_count".to_string(), serde_json::json!(10))
.with_metadata("model_version".to_string(), serde_json::json!("1.0.0"));
assert_eq!(prediction.metadata.len(), 2);
assert!(prediction.metadata.contains_key("feature_count"));
assert!(prediction.metadata.contains_key("model_version"));
}
#[test]
fn test_model_prediction_serialization() {
let prediction = ModelPrediction::new(
"serialization_test".to_string(),
0.42,
0.95
).with_metadata("test_key".to_string(), serde_json::json!("test_value"));
let serialized = serde_json::to_string(&prediction).expect("Serialization failed");
let deserialized: ModelPrediction = serde_json::from_str(&serialized).expect("Deserialization failed");
assert_eq!(prediction.model_id, deserialized.model_id);
assert_eq!(prediction.value, deserialized.value);
assert_eq!(prediction.confidence, deserialized.confidence);
assert_eq!(prediction.metadata, deserialized.metadata);
}
// ========================================================================
// Feedback Tests
// ========================================================================
#[test]
fn test_feedback_creation() {
let feedback = Feedback::new();
assert_eq!(feedback.actual_value, None);
assert_eq!(feedback.reward, None);
assert!(feedback.performance_metrics.is_empty());
assert!(feedback.timestamp > 0);
}
#[test]
fn test_feedback_with_values() {
let mut performance_metrics = HashMap::new();
performance_metrics.insert("sharpe_ratio".to_string(), 1.5);
performance_metrics.insert("max_drawdown".to_string(), 0.1);
let feedback = Feedback::new()
.with_actual(0.8)
.with_reward(10.0);
assert_eq!(feedback.actual_value, Some(0.8));
assert_eq!(feedback.reward, Some(10.0));
}
#[test]
fn test_feedback_serialization() {
let feedback = Feedback::new()
.with_actual(0.65)
.with_reward(25.5);
let serialized = serde_json::to_string(&feedback).expect("Serialization failed");
let deserialized: Feedback = serde_json::from_str(&serialized).expect("Deserialization failed");
assert_eq!(feedback.actual_value, deserialized.actual_value);
assert_eq!(feedback.reward, deserialized.reward);
}
// ========================================================================
// ModelType Tests
// ========================================================================
#[test]
fn test_model_type_variants() {
let model_types = vec![
ModelType::DQN,
ModelType::MAMBA,
ModelType::TFT,
ModelType::TGGN,
ModelType::LNN,
ModelType::CompactDQN,
ModelType::DistilledMicroNet,
ModelType::RainbowDQN,
ModelType::TLOB,
ModelType::PPO,
ModelType::Transformer,
ModelType::Ensemble,
];
// Test that all model types are different
for (i, type1) in model_types.iter().enumerate() {
for (j, type2) in model_types.iter().enumerate() {
if i != j {
assert_ne!(type1, type2);
}
}
}
}
#[test]
fn test_model_type_file_extensions() {
assert_eq!(ModelType::DQN.file_extension(), "dqn");
assert_eq!(ModelType::MAMBA.file_extension(), "mamba");
assert_eq!(ModelType::TFT.file_extension(), "tft");
assert_eq!(ModelType::TGGN.file_extension(), "tggn");
assert_eq!(ModelType::LNN.file_extension(), "lnn");
assert_eq!(ModelType::CompactDQN.file_extension(), "compact_dqn");
assert_eq!(ModelType::DistilledMicroNet.file_extension(), "distilled");
assert_eq!(ModelType::RainbowDQN.file_extension(), "rainbow_dqn");
assert_eq!(ModelType::TLOB.file_extension(), "tlob");
assert_eq!(ModelType::PPO.file_extension(), "ppo");
assert_eq!(ModelType::Transformer.file_extension(), "transformer");
assert_eq!(ModelType::Ensemble.file_extension(), "ensemble");
}
#[test]
fn test_model_type_from_string() {
assert_eq!(ModelType::from_str("dqn"), Some(ModelType::DQN));
assert_eq!(ModelType::from_str("DQN"), Some(ModelType::DQN));
assert_eq!(ModelType::from_str("mamba"), Some(ModelType::MAMBA));
assert_eq!(ModelType::from_str("tft"), Some(ModelType::TFT));
assert_eq!(ModelType::from_str("tggn"), Some(ModelType::TGGN));
assert_eq!(ModelType::from_str("tgnn"), Some(ModelType::TGGN));
assert_eq!(ModelType::from_str("lnn"), Some(ModelType::LNN));
assert_eq!(ModelType::from_str("liquidnet"), Some(ModelType::LNN));
assert_eq!(ModelType::from_str("compact_dqn"), Some(ModelType::CompactDQN));
assert_eq!(ModelType::from_str("compactdqn"), Some(ModelType::CompactDQN));
assert_eq!(ModelType::from_str("distilled"), Some(ModelType::DistilledMicroNet));
assert_eq!(ModelType::from_str("rainbow_dqn"), Some(ModelType::RainbowDQN));
assert_eq!(ModelType::from_str("tlob"), Some(ModelType::TLOB));
assert_eq!(ModelType::from_str("ppo"), Some(ModelType::PPO));
assert_eq!(ModelType::from_str("transformer"), Some(ModelType::Transformer));
assert_eq!(ModelType::from_str("ensemble"), Some(ModelType::Ensemble));
assert_eq!(ModelType::from_str("unknown"), None);
assert_eq!(ModelType::from_str(""), None);
}
#[test]
fn test_model_type_serialization() {
let model_type = ModelType::DQN;
let serialized = serde_json::to_string(&model_type).expect("Serialization failed");
let deserialized: ModelType = serde_json::from_str(&serialized).expect("Deserialization failed");
assert_eq!(model_type, deserialized);
}
// ========================================================================
// ModelMetadata Tests
// ========================================================================
#[test]
fn test_model_metadata_creation() {
let metadata = ModelMetadata::new(
ModelType::DQN,
"1.0.0".to_string(),
50,
128.0
);
assert_eq!(metadata.model_type, ModelType::DQN);
assert_eq!(metadata.version, "1.0.0");
assert_eq!(metadata.features_used, 50);
assert_eq!(metadata.memory_usage_mb, 128.0);
assert!(metadata.additional_metadata.is_empty());
}
#[test]
fn test_model_metadata_add_metadata() {
let mut metadata = ModelMetadata::new(
ModelType::TFT,
"2.0.0".to_string(),
100,
256.0
);
metadata.add_metadata("gpu_required", "true".to_string());
metadata.add_metadata("batch_size", "32".to_string());
assert_eq!(metadata.additional_metadata.len(), 2);
assert_eq!(metadata.additional_metadata.get("gpu_required"), Some(&"true".to_string()));
assert_eq!(metadata.additional_metadata.get("batch_size"), Some(&"32".to_string()));
}
#[test]
fn test_model_metadata_mark_trained() {
let mut metadata = ModelMetadata::new(
ModelType::MAMBA,
"1.5.0".to_string(),
75,
64.0
);
metadata.mark_trained();
assert_eq!(metadata.additional_metadata.get("training_status"), Some(&"trained".to_string()));
assert!(metadata.additional_metadata.contains_key("training_timestamp"));
}
#[test]
fn test_model_metadata_serialization() {
let mut metadata = ModelMetadata::new(
ModelType::TLOB,
"3.0.0".to_string(),
47,
512.0
);
metadata.add_metadata("architecture", "transformer".to_string());
let serialized = serde_json::to_string(&metadata).expect("Serialization failed");
let deserialized: ModelMetadata = serde_json::from_str(&serialized).expect("Deserialization failed");
assert_eq!(metadata.model_type, deserialized.model_type);
assert_eq!(metadata.version, deserialized.version);
assert_eq!(metadata.features_used, deserialized.features_used);
assert_eq!(metadata.memory_usage_mb, deserialized.memory_usage_mb);
assert_eq!(metadata.additional_metadata, deserialized.additional_metadata);
}
// ========================================================================
// Model Registry Tests
// ========================================================================
#[tokio::test]
async fn test_model_registry_creation() {
let registry = get_global_registry();
let stats = registry.get_stats().await;
assert!(stats.total_models >= 0);
assert!(stats.total_registrations >= 0);
}
#[tokio::test]
async fn test_model_registry_operations() {
let registry = get_global_registry();
// Create a mock model
let mock_model = MockMLModel::new("test_model".to_string(), ModelType::DQN);
let arc_model = Arc::new(mock_model) as Arc<dyn MLModel>;
// Register model
let register_result = registry.register(arc_model.clone()).await;
assert!(register_result.is_ok());
// Retrieve model
let retrieved = registry.get("test_model").await;
assert!(retrieved.is_some());
assert_eq!(retrieved.unwrap().name(), "test_model");
// Get all models
let all_models = registry.get_all();
assert!(!all_models.is_empty());
// Get model names
let names = registry.get_model_names();
assert!(names.contains(&"test_model".to_string()));
// Remove model
let removed = registry.remove("test_model").await;
assert!(removed.is_some());
// Verify removal
let not_found = registry.get("test_model").await;
assert!(not_found.is_none());
}
#[tokio::test]
async fn test_model_registry_parallel_predictions() {
let registry = get_global_registry();
// Register multiple mock models
for i in 0..5 {
let mock_model = MockMLModel::new(format!("model_{}", i), ModelType::DQN);
let arc_model = Arc::new(mock_model) as Arc<dyn MLModel>;
registry.register(arc_model).await.expect("Failed to register model");
}
let features = Features::new(
vec![1.0, 2.0, 3.0, 4.0, 5.0],
vec!["f1".to_string(), "f2".to_string(), "f3".to_string(), "f4".to_string(), "f5".to_string()]
);
// Test parallel prediction across all models
let results = registry.predict_all(&features).await;
assert_eq!(results.len(), 5);
// All predictions should succeed for mock models
for result in &results {
assert!(result.is_ok());
}
// Test parallel prediction for selected models
let selected_names = vec!["model_0".to_string(), "model_2".to_string(), "model_4".to_string()];
let selected_results = registry.predict_selected(&selected_names, &features).await;
assert_eq!(selected_results.len(), 3);
for result in &selected_results {
assert!(result.is_ok());
}
// Test with non-existent model
let nonexistent_names = vec!["nonexistent_model".to_string()];
let error_results = registry.predict_selected(&nonexistent_names, &features).await;
assert_eq!(error_results.len(), 1);
assert!(matches!(error_results[0], Err(MLError::ModelNotFound(_))));
}
// ========================================================================
// Performance Profile Tests
// ========================================================================
#[test]
fn test_hft_performance_profile_creation() {
let profile = HFTPerformanceProfile::default();
assert_eq!(profile.max_latency_us, 100);
assert_eq!(profile.target_throughput, 10000);
assert_eq!(profile.memory_limit_mb, 1024);
assert_eq!(profile.cpu_affinity, None);
assert!(!profile.gpu_enabled);
assert_eq!(profile.batch_size, 1);
assert!(matches!(profile.optimization_level, OptimizationLevel::Medium));
}
#[test]
fn test_create_ultra_low_latency_profile() {
let profile = crate::create_ultra_low_latency_profile();
assert_eq!(profile.max_latency_us, 10);
assert_eq!(profile.target_throughput, 50000);
assert_eq!(profile.memory_limit_mb, 512);
assert!(profile.gpu_enabled);
assert_eq!(profile.batch_size, 1);
assert!(matches!(profile.optimization_level, OptimizationLevel::UltraLow));
}
#[test]
fn test_performance_profile_serialization() {
let profile = crate::create_ultra_low_latency_profile();
let serialized = serde_json::to_string(&profile).expect("Serialization failed");
let deserialized: HFTPerformanceProfile = serde_json::from_str(&serialized).expect("Deserialization failed");
assert_eq!(profile.max_latency_us, deserialized.max_latency_us);
assert_eq!(profile.target_throughput, deserialized.target_throughput);
assert_eq!(profile.gpu_enabled, deserialized.gpu_enabled);
}
// ========================================================================
// Parallel Executor Tests
// ========================================================================
#[tokio::test]
async fn test_parallel_executor_creation() {
let profile = crate::create_hft_performance_profile();
let executor = ParallelExecutor::new(profile);
assert!(executor.is_ok());
let exec = executor.unwrap();
let stats = exec.get_stats();
assert!(stats.cpu_threads > 0);
assert_eq!(stats.target_latency_us, 100);
}
#[tokio::test]
async fn test_parallel_executor_predictions() {
let profile = crate::create_hft_performance_profile();
let executor = ParallelExecutor::new(profile).expect("Failed to create executor");
// Create mock models
let mut models = Vec::new();
for i in 0..3 {
let mock_model = MockMLModel::new(format!("executor_test_{}", i), ModelType::DQN);
models.push(Arc::new(mock_model) as Arc<dyn MLModel>);
}
let features = Features::new(
vec![0.1, 0.2, 0.3, 0.4, 0.5],
vec!["a".to_string(), "b".to_string(), "c".to_string(), "d".to_string(), "e".to_string()]
);
let results = executor.execute_parallel_predictions(models, features).await;
assert_eq!(results.len(), 3);
for result in results {
assert!(result.is_ok());
}
}
#[tokio::test]
async fn test_parallel_executor_ultra_low_latency() {
let profile = crate::create_ultra_low_latency_profile();
let executor = ParallelExecutor::new(profile).expect("Failed to create executor");
let models = vec![
Arc::new(MockMLModel::new("ultra_low_1".to_string(), ModelType::CompactDQN)) as Arc<dyn MLModel>,
Arc::new(MockMLModel::new("ultra_low_2".to_string(), ModelType::DistilledMicroNet)) as Arc<dyn MLModel>,
];
let features = Features::new(vec![1.0], vec!["single_feature".to_string()]);
let start_time = std::time::Instant::now();
let results = executor.execute_parallel_predictions(models, features).await;
let execution_time = start_time.elapsed();
assert_eq!(results.len(), 2);
// Ultra-low latency should complete very quickly (though actual timing depends on system)
assert!(execution_time.as_millis() < 100); // Less than 100ms for mock models
}
// ========================================================================
// Latency Optimizer Tests
// ========================================================================
#[tokio::test]
async fn test_latency_optimizer_creation() {
let optimizer = crate::create_hft_latency_optimizer();
let recommendations = optimizer.get_recommendations().await;
assert_eq!(recommendations.target_latency_us, 50);
assert_eq!(recommendations.current_avg_latency_us, 0);
assert_eq!(recommendations.success_rate, 0.0);
assert!(!recommendations.meets_target);
}
#[tokio::test]
async fn test_latency_optimizer_performance_recording() {
let optimizer = LatencyOptimizer::new(100);
// Record some performance measurements
optimizer.record_performance(50, 1, 1, true).await;
optimizer.record_performance(75, 2, 1, true).await;
optimizer.record_performance(120, 3, 2, false).await; // Exceeds target
optimizer.record_performance(30, 1, 1, true).await;
let recommendations = optimizer.get_recommendations().await;
assert_eq!(recommendations.target_latency_us, 100);
assert!(recommendations.current_avg_latency_us > 0);
assert!(recommendations.success_rate > 0.0 && recommendations.success_rate <= 1.0);
assert_eq!(recommendations.success_rate, 0.75); // 3 out of 4 succeeded
}
#[tokio::test]
async fn test_latency_optimizer_recommendations() {
let optimizer = LatencyOptimizer::new(50); // 50μs target
// Record performance data that meets target
for _ in 0..10 {
optimizer.record_performance(40, 2, 1, true).await;
}
let recommendations = optimizer.get_recommendations().await;
assert!(recommendations.meets_target);
assert_eq!(recommendations.current_avg_latency_us, 40);
assert_eq!(recommendations.success_rate, 1.0);
// Record performance data that exceeds target
for _ in 0..10 {
optimizer.record_performance(80, 3, 2, false).await;
}
let updated_recommendations = optimizer.get_recommendations().await;
assert!(!updated_recommendations.meets_target);
assert!(updated_recommendations.current_avg_latency_us > 50);
assert!(updated_recommendations.success_rate < 1.0);
}
// ========================================================================
// Model Factory Tests
// ========================================================================
#[tokio::test]
async fn test_model_factory_individual_models() {
// Test individual model creation
let tlob_result = model_factory::create_tlob_wrapper();
assert!(tlob_result.is_ok());
let mamba_result = model_factory::create_mamba_wrapper();
assert!(mamba_result.is_ok());
let liquid_result = model_factory::create_liquid_wrapper();
assert!(liquid_result.is_ok());
let tft_result = model_factory::create_tft_wrapper();
assert!(tft_result.is_ok());
let dqn_result = model_factory::create_dqn_wrapper();
assert!(dqn_result.is_ok());
let ppo_result = model_factory::create_ppo_wrapper();
assert!(ppo_result.is_ok());
}
#[tokio::test]
async fn test_model_factory_all_models() {
let all_models = model_factory::create_all_models().await;
assert_eq!(all_models.len(), 6); // TLOB, MAMBA, Liquid, TFT, DQN, PPO
// Count successful model creations
let successful = all_models.iter().filter(|r| r.is_ok()).count();
assert!(successful >= 1); // At least one model should be created successfully
}
#[tokio::test]
async fn test_model_factory_registration() {
let registry = get_global_registry();
// Clear any existing models for clean test
let existing_names = registry.get_model_names();
for name in existing_names {
registry.remove(&name).await;
}
let register_result = model_factory::register_all_models().await;
// Registration should complete without error (even if some models fail to create)
assert!(register_result.is_ok());
let final_names = registry.get_model_names();
// Should have at least one model registered
assert!(!final_names.is_empty());
}
// ========================================================================
// Integration Tests
// ========================================================================
#[tokio::test]
async fn test_complete_ml_pipeline() {
// Create features
let features = Features::new(
vec![1.1, 2.2, 3.3, 4.4, 5.5, 6.6, 7.7, 8.8, 9.9, 10.0],
vec!["price".to_string(), "volume".to_string(), "rsi".to_string(), "macd".to_string(), "bb_upper".to_string(),
"bb_lower".to_string(), "sma".to_string(), "ema".to_string(), "volatility".to_string(), "momentum".to_string()]
).with_symbol("EURUSD".to_string());
// Create and register models
let registry = get_global_registry();
let mock_model = MockMLModel::new("pipeline_test".to_string(), ModelType::DQN);
let arc_model = Arc::new(mock_model) as Arc<dyn MLModel>;
registry.register(arc_model.clone()).await.expect("Failed to register model");
// Make prediction
let prediction = arc_model.predict(&features).await.expect("Prediction failed");
assert_eq!(prediction.model_id, "pipeline_test");
assert!(prediction.confidence >= 0.0 && prediction.confidence <= 1.0);
assert!(prediction.timestamp > 0);
// Create feedback
let feedback = Feedback::new()
.with_actual(0.8)
.with_reward(15.0);
// Update model with feedback (should not fail for mock model)
let mut mutable_model = MockMLModel::new("feedback_test".to_string(), ModelType::DQN);
let update_result = mutable_model.update_weights(&feedback).await;
assert!(update_result.is_ok());
}
#[tokio::test]
async fn test_performance_optimization_pipeline() {
let profile = crate::create_ultra_low_latency_profile();
let executor = ParallelExecutor::new(profile).expect("Failed to create executor");
let optimizer = crate::create_hft_latency_optimizer();
// Create multiple models for parallel execution
let mut models = Vec::new();
for i in 0..5 {
let mock_model = MockMLModel::new(format!("perf_test_{}", i), ModelType::DQN);
models.push(Arc::new(mock_model) as Arc<dyn MLModel>);
}
let features = Features::new(
vec![0.1, 0.2, 0.3, 0.4, 0.5],
vec!["a".to_string(), "b".to_string(), "c".to_string(), "d".to_string(), "e".to_string()]
);
// Execute parallel predictions and measure performance
let start_time = std::time::Instant::now();
let results = executor.execute_parallel_predictions(models, features).await;
let execution_time = start_time.elapsed();
assert_eq!(results.len(), 5);
// Record performance in optimizer
optimizer.record_performance(
execution_time.as_micros() as u64,
5,
1,
results.iter().all(|r| r.is_ok())
).await;
let recommendations = optimizer.get_recommendations().await;
assert!(recommendations.current_avg_latency_us > 0);
}
// ========================================================================
// Error Handling and Edge Cases
// ========================================================================
#[tokio::test]
async fn test_model_validation_failures() {
let mock_model = MockMLModel::new("validation_test".to_string(), ModelType::DQN);
let arc_model = Arc::new(mock_model) as Arc<dyn MLModel>;
// Test with empty features
let empty_features = Features::new(vec![], vec![]);
let result = arc_model.predict(&empty_features).await;
assert!(result.is_err());
// Test validation directly
let validation_result = arc_model.validate_features(&empty_features);
assert!(validation_result.is_err());
assert!(matches!(validation_result, Err(MLError::ValidationError { .. })));
}
#[tokio::test]
async fn test_registry_with_not_ready_model() {
let registry = get_global_registry();
let not_ready_model = NotReadyModel::new("not_ready".to_string());
let arc_model = Arc::new(not_ready_model) as Arc<dyn MLModel>;
let register_result = registry.register(arc_model).await;
assert!(register_result.is_err());
assert!(matches!(register_result, Err(MLError::ModelError(_))));
}
#[tokio::test]
async fn test_parallel_executor_with_timeout() {
let mut profile = HFTPerformanceProfile::default();
profile.max_latency_us = 1; // Very short timeout for testing
profile.optimization_level = OptimizationLevel::High; // Conservative mode uses timeouts
let executor = ParallelExecutor::new(profile).expect("Failed to create executor");
let slow_model = SlowModel::new("slow_model".to_string());
let models = vec![Arc::new(slow_model) as Arc<dyn MLModel>];
let features = Features::new(vec![1.0], vec!["test".to_string()]);
let results = executor.execute_parallel_predictions(models, features).await;
assert_eq!(results.len(), 1);
// Should timeout for slow model in conservative mode
// Note: Actual timeout behavior depends on implementation details
}
// ========================================================================
// Stress Tests and Performance Validation
// ========================================================================
#[tokio::test]
async fn test_high_volume_predictions() {
let registry = get_global_registry();
// Register multiple models
for i in 0..10 {
let mock_model = MockMLModel::new(format!("stress_test_{}", i), ModelType::DQN);
let arc_model = Arc::new(mock_model) as Arc<dyn MLModel>;
registry.register(arc_model).await.expect("Failed to register model");
}
// Create many feature sets
let mut feature_sets = Vec::new();
for i in 0..100 {
let features = Features::new(
vec![i as f64 / 100.0, (i as f64 / 100.0) * 2.0],
vec!["feature_1".to_string(), "feature_2".to_string()]
);
feature_sets.push(features);
}
// Execute predictions for all feature sets in parallel
let mut futures = Vec::new();
for features in feature_sets {
let registry_ref = registry.clone();
let future = async move {
registry_ref.predict_all(&features).await
};
futures.push(future);
}
let all_results = futures::future::join_all(futures).await;
assert_eq!(all_results.len(), 100);
// Verify all batches completed successfully
for batch_results in all_results {
assert_eq!(batch_results.len(), 10); // 10 models per batch
for result in batch_results {
assert!(result.is_ok());
}
}
}
#[tokio::test]
async fn test_memory_usage_tracking() {
// Test that models track memory usage correctly
let metadata = ModelMetadata::new(
ModelType::TLOB,
"1.0.0".to_string(),
100,
512.0
);
assert_eq!(metadata.memory_usage_mb, 512.0);
// Test different model types have reasonable memory usage
let models_memory = vec![
(ModelType::CompactDQN, 32.0),
(ModelType::DistilledMicroNet, 16.0),
(ModelType::DQN, 128.0),
(ModelType::MAMBA, 256.0),
(ModelType::TFT, 512.0),
(ModelType::TLOB, 256.0),
];
for (model_type, expected_memory) in models_memory {
let metadata = ModelMetadata::new(model_type, "1.0.0".to_string(), 50, expected_memory);
assert_eq!(metadata.memory_usage_mb, expected_memory);
assert!(metadata.memory_usage_mb > 0.0);
}
}
#[test]
fn test_precision_factor_constant() {
assert_eq!(PRECISION_FACTOR, 100_000_000);
// Test that precision factor provides adequate precision for financial calculations
let price_cents = 12345; // $123.45
let precise_price = price_cents * PRECISION_FACTOR;
let recovered_price = precise_price / PRECISION_FACTOR;
assert_eq!(recovered_price, price_cents);
}
#[test]
fn test_max_inference_latency_constant() {
assert_eq!(MAX_INFERENCE_LATENCY_US, 100);
// Verify the constant is reasonable for HFT requirements
assert!(MAX_INFERENCE_LATENCY_US <= 1000); // Should be sub-millisecond
assert!(MAX_INFERENCE_LATENCY_US >= 1); // Should be at least 1 microsecond
}
}
// ============================================================================
// Mock Models for Testing
// ============================================================================
/// Mock ML model implementation for testing
#[derive(Debug)]
struct MockMLModel {
name: String,
model_type: ModelType,
confidence: f64,
ready: bool,
}
impl MockMLModel {
fn new(name: String, model_type: ModelType) -> Self {
Self {
name,
model_type,
confidence: 0.8,
ready: true,
}
}
}
#[async_trait::async_trait]
impl MLModel for MockMLModel {
fn name(&self) -> &str {
&self.name
}
fn model_type(&self) -> ModelType {
self.model_type
}
async fn predict(&self, features: &Features) -> MLResult<ModelPrediction> {
if !self.ready {
return Err(MLError::ModelError("Model not ready".to_string()));
}
// Validate features
self.validate_features(features)?;
// Simple mock prediction based on feature values
let prediction_value = if !features.values.is_empty() {
features.values.iter().sum::<f64>() / features.values.len() as f64
} else {
0.5 // Default prediction
};
Ok(ModelPrediction::new(
self.name.clone(),
prediction_value,
self.confidence,
))
}
fn get_confidence(&self) -> f64 {
self.confidence
}
async fn update_weights(&mut self, feedback: &Feedback) -> MLResult<()> {
// Mock weight update - adjust confidence based on feedback
if let Some(reward) = feedback.reward {
self.confidence = (self.confidence + reward.signum() * 0.01).clamp(0.0, 1.0);
}
Ok(())
}
fn is_ready(&self) -> bool {
self.ready
}
fn get_metadata(&self) -> ModelMetadata {
ModelMetadata::new(
self.model_type,
"test-1.0.0".to_string(),
10, // features_used
64.0, // memory_usage_mb
)
}
fn validate_features(&self, features: &Features) -> MLResult<()> {
if features.values.is_empty() {
return Err(MLError::ValidationError {
message: "Empty feature vector not allowed".to_string(),
});
}
Ok(())
}
}
/// Mock model that is never ready (for testing error conditions)
#[derive(Debug)]
struct NotReadyModel {
name: String,
}
impl NotReadyModel {
fn new(name: String) -> Self {
Self { name }
}
}
#[async_trait::async_trait]
impl MLModel for NotReadyModel {
fn name(&self) -> &str {
&self.name
}
fn model_type(&self) -> ModelType {
ModelType::DQN
}
async fn predict(&self, _features: &Features) -> MLResult<ModelPrediction> {
Err(MLError::ModelError("Model not ready".to_string()))
}
fn get_confidence(&self) -> f64 {
0.0
}
fn is_ready(&self) -> bool {
false // Never ready
}
fn get_metadata(&self) -> ModelMetadata {
ModelMetadata::new(
ModelType::DQN,
"not-ready-1.0.0".to_string(),
0,
0.0,
)
}
}
/// Mock model that takes a long time to predict (for timeout testing)
#[derive(Debug)]
struct SlowModel {
name: String,
}
impl SlowModel {
fn new(name: String) -> Self {
Self { name }
}
}
#[async_trait::async_trait]
impl MLModel for SlowModel {
fn name(&self) -> &str {
&self.name
}
fn model_type(&self) -> ModelType {
ModelType::DQN
}
async fn predict(&self, _features: &Features) -> MLResult<ModelPrediction> {
// Simulate slow prediction
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
Ok(ModelPrediction::new(
self.name.clone(),
0.5,
0.3,
))
}
fn get_confidence(&self) -> f64 {
0.3
}
fn get_metadata(&self) -> ModelMetadata {
ModelMetadata::new(
ModelType::DQN,
"slow-1.0.0".to_string(),
5,
32.0,
)
}
}
// ============================================================================
// Property-Based Tests for ML Components
// ============================================================================
#[cfg(test)]
mod property_tests {
use super::*;
use proptest::prelude::*;
proptest! {
#[test]
fn test_features_properties(
values in prop::collection::vec(any::<f64>(), 0..100),
names in prop::collection::vec("[a-zA-Z0-9_]+", 0..50)
) {
let features = Features::new(values.clone(), names.clone());
// Basic properties
prop_assert_eq!(features.values, values);
prop_assert_eq!(features.names, names);
prop_assert!(features.timestamp > 0);
// Serialization roundtrip
let serialized = serde_json::to_string(&features).unwrap();
let deserialized: Features = serde_json::from_str(&serialized).unwrap();
prop_assert_eq!(features.values, deserialized.values);
prop_assert_eq!(features.names, deserialized.names);
}
#[test]
fn test_model_prediction_properties(
model_id in "[a-zA-Z0-9_]+",
value in any::<f64>(),
confidence in 0.0..1.0f64
) {
let prediction = ModelPrediction::new(model_id.clone(), value, confidence);
prop_assert_eq!(prediction.model_id, model_id);
prop_assert_eq!(prediction.value, value);
prop_assert_eq!(prediction.confidence, confidence);
prop_assert!(prediction.timestamp > 0);
// Confidence should be in valid range
prop_assert!(confidence >= 0.0 && confidence <= 1.0);
}
#[test]
fn test_feedback_properties(
actual_value in proptest::option::of(any::<f64>()),
reward in proptest::option::of(any::<f64>())
) {
let mut feedback = Feedback::new();
if let Some(actual) = actual_value {
feedback = feedback.with_actual(actual);
}
if let Some(r) = reward {
feedback = feedback.with_reward(r);
}
prop_assert_eq!(feedback.actual_value, actual_value);
prop_assert_eq!(feedback.reward, reward);
prop_assert!(feedback.timestamp > 0);
}
}
}
// ============================================================================
// Benchmark Tests (for manual performance testing)
// ============================================================================
#[cfg(test)]
mod benchmark_tests {
use super::*;
use std::time::Instant;
#[tokio::test]
#[ignore] // Use --ignored to run benchmark tests
async fn benchmark_model_prediction_throughput() {
let mock_model = MockMLModel::new("benchmark_test".to_string(), ModelType::DQN);
let arc_model = Arc::new(mock_model) as Arc<dyn MLModel>;
let features = Features::new(
vec![1.0, 2.0, 3.0, 4.0, 5.0],
vec!["a".to_string(), "b".to_string(), "c".to_string(), "d".to_string(), "e".to_string()]
);
let iterations = 10000;
let start_time = Instant::now();
for _ in 0..iterations {
let _prediction = arc_model.predict(&features).await.expect("Prediction failed");
}
let duration = start_time.elapsed();
let predictions_per_sec = iterations as f64 / duration.as_secs_f64();
println!("Model prediction throughput: {:.0} predictions/sec", predictions_per_sec);
assert!(predictions_per_sec > 1000.0); // Should handle at least 1000 predictions/sec
}
#[tokio::test]
#[ignore] // Use --ignored to run benchmark tests
async fn benchmark_registry_parallel_predictions() {
let registry = get_global_registry();
// Register multiple models
for i in 0..10 {
let mock_model = MockMLModel::new(format!("parallel_bench_{}", i), ModelType::DQN);
let arc_model = Arc::new(mock_model) as Arc<dyn MLModel>;
registry.register(arc_model).await.expect("Failed to register model");
}
let features = Features::new(
vec![1.0, 2.0, 3.0],
vec!["x".to_string(), "y".to_string(), "z".to_string()]
);
let iterations = 1000;
let start_time = Instant::now();
for _ in 0..iterations {
let _results = registry.predict_all(&features).await;
}
let duration = start_time.elapsed();
let parallel_batches_per_sec = iterations as f64 / duration.as_secs_f64();
println!("Parallel prediction batches: {:.0} batches/sec", parallel_batches_per_sec);
println!("Total predictions: {:.0} predictions/sec", parallel_batches_per_sec * 10.0);
assert!(parallel_batches_per_sec > 100.0); // Should handle at least 100 batches/sec
}
#[tokio::test]
#[ignore] // Use --ignored to run benchmark tests
async fn benchmark_executor_latency() {
let profile = crate::create_ultra_low_latency_profile();
let executor = ParallelExecutor::new(profile).expect("Failed to create executor");
let models = vec![
Arc::new(MockMLModel::new("latency_test_1".to_string(), ModelType::CompactDQN)) as Arc<dyn MLModel>,
Arc::new(MockMLModel::new("latency_test_2".to_string(), ModelType::DistilledMicroNet)) as Arc<dyn MLModel>,
];
let features = Features::new(vec![1.0, 2.0], vec!["a".to_string(), "b".to_string()]);
let iterations = 1000;
let mut total_latency_us = 0u64;
for _ in 0..iterations {
let start_time = Instant::now();
let _results = executor.execute_parallel_predictions(models.clone(), features.clone()).await;
let latency = start_time.elapsed();
total_latency_us += latency.as_micros() as u64;
}
let avg_latency_us = total_latency_us / iterations as u64;
println!("Average parallel execution latency: {}μs", avg_latency_us);
// For mock models, should achieve low latency
assert!(avg_latency_us < 1000); // Less than 1ms average
}
}