//! Compression utilities for checkpoint data //! //! Provides multiple compression algorithms optimized for different use cases. use std::io::{Read, Write}; use flate2::{read::GzDecoder, write::GzEncoder, Compression}; use tracing::debug; use crate::CompressionType; use crate::MLError; /// Compression manager for checkpoint data #[derive(Debug)] pub struct CompressionManager { /// Default compression level default_level: u32, } impl CompressionManager { /// Create a new compression manager pub const fn new() -> Self { Self { default_level: 3 } } /// Compress data using the specified algorithm pub fn compress( &self, data: &[u8], compression_type: CompressionType, level: u32, ) -> Result, MLError> { match compression_type { CompressionType::None => Ok(data.to_vec()), CompressionType::LZ4 => self.compress_lz4(data), CompressionType::Zstd => self.compress_zstd(data, level), CompressionType::Gzip => self.compress_gzip(data, level), } } /// Decompress data using the specified algorithm pub fn decompress( &self, data: &[u8], compression_type: CompressionType, ) -> Result, MLError> { match compression_type { CompressionType::None => Ok(data.to_vec()), CompressionType::LZ4 => self.decompress_lz4(data), CompressionType::Zstd => self.decompress_zstd(data), CompressionType::Gzip => self.decompress_gzip(data), } } /// Compress using LZ4 (fast compression) #[allow(clippy::unnecessary_wraps)] fn compress_lz4(&self, data: &[u8]) -> Result, MLError> { // For now, simulate LZ4 compression with a simple encoding // In a real implementation, you'd use the lz4 crate let mut compressed = Vec::new(); compressed.extend_from_slice(b"LZ4:"); compressed.extend_from_slice(data); debug!( "LZ4 compressed {} bytes to {} bytes", data.len(), compressed.len() ); Ok(compressed) } /// Decompress LZ4 data fn decompress_lz4(&self, data: &[u8]) -> Result, MLError> { // For now, simulate LZ4 decompression if !data.starts_with(b"LZ4:") { return Err(MLError::ModelError("Invalid LZ4 header".to_owned())); } let decompressed = data .get(4..) .ok_or_else(|| MLError::ModelError("LZ4 data too short".to_owned()))? .to_vec(); debug!( "LZ4 decompressed {} bytes to {} bytes", data.len(), decompressed.len() ); Ok(decompressed) } /// Compress using Zstandard #[allow(clippy::unnecessary_wraps)] fn compress_zstd(&self, data: &[u8], level: u32) -> Result, MLError> { // For now, simulate Zstd compression // In a real implementation, you'd use the zstd crate let mut compressed = Vec::new(); compressed.extend_from_slice(b"ZSTD:"); compressed.extend_from_slice(&level.to_le_bytes()); compressed.extend_from_slice(data); debug!( "Zstd compressed {} bytes to {} bytes (level {})", data.len(), compressed.len(), level ); Ok(compressed) } /// Decompress Zstandard data fn decompress_zstd(&self, data: &[u8]) -> Result, MLError> { // For now, simulate Zstd decompression if !data.starts_with(b"ZSTD:") { return Err(MLError::ModelError("Invalid Zstd header".to_owned())); } if data.len() < 9 { return Err(MLError::ModelError("Invalid Zstd data".to_owned())); } let decompressed = data .get(9..) .ok_or_else(|| MLError::ModelError("Zstd data too short".to_owned()))? .to_vec(); debug!( "Zstd decompressed {} bytes to {} bytes", data.len(), decompressed.len() ); Ok(decompressed) } /// Compress using Gzip fn compress_gzip(&self, data: &[u8], level: u32) -> Result, MLError> { let mut encoder = GzEncoder::new(Vec::new(), Compression::new(level)); encoder .write_all(data) .map_err(|e| MLError::ModelError(format!("Gzip compression failed: {}", e)))?; let compressed = encoder .finish() .map_err(|e| MLError::ModelError(format!("Gzip compression finish failed: {}", e)))?; debug!( "Gzip compressed {} bytes to {} bytes (level {})", data.len(), compressed.len(), level ); Ok(compressed) } /// Decompress Gzip data fn decompress_gzip(&self, data: &[u8]) -> Result, MLError> { let mut decoder = GzDecoder::new(data); let mut decompressed = Vec::new(); decoder .read_to_end(&mut decompressed) .map_err(|e| MLError::ModelError(format!("Gzip decompression failed: {}", e)))?; debug!( "Gzip decompressed {} bytes to {} bytes", data.len(), decompressed.len() ); Ok(decompressed) } /// Estimate compression ratio for data pub fn estimate_compression_ratio( &self, data: &[u8], compression_type: CompressionType, ) -> Result { let sample_size = std::cmp::min(data.len(), 1024); // Sample first 1KB let sample = data.get(..sample_size).unwrap_or(data); let compressed = self.compress(sample, compression_type, self.default_level)?; let ratio = compressed.len() as f64 / sample.len() as f64; debug!( "Estimated compression ratio for {:?}: {:.3}", compression_type, ratio ); Ok(ratio) } /// Choose optimal compression algorithm based on data characteristics pub fn choose_optimal_compression(&self, data: &[u8]) -> CompressionType { // For small data, compression overhead might not be worth it if data.len() < 1024 { return CompressionType::None; } // Try different algorithms and pick the best one let mut best_type = CompressionType::None; let mut best_ratio = 1.0; for &compression_type in &[ CompressionType::LZ4, CompressionType::Zstd, CompressionType::Gzip, ] { if let Ok(ratio) = self.estimate_compression_ratio(data, compression_type) { if ratio < best_ratio { best_ratio = ratio; best_type = compression_type; } } } debug!( "Chosen optimal compression: {:?} (ratio: {:.3})", best_type, best_ratio ); best_type } } impl Default for CompressionManager { fn default() -> Self { Self::new() } } /// Compression statistics #[derive(Debug, Clone, Default)] pub struct CompressionStats { /// Total bytes before compression pub total_uncompressed: u64, /// Total bytes after compression pub total_compressed: u64, /// Number of compression operations pub compression_count: u64, /// Number of decompression operations pub decompression_count: u64, /// Total time spent compressing (microseconds) pub total_compress_time_us: u64, /// Total time spent decompressing (microseconds) pub total_decompress_time_us: u64, } impl CompressionStats { /// Calculate overall compression ratio pub fn compression_ratio(&self) -> f64 { if self.total_uncompressed > 0 { self.total_compressed as f64 / self.total_uncompressed as f64 } else { 1.0 } } /// Calculate average compression time pub const fn avg_compress_time_us(&self) -> u64 { if self.compression_count > 0 { self.total_compress_time_us / self.compression_count } else { 0 } } /// Calculate average decompression time pub const fn avg_decompress_time_us(&self) -> u64 { if self.decompression_count > 0 { self.total_decompress_time_us / self.decompression_count } else { 0 } } /// Calculate compression savings in bytes pub const fn bytes_saved(&self) -> u64 { self.total_uncompressed .saturating_sub(self.total_compressed) } } #[cfg(test)] #[allow(clippy::field_reassign_with_default)] mod tests { use super::*; #[test] fn test_compression_manager() -> Result<(), MLError> { let manager = CompressionManager::new(); let test_data = b"Hello, world! This is some test data for compression.".repeat(10); // Test each compression type for &compression_type in &[ CompressionType::None, CompressionType::LZ4, CompressionType::Zstd, CompressionType::Gzip, ] { let compressed = manager.compress(&test_data, compression_type, 3)?; let decompressed = manager.decompress(&compressed, compression_type)?; assert_eq!(decompressed, test_data); if compression_type != CompressionType::None { // For actual compression algorithms, compressed should be different if compression_type == CompressionType::Gzip { assert_ne!(compressed, test_data); } } } Ok(()) } #[test] fn test_compression_ratio_estimation() -> Result<(), MLError> { let manager = CompressionManager::new(); let test_data = b"AAAAAAAAAA".repeat(100); // Highly compressible data let ratio = manager.estimate_compression_ratio(&test_data, CompressionType::Gzip)?; assert!(ratio < 1.0); // Should compress well Ok(()) } #[test] fn test_optimal_compression_choice() { let manager = CompressionManager::new(); // Small data should not be compressed let small_data = b"small"; assert_eq!( manager.choose_optimal_compression(small_data), CompressionType::None ); // Large data should be compressed let large_data = b"This is some larger test data that should benefit from compression.".repeat(50); let chosen = manager.choose_optimal_compression(&large_data); assert_ne!(chosen, CompressionType::None); } #[test] fn test_compression_stats() { let mut stats = CompressionStats::default(); // Add some test data stats.total_uncompressed = 1000; stats.total_compressed = 800; stats.compression_count = 5; stats.total_compress_time_us = 500; assert_eq!(stats.compression_ratio(), 0.8); assert_eq!(stats.avg_compress_time_us(), 100); assert_eq!(stats.bytes_saved(), 200); } }