use crate::dict::Dictionary;
use crate::encode::{encode_oneshot, AdvancedOptions};
use crate::error::Error;
use crate::params::{CompressionParameters, Strategy};
use alloc::vec::Vec;
pub const NB_WORKERS_MAX: u32 = 256;
pub const JOB_SIZE_MIN: usize = 512 * 1024;
pub fn default_overlap_log(strategy: Strategy) -> u32 {
match strategy {
Strategy::BtOpt | Strategy::BtUltra | Strategy::BtUltra2 => 9,
_ => 6,
}
}
pub fn overlap_size(window_log: u32, overlap_log: u32, strategy: Strategy) -> usize {
let ov = if overlap_log == 0 {
default_overlap_log(strategy)
} else {
overlap_log.clamp(1, 9)
};
if ov <= 1 {
return 0;
}
let window = 1usize << window_log.min(31);
window >> (9 - ov)
}
pub fn resolve_job_size(requested: usize, window_log: u32, overlap: usize) -> usize {
let window = 1usize << window_log.min(31);
let raw = if requested == 0 {
window.saturating_mul(4).max(1)
} else {
requested.max(1)
};
raw.max(JOB_SIZE_MIN).max(overlap)
}
pub fn default_nb_workers() -> u32 {
let n = std::thread::available_parallelism()
.map(|p| p.get())
.unwrap_or(1);
n.clamp(1, NB_WORKERS_MAX as usize) as u32
}
#[allow(clippy::too_many_arguments)]
pub fn compress_mt(
src: &[u8],
params: CompressionParameters,
checksum: bool,
dict: Option<&Dictionary>,
prefix: &[u8],
write_dict_id: bool,
adv: AdvancedOptions,
) -> Result<Vec<u8>, Error> {
let workers = adv.nb_workers.clamp(1, NB_WORKERS_MAX) as usize;
let overlap = overlap_size(params.window_log, adv.overlap_log, params.strategy);
let job = resolve_job_size(adv.job_size, params.window_log, overlap);
if src.is_empty() || src.len() <= job {
let mut one = adv;
one.nb_workers = 0;
return encode_oneshot(
src,
params,
checksum,
Some(src.len() as u64),
dict,
prefix,
write_dict_id,
one,
);
}
let mut ranges: Vec<(usize, usize)> = Vec::new();
let mut off = 0usize;
while off < src.len() {
let end = (off + job).min(src.len());
ranges.push((off, end));
off = end;
}
let mut job_adv = adv;
job_adv.nb_workers = 0;
let mut parts = Vec::with_capacity(ranges.len());
parts.resize(ranges.len(), Vec::new());
run_jobs(
workers,
&ranges,
|_i, (start, end)| {
let chunk = &src[start..end];
let ov_prefix: &[u8] = if dict.is_some() {
&[]
} else if start == 0 {
prefix
} else {
let ov_from = start.saturating_sub(overlap);
&src[ov_from..start]
};
let mut this_adv = job_adv;
if start != 0 && dict.is_none() {
this_adv.prime_only = true;
}
encode_oneshot(
chunk,
params,
checksum,
Some(chunk.len() as u64),
dict,
ov_prefix,
write_dict_id,
this_adv,
)
},
&mut parts,
)?;
let mut out = Vec::new();
for p in parts {
out.extend_from_slice(&p);
}
Ok(out)
}
fn run_jobs<F>(
workers: usize,
ranges: &[(usize, usize)],
f: F,
parts: &mut [Vec<u8>],
) -> Result<(), Error>
where
F: Fn(usize, (usize, usize)) -> Result<Vec<u8>, Error> + Sync,
{
#[cfg(target_arch = "wasm32")]
{
let _ = workers;
for (i, r) in ranges.iter().copied().enumerate() {
parts[i] = f(i, r)?;
}
Ok(())
}
#[cfg(not(target_arch = "wasm32"))]
{
use core::sync::atomic::{AtomicUsize, Ordering};
let n = ranges.len();
let nthreads = workers.max(1).min(n.max(1));
let next = AtomicUsize::new(0);
let mut err: Option<Error> = None;
std::thread::scope(|s| {
let mut handles = Vec::with_capacity(nthreads);
for _ in 0..nthreads {
let f = &f;
let next = &next;
handles.push(s.spawn(move || -> Result<Vec<(usize, Vec<u8>)>, Error> {
let mut done = Vec::new();
loop {
let idx = next.fetch_add(1, Ordering::Relaxed);
if idx >= n {
break;
}
done.push((idx, f(idx, ranges[idx])?));
}
Ok(done)
}));
}
for h in handles {
match h.join() {
Ok(Ok(done)) => {
for (idx, bytes) in done {
parts[idx] = bytes;
}
}
Ok(Err(e)) => err = Some(e),
Err(_) => err = Some(Error::Corruption),
}
}
});
if let Some(e) = err {
return Err(e);
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::encode::AdvancedOptions;
use crate::inspect_frames;
use crate::params::compression_params;
use crate::{decompress, FrameKind};
fn noise(n: usize) -> Vec<u8> {
let mut s = 0x4D74_u64;
let mut v = vec![0u8; n];
for b in &mut v {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
*b = (s as u8) | 1;
}
v
}
#[test]
fn overlap_size_table() {
assert_eq!(overlap_size(20, 1, Strategy::DFast), 0);
assert_eq!(overlap_size(20, 9, Strategy::DFast), 1 << 20);
assert_eq!(overlap_size(20, 8, Strategy::DFast), 1 << 19);
assert_eq!(overlap_size(20, 6, Strategy::DFast), 1 << 17);
assert_eq!(default_overlap_log(Strategy::BtUltra2), 9);
assert_eq!(default_overlap_log(Strategy::DFast), 6);
}
#[test]
fn job_size_enforces_min() {
let ov = overlap_size(18, 1, Strategy::Fast);
let j = resolve_job_size(1, 18, ov);
assert!(j >= JOB_SIZE_MIN);
}
#[test]
fn mt_two_jobs_roundtrip() {
let src = noise(JOB_SIZE_MIN + JOB_SIZE_MIN / 2);
let params = compression_params(1, Some(src.len() as u64)).unwrap();
let zst = compress_mt(
&src,
params,
true,
None,
&[],
true,
AdvancedOptions {
nb_workers: 2,
job_size: JOB_SIZE_MIN,
overlap_log: 1,
..AdvancedOptions::default()
},
)
.expect("mt");
assert_eq!(decompress(&zst).expect("decode"), src);
let frames = inspect_frames(&zst).expect("list");
let zstd_n = frames
.iter()
.filter(|f| matches!(f.kind, FrameKind::Zstd(_)))
.count();
assert!(zstd_n >= 2, "frames={zstd_n}");
}
#[test]
fn mt_overlap_roundtrip() {
let src = noise(JOB_SIZE_MIN + 64 * 1024);
let params = compression_params(1, Some(src.len() as u64)).unwrap();
let zst = compress_mt(
&src,
params,
true,
None,
&[],
true,
AdvancedOptions {
nb_workers: 2,
job_size: JOB_SIZE_MIN,
overlap_log: 9,
..AdvancedOptions::default()
},
)
.expect("mt ov");
assert_eq!(decompress(&zst).expect("decode"), src);
}
#[test]
fn single_job_is_one_frame() {
let src = noise(1024);
let params = compression_params(1, Some(src.len() as u64)).unwrap();
let zst = compress_mt(
&src,
params,
true,
None,
&[],
true,
AdvancedOptions {
nb_workers: 2,
job_size: JOB_SIZE_MIN,
overlap_log: 1,
..AdvancedOptions::default()
},
)
.unwrap();
let frames = inspect_frames(&zst).unwrap();
assert_eq!(
frames
.iter()
.filter(|f| matches!(f.kind, FrameKind::Zstd(_)))
.count(),
1
);
assert_eq!(decompress(&zst).unwrap(), src);
}
#[test]
fn mt_overlap_repeating_independent_frames() {
let src = b"cli completeness rusty_zstd. ".repeat(25_000);
let params = compression_params(1, Some(src.len() as u64)).unwrap();
let zst = compress_mt(
&src,
params,
true,
None,
&[],
true,
AdvancedOptions {
nb_workers: 2,
job_size: JOB_SIZE_MIN,
overlap_log: 9,
..AdvancedOptions::default()
},
)
.expect("mt ov text");
assert_eq!(decompress(&zst).expect("decode"), src.as_slice());
let n = inspect_frames(&zst)
.unwrap()
.iter()
.filter(|f| matches!(f.kind, FrameKind::Zstd(_)))
.count();
assert!(n >= 2, "frames={n}");
}
}