diff --git a/Cargo.lock b/Cargo.lock index d0a4650..d6fb5bd 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -291,6 +291,7 @@ dependencies = [ "bincode", "bloomfilter", "chrono", + "crc32fast", "criterion", "crossterm", "dotenvy", diff --git a/Cargo.toml b/Cargo.toml index 6127d1d..00359c8 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -42,6 +42,7 @@ api = [] [dependencies] bloomfilter = "3.0" +crc32fast = "1.4" bincode = "1.3" lz4_flex = "0.11" serde = { version = "1.0", features = ["derive"] } diff --git a/src/storage/block.rs b/src/storage/block.rs index 3da0e79..bb7970a 100644 --- a/src/storage/block.rs +++ b/src/storage/block.rs @@ -1,4 +1,5 @@ -use crate::infra::config::StorageConfig; +use crate::infra::{config::StorageConfig, error::LsmError}; +use crc32fast::Hasher; use std::mem::size_of; pub const BLOCK_SIZE: usize = 4096; @@ -76,48 +77,76 @@ impl Block { let num_elements = self.offsets.len() as u32; encoded.extend_from_slice(&num_elements.to_le_bytes()); + // Calculate and append CRC32 checksum (Little Endian) + let mut hasher = Hasher::new(); + hasher.update(&encoded); + let checksum = hasher.finalize(); + encoded.extend_from_slice(&checksum.to_le_bytes()); + encoded } - pub fn decode(data: &[u8]) -> Self { + pub fn decode(data: &[u8]) -> std::result::Result { if data.len() < U32_SIZE { - return Self { - data: Vec::new(), - offsets: Vec::new(), - block_size: BLOCK_SIZE, - }; + return Err(LsmError::CorruptedData( + "Data too short to contain checksum".to_string(), + )); + } + + // Read stored checksum (last 4 bytes) + let checksum_start = data.len() - U32_SIZE; + let stored_checksum = u32::from_le_bytes([ + data[checksum_start], + data[checksum_start + 1], + data[checksum_start + 2], + data[checksum_start + 3], + ]); + + // Extract data without checksum for verification + let data_without_checksum = &data[..checksum_start]; + + // Calculate actual checksum + let mut hasher = Hasher::new(); + hasher.update(data_without_checksum); + let calculated_checksum = hasher.finalize(); + + // Verify checksum + if stored_checksum != calculated_checksum { + return Err(LsmError::CorruptedData( + "CRC32 checksum mismatch: data corruption detected".to_string(), + )); } - let num_elements_start = data.len() - U32_SIZE; + let num_elements_start = data_without_checksum.len() - U32_SIZE; let num_elements = u32::from_le_bytes([ - data[num_elements_start], - data[num_elements_start + 1], - data[num_elements_start + 2], - data[num_elements_start + 3], + data_without_checksum[num_elements_start], + data_without_checksum[num_elements_start + 1], + data_without_checksum[num_elements_start + 2], + data_without_checksum[num_elements_start + 3], ]) as usize; - let offsets_start = data.len() - U32_SIZE - (num_elements * U32_SIZE); - let records_data = data[..offsets_start].to_vec(); + let offsets_start = data_without_checksum.len() - U32_SIZE - (num_elements * U32_SIZE); + let records_data = data_without_checksum[..offsets_start].to_vec(); let mut offsets = Vec::with_capacity(num_elements); let mut offset_pos = offsets_start; for _ in 0..num_elements { let offset = u32::from_le_bytes([ - data[offset_pos], - data[offset_pos + 1], - data[offset_pos + 2], - data[offset_pos + 3], + data_without_checksum[offset_pos], + data_without_checksum[offset_pos + 1], + data_without_checksum[offset_pos + 2], + data_without_checksum[offset_pos + 3], ]); offsets.push(offset); offset_pos += U32_SIZE; } - Self { + Ok(Self { data: records_data, offsets, block_size: BLOCK_SIZE, - } + }) } pub fn len(&self) -> usize { @@ -221,7 +250,7 @@ mod tests { // Verify integrity let encoded = block.encode(); - let decoded = Block::decode(&encoded); + let decoded = Block::decode(&encoded).unwrap(); assert_eq!(decoded.len(), block.len()); assert_eq!(decoded.offsets.len(), block.offsets.len()); @@ -245,7 +274,7 @@ mod tests { fn test_encode_decode_empty_block() { let block = Block::new(BLOCK_SIZE); let encoded = block.encode(); - let decoded = Block::decode(&encoded); + let decoded = Block::decode(&encoded).unwrap(); assert_eq!(decoded.len(), 0); assert!(decoded.is_empty()); } @@ -255,7 +284,7 @@ mod tests { let mut block = Block::new(BLOCK_SIZE); block.add(b"key1", b"value1"); let encoded = block.encode(); - let decoded = Block::decode(&encoded); + let decoded = Block::decode(&encoded).unwrap(); assert_eq!(decoded.len(), 1); assert_eq!(decoded.data_size(), block.data_size()); assert_eq!(decoded.data, block.data); @@ -278,9 +307,85 @@ mod tests { } let encoded = block.encode(); - let decoded = Block::decode(&encoded); + let decoded = Block::decode(&encoded).unwrap(); assert_eq!(decoded.len(), entries.len()); assert_eq!(decoded.data, block.data); assert_eq!(decoded.offsets, block.offsets); } + + #[test] + fn test_crc32_corruption_detected() { + let mut block = Block::new(BLOCK_SIZE); + block.add(b"test_key", b"test_value"); + + let encoded = block.encode(); + let mut corrupted = encoded.clone(); + + // Corrupt a byte in the data section (not the checksum) + corrupted[10] ^= 0xFF; + + // Verify that decode returns a corruption error + let result = Block::decode(&corrupted); + assert!(result.is_err()); + + let err = result.unwrap_err(); + assert!(matches!(err, LsmError::CorruptedData(_))); + assert!(err.to_string().contains("CRC32")); + } + + #[test] + fn test_crc32_valid_checksum() { + let mut block = Block::new(BLOCK_SIZE); + block.add(b"key1", b"value1"); + block.add(b"key2", b"value2"); + + let encoded = block.encode(); + let decoded = Block::decode(&encoded).unwrap(); + + assert_eq!(decoded.len(), 2); + assert_eq!(decoded.data, block.data); + assert_eq!(decoded.offsets, block.offsets); + } + + #[test] + fn test_crc32_checksum_mismatch_single_bit_flip() { + let mut block = Block::new(BLOCK_SIZE); + for i in 0..50 { + let key = format!("key_{:03}", i); + let value = format!("value_{:03}", i); + assert!(block.add(key.as_bytes(), value.as_bytes())); + } + + let encoded = block.encode(); + let corrupted = corrupt_byte(&encoded, 100); + + let result = Block::decode(&corrupted); + assert!(result.is_err()); + + let err = result.unwrap_err(); + assert!(matches!(err, LsmError::CorruptedData(_))); + assert!(err.to_string().contains("mismatch")); + } + + #[test] + fn test_crc32_truncated_file_detected() { + let mut block = Block::new(BLOCK_SIZE); + block.add(b"short_key", b"short_value"); + + let encoded = block.encode(); + // Truncate by removing the checksum bytes + let truncated = &encoded[..encoded.len() - U32_SIZE]; + + let result = Block::decode(truncated); + assert!(result.is_err()); + } + + /// Helper to corrupt a specific byte in the data + fn corrupt_byte(data: &[u8], pos: usize) -> Vec { + let mut corrupted = data.to_vec(); + if pos < corrupted.len() { + corrupted[pos] ^= 0xFF; + } + corrupted + } } diff --git a/src/storage/reader.rs b/src/storage/reader.rs index f433e86..1670df7 100644 --- a/src/storage/reader.rs +++ b/src/storage/reader.rs @@ -128,7 +128,7 @@ impl SstableReader { let block_data = self.read_block(&block_meta)?; // Deserialize block (no lock needed) - let block = Block::decode(&block_data); + let block = Block::decode(&block_data)?; // Linear scan within the block to find the key (no lock needed) Self::search_in_block(&block, key.as_bytes()) @@ -188,7 +188,7 @@ impl SstableReader { for block_meta in &blocks { let block_data = self.read_block(block_meta)?; - let block = Block::decode(&block_data); + let block = Block::decode(&block_data)?; // Access block data through pub(crate) fields for &offset in &block.offsets { @@ -861,4 +861,98 @@ mod tests { handle.join().unwrap(); } } + + #[test] + fn test_sstable_data_corruption_detected() { + use std::fs::File; + use std::io::Read; + use std::io::Write; + + let dir = tempdir().unwrap(); + let path = dir.path().join("corruption_test.sst"); + // Use a small block size to create many blocks in a large file + let config = StorageConfig { + block_size: 128, // Very small block size + ..Default::default() + }; + let _cache = create_test_cache(&config); + + // Write an SSTable with enough data to create many blocks + let mut builder = SstableBuilder::new(path.clone(), config.clone(), 12345).unwrap(); + for i in 0..100 { + let key = format!("key_{:03}", i); + let value = format!("value_{:03}", i); + builder + .add(key.as_bytes(), &create_test_record(&key, value.as_bytes())) + .unwrap(); + } + builder.finish().unwrap(); + + // Open the SSTable to get block metadata + let reader = SstableReader::open(path.clone(), config.clone(), _cache); + let reader = match reader { + Ok(r) => r, + Err(e) => { + // If we can't even open the original SSTable, something is wrong + panic!("Failed to open original SSTable: {:?}", e); + } + }; + + let metadata = reader.metadata(); + + // Corrupt the last block's data + // Get the last block's offset and size from metadata + let last_block = metadata.blocks.last().unwrap(); + let last_block_start = last_block.offset as usize; + let last_block_size = last_block.size as usize; + + // Read the file + let mut file = File::open(&path).unwrap(); + let mut original_data = Vec::new(); + file.read_to_end(&mut original_data).unwrap(); + + // Corrupt a byte in the last compressed block + let corrupt_offset = last_block_start + last_block_size / 2; + + if corrupt_offset < original_data.len() - 8 { + // Keep footer intact + original_data[corrupt_offset] ^= 0xFF; + + // Write the corrupted data back + let mut file = File::create(&path).unwrap(); + file.write_all(&original_data).unwrap(); + drop(file); + + // Re-open reader with a fresh cache + let fresh_cache = create_test_cache(&config); + let reader = SstableReader::open(path, config, fresh_cache).unwrap(); + + // Try to read the last block which should trigger the CRC32 check + // Use the key from the last block (we'll use a key we know exists) + // For simplicity, get the last key by reading metadata.max_key + let result = reader.get(&String::from_utf8_lossy(&metadata.max_key)); + + // The corruption should cause CRC32 verification to fail + assert!( + result.is_err(), + "Should fail to read from corrupted SSTable, got: {:?}", + result + ); + + let err = result.unwrap_err(); + assert!( + matches!(err, LsmError::CorruptedData(_)), + "Expected CorruptedData error, got: {:?}", + err + ); + assert!( + err.to_string().contains("CRC32"), + "Error should mention CRC32" + ); + } else { + // Fallback: if corruption position is invalid, just verify the block tests still work + // This shouldn't happen in practice + panic!("Corruption position calculation failed"); + } + } } diff --git a/src/storage/sst_iterator.rs b/src/storage/sst_iterator.rs index 77ab641..eef4b4a 100644 --- a/src/storage/sst_iterator.rs +++ b/src/storage/sst_iterator.rs @@ -92,7 +92,7 @@ impl SstableIterator { } let block_meta = meta.blocks[block_idx].clone(); let raw = self.reader.read_block(&block_meta)?; - self.current_block = Some(Block::decode(&raw)); + self.current_block = Some(Block::decode(&raw)?); self.block_index = block_idx; self.offset_index = 0; Ok(())