use crate::error::{Result, SZipError};
use crate::writer::CompressionMethod;
use async_compression::tokio::bufread::DeflateEncoder;
use std::io::Cursor;
use std::path::PathBuf;
use std::sync::Arc;
use tokio::io::AsyncReadExt;
use tokio::sync::{mpsc, Semaphore};
#[derive(Debug, Clone)]
pub struct ParallelConfig {
pub max_concurrent: usize,
pub compression_level: u32,
pub compression_method: CompressionMethod,
}
impl Default for ParallelConfig {
fn default() -> Self {
Self {
max_concurrent: 4,
compression_level: 6,
compression_method: CompressionMethod::Deflate,
}
}
}
impl ParallelConfig {
pub fn conservative() -> Self {
Self {
max_concurrent: 2,
compression_level: 6,
compression_method: CompressionMethod::Deflate,
}
}
pub fn balanced() -> Self {
Self::default()
}
pub fn aggressive() -> Self {
Self {
max_concurrent: 8,
compression_level: 6,
compression_method: CompressionMethod::Deflate,
}
}
pub fn with_max_concurrent(mut self, max: usize) -> crate::error::Result<Self> {
if max == 0 {
return Err(crate::error::SZipError::InvalidFormat(
"max_concurrent must be at least 1".to_string(),
));
}
if max > 16 {
return Err(crate::error::SZipError::InvalidFormat(
"max_concurrent must not exceed 16".to_string(),
));
}
self.max_concurrent = max;
Ok(self)
}
pub fn with_compression_level(mut self, level: u32) -> Self {
self.compression_level = level;
self
}
pub fn estimated_peak_memory_mb(&self) -> usize {
self.max_concurrent * 4
}
}
pub struct ParallelEntry {
pub name: String,
pub path: PathBuf,
}
impl ParallelEntry {
pub fn new(name: impl Into<String>, path: impl Into<PathBuf>) -> Self {
Self {
name: name.into(),
path: path.into(),
}
}
}
pub(crate) struct CompressedEntry {
pub name: String,
pub data: Vec<u8>,
pub uncompressed_size: u64,
pub crc32: u32,
}
async fn compress_file_deflate(path: PathBuf, level: u32) -> Result<(Vec<u8>, u64, u32)> {
let data = tokio::fs::read(&path).await?;
let uncompressed_size = data.len() as u64;
let crc32 = crc32fast::hash(&data);
let cursor = Cursor::new(data);
let mut encoder =
DeflateEncoder::with_quality(cursor, async_compression::Level::Precise(level as i32));
let mut compressed = Vec::new();
encoder.read_to_end(&mut compressed).await?;
Ok((compressed, uncompressed_size, crc32))
}
pub(crate) async fn compress_entries_parallel(
entries: Vec<ParallelEntry>,
config: ParallelConfig,
) -> Result<Vec<CompressedEntry>> {
let semaphore = Arc::new(Semaphore::new(config.max_concurrent));
let (tx, mut rx) = mpsc::channel(config.max_concurrent);
let handles: Vec<_> = entries
.into_iter()
.enumerate()
.map(|(index, entry)| {
let semaphore = semaphore.clone();
let tx = tx.clone();
let config = config.clone();
tokio::task::spawn(async move {
let _permit = semaphore
.acquire()
.await
.map_err(|_e| SZipError::InvalidFormat("Semaphore error".to_string()))?;
let (compressed, uncompressed_size, crc32) = match config.compression_method {
CompressionMethod::Deflate => {
compress_file_deflate(entry.path, config.compression_level).await?
}
_ => {
return Err(SZipError::InvalidFormat(
"Only DEFLATE supported in parallel compression".to_string(),
));
}
};
let result = CompressedEntry {
name: entry.name,
data: compressed,
uncompressed_size,
crc32,
};
tx.send((index, result))
.await
.map_err(|_e| SZipError::InvalidFormat("Channel send error".to_string()))?;
Ok::<_, SZipError>(())
})
})
.collect();
drop(tx);
let mut results = Vec::new();
while let Some((index, entry)) = rx.recv().await {
results.push((index, entry));
}
for handle in handles {
handle
.await
.map_err(|_e| SZipError::InvalidFormat("Task join error".to_string()))??;
}
results.sort_by_key(|(index, _)| *index);
Ok(results.into_iter().map(|(_, entry)| entry).collect())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_config_defaults() {
let config = ParallelConfig::default();
assert_eq!(config.max_concurrent, 4);
assert_eq!(config.compression_level, 6);
}
#[test]
fn test_config_presets() {
let conservative = ParallelConfig::conservative();
assert_eq!(conservative.max_concurrent, 2);
let aggressive = ParallelConfig::aggressive();
assert_eq!(aggressive.max_concurrent, 8);
}
#[test]
fn test_memory_estimation() {
let config = ParallelConfig::balanced();
let estimated = config.estimated_peak_memory_mb();
assert_eq!(estimated, 16); }
#[test]
fn test_invalid_max_concurrent_zero() {
let result = ParallelConfig::default().with_max_concurrent(0);
assert!(result.is_err(), "Expected error for max_concurrent=0");
let msg = result.unwrap_err().to_string();
assert!(
msg.contains("at least 1"),
"Error should mention minimum: {}",
msg
);
}
#[test]
fn test_invalid_max_concurrent_too_high() {
let result = ParallelConfig::default().with_max_concurrent(20);
assert!(result.is_err(), "Expected error for max_concurrent=20");
let msg = result.unwrap_err().to_string();
assert!(msg.contains("16"), "Error should mention maximum: {}", msg);
}
#[test]
fn test_valid_max_concurrent() {
let config = ParallelConfig::default().with_max_concurrent(8).unwrap();
assert_eq!(config.max_concurrent, 8);
}
}