Files
foxhunt/ml/src/checkpoint/validation.rs
jgrusewski 6093eac7bf 🔧 Tonic 0.14 Upgrade: Auto-generated and build system changes
Wave 64-65 cleanup: Proto regeneration and build system updates from Tonic 0.12→0.14 upgrade

Files updated:
- Cargo.lock: Dependency resolution for Tonic 0.14.2
- All build.rs: Updated for tonic-prost-build
- Proto files: Regenerated with tonic-prost 0.14
- Examples/tests: Updated for new gRPC API

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-Authored-By: Claude <noreply@anthropic.com>
2025-10-03 07:34:26 +02:00

526 lines
16 KiB
Rust

//! Validation utilities for checkpoint integrity
//!
//! Provides checksum validation and corruption detection for checkpoints.
use std::collections::HashMap;
use sha2::{Digest, Sha256};
use tracing::{debug, error, warn};
use super::{CheckpointMetadata, ModelType};
use crate::MLError;
/// 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_string(),
));
}
if metadata.model_name.is_empty() {
return Err(MLError::ModelError(
"Model name cannot be empty".to_string(),
));
}
if metadata.version.is_empty() {
return Err(MLError::ModelError(
"Model version cannot be empty".to_string(),
));
}
// 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_string(),
));
}
}
// Validate metrics ranges
if let Some(accuracy) = metadata.accuracy {
if accuracy < 0.0 || accuracy > 1.0 {
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::<u32>()
.map_err(|_| MLError::ModelError(format!("Invalid major version: {}", parts[0])))?;
let minor = parts[1]
.parse::<u32>()
.map_err(|_| MLError::ModelError(format!("Invalid minor version: {}", parts[1])))?;
let patch = parts[2]
.parse::<u32>()
.map_err(|_| MLError::ModelError(format!("Invalid patch version: {}", 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<ValidationReport, MLError> {
let mut report = ValidationReport::new();
// Checksum validation
if let Err(e) = self.validate_checksum(data, &metadata.checksum) {
report.add_error("checksum".to_string(), e.to_string());
} else {
report.add_success("checksum".to_string());
}
// Metadata validation
if let Err(e) = self.validate_metadata(metadata) {
report.add_error("metadata".to_string(), e.to_string());
} else {
report.add_success("metadata".to_string());
}
// Model compatibility validation
if let Err(e) = self.validate_model_compatibility(expected_model_type, metadata) {
report.add_error("model_compatibility".to_string(), e.to_string());
} else {
report.add_success("model_compatibility".to_string());
}
// Version compatibility validation
if let Err(e) = self.validate_version_compatibility(current_version, &metadata.version) {
report.add_warning("version_compatibility".to_string(), e.to_string());
} else {
report.add_success("version_compatibility".to_string());
}
// Data size validation
if data.len() as u64 != metadata.file_size {
report.add_error(
"data_size".to_string(),
format!(
"Data size mismatch: expected {}, got {}",
metadata.file_size,
data.len()
),
);
} else {
report.add_success("data_size".to_string());
}
debug!(
"Comprehensive validation completed with {} errors, {} warnings",
report.errors.len(),
report.warnings.len()
);
Ok(report)
}
/// Get validation statistics
pub 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<String>,
/// Validation warnings (non-critical issues)
pub warnings: HashMap<String, String>,
/// Validation errors (critical issues)
pub errors: HashMap<String, String>,
}
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)]
mod tests {
use super::*;
use chrono::Utc;
// use crate::safe_operations; // DISABLED - module not found
#[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_string(),
model_type: ModelType::DQN,
model_name: "test_model".to_string(),
version: "1.2.3".to_string(),
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_string();
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_string();
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_string(),
model_type: ModelType::TFT,
model_name: "test_model".to_string(),
version: "1.0.0".to_string(),
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_string());
report.add_warning("version".to_string(), "Minor version mismatch".to_string());
report.add_error("metadata".to_string(), "Invalid field".to_string());
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"));
}
}