use crate::{lz4, Error, Level, Result};
#[cfg(feature = "parallel")]
use rayon::prelude::*;
const BLOCK_SIZE: usize = 64 * 1024;
const MAGIC: &[u8; 4] = b"AVPZ"; const VERSION: u8 = 1;
#[cfg(feature = "parallel")]
pub fn compress_parallel(data: &[u8], num_threads: usize) -> Result<Vec<u8>> {
compress_parallel_with_level(data, num_threads, Level::default())
}
#[cfg(feature = "parallel")]
pub fn compress_parallel_with_level(
data: &[u8],
num_threads: usize,
level: Level,
) -> Result<Vec<u8>> {
if data.is_empty() {
return Ok(create_empty_parallel_header());
}
let pool = if num_threads > 0 {
rayon::ThreadPoolBuilder::new()
.num_threads(num_threads)
.build()
.map_err(|e| Error::InvalidInput(format!("Thread pool error: {}", e)))?
} else {
rayon::ThreadPoolBuilder::new()
.build()
.map_err(|e| Error::InvalidInput(format!("Thread pool error: {}", e)))?
};
let blocks: Vec<&[u8]> = data.chunks(BLOCK_SIZE).collect();
let num_blocks = blocks.len();
let compressed_blocks: Vec<Vec<u8>> = pool.install(|| {
blocks
.par_iter()
.map(|block| lz4::compress_with_level(block, level))
.collect::<Result<Vec<_>>>()
})?;
let mut output = Vec::new();
output.extend_from_slice(MAGIC); output.push(VERSION); output.extend_from_slice(&(data.len() as u64).to_le_bytes()); output.extend_from_slice(&(num_blocks as u32).to_le_bytes());
for block in &compressed_blocks {
output.extend_from_slice(&(block.len() as u32).to_le_bytes());
}
for block in compressed_blocks {
output.extend_from_slice(&block);
}
Ok(output)
}
#[cfg(feature = "parallel")]
pub fn decompress_parallel(data: &[u8], num_threads: usize) -> Result<Vec<u8>> {
if data.len() < 17 {
return Err(Error::InvalidInput("Data too short for parallel format".to_string()));
}
if &data[0..4] != MAGIC {
return Err(Error::InvalidInput("Invalid parallel format magic".to_string()));
}
let version = data[4];
if version != VERSION {
return Err(Error::InvalidInput(format!(
"Unsupported version: {}",
version
)));
}
let original_size = u64::from_le_bytes([
data[5], data[6], data[7], data[8], data[9], data[10], data[11], data[12],
]) as usize;
if original_size == 0 {
return Ok(Vec::new());
}
let num_blocks = u32::from_le_bytes([data[13], data[14], data[15], data[16]]) as usize;
let mut pos = 17;
let mut block_sizes = Vec::with_capacity(num_blocks);
for _ in 0..num_blocks {
if pos + 4 > data.len() {
return Err(Error::CorruptedData("Truncated block sizes".to_string()));
}
let size = u32::from_le_bytes([data[pos], data[pos + 1], data[pos + 2], data[pos + 3]]) as usize;
block_sizes.push(size);
pos += 4;
}
let mut compressed_blocks = Vec::with_capacity(num_blocks);
for &size in &block_sizes {
if pos + size > data.len() {
return Err(Error::CorruptedData("Truncated block data".to_string()));
}
compressed_blocks.push(&data[pos..pos + size]);
pos += size;
}
let pool = if num_threads > 0 {
rayon::ThreadPoolBuilder::new()
.num_threads(num_threads)
.build()
.map_err(|e| Error::InvalidInput(format!("Thread pool error: {}", e)))?
} else {
rayon::ThreadPoolBuilder::new()
.build()
.map_err(|e| Error::InvalidInput(format!("Thread pool error: {}", e)))?
};
let decompressed_blocks: Vec<Vec<u8>> = pool.install(|| {
compressed_blocks
.par_iter()
.map(|block| lz4::decompress(block))
.collect::<Result<Vec<_>>>()
})?;
let mut output = Vec::with_capacity(original_size);
for block in decompressed_blocks {
output.extend_from_slice(&block);
}
if output.len() != original_size {
return Err(Error::CorruptedData(format!(
"Size mismatch: expected {}, got {}",
original_size,
output.len()
)));
}
Ok(output)
}
fn create_empty_parallel_header() -> Vec<u8> {
let mut output = Vec::new();
output.extend_from_slice(MAGIC);
output.push(VERSION);
output.extend_from_slice(&0u64.to_le_bytes()); output.extend_from_slice(&0u32.to_le_bytes()); output
}
#[cfg(not(feature = "parallel"))]
pub fn compress_parallel(data: &[u8], _num_threads: usize) -> Result<Vec<u8>> {
Err(Error::InvalidInput(
"Parallel compression requires 'parallel' feature".to_string(),
))
}
#[cfg(not(feature = "parallel"))]
pub fn compress_parallel_with_level(
data: &[u8],
_num_threads: usize,
_level: Level,
) -> Result<Vec<u8>> {
Err(Error::InvalidInput(
"Parallel compression requires 'parallel' feature".to_string(),
))
}
#[cfg(not(feature = "parallel"))]
pub fn decompress_parallel(data: &[u8], _num_threads: usize) -> Result<Vec<u8>> {
Err(Error::InvalidInput(
"Parallel decompression requires 'parallel' feature".to_string(),
))
}
#[cfg(all(test, feature = "parallel"))]
mod tests {
use super::*;
#[test]
fn test_parallel_compress_decompress() {
let data = vec![b'A'; 100000];
let compressed = compress_parallel(&data, 4).unwrap();
let decompressed = decompress_parallel(&compressed, 4).unwrap();
assert_eq!(data, decompressed);
}
#[test]
fn test_parallel_empty() {
let data = vec![];
let compressed = compress_parallel(&data, 4).unwrap();
let decompressed = decompress_parallel(&compressed, 4).unwrap();
assert_eq!(data, decompressed);
}
#[test]
fn test_parallel_large() {
let data = (0..1000000).map(|i| (i % 256) as u8).collect::<Vec<u8>>();
let compressed = compress_parallel(&data, 8).unwrap();
assert!(compressed.len() < data.len());
let decompressed = decompress_parallel(&compressed, 8).unwrap();
assert_eq!(data, decompressed);
}
#[test]
fn test_parallel_all_levels() {
let data = vec![b'X'; 200000];
for level in [Level::Fast, Level::Balanced, Level::Best] {
let compressed = compress_parallel_with_level(&data, 4, level).unwrap();
let decompressed = decompress_parallel(&compressed, 4).unwrap();
assert_eq!(data, decompressed);
}
}
}