#![allow( clippy::assertions_on_constants, clippy::assertions_on_result_states, clippy::clone_on_copy, clippy::decimal_literal_representation, clippy::doc_markdown, clippy::empty_line_after_doc_comments, clippy::field_reassign_with_default, clippy::get_unwrap, clippy::identity_op, clippy::inconsistent_digit_grouping, clippy::indexing_slicing, clippy::integer_division, clippy::len_zero, clippy::let_underscore_must_use, clippy::manual_div_ceil, clippy::manual_let_else, clippy::manual_range_contains, clippy::modulo_arithmetic, clippy::needless_range_loop, clippy::non_ascii_literal, clippy::redundant_clone, clippy::shadow_reuse, clippy::shadow_same, clippy::shadow_unrelated, clippy::single_match_else, clippy::str_to_string, clippy::string_slice, clippy::tests_outside_test_module, clippy::too_many_lines, clippy::unnecessary_wraps, clippy::unseparated_literal_suffix, clippy::use_debug, clippy::useless_vec, clippy::wildcard_enum_match_arm, clippy::else_if_without_else, clippy::expect_used, clippy::missing_const_for_fn, clippy::similar_names, clippy::type_complexity, clippy::collapsible_else_if, clippy::doc_lazy_continuation, clippy::items_after_test_module, clippy::map_clone, clippy::multiple_unsafe_ops_per_block, clippy::unwrap_or_default, clippy::assign_op_pattern, clippy::needless_borrow, clippy::println_empty_string, clippy::unnecessary_cast, clippy::used_underscore_binding, clippy::create_dir, clippy::implicit_saturating_sub, clippy::exit, clippy::expect_fun_call, clippy::too_many_arguments, clippy::unnecessary_map_or, clippy::unwrap_used, dead_code, unused_imports, unused_variables, clippy::cloned_ref_to_slice_refs, clippy::neg_multiply, clippy::while_let_loop, clippy::bool_assert_comparison, clippy::excessive_precision, clippy::trivially_copy_pass_by_ref, clippy::op_ref, clippy::redundant_closure, clippy::unnecessary_lazy_evaluations, clippy::if_then_some_else_none, clippy::unnecessary_to_owned, clippy::single_component_path_imports, )] //! Model Registry Checkpoint Integration Tests //! //! TDD tests for checkpoint versioning, metadata tracking, and production model registration. //! Wave 10 Agent 10.8 - Training → Paper Trading Integration use chrono::Utc; use ml::model_registry::{ModelRegistry, ModelVersionMetadata}; use ml::{MLResult, ModelType}; use std::path::PathBuf; // Test database URL const TEST_DB_URL: &str = "postgresql://foxhunt:foxhunt_dev_password@localhost:5432/foxhunt"; const TEST_S3_PATH: &str = "s3://foxhunt-ml-models-test/"; /// Test 1: Register trained DQN model with checkpoint path #[tokio::test] #[ignore = "Requires PostgreSQL"] async fn test_register_dqn_checkpoint() -> MLResult<()> { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?; // Find latest DQN checkpoint let checkpoint_path = PathBuf::from( "/home/jgrusewski/Work/foxhunt/ml/trained_models/production/dqn/dqn_epoch_30.safetensors", ); let mut metadata = ModelVersionMetadata::new( "dqn-production-v1.0.0".to_string(), ModelType::DQN, "1.0.0".to_string(), "ES.FUT_2024_Q4".to_string(), "s3://foxhunt-ml-models/dqn/1.0.0/".to_string(), ); // Add checkpoint metadata metadata.add_hyperparameter("epochs", serde_json::json!(30)); metadata.add_hyperparameter("batch_size", serde_json::json!(128)); metadata.add_hyperparameter("learning_rate", serde_json::json!(0.0001)); metadata.add_metric("final_loss", serde_json::json!(0.0342)); metadata.add_metric("validation_loss", serde_json::json!(0.0356)); metadata.add_metadata( "checkpoint_path", checkpoint_path.to_string_lossy().to_string(), ); metadata.add_metadata("training_duration_hours", "2.5".to_string()); metadata.set_checksum("sha256:dqn_epoch_30_checksum".to_string()); // Register registry.register_version(&metadata).await?; // Verify retrieval let retrieved = registry .get_model_by_version("dqn-production-v1.0.0") .await?; assert_eq!(retrieved.model_id, "dqn-production-v1.0.0"); assert_eq!(retrieved.model_type, ModelType::DQN); assert_eq!(retrieved.version, "1.0.0"); assert!(retrieved.metadata.contains_key("checkpoint_path")); Ok(()) } /// Test 2: Register trained PPO model with actor-critic checkpoints #[tokio::test] #[ignore = "Requires PostgreSQL"] async fn test_register_ppo_checkpoint() -> MLResult<()> { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?; let actor_checkpoint = PathBuf::from("/home/jgrusewski/Work/foxhunt/ml/trained_models/production/ppo/ppo_actor_epoch_420.safetensors"); let critic_checkpoint = PathBuf::from("/home/jgrusewski/Work/foxhunt/ml/trained_models/production/ppo/ppo_critic_epoch_420.safetensors"); let mut metadata = ModelVersionMetadata::new( "ppo-production-v1.0.0".to_string(), ModelType::PPO, "1.0.0".to_string(), "ES.FUT_2024_Q4".to_string(), "s3://foxhunt-ml-models/ppo/1.0.0/".to_string(), ); // Add PPO-specific hyperparameters metadata.add_hyperparameter("epochs", serde_json::json!(420)); metadata.add_hyperparameter("batch_size", serde_json::json!(64)); metadata.add_hyperparameter("learning_rate", serde_json::json!(0.0003)); metadata.add_hyperparameter("gamma", serde_json::json!(0.99)); metadata.add_hyperparameter("gae_lambda", serde_json::json!(0.95)); metadata.add_metric("final_actor_loss", serde_json::json!(0.0152)); metadata.add_metric("final_critic_loss", serde_json::json!(0.0089)); metadata.add_metric("avg_reward", serde_json::json!(45.3)); metadata.add_metadata( "actor_checkpoint_path", actor_checkpoint.to_string_lossy().to_string(), ); metadata.add_metadata( "critic_checkpoint_path", critic_checkpoint.to_string_lossy().to_string(), ); metadata.set_checksum("sha256:ppo_epoch_420_checksum".to_string()); registry.register_version(&metadata).await?; let retrieved = registry .get_model_by_version("ppo-production-v1.0.0") .await?; assert_eq!(retrieved.model_type, ModelType::PPO); assert!(retrieved.metadata.contains_key("actor_checkpoint_path")); assert!(retrieved.metadata.contains_key("critic_checkpoint_path")); Ok(()) } /// Test 3: Register trained MAMBA-2 model with training metrics #[tokio::test] #[ignore = "Requires PostgreSQL"] async fn test_register_mamba2_checkpoint() -> MLResult<()> { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?; let mut metadata = ModelVersionMetadata::new( "mamba2-production-v1.0.0".to_string(), ModelType::MAMBA, "1.0.0".to_string(), "ES.FUT_2024_Q4".to_string(), "s3://foxhunt-ml-models/mamba2/1.0.0/".to_string(), ); // Add MAMBA-2 hyperparameters metadata.add_hyperparameter("epochs", serde_json::json!(24)); metadata.add_hyperparameter("batch_size", serde_json::json!(32)); metadata.add_hyperparameter("learning_rate", serde_json::json!(0.0001)); metadata.add_hyperparameter("d_model", serde_json::json!(256)); metadata.add_hyperparameter("n_layers", serde_json::json!(6)); metadata.add_hyperparameter("state_size", serde_json::json!(16)); metadata.add_metric("best_val_loss", serde_json::json!(1.4318895660848898)); metadata.add_metric("best_epoch", serde_json::json!(3)); metadata.add_metric("final_perplexity", serde_json::json!(4.1866025848353)); metadata.add_metadata( "checkpoint_path", "/home/jgrusewski/Work/foxhunt/ml/checkpoints/mamba2_dbn/".to_string(), ); metadata.add_metadata("training_duration_hours", "0.031".to_string()); metadata.set_checksum("sha256:mamba2_epoch_24_checksum".to_string()); registry.register_version(&metadata).await?; let retrieved = registry .get_model_by_version("mamba2-production-v1.0.0") .await?; assert_eq!(retrieved.model_type, ModelType::MAMBA); // Verify metrics let metrics = retrieved.metrics.as_object().unwrap(); assert!(metrics.contains_key("best_val_loss")); assert!(metrics.contains_key("final_perplexity")); Ok(()) } /// Test 4: Register trained TFT model with multiple checkpoints #[tokio::test] #[ignore = "Requires PostgreSQL"] async fn test_register_tft_checkpoint() -> MLResult<()> { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?; let checkpoint_path = PathBuf::from( "/home/jgrusewski/Work/foxhunt/ml/trained_models/production/tft/tft_epoch_100.safetensors", ); let mut metadata = ModelVersionMetadata::new( "tft-production-v1.0.0".to_string(), ModelType::TFT, "1.0.0".to_string(), "ES.FUT_2024_Q4".to_string(), "s3://foxhunt-ml-models/tft/1.0.0/".to_string(), ); // Add TFT hyperparameters metadata.add_hyperparameter("epochs", serde_json::json!(100)); metadata.add_hyperparameter("batch_size", serde_json::json!(256)); metadata.add_hyperparameter("learning_rate", serde_json::json!(0.0001)); metadata.add_hyperparameter("hidden_size", serde_json::json!(256)); metadata.add_hyperparameter("num_attention_heads", serde_json::json!(8)); metadata.add_metric("final_loss", serde_json::json!(0.0198)); metadata.add_metric("validation_loss", serde_json::json!(0.0213)); metadata.add_metric("sharpe_ratio", serde_json::json!(2.4)); metadata.add_metadata( "checkpoint_path", checkpoint_path.to_string_lossy().to_string(), ); metadata.set_checksum("sha256:tft_epoch_100_checksum".to_string()); registry.register_version(&metadata).await?; let retrieved = registry .get_model_by_version("tft-production-v1.0.0") .await?; assert_eq!(retrieved.model_type, ModelType::TFT); // Verify hyperparameters let hyperparams = retrieved.hyperparameters.as_object().unwrap(); assert_eq!(hyperparams.get("epochs").unwrap(), &serde_json::json!(100)); Ok(()) } /// Test 5: Register TFT-INT8 quantized model #[tokio::test] #[ignore = "Requires PostgreSQL"] async fn test_register_tft_int8_checkpoint() -> MLResult<()> { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?; let mut metadata = ModelVersionMetadata::new( "tft-int8-production-v1.0.0".to_string(), ModelType::TFT, "1.0.0-int8".to_string(), "ES.FUT_2024_Q4".to_string(), "s3://foxhunt-ml-models/tft-int8/1.0.0/".to_string(), ); metadata.add_hyperparameter("quantization", serde_json::json!("int8")); metadata.add_hyperparameter("epochs", serde_json::json!(100)); metadata.add_metric("inference_latency_ms", serde_json::json!(3.2)); metadata.add_metric("model_size_mb", serde_json::json!(128)); metadata.add_metadata("quantization_method", "static_int8".to_string()); metadata.add_metadata("optimization_level", "production".to_string()); metadata.set_checksum("sha256:tft_int8_checksum".to_string()); registry.register_version(&metadata).await?; let retrieved = registry .get_model_by_version("tft-int8-production-v1.0.0") .await?; assert_eq!(retrieved.version, "1.0.0-int8"); assert!(retrieved.metadata.contains_key("quantization_method")); Ok(()) } /// Test 6: Version increment handling #[tokio::test] #[ignore = "Requires PostgreSQL"] async fn test_version_increment() -> MLResult<()> { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?; // Register v1.0.0 let mut metadata_v1 = ModelVersionMetadata::new( "dqn-version-test-v1.0.0".to_string(), ModelType::DQN, "1.0.0".to_string(), "test_data".to_string(), "s3://test/dqn/1.0.0/".to_string(), ); metadata_v1.add_metric("loss", serde_json::json!(0.05)); registry.register_version(&metadata_v1).await?; // Register v1.1.0 (improvement) let mut metadata_v1_1 = ModelVersionMetadata::new( "dqn-version-test-v1.1.0".to_string(), ModelType::DQN, "1.1.0".to_string(), "test_data".to_string(), "s3://test/dqn/1.1.0/".to_string(), ); metadata_v1_1.add_metric("loss", serde_json::json!(0.03)); registry.register_version(&metadata_v1_1).await?; // Register v2.0.0 (major update) let mut metadata_v2 = ModelVersionMetadata::new( "dqn-version-test-v2.0.0".to_string(), ModelType::DQN, "2.0.0".to_string(), "test_data".to_string(), "s3://test/dqn/2.0.0/".to_string(), ); metadata_v2.add_metric("loss", serde_json::json!(0.01)); registry.register_version(&metadata_v2).await?; // Verify all versions exist let v1 = registry .get_model_by_version("dqn-version-test-v1.0.0") .await?; assert_eq!(v1.version, "1.0.0"); let v1_1 = registry .get_model_by_version("dqn-version-test-v1.1.0") .await?; assert_eq!(v1_1.version, "1.1.0"); let v2 = registry .get_model_by_version("dqn-version-test-v2.0.0") .await?; assert_eq!(v2.version, "2.0.0"); Ok(()) } /// Test 7: Checkpoint path validation #[tokio::test] #[ignore = "Requires PostgreSQL"] async fn test_checkpoint_path_metadata() -> MLResult<()> { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?; let checkpoint_path = PathBuf::from( "/home/jgrusewski/Work/foxhunt/ml/trained_models/production/dqn/dqn_epoch_30.safetensors", ); let mut metadata = ModelVersionMetadata::new( "dqn-checkpoint-path-test".to_string(), ModelType::DQN, "1.0.0".to_string(), "test_data".to_string(), "s3://test/dqn/1.0.0/".to_string(), ); metadata.add_metadata( "checkpoint_path", checkpoint_path.to_string_lossy().to_string(), ); metadata.add_metadata("checkpoint_format", "safetensors".to_string()); metadata.add_metadata("checkpoint_size_mb", "256".to_string()); registry.register_version(&metadata).await?; let retrieved = registry .get_model_by_version("dqn-checkpoint-path-test") .await?; assert!(retrieved.metadata.contains_key("checkpoint_path")); assert_eq!( retrieved.metadata.get("checkpoint_format").unwrap(), "safetensors" ); Ok(()) } /// Test 8: Multi-model registry query #[tokio::test] #[ignore = "Requires PostgreSQL"] async fn test_multi_model_registry_query() -> MLResult<()> { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?; // Register multiple models let model_types = vec![ (ModelType::DQN, "dqn-multi-test"), (ModelType::PPO, "ppo-multi-test"), (ModelType::MAMBA, "mamba-multi-test"), (ModelType::TFT, "tft-multi-test"), ]; for (model_type, model_id) in model_types { let metadata = ModelVersionMetadata::new( model_id.to_string(), model_type, "1.0.0".to_string(), "test_data".to_string(), format!("s3://test/{}/1.0.0/", model_id), ); registry.register_version(&metadata).await?; } // Query by type let dqn_models = registry.get_models_by_type(ModelType::DQN).await?; assert!(dqn_models.iter().any(|m| m.model_id == "dqn-multi-test")); let ppo_models = registry.get_models_by_type(ModelType::PPO).await?; assert!(ppo_models.iter().any(|m| m.model_id == "ppo-multi-test")); Ok(()) } /// Test 9: Production model promotion workflow #[tokio::test] #[ignore = "Requires PostgreSQL"] async fn test_production_promotion_workflow() -> MLResult<()> { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?; let metadata = ModelVersionMetadata::new( "dqn-promotion-test".to_string(), ModelType::DQN, "1.0.0".to_string(), "test_data".to_string(), "s3://test/dqn/1.0.0/".to_string(), ); // Start as experimental assert!(metadata.is_experimental); assert!(!metadata.is_production); registry.register_version(&metadata).await?; // Promote to production registry.mark_production("dqn-promotion-test").await?; // Verify production status let retrieved = registry.get_model_by_version("dqn-promotion-test").await?; assert!(retrieved.is_production); assert!(!retrieved.is_experimental); // Verify in production query let production_models = registry.get_production_models().await?; assert!(production_models .iter() .any(|m| m.model_id == "dqn-promotion-test")); Ok(()) } /// Test 10: Training metrics metadata #[tokio::test] #[ignore = "Requires PostgreSQL"] async fn test_training_metrics_metadata() -> MLResult<()> { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?; let mut metadata = ModelVersionMetadata::new( "dqn-metrics-test".to_string(), ModelType::DQN, "1.0.0".to_string(), "test_data".to_string(), "s3://test/dqn/1.0.0/".to_string(), ); // Add comprehensive metrics metadata.add_metric("final_loss", serde_json::json!(0.0342)); metadata.add_metric("validation_loss", serde_json::json!(0.0356)); metadata.add_metric("best_epoch", serde_json::json!(28)); metadata.add_metric("total_epochs", serde_json::json!(30)); metadata.add_metric("training_duration_hours", serde_json::json!(2.5)); metadata.add_metric("gpu_memory_used_gb", serde_json::json!(3.2)); metadata.add_metric("avg_epoch_time_seconds", serde_json::json!(300)); registry.register_version(&metadata).await?; let retrieved = registry.get_model_by_version("dqn-metrics-test").await?; let metrics = retrieved.metrics.as_object().unwrap(); assert_eq!( metrics.get("final_loss").unwrap(), &serde_json::json!(0.0342) ); assert_eq!(metrics.get("best_epoch").unwrap(), &serde_json::json!(28)); assert!(metrics.contains_key("gpu_memory_used_gb")); Ok(()) } /// Test 11: List all checkpoints for a model type #[tokio::test] #[ignore = "Requires PostgreSQL"] async fn test_list_checkpoints_by_type() -> MLResult<()> { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?; // Register multiple DQN checkpoints for epoch in [10, 20, 30] { let mut metadata = ModelVersionMetadata::new( format!("dqn-checkpoint-epoch-{}", epoch), ModelType::DQN, format!("1.0.{}", epoch), "test_data".to_string(), format!("s3://test/dqn/1.0.{}/", epoch), ); metadata.add_metadata("epoch", epoch.to_string()); registry.register_version(&metadata).await?; } let dqn_models = registry.get_models_by_type(ModelType::DQN).await?; let checkpoint_models: Vec<_> = dqn_models .iter() .filter(|m| m.model_id.starts_with("dqn-checkpoint-epoch-")) .collect(); assert!(checkpoint_models.len() >= 3); Ok(()) } /// Test 12: Checkpoint metadata completeness #[tokio::test] #[ignore = "Requires PostgreSQL"] async fn test_checkpoint_metadata_completeness() -> MLResult<()> { let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?; let mut metadata = ModelVersionMetadata::new( "complete-metadata-test".to_string(), ModelType::DQN, "1.0.0".to_string(), "ES.FUT_2024_Q4".to_string(), "s3://test/dqn/1.0.0/".to_string(), ); // Add comprehensive metadata metadata.add_hyperparameter("epochs", serde_json::json!(30)); metadata.add_hyperparameter("batch_size", serde_json::json!(128)); metadata.add_hyperparameter("learning_rate", serde_json::json!(0.0001)); metadata.add_hyperparameter("gamma", serde_json::json!(0.99)); metadata.add_hyperparameter("epsilon_start", serde_json::json!(1.0)); metadata.add_hyperparameter("epsilon_end", serde_json::json!(0.01)); metadata.add_metric("final_loss", serde_json::json!(0.0342)); metadata.add_metric("validation_loss", serde_json::json!(0.0356)); metadata.add_metric("sharpe_ratio", serde_json::json!(2.1)); metadata.add_metric("max_drawdown", serde_json::json!(0.12)); metadata.add_metadata( "checkpoint_path", "/path/to/checkpoint.safetensors".to_string(), ); metadata.add_metadata("training_date", Utc::now().to_rfc3339()); metadata.add_metadata("cuda_version", "12.1".to_string()); metadata.add_metadata("pytorch_version", "2.0.0".to_string()); metadata.set_checksum("sha256:complete_metadata_checksum".to_string()); registry.register_version(&metadata).await?; let retrieved = registry .get_model_by_version("complete-metadata-test") .await?; // Verify hyperparameters let hyperparams = retrieved.hyperparameters.as_object().unwrap(); assert_eq!(hyperparams.len(), 6); // Verify metrics let metrics = retrieved.metrics.as_object().unwrap(); assert_eq!(metrics.len(), 4); // Verify metadata assert_eq!(retrieved.metadata.len(), 4); assert!(retrieved.metadata.contains_key("checkpoint_path")); assert!(retrieved.metadata.contains_key("cuda_version")); Ok(()) }