//! Model Registry API Example //! //! This example demonstrates how to use the model registry system //! for tracking ML model versions, metadata, and production deployments. //! //! # Usage //! //! ```bash //! # Start PostgreSQL (via docker-compose) //! docker-compose up -d postgres //! //! # Run the example //! cargo run --example model_registry_api //! ``` use ml::model_registry::{ModelRegistry, ModelVersionMetadata, RegistryStatistics}; use ml::ModelType; use std::error::Error; #[tokio::main] async fn main() -> Result<(), Box> { // Initialize tracing tracing_subscriber::fmt::init(); println!("šŸš€ Model Registry API Example"); println!("===============================\n"); // Initialize registry let database_url = std::env::var("DATABASE_URL") .unwrap_or_else(|_| "postgresql://foxhunt:foxhunt_dev_password@localhost:5432/foxhunt".to_string()); let s3_base_path = "s3://foxhunt-ml-models/"; println!("šŸ“Š Connecting to database: {}", database_url); let registry = ModelRegistry::new(&database_url, s3_base_path).await?; println!("āœ… Registry initialized\n"); // Example 1: Register a DQN model println!("šŸ“ Example 1: Registering DQN model v1.0.0"); println!("-------------------------------------------"); let mut dqn_metadata = ModelVersionMetadata::new( "dqn-v1.0.0".to_string(), ModelType::DQN, "1.0.0".to_string(), "databento_2024_Q4".to_string(), "s3://foxhunt-ml-models/dqn/1.0.0/".to_string(), ); // Add hyperparameters dqn_metadata.add_hyperparameter("epochs", serde_json::json!(500)); dqn_metadata.add_hyperparameter("batch_size", serde_json::json!(128)); dqn_metadata.add_hyperparameter("learning_rate", serde_json::json!(0.0001)); dqn_metadata.add_hyperparameter("gamma", serde_json::json!(0.99)); // Add training metrics dqn_metadata.add_metric("final_loss", serde_json::json!(0.001)); dqn_metadata.add_metric("best_epoch", serde_json::json!(487)); dqn_metadata.add_metric("training_time_seconds", serde_json::json!(168)); dqn_metadata.add_metric("sharpe_ratio", serde_json::json!(2.3)); // Set checksum dqn_metadata.set_checksum("sha256:abc123def456...".to_string()); // Add custom metadata dqn_metadata.add_metadata("trainer", "ml_training_service"); dqn_metadata.add_metadata("gpu_type", "RTX 3050 Ti"); dqn_metadata.add_metadata("dataset_size", "10M samples"); // Register model registry.register_version(&dqn_metadata).await?; println!("āœ… DQN v1.0.0 registered as experimental\n"); // Example 2: Register a MAMBA model println!("šŸ“ Example 2: Registering MAMBA model v1.0.0"); println!("----------------------------------------------"); let mut mamba_metadata = ModelVersionMetadata::new( "mamba-v1.0.0".to_string(), ModelType::MAMBA, "1.0.0".to_string(), "databento_2024_Q4".to_string(), "s3://foxhunt-ml-models/mamba/1.0.0/".to_string(), ); mamba_metadata.add_hyperparameter("state_size", serde_json::json!(16)); mamba_metadata.add_hyperparameter("seq_len", serde_json::json!(100)); mamba_metadata.add_metric("final_loss", serde_json::json!(0.0008)); mamba_metadata.add_metric("sharpe_ratio", serde_json::json!(2.5)); mamba_metadata.set_checksum("sha256:mamba123...".to_string()); registry.register_version(&mamba_metadata).await?; println!("āœ… MAMBA v1.0.0 registered as experimental\n"); // Example 3: Mark DQN as production println!("šŸ“ Example 3: Promoting DQN to production"); println!("------------------------------------------"); registry.mark_production("dqn-v1.0.0").await?; println!("āœ… DQN v1.0.0 promoted to production\n"); // Example 4: Query models println!("šŸ“ Example 4: Querying models"); println!("-----------------------------"); // Get production models let production_models = registry.get_production_models().await?; println!("šŸ­ Production models: {}", production_models.len()); for model in &production_models { println!(" - {} ({})", model.model_id, format!("{:?}", model.model_type)); println!(" Version: {}", model.version); println!(" Trained: {}", model.training_date.format("%Y-%m-%d %H:%M:%S")); println!(" S3: {}", model.s3_location); } println!(); // Get experimental models let experimental_models = registry.get_experimental_models().await?; println!("šŸ”¬ Experimental models: {}", experimental_models.len()); for model in &experimental_models { println!(" - {} ({})", model.model_id, format!("{:?}", model.model_type)); } println!(); // Get models by type let dqn_models = registry.get_models_by_type(ModelType::DQN).await?; println!("šŸŽÆ DQN models: {}", dqn_models.len()); for model in &dqn_models { println!(" - {} (status: {})", model.model_id, if model.is_production { "production" } else if model.is_experimental { "experimental" } else { "unknown" } ); } println!(); // Example 5: Retrieve specific model println!("šŸ“ Example 5: Retrieving specific model"); println!("---------------------------------------"); let retrieved = registry.get_model_by_version("dqn-v1.0.0").await?; println!("šŸ“¦ Model: {}", retrieved.model_id); println!(" Type: {:?}", retrieved.model_type); println!(" Version: {}", retrieved.version); println!(" Training Date: {}", retrieved.training_date.format("%Y-%m-%d %H:%M:%S")); println!(" Data Source: {}", retrieved.data_source); println!(" S3 Location: {}", retrieved.s3_location); println!(" Checksum: {}", retrieved.checksum); println!(" Production: {}", retrieved.is_production); println!(" Experimental: {}", retrieved.is_experimental); println!("\n Hyperparameters:"); if let Some(obj) = retrieved.hyperparameters.as_object() { for (key, value) in obj { println!(" - {}: {}", key, value); } } println!("\n Metrics:"); if let Some(obj) = retrieved.metrics.as_object() { for (key, value) in obj { println!(" - {}: {}", key, value); } } println!("\n Metadata:"); for (key, value) in &retrieved.metadata { println!(" - {}: {}", key, value); } println!(); // Example 6: Get registry statistics println!("šŸ“ Example 6: Registry statistics"); println!("---------------------------------"); let stats: RegistryStatistics = registry.get_statistics().await?; println!("šŸ“Š Registry Statistics:"); println!(" Total models: {}", stats.total_count); println!(" Production models: {}", stats.production_count); println!(" Experimental models: {}", stats.experimental_count); println!(" Archived models: {}", stats.archived_count); println!(" Model types: {}", stats.model_types_count); if let Some(latest) = stats.latest_training_date { println!(" Latest training: {}", latest.format("%Y-%m-%d %H:%M:%S")); } if let Some(earliest) = stats.earliest_training_date { println!(" Earliest training: {}", earliest.format("%Y-%m-%d %H:%M:%S")); } println!(); // Example 7: Query by date range println!("šŸ“ Example 7: Querying by date range"); println!("------------------------------------"); 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?; println!("šŸ“… Models trained in last 24 hours: {}", recent_models.len()); for model in &recent_models { println!(" - {} (trained {})", model.model_id, model.training_date.format("%Y-%m-%d %H:%M:%S") ); } println!(); // Example 8: Archive old model println!("šŸ“ Example 8: Archiving model"); println!("-----------------------------"); // Register a model to archive let mut old_metadata = ModelVersionMetadata::new( "dqn-v0.9.0".to_string(), ModelType::DQN, "0.9.0".to_string(), "databento_2024_Q3".to_string(), "s3://foxhunt-ml-models/dqn/0.9.0/".to_string(), ); old_metadata.set_checksum("sha256:old123...".to_string()); registry.register_version(&old_metadata).await?; // Archive it registry.archive_model("dqn-v0.9.0").await?; println!("āœ… DQN v0.9.0 archived\n"); // Example 9: Error handling println!("šŸ“ Example 9: Error handling"); println!("----------------------------"); match registry.get_model_by_version("nonexistent-model").await { Ok(_) => println!("āŒ Should have failed!"), Err(e) => println!("āœ… Correctly handled missing model: {}", e), } println!(); println!("šŸŽ‰ All examples completed successfully!"); println!("\nšŸ’” Key Features Demonstrated:"); println!(" āœ“ Model registration with metadata"); println!(" āœ“ Hyperparameter and metric tracking"); println!(" āœ“ Production/experimental tagging"); println!(" āœ“ Version queries (by ID, type, date)"); println!(" āœ“ Model archival and lifecycle management"); println!(" āœ“ Registry statistics and monitoring"); println!(" āœ“ Error handling and validation"); Ok(()) }