//! Validation utilities for checkpoint integrity //! //! Provides checksum validation and corruption detection for checkpoints. use std::collections::HashMap; use std::path::Path; use sha2::{Digest, Sha256}; use tracing::{debug, error, warn}; use super::{CheckpointMetadata, ModelType}; use crate::MLError; /// Write SHA-256 checksum sidecar file alongside a safetensors checkpoint. /// Creates `{path}.sha256` containing the hex digest. pub fn write_checksum(safetensors_path: &Path) -> Result<(), MLError> { let bytes = std::fs::read(safetensors_path).map_err(|e| { MLError::CheckpointError(format!("Failed to read file for checksum: {}", e)) })?; let hash = Sha256::digest(&bytes); let hex = format!("{:x}", hash); let checksum_path = safetensors_path.with_extension("sha256"); std::fs::write(&checksum_path, hex.as_bytes()).map_err(|e| { MLError::CheckpointError(format!("Failed to write checksum: {}", e)) })?; Ok(()) } /// Verify SHA-256 checksum of a safetensors file against its `.sha256` sidecar. /// Returns `Ok(true)` if valid, `Ok(false)` if mismatch. /// Returns `Ok(true)` if no sidecar exists (backwards compatible). pub fn verify_checksum(safetensors_path: &Path) -> Result { let checksum_path = safetensors_path.with_extension("sha256"); if !checksum_path.exists() { warn!( "No checksum file for {}, skipping integrity check", safetensors_path.display() ); return Ok(true); } let expected = std::fs::read_to_string(&checksum_path).map_err(|e| { MLError::CheckpointError(format!("Failed to read checksum: {}", e)) })?; let bytes = std::fs::read(safetensors_path).map_err(|e| { MLError::CheckpointError(format!("Failed to read file for verification: {}", e)) })?; let actual = format!("{:x}", Sha256::digest(&bytes)); Ok(actual.trim() == expected.trim()) } /// Validation manager for checkpoint integrity #[derive(Debug)] pub struct ValidationManager { /// Validation statistics stats: ValidationStats, } impl ValidationManager { /// Create a new validation manager pub fn new() -> Self { Self { stats: ValidationStats::default(), } } /// Calculate SHA-256 checksum of data pub fn calculate_checksum(&self, data: &[u8]) -> String { let mut hasher = Sha256::new(); hasher.update(data); format!("{:x}", hasher.finalize()) } /// Validate checksum of data pub fn validate_checksum(&self, data: &[u8], expected_checksum: &str) -> Result<(), MLError> { let calculated_checksum = self.calculate_checksum(data); if calculated_checksum != expected_checksum { error!( "Checksum mismatch: expected {}, got {}", expected_checksum, calculated_checksum ); return Err(MLError::ModelError(format!( "Checkpoint corruption detected: checksum mismatch (expected: {}, got: {})", expected_checksum, calculated_checksum ))); } debug!("Checksum validation passed: {}", calculated_checksum); Ok(()) } /// Validate metadata consistency pub fn validate_metadata(&self, metadata: &CheckpointMetadata) -> Result<(), MLError> { // Check required fields if metadata.checkpoint_id.is_empty() { return Err(MLError::ModelError( "Checkpoint ID cannot be empty".to_owned(), )); } if metadata.model_name.is_empty() { return Err(MLError::ModelError( "Model name cannot be empty".to_owned(), )); } if metadata.version.is_empty() { return Err(MLError::ModelError( "Model version cannot be empty".to_owned(), )); } // Validate version format (basic semantic versioning) if !self.is_valid_version(&metadata.version) { return Err(MLError::ModelError(format!( "Invalid version format: {}", metadata.version ))); } // Check file size consistency if metadata.file_size == 0 { warn!("Checkpoint has zero file size: {}", metadata.checkpoint_id); } if let Some(compressed_size) = metadata.compressed_size { if compressed_size > metadata.file_size { return Err(MLError::ModelError( "Compressed size cannot be larger than original size".to_owned(), )); } } // Validate metrics ranges if let Some(accuracy) = metadata.accuracy { if !(0.0..=1.0).contains(&accuracy) { return Err(MLError::ModelError(format!( "Accuracy must be between 0 and 1, got: {}", accuracy ))); } } debug!( "Metadata validation passed for checkpoint: {}", metadata.checkpoint_id ); Ok(()) } /// Validate model type compatibility pub fn validate_model_compatibility( &self, expected_type: ModelType, metadata: &CheckpointMetadata, ) -> Result<(), MLError> { if metadata.model_type != expected_type { return Err(MLError::ModelError(format!( "Model type mismatch: expected {:?}, got {:?}", expected_type, metadata.model_type ))); } debug!("Model compatibility validation passed"); Ok(()) } /// Validate version compatibility pub fn validate_version_compatibility( &self, current_version: &str, checkpoint_version: &str, ) -> Result<(), MLError> { let current_parts = self.parse_version(current_version)?; let checkpoint_parts = self.parse_version(checkpoint_version)?; // Check major version compatibility if current_parts.0 != checkpoint_parts.0 { return Err(MLError::ModelError(format!( "Major version incompatibility: current {}, checkpoint {}", current_version, checkpoint_version ))); } // Warn about minor version differences if current_parts.1 != checkpoint_parts.1 { warn!( "Minor version difference: current {}, checkpoint {}", current_version, checkpoint_version ); } debug!("Version compatibility validation passed"); Ok(()) } /// Check if version string is valid fn is_valid_version(&self, version: &str) -> bool { self.parse_version(version).is_ok() } /// Parse semantic version string fn parse_version(&self, version: &str) -> Result<(u32, u32, u32), MLError> { let parts: Vec<&str> = version.split('.').collect(); if parts.len() != 3 { return Err(MLError::ModelError(format!( "Invalid version format: {} (expected major.minor.patch)", version ))); } let major = parts[0] .parse::() .map_err(|e| MLError::ModelError(format!("Invalid major version '{}': {e}", parts[0])))?; let minor = parts[1] .parse::() .map_err(|e| MLError::ModelError(format!("Invalid minor version '{}': {e}", parts[1])))?; let patch = parts[2] .parse::() .map_err(|e| MLError::ModelError(format!("Invalid patch version '{}': {e}", parts[2])))?; Ok((major, minor, patch)) } /// Perform comprehensive validation pub fn comprehensive_validation( &self, data: &[u8], metadata: &CheckpointMetadata, expected_model_type: ModelType, current_version: &str, ) -> Result { let mut report = ValidationReport::new(); // Checksum validation if let Err(e) = self.validate_checksum(data, &metadata.checksum) { report.add_error("checksum".to_owned(), e.to_string()); } else { report.add_success("checksum".to_owned()); } // Metadata validation if let Err(e) = self.validate_metadata(metadata) { report.add_error("metadata".to_owned(), e.to_string()); } else { report.add_success("metadata".to_owned()); } // Model compatibility validation if let Err(e) = self.validate_model_compatibility(expected_model_type, metadata) { report.add_error("model_compatibility".to_owned(), e.to_string()); } else { report.add_success("model_compatibility".to_owned()); } // Version compatibility validation if let Err(e) = self.validate_version_compatibility(current_version, &metadata.version) { report.add_warning("version_compatibility".to_owned(), e.to_string()); } else { report.add_success("version_compatibility".to_owned()); } // Data size validation if data.len() as u64 != metadata.file_size { report.add_error( "data_size".to_owned(), format!( "Data size mismatch: expected {}, got {}", metadata.file_size, data.len() ), ); } else { report.add_success("data_size".to_owned()); } debug!( "Comprehensive validation completed with {} errors, {} warnings", report.errors.len(), report.warnings.len() ); Ok(report) } /// Get validation statistics pub const fn get_stats(&self) -> &ValidationStats { &self.stats } } impl Default for ValidationManager { fn default() -> Self { Self::new() } } /// Validation statistics #[derive(Debug, Clone, Default)] pub struct ValidationStats { /// Total validations performed pub total_validations: u64, /// Successful validations pub successful_validations: u64, /// Failed validations pub failed_validations: u64, /// Checksum mismatches detected pub checksum_failures: u64, /// Metadata validation failures pub metadata_failures: u64, /// Version compatibility issues pub version_issues: u64, } impl ValidationStats { /// Calculate success rate pub fn success_rate(&self) -> f64 { if self.total_validations > 0 { self.successful_validations as f64 / self.total_validations as f64 } else { 0.0 } } } /// Validation report containing results of comprehensive validation #[derive(Debug, Clone)] pub struct ValidationReport { /// Successful validation checks pub successes: Vec, /// Validation warnings (non-critical issues) pub warnings: HashMap, /// Validation errors (critical issues) pub errors: HashMap, } impl ValidationReport { /// Create a new validation report pub fn new() -> Self { Self { successes: Vec::new(), warnings: HashMap::new(), errors: HashMap::new(), } } /// Add a successful check pub fn add_success(&mut self, check: String) { self.successes.push(check); } /// Add a warning pub fn add_warning(&mut self, check: String, message: String) { self.warnings.insert(check, message); } /// Add an error pub fn add_error(&mut self, check: String, message: String) { self.errors.insert(check, message); } /// Check if validation passed (no errors) pub fn is_valid(&self) -> bool { self.errors.is_empty() } /// Check if there are warnings pub fn has_warnings(&self) -> bool { !self.warnings.is_empty() } /// Get summary of validation results pub fn summary(&self) -> String { if self.is_valid() { if self.has_warnings() { format!( "Validation passed with {} warnings ({} successful checks)", self.warnings.len(), self.successes.len() ) } else { format!( "Validation passed successfully ({} checks)", self.successes.len() ) } } else { format!( "Validation failed with {} errors and {} warnings", self.errors.len(), self.warnings.len() ) } } } impl Default for ValidationReport { fn default() -> Self { Self::new() } } #[cfg(test)] #[allow( clippy::assertions_on_result_states, clippy::len_zero, clippy::redundant_clone )] mod tests { use super::*; use chrono::Utc; #[test] fn test_checksum_validation() { let validator = ValidationManager::new(); let data = b"test data for checksum"; let checksum = validator.calculate_checksum(data); assert!(validator.validate_checksum(data, &checksum).is_ok()); // Test with wrong checksum let wrong_checksum = "wrong_checksum"; assert!(validator.validate_checksum(data, wrong_checksum).is_err()); } #[test] fn test_metadata_validation() { let validator = ValidationManager::new(); // Valid metadata let valid_metadata = CheckpointMetadata { checkpoint_id: "test_id".to_owned(), model_type: ModelType::DQN, model_name: "test_model".to_owned(), version: "1.2.3".to_owned(), created_at: Utc::now(), file_size: 1000, compressed_size: Some(800), accuracy: Some(0.95), ..Default::default() }; assert!(validator.validate_metadata(&valid_metadata).is_ok()); // Invalid metadata - empty model name let mut invalid_metadata = valid_metadata.clone(); invalid_metadata.model_name = "".to_owned(); assert!(validator.validate_metadata(&invalid_metadata).is_err()); // Invalid metadata - invalid version let mut invalid_version = valid_metadata.clone(); invalid_version.version = "invalid_version".to_owned(); assert!(validator.validate_metadata(&invalid_version).is_err()); // Invalid metadata - accuracy out of range let mut invalid_accuracy = valid_metadata.clone(); invalid_accuracy.accuracy = Some(1.5); assert!(validator.validate_metadata(&invalid_accuracy).is_err()); } #[test] fn test_version_parsing() -> Result<(), MLError> { let validator = ValidationManager::new(); assert_eq!(validator.parse_version("1.2.3")?, (1, 2, 3)); assert_eq!(validator.parse_version("0.1.0")?, (0, 1, 0)); assert!(validator.parse_version("1.2").is_err()); assert!(validator.parse_version("1.2.3.4").is_err()); assert!(validator.parse_version("a.b.c").is_err()); Ok(()) } #[test] fn test_version_compatibility() { let validator = ValidationManager::new(); // Same major version should be compatible assert!(validator .validate_version_compatibility("1.2.3", "1.3.0") .is_ok()); // Different major version should be incompatible assert!(validator .validate_version_compatibility("1.0.0", "2.0.0") .is_err()); // Invalid versions should return error assert!(validator .validate_version_compatibility("invalid", "1.0.0") .is_err()); } #[test] fn test_model_compatibility() { let validator = ValidationManager::new(); let metadata = CheckpointMetadata { model_type: ModelType::DQN, ..Default::default() }; // Same model type should be compatible assert!(validator .validate_model_compatibility(ModelType::DQN, &metadata) .is_ok()); // Different model type should be incompatible assert!(validator .validate_model_compatibility(ModelType::MAMBA, &metadata) .is_err()); } #[test] fn test_comprehensive_validation() -> Result<(), MLError> { let validator = ValidationManager::new(); let data = b"test checkpoint data"; let metadata = CheckpointMetadata { checkpoint_id: "test_id".to_owned(), model_type: ModelType::TFT, model_name: "test_model".to_owned(), version: "1.0.0".to_owned(), created_at: Utc::now(), file_size: data.len() as u64, checksum: validator.calculate_checksum(data), ..Default::default() }; let report = validator.comprehensive_validation(data, &metadata, ModelType::TFT, "1.0.0")?; assert!(report.is_valid()); assert!(!report.has_warnings()); assert!(report.successes.len() > 0); Ok(()) } #[test] fn test_validation_report() { let mut report = ValidationReport::new(); report.add_success("checksum".to_owned()); report.add_warning("version".to_owned(), "Minor version mismatch".to_owned()); report.add_error("metadata".to_owned(), "Invalid field".to_owned()); assert!(!report.is_valid()); assert!(report.has_warnings()); assert_eq!(report.successes.len(), 1); assert_eq!(report.warnings.len(), 1); assert_eq!(report.errors.len(), 1); let summary = report.summary(); assert!(summary.contains("failed")); assert!(summary.contains("1 errors")); assert!(summary.contains("1 warnings")); } #[test] fn test_sidecar_checksum_roundtrip() -> Result<(), Box> { let dir = tempfile::TempDir::new()?; let path = dir.path().join("test_model.safetensors"); std::fs::write(&path, b"fake model data")?; write_checksum(&path)?; let checksum_path = path.with_extension("sha256"); assert!(checksum_path.exists()); assert!(verify_checksum(&path)?); Ok(()) } #[test] fn test_sidecar_checksum_detects_corruption() -> Result<(), Box> { let dir = tempfile::TempDir::new()?; let path = dir.path().join("test_model.safetensors"); std::fs::write(&path, b"original data")?; write_checksum(&path)?; // Corrupt the file std::fs::write(&path, b"corrupted data")?; assert!(!verify_checksum(&path)?); Ok(()) } #[test] fn test_verify_without_sidecar_returns_true() -> Result<(), Box> { let dir = tempfile::TempDir::new()?; let path = dir.path().join("test_model.safetensors"); std::fs::write(&path, b"no checksum file")?; // No .sha256 file exists -- should return true (backwards compatible) assert!(verify_checksum(&path)?); Ok(()) } }