//! Comprehensive tests for the core InferenceEngine //! //! This test suite validates the complete inference engine functionality //! including model loading, graph optimization, batch processing, and //! integration with all subsystems. use futures::future; use rtx_inference::*; use rtx_tensor::{DType, Device, Shape, Tensor}; use std::collections::HashMap; use std::sync::Arc; use std::time::{Duration, Instant}; use tokio::sync::Mutex; use tokio_stream::StreamExt; use uuid::Uuid; /// Mock model weights for testing #[derive(Debug, Clone)] struct MockModelWeights { layers: HashMap, config: ModelConfig, } fn create_test_config() -> InferenceEngineConfig { InferenceEngineConfig { model_path: "test_model.safetensors".to_string(), device: Device::cpu(), max_batch_size: 32, max_sequence_length: 2048, memory_pool_size: 1024 * 1024 * 1024, // 1GB optimization_level: 2, enable_caching: true, enable_quantization: false, request_manager_config: RequestManagerConfig::default(), scheduler_config: BatchSchedulerConfig::default(), cache_config: KvCacheConfig::default(), } } impl MockModelWeights { fn new() -> InferenceResult { let config = ModelConfig { vocab_size: 32000, hidden_size: 4096, num_layers: 32, num_heads: 32, max_position_embeddings: 2048, layer_norm_epsilon: 1e-6, }; let mut layers = HashMap::new(); // Create mock embedding layer let embedding_weight = Tensor::randn(&[config.vocab_size, config.hidden_size], &Device::cpu())?; layers.insert("embedding.weight".to_string(), embedding_weight); // Create mock transformer layers for layer_idx in 0..config.num_layers { // Attention weights let attn_q = Tensor::randn(&[config.hidden_size, config.hidden_size], &Device::cpu())?; let attn_k = Tensor::randn(&[config.hidden_size, config.hidden_size], &Device::cpu())?; let attn_v = Tensor::randn(&[config.hidden_size, config.hidden_size], &Device::cpu())?; let attn_o = Tensor::randn(&[config.hidden_size, config.hidden_size], &Device::cpu())?; layers.insert( format!("layers.{}.attention.q_proj.weight", layer_idx), attn_q, ); layers.insert( format!("layers.{}.attention.k_proj.weight", layer_idx), attn_k, ); layers.insert( format!("layers.{}.attention.v_proj.weight", layer_idx), attn_v, ); layers.insert( format!("layers.{}.attention.o_proj.weight", layer_idx), attn_o, ); // MLP weights let mlp_gate = Tensor::randn( &[config.hidden_size, config.hidden_size * 4], &Device::cpu(), )?; let mlp_up = Tensor::randn( &[config.hidden_size, config.hidden_size * 4], &Device::cpu(), )?; let mlp_down = Tensor::randn( &[config.hidden_size * 4, config.hidden_size], &Device::cpu(), )?; layers.insert( format!("layers.{}.mlp.gate_proj.weight", layer_idx), mlp_gate, ); layers.insert(format!("layers.{}.mlp.up_proj.weight", layer_idx), mlp_up); layers.insert( format!("layers.{}.mlp.down_proj.weight", layer_idx), mlp_down, ); // Layer norm weights let input_layernorm = Tensor::ones(&[config.hidden_size], &Device::cpu())?; let post_attention_layernorm = Tensor::ones(&[config.hidden_size], &Device::cpu())?; layers.insert( format!("layers.{}.input_layernorm.weight", layer_idx), input_layernorm, ); layers.insert( format!("layers.{}.post_attention_layernorm.weight", layer_idx), post_attention_layernorm, ); } // Output layer let lm_head = Tensor::randn(&[config.hidden_size, config.vocab_size], &Device::cpu())?; layers.insert("lm_head.weight".to_string(), lm_head); Ok(Self { layers, config }) } fn get_layer(&self, name: &str) -> Option<&Tensor> { self.layers.get(name) } fn layer_count(&self) -> usize { self.layers.len() } } // Core InferenceEngine tests - these should fail initially #[tokio::test] async fn test_inference_engine_creation() { let config = create_test_config(); let engine = InferenceEngine::new(config.clone()).await; assert!( engine.is_ok(), "Failed to create inference engine: {:?}", engine.err() ); let engine = engine.unwrap(); assert_eq!(engine.device(), &config.device); assert_eq!(engine.max_batch_size(), config.max_batch_size); } #[tokio::test] async fn test_model_loading_and_registration() { let config = create_test_config(); let mut engine = InferenceEngine::new(config).await.unwrap(); // Create mock model weights let mock_weights = MockModelWeights::new().unwrap(); // Load model should fail initially (not implemented) let result = engine .load_model("test_model", &mock_weights.layers, &mock_weights.config) .await; assert!(result.is_ok(), "Model loading should succeed"); // Check model is registered let models = engine.list_models().await; assert!(models.contains(&"test_model".to_string())); // Get model info let info = engine.get_model_info("test_model").await.unwrap(); assert_eq!(info.name, "test_model"); assert_eq!(info.layer_count, mock_weights.layer_count()); assert!(info.memory_usage > 0); } #[tokio::test] async fn test_single_inference_request() { let config = create_test_config(); let mut engine = InferenceEngine::new(config).await.unwrap(); // Load model let mock_weights = MockModelWeights::new().unwrap(); engine .load_model("test_model", &mock_weights.layers, &mock_weights.config) .await .unwrap(); // Create inference request let input_tokens = vec![1, 2, 3, 4, 5]; let request = InferenceRequest::new("test_model".to_string(), input_tokens.clone(), 10); let request_id = request.id.clone(); // Run inference let start_time = Instant::now(); let result = engine.infer(request).await; let inference_time = start_time.elapsed(); assert!(result.is_ok(), "Inference should succeed"); let result = result.unwrap(); // Validate output assert_eq!(result.request_id, request_id); assert!(!result.output_tokens.is_empty()); assert!(result.output_tokens.len() <= 10); assert!(matches!( result.finish_reason, FinishReason::Length | FinishReason::EndOfSequence )); // Validate performance metrics let metrics = result.metrics.unwrap(); assert!(metrics.total_time > Duration::ZERO); assert!(metrics.generation_time > Duration::ZERO); assert_eq!(metrics.input_token_count, input_tokens.len()); assert_eq!(metrics.output_token_count, result.output_tokens.len()); assert!(metrics.tokens_per_second > 0.0); println!( "Inference completed in {:?}, {:.2} tokens/sec", inference_time, metrics.tokens_per_second ); } #[tokio::test] async fn test_batch_inference() { let mut config = create_test_config(); config.max_batch_size = 4; let mut engine = InferenceEngine::new(config).await.unwrap(); // Load model let mock_weights = MockModelWeights::new().unwrap(); engine .load_model("test_model", &mock_weights.layers, &mock_weights.config) .await .unwrap(); // Create multiple requests let requests = vec![ InferenceRequest::new("test_model".to_string(), vec![1, 2, 3], 5), InferenceRequest::new("test_model".to_string(), vec![4, 5, 6, 7], 5), InferenceRequest::new("test_model".to_string(), vec![8, 9], 5), InferenceRequest::new("test_model".to_string(), vec![10, 11, 12, 13, 14], 5), ]; // Submit all requests concurrently let mut handles = vec![]; for request in requests { let engine_clone = engine.clone(); let handle = tokio::spawn(async move { engine_clone.infer(request).await }); handles.push(handle); } // Wait for all results let results: Vec<_> = future::join_all(handles).await; // Validate all succeeded for result in results { let inference_result = result.unwrap().unwrap(); assert!(!inference_result.output_tokens.is_empty()); assert!(inference_result.metrics.is_some()); } } #[tokio::test] async fn test_streaming_inference() { let config = create_test_config(); let mut engine = InferenceEngine::new(config).await.unwrap(); // Load model let mock_weights = MockModelWeights::new().unwrap(); engine .load_model("test_model", &mock_weights.layers, &mock_weights.config) .await .unwrap(); // Create streaming request let mut request = InferenceRequest::new("test_model".to_string(), vec![1, 2, 3], 20); request.stream = true; // Start streaming inference let mut stream = engine.infer_stream(request).await.unwrap(); let mut token_count = 0; let mut total_time = Duration::ZERO; // Collect streaming tokens while let Some(token_result) = stream.next().await { assert!(token_result.is_ok()); let token = token_result.unwrap(); assert!(token.id > 0); assert!(!token.text.is_empty()); assert!(token.probability > 0.0); token_count += 1; total_time += token.generation_time; if token_count >= 20 { break; } } assert!(token_count > 0); assert!(total_time > Duration::ZERO); println!( "Generated {} tokens via streaming in {:?}", token_count, total_time ); } #[tokio::test] async fn test_memory_management() { let mut config = create_test_config(); config.memory_pool_size = 100 * 1024 * 1024; // 100MB let mut engine = InferenceEngine::new(config).await.unwrap(); // Load model let mock_weights = MockModelWeights::new().unwrap(); engine .load_model("test_model", &mock_weights.layers, &mock_weights.config) .await .unwrap(); // Get initial memory stats let initial_stats = engine.get_memory_stats().await; assert!(initial_stats.total_allocated > 0); assert!(initial_stats.total_allocated <= 100 * 1024 * 1024); // Run inference to allocate memory let request = InferenceRequest::new("test_model".to_string(), vec![1, 2, 3], 10); let _result = engine.infer(request).await.unwrap(); // Check memory usage increased let after_inference_stats = engine.get_memory_stats().await; assert!(after_inference_stats.total_allocated >= initial_stats.total_allocated); // Force garbage collection engine.gc().await.unwrap(); // Memory should be reclaimed let after_gc_stats = engine.get_memory_stats().await; assert!(after_gc_stats.total_allocated <= after_inference_stats.total_allocated); } #[tokio::test] async fn test_graph_optimization() { let mut config = create_test_config(); config.optimization_level = 3; // Maximum optimization let mut engine = InferenceEngine::new(config).await.unwrap(); // Load model let mock_weights = MockModelWeights::new().unwrap(); engine .load_model("test_model", &mock_weights.layers, &mock_weights.config) .await .unwrap(); // Get optimization stats let stats = engine.get_optimization_stats("test_model").await.unwrap(); assert!(stats.optimizations_applied > 0); assert!(stats.nodes_eliminated > 0); assert!(stats.memory_saved > 0); assert!(stats.estimated_speedup > 1.0); // Compare optimized vs unoptimized performance let request = InferenceRequest::new("test_model".to_string(), vec![1, 2, 3, 4, 5], 10); // Measure with optimization let start = Instant::now(); let _optimized_result = engine.infer(request.clone()).await.unwrap(); let optimized_time = start.elapsed(); // Disable optimization temporarily engine .set_optimization_level("test_model", 0) .await .unwrap(); // Measure without optimization let start = Instant::now(); let _unoptimized_result = engine.infer(request.clone()).await.unwrap(); let unoptimized_time = start.elapsed(); // Optimized should be faster assert!( optimized_time < unoptimized_time, "Optimized inference ({:?}) should be faster than unoptimized ({:?})", optimized_time, unoptimized_time ); } #[tokio::test] async fn test_quantization_integration() { let mut config = create_test_config(); config.enable_quantization = true; let mut engine = InferenceEngine::new(config).await.unwrap(); // Load model let mock_weights = MockModelWeights::new().unwrap(); engine .load_model("test_model", &mock_weights.layers, &mock_weights.config) .await .unwrap(); // Apply quantization let quant_config = QuantizationConfig::new() .with_scheme(QuantizationScheme::INT8) .with_calibration_samples(100); engine .quantize_model("test_model", quant_config) .await .unwrap(); // Check quantization was applied let model_info = engine.get_model_info("test_model").await.unwrap(); assert!(model_info.is_quantized); assert!(model_info.memory_usage < model_info.original_memory_usage); // Inference should still work let request = InferenceRequest::new("test_model".to_string(), vec![1, 2, 3], 5); let result = engine.infer(request).await.unwrap(); assert!(!result.output_tokens.is_empty()); } #[tokio::test] async fn test_performance_monitoring() { let config = create_test_config(); let mut engine = InferenceEngine::new(config).await.unwrap(); // Load model let mock_weights = MockModelWeights::new().unwrap(); engine .load_model("test_model", &mock_weights.layers, &mock_weights.config) .await .unwrap(); // Run multiple inferences to generate metrics for i in 0..10 { let request = InferenceRequest::new("test_model".to_string(), vec![1, 2, i as i32], 5); let _result = engine.infer(request).await.unwrap(); } // Get performance metrics let metrics = engine.get_performance_metrics().await; assert!(metrics.total_requests >= 10); assert!(metrics.total_tokens_generated > 0); assert!(metrics.average_latency > Duration::ZERO); assert!(metrics.average_throughput > 0.0); assert!(metrics.success_rate >= 0.9); // Should be high // Get model-specific metrics let model_metrics = engine.get_model_metrics("test_model").await.unwrap(); assert!(model_metrics.inference_count >= 10); assert!(model_metrics.total_inference_time > Duration::ZERO); assert!(model_metrics.average_tokens_per_request > 0.0); } #[tokio::test] async fn test_error_handling() { let config = create_test_config(); let mut engine = InferenceEngine::new(config).await.unwrap(); // Test inference without model let request = InferenceRequest::new("nonexistent_model".to_string(), vec![1, 2, 3], 5); let result = engine.infer(request).await; assert!(result.is_err()); assert!(matches!( result.unwrap_err(), InferenceError::ModelNotFound { .. } )); // Load model let mock_weights = MockModelWeights::new().unwrap(); engine .load_model("test_model", &mock_weights.layers, &mock_weights.config) .await .unwrap(); // Test with invalid input let invalid_request = InferenceRequest::new("test_model".to_string(), vec![], 5); // Empty input let result = engine.infer(invalid_request).await; assert!(result.is_err()); assert!(matches!( result.unwrap_err(), InferenceError::InvalidRequest { .. } )); // Test with too large input let large_input = vec![1i32; 10000]; // Exceed max sequence length let large_request = InferenceRequest::new("test_model".to_string(), large_input, 5); let result = engine.infer(large_request).await; assert!(result.is_err()); assert!(matches!( result.unwrap_err(), InferenceError::InvalidRequest { .. } )); } #[tokio::test] async fn test_concurrent_inference() { let mut config = create_test_config(); config.max_batch_size = 8; let mut engine_temp = InferenceEngine::new(config).await.unwrap(); // Load model let mock_weights = MockModelWeights::new().unwrap(); engine_temp .load_model("test_model", &mock_weights.layers, &mock_weights.config) .await .unwrap(); let engine = Arc::new(engine_temp); // Run concurrent inferences let mut handles = vec![]; for i in 0..20 { let engine_clone = engine.clone(); let handle = tokio::spawn(async move { let request = InferenceRequest::new("test_model".to_string(), vec![i, i + 1, i + 2], 3); engine_clone.infer(request).await }); handles.push(handle); } // Wait for all to complete let results = future::join_all(handles).await; // Check all succeeded let mut success_count = 0; for result in results { if result.unwrap().is_ok() { success_count += 1; } } assert_eq!( success_count, 20, "All concurrent inferences should succeed" ); // Check performance metrics show concurrent processing let metrics = engine.get_performance_metrics().await; assert!(metrics.concurrent_requests_handled > 0); assert!(metrics.peak_batch_size > 1); } #[tokio::test] async fn test_model_version_management() { let config = create_test_config(); let mut engine = InferenceEngine::new(config).await.unwrap(); // Load model version 1 let mock_weights_v1 = MockModelWeights::new().unwrap(); engine .load_model_version( "test_model", "v1.0", &mock_weights_v1.layers, &mock_weights_v1.config, ) .await .unwrap(); // Load model version 2 let mock_weights_v2 = MockModelWeights::new().unwrap(); engine .load_model_version( "test_model", "v2.0", &mock_weights_v2.layers, &mock_weights_v2.config, ) .await .unwrap(); // List versions let versions = engine.list_model_versions("test_model").await.unwrap(); assert!(versions.contains(&"v1.0".to_string())); assert!(versions.contains(&"v2.0".to_string())); // Switch to specific version engine .switch_model_version("test_model", "v1.0") .await .unwrap(); let current_version = engine .get_current_model_version("test_model") .await .unwrap(); assert_eq!(current_version, "v1.0"); // Test inference works with specific version let request = InferenceRequest::new("test_model".to_string(), vec![1, 2, 3], 5); let result = engine.infer(request).await.unwrap(); assert!(!result.output_tokens.is_empty()); } #[tokio::test] async fn test_health_checks() { let config = create_test_config(); let mut engine = InferenceEngine::new(config).await.unwrap(); // Initial health should be ok let health = engine.health_check().await; assert!(health.is_healthy); assert_eq!(health.status, "Ready"); // Load model let mock_weights = MockModelWeights::new().unwrap(); engine .load_model("test_model", &mock_weights.layers, &mock_weights.config) .await .unwrap(); // Health should still be good let health = engine.health_check().await; assert!(health.is_healthy); assert!(health.models_loaded > 0); // Detailed health check let detailed_health = engine.detailed_health_check().await; assert!(detailed_health.is_healthy); assert!(detailed_health.memory_usage.total_allocated > 0); assert!(detailed_health.memory_usage.utilization < 1.0); assert!(detailed_health.model_health.contains_key("test_model")); } // Integration tests combining multiple features #[tokio::test] async fn test_full_pipeline_integration() { let mut config = create_test_config(); config.max_batch_size = 4; config.enable_caching = true; config.enable_quantization = true; config.optimization_level = 2; let mut engine = InferenceEngine::new(config).await.unwrap(); // Load and optimize model let mock_weights = MockModelWeights::new().unwrap(); engine .load_model("test_model", &mock_weights.layers, &mock_weights.config) .await .unwrap(); // Apply quantization let quant_config = QuantizationConfig::new().with_scheme(QuantizationScheme::INT8); engine .quantize_model("test_model", quant_config) .await .unwrap(); // Run inference workload let mut results = vec![]; for i in 0..10 { let request = InferenceRequest::new("test_model".to_string(), vec![i, i + 1, i + 2], 8); let result = engine.infer(request).await.unwrap(); results.push(result); } // Validate all succeeded assert_eq!(results.len(), 10); for result in &results { assert!(!result.output_tokens.is_empty()); assert!(result.metrics.is_some()); } // Check performance metrics let metrics = engine.get_performance_metrics().await; assert!(metrics.total_requests >= 10); assert!(metrics.success_rate >= 0.9); // Validate memory efficiency let memory_stats = engine.get_memory_stats().await; assert!(memory_stats.fragmentation < 0.2); // Low fragmentation assert!(memory_stats.utilization > 0.1); // Some memory used } // Note: TensorError -> InferenceError conversion is handled internally by the crate