//! Checkpoint Uploader - Upload trained model checkpoints to S3 //! //! This utility uploads all trained model checkpoints from local storage //! to S3-compatible storage (MinIO in development) for archival and //! production deployment. //! //! Usage: //! cargo run --example checkpoint_uploader -- --source-dir ml/trained_models/production use std::path::{Path, PathBuf}; use std::time::Instant; use storage::{ObjectStoreBackend, Storage}; use config::schemas::S3Config; use clap::Parser; use tracing::{info, warn, error}; #[derive(Parser, Debug)] #[clap(name = "checkpoint_uploader")] #[clap(about = "Upload trained model checkpoints to S3")] struct Args { /// Source directory containing checkpoints #[clap(short, long, default_value = "ml/trained_models/production")] source_dir: PathBuf, /// S3 bucket name #[clap(short, long, default_value = "foxhunt-ml-models")] bucket: String, /// Dry run - don't actually upload #[clap(short, long)] dry_run: bool, } #[derive(Debug)] struct UploadStats { total_files: usize, uploaded_files: usize, failed_files: usize, total_bytes: u64, duration_secs: f64, } impl UploadStats { fn new() -> Self { Self { total_files: 0, uploaded_files: 0, failed_files: 0, total_bytes: 0, duration_secs: 0.0, } } fn throughput_mbps(&self) -> f64 { if self.duration_secs > 0.0 { (self.total_bytes as f64) / (1024.0 * 1024.0 * self.duration_secs) } else { 0.0 } } } /// Parse checkpoint filename to extract model name and version fn parse_checkpoint_filename(filename: &str) -> Option<(String, String, String)> { // Expected formats: // - dqn_epoch_100.safetensors -> (dqn, epoch_100, .safetensors) // - ppo_checkpoint_epoch_200.safetensors -> (ppo, epoch_200, .safetensors) // - dqn_final_epoch500.safetensors -> (dqn, epoch500, .safetensors) if !filename.ends_with(".safetensors") { return None; } let name_without_ext = filename.trim_end_matches(".safetensors"); // Try to extract model name and epoch if let Some(pos) = name_without_ext.find("_epoch") { let model_name = &name_without_ext[..pos]; let model_clean = model_name.trim_end_matches("_checkpoint").trim_end_matches("_final"); let epoch_part = &name_without_ext[pos..]; return Some(( model_clean.to_string(), epoch_part.to_string(), ".safetensors".to_string() )); } None } /// Generate S3 path for checkpoint fn get_s3_path(model_name: &str, version: &str, filename: &str) -> String { format!("{}/{}/checkpoints/{}", model_name, version, filename) } async fn upload_checkpoint( backend: &ObjectStoreBackend, source_path: &Path, s3_path: &str, dry_run: bool, ) -> Result> { let file_size = tokio::fs::metadata(source_path).await?.len(); if dry_run { info!("DRY RUN: Would upload {} ({} bytes) -> {}", source_path.display(), file_size, s3_path); return Ok(file_size); } info!("Uploading {} ({} bytes) -> {}", source_path.display(), file_size, s3_path); // Read file contents let data = tokio::fs::read(source_path).await?; // Upload to S3 backend.store(s3_path, &data).await?; info!("Successfully uploaded: {}", s3_path); Ok(file_size) } #[tokio::main] async fn main() -> Result<(), Box> { // Initialize logging tracing_subscriber::fmt() .with_env_filter( tracing_subscriber::EnvFilter::from_default_env() .add_directive("checkpoint_uploader=info".parse()?) .add_directive("storage=info".parse()?) ) .init(); let args = Args::parse(); info!("Checkpoint Uploader"); info!(" Source directory: {}", args.source_dir.display()); info!(" S3 bucket: {}", args.bucket); info!(" Dry run: {}", args.dry_run); // Configure S3 backend for MinIO let s3_config = S3Config { bucket_name: args.bucket.clone(), region: "us-east-1".to_string(), access_key_id: Some("foxhunt".to_string()), secret_access_key: Some("foxhunt_dev_password".to_string()), session_token: None, endpoint_url: Some("http://localhost:9000".to_string()), force_path_style: true, timeout: std::time::Duration::from_secs(30), max_retry_attempts: 3, use_ssl: false, }; info!("Initializing S3 backend..."); let backend = ObjectStoreBackend::new(s3_config, None).await?; info!("S3 backend initialized successfully"); // Scan source directory for checkpoints info!("Scanning directory: {}", args.source_dir.display()); let mut entries = tokio::fs::read_dir(&args.source_dir).await?; let mut checkpoints = Vec::new(); while let Some(entry) = entries.next_entry().await? { let path = entry.path(); if path.is_file() { if let Some(filename) = path.file_name().and_then(|n| n.to_str()) { if filename.ends_with(".safetensors") { checkpoints.push(path); } } } } info!("Found {} checkpoint files", checkpoints.len()); // Upload checkpoints let start = Instant::now(); let mut stats = UploadStats::new(); stats.total_files = checkpoints.len(); for checkpoint_path in checkpoints { let filename = checkpoint_path.file_name() .and_then(|n| n.to_str()) .unwrap_or("unknown"); // Parse filename to determine model and version let (model_name, version) = if let Some((model, ver, _)) = parse_checkpoint_filename(filename) { (model, ver) } else { warn!("Could not parse checkpoint filename: {}, using defaults", filename); ("unknown".to_string(), "v1.0".to_string()) }; // Generate S3 path let s3_path = get_s3_path(&model_name, &version, filename); // Upload checkpoint match upload_checkpoint(&backend, &checkpoint_path, &s3_path, args.dry_run).await { Ok(size) => { stats.uploaded_files += 1; stats.total_bytes += size; } Err(e) => { error!("Failed to upload {}: {}", filename, e); stats.failed_files += 1; } } } stats.duration_secs = start.elapsed().as_secs_f64(); // Print summary println!("\n╔══════════════════════════════════════════════════════════╗"); println!("║ Checkpoint Upload Summary ║"); println!("╠══════════════════════════════════════════════════════════╣"); println!("║ Total files: {:>6} ║", stats.total_files); println!("║ Uploaded: {:>6} ║", stats.uploaded_files); println!("║ Failed: {:>6} ║", stats.failed_files); println!("║ Total size: {:>6} MB ║", stats.total_bytes / (1024 * 1024)); println!("║ Duration: {:>6.2} seconds ║", stats.duration_secs); println!("║ Throughput: {:>6.2} MB/s ║", stats.throughput_mbps()); println!("╚══════════════════════════════════════════════════════════╝"); if args.dry_run { println!("\nDRY RUN COMPLETE - No files were actually uploaded"); } else { println!("\nUpload complete!"); // Verify uploads by listing S3 bucket info!("Verifying uploads..."); let uploaded_objects = backend.list("").await?; println!("S3 bucket now contains {} objects", uploaded_objects.len()); } Ok(()) } #[cfg(test)] mod tests { use super::*; #[test] fn test_parse_checkpoint_filename() { let test_cases = vec![ ("dqn_epoch_100.safetensors", Some(("dqn".to_string(), "_epoch_100".to_string(), ".safetensors".to_string()))), ("ppo_checkpoint_epoch_200.safetensors", Some(("ppo".to_string(), "_epoch_200".to_string(), ".safetensors".to_string()))), ("dqn_final_epoch500.safetensors", Some(("dqn".to_string(), "_epoch500".to_string(), ".safetensors".to_string()))), ("invalid.txt", None), ]; for (input, expected) in test_cases { let result = parse_checkpoint_filename(input); assert_eq!(result, expected, "Failed for input: {}", input); } } #[test] fn test_get_s3_path() { let path = get_s3_path("dqn", "epoch_100", "dqn_epoch_100.safetensors"); assert_eq!(path, "dqn/epoch_100/checkpoints/dqn_epoch_100.safetensors"); } }