use crate::config::{CompressionConfig, CompressionMode};
use crate::error::{Error, Result};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum Algorithm {
None = 0,
Lz4 = 1,
Zstd = 2,
Pcm = 3,
}
impl Default for Algorithm {
fn default() -> Self {
if cfg!(feature = "zstd-codec") {
Algorithm::Zstd
} else if cfg!(feature = "lz4-codec") {
Algorithm::Lz4
} else {
Algorithm::None
}
}
}
impl Algorithm {
pub fn from_u8(v: u8) -> Result<Self> {
match v {
0 => Ok(Algorithm::None),
1 => Ok(Algorithm::Lz4),
2 => Ok(Algorithm::Zstd),
3 => Ok(Algorithm::Pcm),
other => Err(Error::Compress(format!("unknown algorithm id {other}"))),
}
}
pub fn available(self) -> bool {
match self {
Algorithm::None => true,
Algorithm::Lz4 => cfg!(feature = "lz4-codec"),
Algorithm::Zstd => cfg!(feature = "zstd-codec"),
Algorithm::Pcm => true,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Encoded {
pub algorithm: Algorithm,
pub raw_len: usize,
}
pub struct Codec {
#[cfg_attr(not(feature = "zstd-codec"), allow(dead_code))]
level: i32,
scratch: Vec<u8>,
#[cfg(feature = "zstd-codec")]
zc: Option<zstd::bulk::Compressor<'static>>,
#[cfg(feature = "zstd-codec")]
zd: Option<zstd::bulk::Decompressor<'static>>,
}
impl Default for Codec {
fn default() -> Self {
Self::new()
}
}
impl Codec {
pub fn new() -> Self {
Self {
level: i32::MIN,
scratch: Vec::new(),
#[cfg(feature = "zstd-codec")]
zc: None,
#[cfg(feature = "zstd-codec")]
zd: None,
}
}
pub fn compress_into(
&mut self,
cfg: &CompressionConfig,
hint: FileHint,
input: &[u8],
out: &mut Vec<u8>,
) -> Result<Encoded> {
let raw_len = input.len();
let algo = select_algorithm(cfg, hint, input);
if algo == Algorithm::None {
out.extend_from_slice(input);
return Ok(Encoded {
algorithm: Algorithm::None,
raw_len,
});
}
let produced = self.run(algo, input, cfg.level)?;
if !worth_it(produced, raw_len, cfg.min_gain) {
out.extend_from_slice(input);
return Ok(Encoded {
algorithm: Algorithm::None,
raw_len,
});
}
out.extend_from_slice(&self.scratch[..produced]);
Ok(Encoded {
algorithm: algo,
raw_len,
})
}
pub fn compress_in_place(
&mut self,
cfg: &CompressionConfig,
hint: FileHint,
buf: &mut Vec<u8>,
prefix: usize,
) -> Result<Encoded> {
let raw_len = buf.len() - prefix;
if let Some(fmt) = hint
.audio
.filter(|_| cfg.audio_codec && cfg.mode != CompressionMode::Off)
{
self.scratch.clear();
if let Some(n) =
super::pcm::encode(&fmt, hint.chunk_offset, &buf[prefix..], &mut self.scratch)
{
if worth_it(n, raw_len, cfg.min_gain) {
buf.truncate(prefix);
buf.extend_from_slice(&self.scratch[..n]);
return Ok(Encoded {
algorithm: Algorithm::Pcm,
raw_len,
});
}
}
}
let algo = select_algorithm(cfg, hint, &buf[prefix..]);
if algo == Algorithm::None {
return Ok(Encoded {
algorithm: Algorithm::None,
raw_len,
});
}
let produced = self.run_from(algo, buf, prefix, cfg.level)?;
if !worth_it(produced, raw_len, cfg.min_gain) {
return Ok(Encoded {
algorithm: Algorithm::None,
raw_len,
});
}
buf.truncate(prefix);
buf.extend_from_slice(&self.scratch[..produced]);
Ok(Encoded {
algorithm: algo,
raw_len,
})
}
pub fn decompress_into(
&mut self,
algo: Algorithm,
raw_len: usize,
input: &[u8],
out: &mut Vec<u8>,
) -> Result<()> {
match algo {
Algorithm::None => {
if input.len() != raw_len {
return Err(Error::Compress(format!(
"raw chunk length {} does not match declared {}",
input.len(),
raw_len
)));
}
out.extend_from_slice(input);
Ok(())
}
Algorithm::Zstd => self.decompress_zstd(input, raw_len, out),
Algorithm::Lz4 => self.decompress_lz4(input, raw_len, out),
Algorithm::Pcm => {
let before = out.len();
super::pcm::decode(input, out)?;
if out.len() - before != raw_len {
out.truncate(before);
return Err(Error::Compress(format!(
"pcm produced {} bytes, header declared {raw_len}",
out.len() - before
)));
}
Ok(())
}
}
}
#[cfg_attr(not(feature = "lz4-codec"), allow(dead_code))]
fn scratch_at_least(&mut self, n: usize) -> &mut [u8] {
if self.scratch.len() < n {
self.scratch.resize(n, 0);
}
&mut self.scratch[..n]
}
fn run(&mut self, algo: Algorithm, input: &[u8], level: i32) -> Result<usize> {
match algo {
Algorithm::Zstd => self.compress_zstd(input, level),
Algorithm::Lz4 => self.compress_lz4(input),
Algorithm::Pcm | Algorithm::None => unreachable!("handled by the caller"),
}
}
fn run_from(
&mut self,
algo: Algorithm,
buf: &[u8],
prefix: usize,
level: i32,
) -> Result<usize> {
self.run(algo, &buf[prefix..], level)
}
}
#[inline]
fn worth_it(produced: usize, raw_len: usize, min_gain: f32) -> bool {
let gain = 1.0 - (produced as f32 / raw_len.max(1) as f32);
gain >= min_gain
}
thread_local! {
static TLS_CODEC: std::cell::RefCell<Codec> = std::cell::RefCell::new(Codec::new());
}
pub fn compress_into(
cfg: &CompressionConfig,
hint: FileHint,
input: &[u8],
out: &mut Vec<u8>,
) -> Result<Encoded> {
TLS_CODEC.with(|c| c.borrow_mut().compress_into(cfg, hint, input, out))
}
pub fn decompress_into(
algo: Algorithm,
raw_len: usize,
input: &[u8],
out: &mut Vec<u8>,
) -> Result<()> {
TLS_CODEC.with(|c| c.borrow_mut().decompress_into(algo, raw_len, input, out))
}
pub fn with_codec<R>(f: impl FnOnce(&mut Codec) -> R) -> R {
TLS_CODEC.with(|c| f(&mut c.borrow_mut()))
}
#[derive(Debug, Clone, Copy, Default)]
pub struct FileHint {
pub known_incompressible: bool,
pub audio: Option<super::pcm::AudioFormat>,
pub chunk_offset: u64,
}
fn select_algorithm(cfg: &CompressionConfig, hint: FileHint, input: &[u8]) -> Algorithm {
match cfg.mode {
CompressionMode::Off => return Algorithm::None,
CompressionMode::Always => {
return if cfg.algorithm.available() {
cfg.algorithm
} else {
Algorithm::None
}
}
CompressionMode::Adaptive => {}
}
if !cfg.algorithm.available() || hint.known_incompressible {
return Algorithm::None;
}
if input.len() < 1024 {
return Algorithm::None;
}
if looks_incompressible(input, cfg.probe_bytes) {
return Algorithm::None;
}
cfg.algorithm
}
fn looks_incompressible(input: &[u8], probe_bytes: usize) -> bool {
let n = probe_bytes.min(input.len());
if n < 256 {
return false;
}
let sample = &input[..n];
let mut hist = [0u32; 256];
for &b in sample {
hist[b as usize] += 1;
}
let len = n as f32;
let mut entropy = 0.0f32;
for &c in hist.iter() {
if c != 0 {
let p = c as f32 / len;
entropy -= p * p.log2();
}
}
entropy > 7.8
}
pub fn is_incompressible_extension(cfg: &CompressionConfig, path: &str) -> bool {
let ext = match path.rsplit_once('.') {
Some((_, e)) if !e.is_empty() && e.len() <= 12 => e,
_ => return false,
};
let lower = ext.to_ascii_lowercase();
cfg.incompressible_extensions.contains(&lower)
}
impl Codec {
#[cfg(feature = "zstd-codec")]
fn compress_zstd(&mut self, input: &[u8], level: i32) -> Result<usize> {
if self.zc.is_none() {
self.zc = Some(
zstd::bulk::Compressor::new(level)
.map_err(|e| Error::Compress(format!("zstd context: {e}")))?,
);
self.level = level;
}
if self.level != level {
self.zc
.as_mut()
.expect("just built")
.set_compression_level(level)
.map_err(|e| Error::Compress(format!("zstd level: {e}")))?;
self.level = level;
}
let bound = zstd::zstd_safe::compress_bound(input.len());
if self.scratch.len() < bound {
self.scratch.resize(bound, 0);
}
let (zc, scratch) = (
self.zc.as_mut().expect("just built"),
&mut self.scratch[..bound],
);
zc.compress_to_buffer(input, scratch)
.map_err(|e| Error::Compress(format!("zstd: {e}")))
}
#[cfg(not(feature = "zstd-codec"))]
fn compress_zstd(&mut self, _input: &[u8], _level: i32) -> Result<usize> {
Err(Error::Compress("zstd support not compiled in".into()))
}
#[cfg(feature = "zstd-codec")]
fn decompress_zstd(&mut self, input: &[u8], raw_len: usize, out: &mut Vec<u8>) -> Result<()> {
if self.zd.is_none() {
self.zd = Some(
zstd::bulk::Decompressor::new()
.map_err(|e| Error::Compress(format!("zstd context: {e}")))?,
);
}
let before = out.len();
out.resize(before + raw_len, 0);
let written = self
.zd
.as_mut()
.expect("just built")
.decompress_to_buffer(input, &mut out[before..])
.map_err(|e| Error::Compress(format!("zstd decode: {e}")))?;
if written != raw_len {
out.truncate(before);
return Err(Error::Compress(format!(
"zstd produced {written} bytes, header declared {raw_len}"
)));
}
Ok(())
}
#[cfg(not(feature = "zstd-codec"))]
fn decompress_zstd(
&mut self,
_input: &[u8],
_raw_len: usize,
_out: &mut Vec<u8>,
) -> Result<()> {
Err(Error::Compress(
"peer used zstd but zstd support is not compiled in".into(),
))
}
#[cfg(feature = "lz4-codec")]
fn compress_lz4(&mut self, input: &[u8]) -> Result<usize> {
let bound = lz4_flex::block::get_maximum_output_size(input.len());
let dst = self.scratch_at_least(bound);
lz4_flex::block::compress_into(input, dst).map_err(|e| Error::Compress(format!("lz4: {e}")))
}
#[cfg(not(feature = "lz4-codec"))]
fn compress_lz4(&mut self, _input: &[u8]) -> Result<usize> {
Err(Error::Compress("lz4 support not compiled in".into()))
}
#[cfg(feature = "lz4-codec")]
fn decompress_lz4(&mut self, input: &[u8], raw_len: usize, out: &mut Vec<u8>) -> Result<()> {
let before = out.len();
out.resize(before + raw_len, 0);
let written = lz4_flex::block::decompress_into(input, &mut out[before..])
.map_err(|e| Error::Compress(format!("lz4 decode: {e}")))?;
if written != raw_len {
out.truncate(before);
return Err(Error::Compress(format!(
"lz4 produced {written} bytes, header declared {raw_len}"
)));
}
Ok(())
}
#[cfg(not(feature = "lz4-codec"))]
fn decompress_lz4(&mut self, _input: &[u8], _raw_len: usize, _out: &mut Vec<u8>) -> Result<()> {
Err(Error::Compress(
"peer used lz4 but lz4 support is not compiled in".into(),
))
}
}
#[cfg(test)]
mod tests {
use super::*;
fn text_chunk() -> Vec<u8> {
"the quick brown fox jumps over the lazy dog. "
.repeat(4000)
.into_bytes()
}
fn random_chunk(n: usize) -> Vec<u8> {
let mut s = 0x2545F4914F6CDD1Du64;
(0..n)
.map(|_| {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
(s >> 24) as u8
})
.collect()
}
#[test]
fn roundtrip_all_algorithms() {
let data = text_chunk();
for algo in [Algorithm::None, Algorithm::Lz4, Algorithm::Zstd] {
if !algo.available() {
continue;
}
let cfg = CompressionConfig {
mode: CompressionMode::Always,
algorithm: algo,
..Default::default()
};
let mut enc = Vec::new();
let e = compress_into(&cfg, FileHint::default(), &data, &mut enc).unwrap();
let mut dec = Vec::new();
decompress_into(e.algorithm, e.raw_len, &enc, &mut dec).unwrap();
assert_eq!(dec, data, "roundtrip failed for {algo:?}");
}
}
#[test]
fn adaptive_skips_high_entropy_data() {
let cfg = CompressionConfig::default();
let data = random_chunk(256 * 1024);
let mut enc = Vec::new();
let e = compress_into(&cfg, FileHint::default(), &data, &mut enc).unwrap();
assert_eq!(e.algorithm, Algorithm::None);
assert_eq!(enc.len(), data.len());
}
#[test]
fn adaptive_compresses_text() {
let cfg = CompressionConfig::default();
let data = text_chunk();
let mut enc = Vec::new();
let e = compress_into(&cfg, FileHint::default(), &data, &mut enc).unwrap();
assert_ne!(e.algorithm, Algorithm::None);
assert!(enc.len() < data.len() / 2);
}
#[test]
fn extension_hint_forces_raw() {
let cfg = CompressionConfig::default();
let data = text_chunk();
let hint = FileHint {
known_incompressible: true,
..Default::default()
};
let mut enc = Vec::new();
let e = compress_into(&cfg, hint, &data, &mut enc).unwrap();
assert_eq!(e.algorithm, Algorithm::None);
}
#[test]
fn extension_matching() {
let cfg = CompressionConfig::default();
assert!(is_incompressible_extension(&cfg, "song.FLAC"));
assert!(is_incompressible_extension(&cfg, "a/b/movie.mkv"));
assert!(!is_incompressible_extension(&cfg, "master.wav"));
assert!(!is_incompressible_extension(&cfg, "notes.txt"));
assert!(!is_incompressible_extension(&cfg, "no_extension"));
}
#[test]
fn decompress_rejects_length_mismatch() {
let cfg = CompressionConfig {
mode: CompressionMode::Always,
..Default::default()
};
let data = text_chunk();
let mut enc = Vec::new();
let e = compress_into(&cfg, FileHint::default(), &data, &mut enc).unwrap();
if e.algorithm == Algorithm::None {
return;
}
let mut dec = Vec::new();
assert!(decompress_into(e.algorithm, e.raw_len / 2, &enc, &mut dec).is_err());
}
}