#![allow(clippy::module_name_repetitions)] #![allow(clippy::integer_division)] #![allow(clippy::shadow_reuse, clippy::shadow_same, clippy::shadow_unrelated)] // Tensor ops: let x = x.relu() is idiomatic #![allow(clippy::non_ascii_literal)] // Math symbols in ML documentation and error messages #![allow(clippy::partial_pub_fields)] // ML config structs: some fields are pub API, some internal #![allow(clippy::same_name_method)] // Intentional: inherent methods shadow trait defaults for ML-specific behavior #![allow(clippy::indexing_slicing)] // Tensor/matrix indexing with bounds guaranteed by construction //! # Unified Model Weight Persistence System //! //! Comprehensive checkpoint system for all AI models //! with versioning, metadata, compression, and validation. //! //! ## Key Features //! //! - **Unified Interface**: Single API for all model checkpointing //! - **Model Versioning**: Semantic versioning with compatibility checks //! - **Metadata Management**: Training metrics, hyperparameters, performance stats //! - **Compression**: Optional LZ4/Zstd compression for large models //! - **Validation**: Checksum verification and corruption detection //! - **Incremental Saves**: Delta checkpoints for memory efficiency //! - **Async I/O**: Non-blocking checkpoint operations //! - **Multi-format**: Binary, JSON, and custom formats use std::collections::HashMap; use std::fs::{self}; use std::path::PathBuf; use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::Arc; #[cfg(test)] use async_trait::async_trait; use chrono::{DateTime, Utc}; use dashmap::DashMap; use serde::{Deserialize, Serialize}; use sha2::{Digest, Sha256}; use tokio::sync::RwLock; use tracing::{info, instrument, warn}; use uuid::Uuid; // Re-export foundation types so sub-modules can use `crate::MLError` / `crate::ModelType` pub use ml_core::MLError; pub use ml_core::ModelType; pub mod compression; pub mod signer; pub mod storage; pub mod validation; pub mod versioning; #[cfg(test)] pub mod integration_tests; // Re-export key types for external usage pub use compression::CompressionManager; pub use signer::{CheckpointSigner, SignatureInfo}; #[cfg(feature = "s3-storage")] pub use storage::S3CheckpointStorage; pub use storage::{CheckpointStorage, FileSystemStorage, MemoryStorage, StorageStats}; pub use validation::{verify_checksum, write_checksum, ValidationManager}; pub use versioning::VersionManager; /// Checkpoint format options #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] pub enum CheckpointFormat { /// Binary format (fastest) Binary, /// JSON format (human-readable) JSON, /// `MessagePack` format (compact) MessagePack, /// Custom optimized format Custom, } /// Compression algorithm options #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] pub enum CompressionType { /// No compression None, /// LZ4 - fast compression LZ4, /// Zstandard - balanced compression Zstd, /// Gzip - high compression Gzip, } /// Checkpoint metadata containing model information and training state #[derive(Debug, Clone, Serialize, Deserialize)] pub struct CheckpointMetadata { /// Unique checkpoint identifier pub checkpoint_id: String, /// Model type pub model_type: ModelType, /// Model name/identifier pub model_name: String, /// Model version (semantic versioning) pub version: String, /// Creation timestamp pub created_at: DateTime, /// Training epoch when checkpoint was created pub epoch: Option, /// Training step when checkpoint was created pub step: Option, /// Training loss at checkpoint time pub loss: Option, /// Validation accuracy at checkpoint time pub accuracy: Option, /// Model hyperparameters pub hyperparameters: HashMap, /// Training metrics and statistics pub metrics: HashMap, /// Model architecture information pub architecture: HashMap, /// Checkpoint file format pub format: CheckpointFormat, /// Compression algorithm used pub compression: CompressionType, /// File size in bytes pub file_size: u64, /// Compressed file size (if compressed) pub compressed_size: Option, /// SHA-256 checksum for validation pub checksum: String, /// Tags for organizing checkpoints pub tags: Vec, /// Additional custom metadata pub custom_metadata: HashMap, // Security fields (Agent 122 - SEC-001 fix) /// HMAC-SHA256 signature (hex-encoded) pub signature: Option, /// Signing algorithm identifier pub signature_algorithm: String, /// Signing key ID for rotation support pub signing_key_id: String, /// Signature timestamp pub signed_at: Option>, } impl Default for CheckpointMetadata { fn default() -> Self { Self { checkpoint_id: Uuid::new_v4().to_string(), model_type: ModelType::DQN, // Default to first variant model_name: String::new(), version: "1.0.0".to_owned(), created_at: Utc::now(), epoch: None, step: None, loss: None, accuracy: None, hyperparameters: HashMap::new(), metrics: HashMap::new(), architecture: HashMap::new(), format: CheckpointFormat::Binary, compression: CompressionType::None, file_size: 0, compressed_size: None, checksum: String::new(), tags: Vec::new(), custom_metadata: HashMap::new(), signature: None, signature_algorithm: String::new(), signing_key_id: String::new(), signed_at: None, } } } impl CheckpointMetadata { /// Create new checkpoint metadata pub fn new(model_type: ModelType, model_name: String, version: String) -> Self { Self { checkpoint_id: Uuid::new_v4().to_string(), model_type, model_name, version, created_at: Utc::now(), epoch: None, step: None, loss: None, accuracy: None, hyperparameters: HashMap::new(), metrics: HashMap::new(), architecture: HashMap::new(), format: CheckpointFormat::Binary, compression: CompressionType::None, file_size: 0, compressed_size: None, checksum: String::new(), tags: Vec::new(), custom_metadata: HashMap::new(), signature: None, signature_algorithm: String::new(), signing_key_id: String::new(), signed_at: None, } } /// Add training state information pub const fn with_training_state( mut self, epoch: Option, step: Option, loss: Option, accuracy: Option, ) -> Self { self.epoch = epoch; self.step = step; self.loss = loss; self.accuracy = accuracy; self } /// Add hyperparameters pub fn with_hyperparameters(mut self, hyperparams: HashMap) -> Self { self.hyperparameters = hyperparams; self } /// Add metrics pub fn with_metrics(mut self, metrics: HashMap) -> Self { self.metrics = metrics; self } /// Add tags pub fn with_tags(mut self, tags: Vec) -> Self { self.tags = tags; self } /// Check if this is a newer version than another metadata pub fn is_newer_than(&self, other: &CheckpointMetadata) -> bool { if self.model_type != other.model_type || self.model_name != other.model_name { return false; } // Compare by epoch if available if let (Some(self_epoch), Some(other_epoch)) = (self.epoch, other.epoch) { return self_epoch > other_epoch; } // Compare by step if available if let (Some(self_step), Some(other_step)) = (self.step, other.step) { return self_step > other_step; } // Compare by timestamp self.created_at > other.created_at } /// Generate filename for this checkpoint pub fn generate_filename(&self) -> String { let timestamp = self.created_at.format("%Y%m%d_%H%M%S%.3f"); let epoch_str = self.epoch.map(|e| format!("_e{}", e)).unwrap_or_default(); let step_str = self.step.map(|s| format!("_s{}", s)).unwrap_or_default(); let ext = self.model_type.file_extension(); format!( "{}_{}_v{}{}{}_{}.{}", self.model_type.file_extension(), self.model_name, self.version, epoch_str, step_str, timestamp, ext ) } } /// Configuration for checkpoint operations #[derive(Debug, Clone, Serialize, Deserialize)] pub struct CheckpointConfig { /// Base directory for checkpoints pub base_dir: PathBuf, /// Default compression type pub compression: CompressionType, /// Default checkpoint format pub format: CheckpointFormat, /// Maximum number of checkpoints to keep per model pub max_checkpoints_per_model: usize, /// Automatic cleanup of old checkpoints pub auto_cleanup: bool, /// Enable checksum validation pub validate_checksums: bool, /// Enable incremental checkpoints (delta saves) pub incremental_checkpoints: bool, /// Compression level (0-9, algorithm dependent) pub compression_level: u32, /// Enable async I/O operations pub async_io: bool, /// Buffer size for I/O operations pub buffer_size: usize, } impl Default for CheckpointConfig { fn default() -> Self { Self { base_dir: PathBuf::from("./checkpoints"), compression: CompressionType::LZ4, format: CheckpointFormat::Binary, max_checkpoints_per_model: 10, auto_cleanup: true, validate_checksums: true, incremental_checkpoints: true, compression_level: 3, async_io: true, buffer_size: 64 * 1024, // 64KB } } } /// Statistics for checkpoint operations #[derive(Debug, Default, Serialize, Deserialize)] pub struct CheckpointStats { /// Total checkpoints saved pub total_saved: AtomicU64, /// Total checkpoints loaded pub total_loaded: AtomicU64, /// Total bytes saved pub total_bytes_saved: AtomicU64, /// Total bytes loaded pub total_bytes_loaded: AtomicU64, /// Total compression savings (bytes) pub compression_savings: AtomicU64, /// Average save time (microseconds) pub avg_save_time_us: AtomicU64, /// Average load time (microseconds) pub avg_load_time_us: AtomicU64, /// Failed operations pub failed_operations: AtomicU64, } impl CheckpointStats { /// Record a save operation pub fn record_save(&self, bytes_saved: u64, save_time_us: u64, compression_savings: u64) { self.total_saved.fetch_add(1, Ordering::Relaxed); self.total_bytes_saved .fetch_add(bytes_saved, Ordering::Relaxed); self.compression_savings .fetch_add(compression_savings, Ordering::Relaxed); // Update moving average let count = self.total_saved.load(Ordering::Relaxed); let current_avg = self.avg_save_time_us.load(Ordering::Relaxed); let new_avg = ((current_avg * (count - 1)) + save_time_us) / count; self.avg_save_time_us.store(new_avg, Ordering::Relaxed); } /// Record a load operation pub fn record_load(&self, bytes_loaded: u64, load_time_us: u64) { self.total_loaded.fetch_add(1, Ordering::Relaxed); self.total_bytes_loaded .fetch_add(bytes_loaded, Ordering::Relaxed); // Update moving average let count = self.total_loaded.load(Ordering::Relaxed); let current_avg = self.avg_load_time_us.load(Ordering::Relaxed); let new_avg = ((current_avg * (count - 1)) + load_time_us) / count; self.avg_load_time_us.store(new_avg, Ordering::Relaxed); } /// Record a failed operation pub fn record_failure(&self) { self.failed_operations.fetch_add(1, Ordering::Relaxed); } /// Get statistics as a map pub fn to_map(&self) -> HashMap { let mut map = HashMap::new(); map.insert( "total_saved".to_owned(), self.total_saved.load(Ordering::Relaxed), ); map.insert( "total_loaded".to_owned(), self.total_loaded.load(Ordering::Relaxed), ); map.insert( "total_bytes_saved".to_owned(), self.total_bytes_saved.load(Ordering::Relaxed), ); map.insert( "total_bytes_loaded".to_owned(), self.total_bytes_loaded.load(Ordering::Relaxed), ); map.insert( "compression_savings".to_owned(), self.compression_savings.load(Ordering::Relaxed), ); map.insert( "avg_save_time_us".to_owned(), self.avg_save_time_us.load(Ordering::Relaxed), ); map.insert( "avg_load_time_us".to_owned(), self.avg_load_time_us.load(Ordering::Relaxed), ); map.insert( "failed_operations".to_owned(), self.failed_operations.load(Ordering::Relaxed), ); map } } // Checkpointable trait is defined in ml-core to avoid orphan rule issues pub use ml_core::Checkpointable; /// Main checkpoint manager for all AI models #[derive(Debug)] pub struct CheckpointManager { /// Configuration config: CheckpointConfig, /// Storage backend storage: Arc, /// Checkpoint metadata index metadata_index: Arc>>, /// Statistics stats: Arc, /// Compression manager compression_manager: Arc, /// Validation manager validation_manager: Arc, } impl CheckpointManager { /// Create a new checkpoint manager pub fn new(config: CheckpointConfig) -> Result { // Create base directory if it doesn't exist if !config.base_dir.exists() { fs::create_dir_all(&config.base_dir).map_err(|e| { MLError::ModelError(format!("Failed to create checkpoint directory: {}", e)) })?; } let storage: Arc = Arc::new(FileSystemStorage::new(config.base_dir.clone())); let metadata_index = Arc::new(RwLock::new(DashMap::new())); let stats = Arc::new(CheckpointStats::default()); let compression_manager = Arc::new(CompressionManager::new()); let validation_manager = Arc::new(ValidationManager::new()); Ok(Self { config, storage, metadata_index, stats, compression_manager, validation_manager, }) } /// Save a checkpoint for a model #[instrument(skip(self, model))] pub async fn save_checkpoint( &self, model: &M, tags: Option>, ) -> Result { let start_time = std::time::Instant::now(); info!("Saving checkpoint for model: {}", model.model_name()); // Create metadata let (epoch, step, loss, accuracy) = model.get_training_state(); let mut metadata = CheckpointMetadata::new( model.model_type(), model.model_name().to_owned(), model.model_version().to_owned(), ) .with_training_state(epoch, step, loss, accuracy) .with_hyperparameters(model.get_hyperparameters()) .with_metrics(model.get_metrics()); if let Some(tags) = tags { metadata = metadata.with_tags(tags); } metadata.format = self.config.format; metadata.compression = self.config.compression; // Add architecture info metadata.architecture = model.get_architecture_info(); // Serialize model state let model_data = model.serialize_state().await?; let original_size = model_data.len() as u64; // Compress if needed let (final_data, compressed_size) = if self.config.compression != CompressionType::None { let compressed = self.compression_manager.compress( &model_data, self.config.compression, self.config.compression_level, )?; let comp_size = compressed.len() as u64; (compressed, Some(comp_size)) } else { (model_data, None) }; // Calculate checksum let mut hasher = Sha256::new(); hasher.update(&final_data); let checksum = format!("{:x}", hasher.finalize()); // Update metadata metadata.file_size = original_size; metadata.compressed_size = compressed_size; metadata.checksum = checksum; // Generate filename let filename = metadata.generate_filename(); // Save to storage self.storage .save_checkpoint(&filename, &final_data, &metadata) .await?; // Update index { let index = self.metadata_index.write().await; index.insert(metadata.checkpoint_id.clone(), metadata.clone()); } // Cleanup old checkpoints if needed if self.config.auto_cleanup { self.cleanup_old_checkpoints(model.model_type(), model.model_name()) .await?; } // Record statistics let save_time_us = start_time.elapsed().as_micros() as u64; let compression_savings = compressed_size .map(|cs| original_size.saturating_sub(cs)) .unwrap_or(0); self.stats .record_save(original_size, save_time_us, compression_savings); info!( "Checkpoint saved: {} ({} bytes, {}\u{b5}s, {:.1}% compression)", filename, final_data.len(), save_time_us, if compressed_size.is_some() { (compression_savings as f64 / original_size as f64) * 100.0 } else { 0.0 } ); Ok(metadata.checkpoint_id) } /// Load a checkpoint into a model #[instrument(skip(self, model))] pub async fn load_checkpoint( &self, model: &mut M, checkpoint_id: &str, ) -> Result { let start_time = std::time::Instant::now(); info!("Loading checkpoint: {}", checkpoint_id); // Get metadata let metadata = { let index = self.metadata_index.read().await; index.get(checkpoint_id).map(|entry| entry.clone()) }; let metadata = metadata.ok_or_else(|| { MLError::ModelError(format!("Checkpoint not found: {}", checkpoint_id)) })?; // Verify model compatibility if metadata.model_type != model.model_type() { return Err(MLError::ModelError(format!( "Model type mismatch: expected {:?}, got {:?}", model.model_type(), metadata.model_type ))); } // Load data from storage let filename = metadata.generate_filename(); let data = self.storage.load_checkpoint(&filename).await?; // Validate checksum if enabled if self.config.validate_checksums { self.validation_manager .validate_checksum(&data, &metadata.checksum)?; } // Decompress if needed let final_data = if metadata.compression != CompressionType::None { self.compression_manager .decompress(&data, metadata.compression)? } else { data }; // Deserialize into model model.deserialize_state(&final_data).await?; // Validate loaded state model.validate_loaded_state()?; // Record statistics let load_time_us = start_time.elapsed().as_micros() as u64; self.stats .record_load(final_data.len() as u64, load_time_us); info!( "Checkpoint loaded: {} ({} bytes, {}\u{b5}s)", filename, final_data.len(), load_time_us ); Ok(metadata) } /// Load the latest checkpoint for a model pub async fn load_latest_checkpoint( &self, model: &mut M, ) -> Result, MLError> { let model_type = model.model_type(); let model_name = model.model_name(); // Find latest checkpoint let latest_checkpoint = { let index = self.metadata_index.read().await; index .iter() .filter(|entry| { let metadata = entry.value(); metadata.model_type == model_type && metadata.model_name == model_name }) .max_by(|a, b| a.value().created_at.cmp(&b.value().created_at)) .map(|entry| entry.value().clone()) }; if let Some(metadata) = latest_checkpoint { let loaded_metadata = self.load_checkpoint(model, &metadata.checkpoint_id).await?; Ok(Some(loaded_metadata)) } else { Ok(None) } } /// List all checkpoints for a model pub async fn list_checkpoints( &self, model_type: ModelType, model_name: &str, ) -> Vec { let index = self.metadata_index.read().await; let mut checkpoints: Vec<_> = index .iter() .filter(|entry| { let metadata = entry.value(); // Empty string matches all model names (wildcard behavior) metadata.model_type == model_type && (model_name.is_empty() || metadata.model_name == model_name) }) .map(|entry| entry.value().clone()) .collect(); // Sort by creation time (newest first) checkpoints.sort_by(|a, b| b.created_at.cmp(&a.created_at)); checkpoints } /// Delete a checkpoint pub async fn delete_checkpoint(&self, checkpoint_id: &str) -> Result<(), MLError> { // Get metadata let metadata = { let index = self.metadata_index.read().await; index.get(checkpoint_id).map(|entry| entry.clone()) }; let metadata = metadata.ok_or_else(|| { MLError::ModelError(format!("Checkpoint not found: {}", checkpoint_id)) })?; // Delete from storage let filename = metadata.generate_filename(); self.storage.delete_checkpoint(&filename).await?; // Remove from index { let index = self.metadata_index.write().await; index.remove(checkpoint_id); } info!("Checkpoint deleted: {}", checkpoint_id); Ok(()) } /// Cleanup old checkpoints for a model async fn cleanup_old_checkpoints( &self, model_type: ModelType, model_name: &str, ) -> Result<(), MLError> { let mut checkpoints = self.list_checkpoints(model_type, model_name).await; if checkpoints.len() <= self.config.max_checkpoints_per_model { return Ok(()); } // Remove oldest checkpoints checkpoints.sort_by(|a, b| a.created_at.cmp(&b.created_at)); let to_remove = checkpoints.len() - self.config.max_checkpoints_per_model; for checkpoint in checkpoints.into_iter().take(to_remove) { if let Err(e) = self.delete_checkpoint(&checkpoint.checkpoint_id).await { warn!( "Failed to delete old checkpoint {}: {}", checkpoint.checkpoint_id, e ); } } info!( "Cleaned up {} old checkpoints for {}:{}", to_remove, model_type.file_extension(), model_name ); Ok(()) } /// Get checkpoint statistics pub fn get_stats(&self) -> HashMap { self.stats.to_map() } /// Refresh metadata index from storage pub async fn refresh_index(&self) -> Result<(), MLError> { let all_metadata = self.storage.list_all_checkpoints().await?; let index = self.metadata_index.write().await; index.clear(); for metadata in all_metadata { index.insert(metadata.checkpoint_id.clone(), metadata); } info!( "Refreshed checkpoint index with {} checkpoints", index.len() ); Ok(()) } /// Get checkpoint by tags pub async fn find_checkpoints_by_tags(&self, tags: &[String]) -> Vec { let index = self.metadata_index.read().await; index .iter() .filter(|entry| { let metadata = entry.value(); tags.iter().all(|tag| metadata.tags.contains(tag)) }) .map(|entry| entry.value().clone()) .collect() } } #[cfg(test)] #[allow(clippy::assertions_on_result_states)] mod tests { use super::*; use std::sync::atomic::AtomicBool; // Mock model for testing struct MockModel { model_type: ModelType, model_name: String, version: String, data: Vec, trained: AtomicBool, } impl MockModel { fn new(model_type: ModelType, name: String, version: String) -> Self { Self { model_type, model_name: name, version, data: vec![1, 2, 3, 4, 5], trained: AtomicBool::new(false), } } } #[async_trait] impl Checkpointable for MockModel { fn model_type(&self) -> ModelType { self.model_type } fn model_name(&self) -> &str { &self.model_name } fn model_version(&self) -> &str { &self.version } async fn serialize_state(&self) -> Result, MLError> { Ok(self.data.clone()) } async fn deserialize_state(&mut self, data: &[u8]) -> Result<(), MLError> { self.data = data.to_vec(); self.trained.store(true, Ordering::Relaxed); Ok(()) } fn get_training_state(&self) -> (Option, Option, Option, Option) { (Some(10), Some(1000), Some(0.1), Some(0.95)) } } #[tokio::test] async fn test_checkpoint_save_load() -> Result<(), MLError> { let temp_dir = tempfile::tempdir()?; let config = CheckpointConfig { base_dir: temp_dir.path().to_path_buf(), compression: CompressionType::None, ..Default::default() }; let manager = CheckpointManager::new(config)?; let mut model = MockModel::new( ModelType::DQN, "test_model".to_owned(), "1.0.0".to_owned(), ); // Save checkpoint let checkpoint_id = manager .save_checkpoint(&model, Some(vec!["test".to_owned()])) .await?; // Modify model data model.data = vec![9, 8, 7, 6, 5]; // Load checkpoint let metadata = manager.load_checkpoint(&mut model, &checkpoint_id).await?; // Verify data was restored assert_eq!(model.data, vec![1, 2, 3, 4, 5]); assert_eq!(metadata.model_type, ModelType::DQN); assert_eq!(metadata.model_name, "test_model"); assert_eq!(metadata.tags, vec!["test".to_owned()]); assert!(model.trained.load(Ordering::Relaxed)); Ok(()) } #[tokio::test] async fn test_checkpoint_compression() -> Result<(), MLError> { let temp_dir = tempfile::tempdir()?; let config = CheckpointConfig { base_dir: temp_dir.path().to_path_buf(), compression: CompressionType::LZ4, ..Default::default() }; let manager = CheckpointManager::new(config)?; let mut model = MockModel::new( ModelType::MAMBA, "test_model".to_owned(), "1.0.0".to_owned(), ); // Create larger data for compression test model.data = vec![42; 1000]; let checkpoint_id = manager.save_checkpoint(&model, None).await?; // Clear data model.data.clear(); // Load and verify manager.load_checkpoint(&mut model, &checkpoint_id).await?; assert_eq!(model.data, vec![42; 1000]); Ok(()) } #[tokio::test] async fn test_checkpoint_metadata() -> Result<(), MLError> { let metadata = CheckpointMetadata::new( ModelType::TFT, "transformer_model".to_owned(), "2.1.0".to_owned(), ) .with_training_state(Some(50), Some(5000), Some(0.05), Some(0.98)) .with_tags(vec!["production".to_owned(), "validated".to_owned()]); assert_eq!(metadata.model_type, ModelType::TFT); assert_eq!(metadata.epoch, Some(50)); assert_eq!(metadata.accuracy, Some(0.98)); assert!(metadata.tags.contains(&"production".to_owned())); let filename = metadata.generate_filename(); assert!(filename.contains("tft")); assert!(filename.contains("transformer_model")); assert!(filename.contains("v2.1.0")); assert!(filename.contains("e50")); assert!(filename.contains("s5000")); Ok(()) } #[tokio::test] async fn test_list_and_cleanup_checkpoints() -> Result<(), MLError> { let temp_dir = tempfile::tempdir()?; let config = CheckpointConfig { base_dir: temp_dir.path().to_path_buf(), max_checkpoints_per_model: 2, auto_cleanup: false, ..Default::default() }; let manager = CheckpointManager::new(config)?; let model = MockModel::new( ModelType::TGGN, "graph_model".to_owned(), "1.0.0".to_owned(), ); // Save multiple checkpoints let id1 = manager.save_checkpoint(&model, None).await?; tokio::time::sleep(tokio::time::Duration::from_millis(10)).await; let id2 = manager.save_checkpoint(&model, None).await?; tokio::time::sleep(tokio::time::Duration::from_millis(10)).await; let id3 = manager.save_checkpoint(&model, None).await?; // List checkpoints let checkpoints = manager .list_checkpoints(ModelType::TGGN, "graph_model") .await; assert_eq!(checkpoints.len(), 3); // Manual cleanup manager .cleanup_old_checkpoints(ModelType::TGGN, "graph_model") .await?; // Should have only 2 checkpoints now let checkpoints = manager .list_checkpoints(ModelType::TGGN, "graph_model") .await; assert_eq!(checkpoints.len(), 2); // The oldest checkpoint should be gone assert!(manager .load_checkpoint( &mut MockModel::new( ModelType::TGGN, "graph_model".to_owned(), "1.0.0".to_owned() ), &id1 ) .await .is_err()); assert!(manager .load_checkpoint( &mut MockModel::new( ModelType::TGGN, "graph_model".to_owned(), "1.0.0".to_owned() ), &id2 ) .await .is_ok()); assert!(manager .load_checkpoint( &mut MockModel::new( ModelType::TGGN, "graph_model".to_owned(), "1.0.0".to_owned() ), &id3 ) .await .is_ok()); Ok(()) } }