use crate::codec::{Codec, CodecTunables, PerCodecTunables};
use crate::error::CoreError;
#[derive(Clone, Debug)]
pub struct ZstdTunables {
pub quality: u8,
}
impl Default for ZstdTunables {
fn default() -> Self {
Self { quality: 2 }
}
}
fn level_for_quality(quality: u8) -> omnizip_zstd::ZstdLevel {
match quality {
0..=2 => omnizip_zstd::ZstdLevel::Fastest,
3..=5 => omnizip_zstd::ZstdLevel::Fast,
6..=11 => omnizip_zstd::ZstdLevel::Default,
12..=21 => omnizip_zstd::ZstdLevel::Better,
_ => omnizip_zstd::ZstdLevel::Best,
}
}
pub struct ZstdCodec;
const MT_WHOLE_FILE_THRESHOLD: usize = 4 * 1024 * 1024;
fn compress_verified(
plaintext: &[u8],
level: omnizip_zstd::ZstdLevel,
) -> Result<Vec<u8>, CoreError> {
let out = if plaintext.len() >= MT_WHOLE_FILE_THRESHOLD {
let threads = std::thread::available_parallelism().map_or(2, |n| std::cmp::max(2, n.get()));
omnizip_zstd::compress_mt(plaintext, level, threads)
} else {
omnizip_zstd::compress(plaintext, level)
};
out.map_err(|e| CoreError::Corrupt {
reason: format!("zstd compress (level {level}) failed: {e}"),
})
}
impl Codec for ZstdCodec {
fn id(&self) -> u8 {
super::CODEC_ZSTD
}
fn name(&self) -> &'static str {
"zstd"
}
fn compress(&self, plaintext: &[u8]) -> Result<Vec<u8>, CoreError> {
compress(plaintext)
}
fn decompress(&self, compressed: &[u8], expected_len: u32) -> Result<Vec<u8>, CoreError> {
let result =
omnizip_zstd::decompress(compressed, expected_len).map_err(|e| CoreError::Corrupt {
reason: format!("zstd decompress failed: {e}"),
})?;
let expected_us = usize::try_from(expected_len).map_err(|_| CoreError::Corrupt {
reason: format!("decompress: expected_len {expected_len} exceeds usize"),
})?;
if result.len() != expected_us {
return Err(CoreError::Corrupt {
reason: format!(
"zstd decompress: result length {} does not match plaintext_len {expected_us}",
result.len()
),
});
}
Ok(result)
}
fn compress_with_tunables(
&self,
plaintext: &[u8],
t: &CodecTunables,
) -> Result<Vec<u8>, CoreError> {
let quality = if t.zstd_quality > 0 {
t.zstd_quality
} else {
2
};
let level = level_for_quality(quality);
compress_verified(plaintext, level)
}
}
impl PerCodecTunables for ZstdCodec {
type Tunables = ZstdTunables;
fn compress_with_owned_tunables(
&self,
plaintext: &[u8],
t: &Self::Tunables,
) -> Result<Vec<u8>, CoreError> {
let level = level_for_quality(t.quality);
compress_verified(plaintext, level)
}
}
pub(crate) fn compress(plaintext: &[u8]) -> Result<Vec<u8>, CoreError> {
compress_verified(plaintext, omnizip_zstd::ZstdLevel::Fastest)
}
#[allow(dead_code)]
pub(crate) fn compress_at_level(
plaintext: &[u8],
level: omnizip_zstd::ZstdLevel,
) -> Result<Vec<u8>, CoreError> {
compress_verified(plaintext, level)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn quality_to_level_band_map_is_pinned() {
use omnizip_zstd::ZstdLevel;
assert_eq!(
level_for_quality(ZstdTunables::default().quality),
ZstdLevel::Fastest
);
for q in 0..=2 {
assert_eq!(level_for_quality(q), ZstdLevel::Fastest, "q{q}");
}
for q in 3..=5 {
assert_eq!(level_for_quality(q), ZstdLevel::Fast, "q{q}");
}
for q in 6..=11 {
assert_eq!(level_for_quality(q), ZstdLevel::Default, "q{q}");
}
for q in 12..=21 {
assert_eq!(level_for_quality(q), ZstdLevel::Better, "q{q}");
}
for q in 22..=255 {
assert_eq!(level_for_quality(q), ZstdLevel::Best, "q{q}");
}
}
#[test]
fn omnizip_315_blob_round_trips_at_all_levels() {
const B64: &str = "AgAAAAIAAAAAAAAApIEAAAAAAAAAAAAAjfBCgS8IzhiN8EKBLwjOGAEAAAAEpQAAAGR1cGxpY2F0ZSBpbmxpbmUgY29udGVudDogdGhlIHNhbWUgMjAwLWlzaCBieXRlcyBpbiB0aHJlZSBmaWxlcywgc28gdGhlIHdyaXRlcidzIGlubGluZSBkZWR1cCBmaWxlcyBvbiBldmVyeSByZWFsaXN0aWMgdHJlZS4gUGFkZGluZyBwYWRkaW5nIHBhZGRpbmcgcGFkZGluZyBwYWRkaW5nIQEAAAAAAAAA7UEAAAAAAAAAAAAABelAgS8IzhgF6UCBLwjOGAEAAAAA0n/vT8wNhb/EicVbOmpyaI3ka3H9+fam7ksII2Ipyd4BAAAAAQEAAAAJAAAAZHVwLWEudHh0AgAAAAAAAAAB";
let blob = b64(B64);
assert_eq!(blob.len(), 318);
for level in [
omnizip_zstd::ZstdLevel::Fastest,
omnizip_zstd::ZstdLevel::Fast,
omnizip_zstd::ZstdLevel::Default,
omnizip_zstd::ZstdLevel::Better,
omnizip_zstd::ZstdLevel::Best,
] {
let frame =
compress_verified(&blob, level).unwrap_or_else(|e| panic!("{level:?}: {e}"));
let back = omnizip_zstd::decompress(&frame, 318)
.unwrap_or_else(|e| panic!("{level:?} decode: {e}"));
assert_eq!(back, blob, "{level:?} must round-trip");
}
}
#[test]
fn assorted_inputs_round_trip_at_all_levels() {
let mut state = 0x0dd_ba11_5eed_u64;
let mut next = move || {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
state
};
for len in [0usize, 1, 17, 511, 318, 4096, 65536] {
let data: Vec<u8> = (0..len)
.map(|i| ((next() >> 32) ^ i as u64) as u8)
.collect();
let frame = compress(&data).expect("compress");
if data.is_empty() || frame.len() >= data.len() {
continue; }
let back = omnizip_zstd::decompress(&frame, data.len() as u32).expect("decode");
assert_eq!(back, data, "len {len} must round-trip");
}
}
fn b64(s: &str) -> Vec<u8> {
let mut out = Vec::new();
let (mut buf, mut bits) = (0u32, 0u32);
for c in s.chars() {
let v = match c {
'A'..='Z' => u32::from(c) - 65,
'a'..='z' => u32::from(c) - 71,
'0'..='9' => u32::from(c) + 4,
'+' => 62,
'/' => 63,
_ => continue,
};
buf = (buf << 6) | v;
bits += 6;
if bits >= 8 {
bits -= 8;
out.push(((buf >> bits) & 0xFF) as u8);
}
}
out
}
}
#[cfg(test)]
mod mt_tests {
use super::*;
#[test]
fn multi_thread_frames_round_trip_and_are_thread_deterministic() {
let mut payload = Vec::with_capacity(9 * 1024 * 1024);
let mut state = 0x5EED_F00Du64;
while payload.len() < 9 * 1024 * 1024 {
state ^= state << 13;
state ^= state >> 7;
payload.extend_from_slice(&state.to_le_bytes());
payload.extend_from_slice(b"mt-frame filler line\n");
}
let a =
compress_verified(&payload, omnizip_zstd::ZstdLevel::Fastest).expect("mt compress a");
let b = omnizip_zstd::compress_mt(&payload, omnizip_zstd::ZstdLevel::Fastest, 2)
.expect("two threads");
let c = omnizip_zstd::compress_mt(&payload, omnizip_zstd::ZstdLevel::Fastest, 8)
.expect("eight threads");
assert_eq!(a, b, "registry path matches explicit 2 threads");
assert_eq!(b, c, "output must not depend on thread count (>= 2)");
let back = omnizip_zstd::decompress(&a, payload.len() as u32).expect("decode mt");
assert_eq!(back, payload, "mt frames round-trip");
}
}