//! Model Registry Integration Tests //! //! Comprehensive tests for the ML model versioning and registry system. use ml::model_registry::{ModelRegistry, ModelVersionMetadata}; use ml::ModelType; // Test database URL (requires PostgreSQL running) const TEST_DB_URL: &str = "postgresql://foxhunt:foxhunt_dev_password@localhost:5432/foxhunt"; const TEST_S3_PATH: &str = "s3://foxhunt-ml-models-test/"; #[tokio::test] #[ignore] // Requires PostgreSQL async fn test_registry_initialization() { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await; assert!(registry.is_ok(), "Failed to initialize registry: {:?}", registry.err()); } #[tokio::test] #[ignore] // Requires PostgreSQL async fn test_register_and_retrieve_model() { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await.unwrap(); // Create metadata let mut metadata = ModelVersionMetadata::new( format!("dqn-test-{}", uuid::Uuid::new_v4()), ModelType::DQN, "1.0.0".to_string(), "test_data".to_string(), "s3://test/dqn/1.0.0/".to_string(), ); metadata.add_hyperparameter("epochs", serde_json::json!(500)); metadata.add_metric("final_loss", serde_json::json!(0.001)); metadata.set_checksum("sha256:test123".to_string()); let model_id = metadata.model_id.clone(); // Register registry.register_version(&metadata).await.unwrap(); // Retrieve let retrieved = registry.get_model_by_version(&model_id).await.unwrap(); assert_eq!(retrieved.model_id, model_id); assert_eq!(retrieved.version, "1.0.0"); assert_eq!(retrieved.data_source, "test_data"); assert_eq!(retrieved.checksum, "sha256:test123"); } #[tokio::test] #[ignore] // Requires PostgreSQL async fn test_hyperparameters_and_metrics() { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await.unwrap(); let mut metadata = ModelVersionMetadata::new( format!("tft-test-{}", uuid::Uuid::new_v4()), ModelType::TFT, "2.0.0".to_string(), "test_data".to_string(), "s3://test/tft/2.0.0/".to_string(), ); // Add multiple hyperparameters metadata.add_hyperparameter("epochs", serde_json::json!(1000)); metadata.add_hyperparameter("batch_size", serde_json::json!(256)); metadata.add_hyperparameter("learning_rate", serde_json::json!(0.0001)); metadata.add_hyperparameter("dropout", serde_json::json!(0.2)); // Add multiple metrics metadata.add_metric("final_loss", serde_json::json!(0.0005)); metadata.add_metric("validation_loss", serde_json::json!(0.0008)); metadata.add_metric("sharpe_ratio", serde_json::json!(2.5)); metadata.add_metric("max_drawdown", serde_json::json!(0.15)); metadata.set_checksum("sha256:tft456".to_string()); let model_id = metadata.model_id.clone(); // Register registry.register_version(&metadata).await.unwrap(); // Retrieve and verify let retrieved = registry.get_model_by_version(&model_id).await.unwrap(); // Verify hyperparameters let hyperparams = retrieved.hyperparameters.as_object().unwrap(); assert_eq!(hyperparams.get("epochs").unwrap(), &serde_json::json!(1000)); assert_eq!(hyperparams.get("batch_size").unwrap(), &serde_json::json!(256)); // Verify metrics let metrics = retrieved.metrics.as_object().unwrap(); assert_eq!(metrics.get("final_loss").unwrap(), &serde_json::json!(0.0005)); assert_eq!(metrics.get("sharpe_ratio").unwrap(), &serde_json::json!(2.5)); } #[tokio::test] #[ignore] // Requires PostgreSQL async fn test_production_tagging() { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await.unwrap(); let metadata = ModelVersionMetadata::new( format!("mamba-test-{}", uuid::Uuid::new_v4()), ModelType::MAMBA, "1.0.0".to_string(), "test_data".to_string(), "s3://test/mamba/1.0.0/".to_string(), ); let model_id = metadata.model_id.clone(); // Register as experimental registry.register_version(&metadata).await.unwrap(); // Verify experimental let retrieved = registry.get_model_by_version(&model_id).await.unwrap(); assert!(retrieved.is_experimental); assert!(!retrieved.is_production); // Promote to production registry.mark_production(&model_id).await.unwrap(); // Verify production let retrieved = registry.get_model_by_version(&model_id).await.unwrap(); assert!(retrieved.is_production); assert!(!retrieved.is_experimental); } #[tokio::test] #[ignore] // Requires PostgreSQL async fn test_get_production_models() { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await.unwrap(); // Register a production model let mut metadata = ModelVersionMetadata::new( format!("dqn-prod-{}", uuid::Uuid::new_v4()), ModelType::DQN, "1.0.0".to_string(), "test_data".to_string(), "s3://test/dqn/1.0.0/".to_string(), ); metadata.set_checksum("sha256:prod123".to_string()); let model_id = metadata.model_id.clone(); registry.register_version(&metadata).await.unwrap(); registry.mark_production(&model_id).await.unwrap(); // Query production models let production_models = registry.get_production_models().await.unwrap(); // Verify at least one production model exists assert!(!production_models.is_empty()); // Verify all returned models are production for model in &production_models { assert!(model.is_production); assert!(!model.is_archived); } } #[tokio::test] #[ignore] // Requires PostgreSQL async fn test_get_models_by_type() { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await.unwrap(); // Register multiple PPO models for i in 0..3 { let metadata = ModelVersionMetadata::new( format!("ppo-test-{}-{}", i, uuid::Uuid::new_v4()), ModelType::PPO, format!("1.0.{}", i), "test_data".to_string(), format!("s3://test/ppo/1.0.{}/", i), ); registry.register_version(&metadata).await.unwrap(); } // Query PPO models let ppo_models = registry.get_models_by_type(ModelType::PPO).await.unwrap(); // Verify at least 3 PPO models exist assert!(ppo_models.len() >= 3); // Verify all are PPO for model in &ppo_models { assert_eq!(model.model_type, ModelType::PPO); } } #[tokio::test] #[ignore] // Requires PostgreSQL async fn test_archive_model() { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await.unwrap(); let metadata = ModelVersionMetadata::new( format!("tlob-archive-{}", uuid::Uuid::new_v4()), ModelType::TLOB, "0.9.0".to_string(), "test_data".to_string(), "s3://test/tlob/0.9.0/".to_string(), ); let model_id = metadata.model_id.clone(); // Register registry.register_version(&metadata).await.unwrap(); // Archive registry.archive_model(&model_id).await.unwrap(); // Verify archived let retrieved = registry.get_model_by_version(&model_id).await.unwrap(); assert!(retrieved.is_archived); } #[tokio::test] #[ignore] // Requires PostgreSQL async fn test_get_registry_statistics() { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await.unwrap(); // Get statistics let stats = registry.get_statistics().await.unwrap(); // Verify basic stats structure assert!(stats.total_count >= 0); assert!(stats.production_count <= stats.total_count); assert!(stats.experimental_count <= stats.total_count); assert!(stats.archived_count <= stats.total_count); assert!(stats.model_types_count >= 0); } #[tokio::test] #[ignore] // Requires PostgreSQL async fn test_date_range_query() { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await.unwrap(); // Register a model let metadata = ModelVersionMetadata::new( format!("transformer-test-{}", uuid::Uuid::new_v4()), ModelType::Transformer, "1.0.0".to_string(), "test_data".to_string(), "s3://test/transformer/1.0.0/".to_string(), ); registry.register_version(&metadata).await.unwrap(); // Query last 24 hours let now = chrono::Utc::now(); let one_day_ago = now - chrono::Duration::days(1); let recent_models = registry.get_models_by_date_range(one_day_ago, now).await.unwrap(); // Should find at least the model we just registered assert!(!recent_models.is_empty()); } #[tokio::test] #[ignore] // Requires PostgreSQL async fn test_model_not_found_error() { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await.unwrap(); // Try to retrieve non-existent model let result = registry.get_model_by_version("nonexistent-model-xyz").await; assert!(result.is_err()); } #[tokio::test] #[ignore] // Requires PostgreSQL async fn test_update_model_metadata() { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await.unwrap(); // Register initial version let mut metadata = ModelVersionMetadata::new( format!("ensemble-test-{}", uuid::Uuid::new_v4()), ModelType::Ensemble, "1.0.0".to_string(), "test_data_v1".to_string(), "s3://test/ensemble/1.0.0/".to_string(), ); metadata.add_metric("accuracy", serde_json::json!(0.85)); metadata.set_checksum("sha256:v1".to_string()); let model_id = metadata.model_id.clone(); registry.register_version(&metadata).await.unwrap(); // Update with new data let mut updated_metadata = metadata.clone(); updated_metadata.data_source = "test_data_v2".to_string(); updated_metadata.add_metric("accuracy", serde_json::json!(0.90)); updated_metadata.set_checksum("sha256:v2".to_string()); registry.register_version(&updated_metadata).await.unwrap(); // Verify update let retrieved = registry.get_model_by_version(&model_id).await.unwrap(); assert_eq!(retrieved.data_source, "test_data_v2"); assert_eq!(retrieved.checksum, "sha256:v2"); } #[tokio::test] #[ignore] // Requires PostgreSQL async fn test_multiple_model_types() { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await.unwrap(); let model_types = vec![ ModelType::DQN, ModelType::MAMBA, ModelType::TFT, ModelType::PPO, ModelType::TLOB, ModelType::Transformer, ]; // Register one of each type for model_type in model_types { let metadata = ModelVersionMetadata::new( format!("{:?}-multi-{}", model_type, uuid::Uuid::new_v4()), model_type, "1.0.0".to_string(), "test_data".to_string(), format!("s3://test/{:?}/1.0.0/", model_type), ); registry.register_version(&metadata).await.unwrap(); } // Verify statistics let stats = registry.get_statistics().await.unwrap(); assert!(stats.model_types_count >= 6); } #[tokio::test] #[ignore] // Requires PostgreSQL async fn test_cache_functionality() { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await.unwrap(); let metadata = ModelVersionMetadata::new( format!("cache-test-{}", uuid::Uuid::new_v4()), ModelType::DQN, "1.0.0".to_string(), "test_data".to_string(), "s3://test/cache/1.0.0/".to_string(), ); let model_id = metadata.model_id.clone(); registry.register_version(&metadata).await.unwrap(); // First retrieval (from database) let start1 = std::time::Instant::now(); let _ = registry.get_model_by_version(&model_id).await.unwrap(); let duration1 = start1.elapsed(); // Second retrieval (from cache, should be faster) let start2 = std::time::Instant::now(); let _ = registry.get_model_by_version(&model_id).await.unwrap(); let duration2 = start2.elapsed(); // Cache should be faster (not guaranteed but likely) println!("First retrieval: {:?}", duration1); println!("Second retrieval (cached): {:?}", duration2); }