1use crate::config::{CompressionConfig, CompressionMode};
9use crate::error::{Error, Result};
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq)]
13#[repr(u8)]
14pub enum Algorithm {
15 None = 0,
16 Lz4 = 1,
19 Zstd = 2,
21 Pcm = 3,
24}
25
26impl Default for Algorithm {
27 fn default() -> Self {
28 if cfg!(feature = "zstd-codec") {
29 Algorithm::Zstd
30 } else if cfg!(feature = "lz4-codec") {
31 Algorithm::Lz4
32 } else {
33 Algorithm::None
34 }
35 }
36}
37
38impl Algorithm {
39 pub fn from_u8(v: u8) -> Result<Self> {
40 match v {
41 0 => Ok(Algorithm::None),
42 1 => Ok(Algorithm::Lz4),
43 2 => Ok(Algorithm::Zstd),
44 3 => Ok(Algorithm::Pcm),
45 other => Err(Error::Compress(format!("unknown algorithm id {other}"))),
46 }
47 }
48
49 pub fn available(self) -> bool {
50 match self {
51 Algorithm::None => true,
52 Algorithm::Lz4 => cfg!(feature = "lz4-codec"),
53 Algorithm::Zstd => cfg!(feature = "zstd-codec"),
54 Algorithm::Pcm => true,
56 }
57 }
58}
59
60#[derive(Debug, Clone, Copy, PartialEq, Eq)]
62pub struct Encoded {
63 pub algorithm: Algorithm,
64 pub raw_len: usize,
66}
67
68pub struct Codec {
79 #[cfg_attr(not(feature = "zstd-codec"), allow(dead_code))]
81 level: i32,
82 scratch: Vec<u8>,
83 #[cfg(feature = "zstd-codec")]
84 zc: Option<zstd::bulk::Compressor<'static>>,
85 #[cfg(feature = "zstd-codec")]
86 zd: Option<zstd::bulk::Decompressor<'static>>,
87}
88
89impl Default for Codec {
90 fn default() -> Self {
91 Self::new()
92 }
93}
94
95impl Codec {
96 pub fn new() -> Self {
97 Self {
98 level: i32::MIN,
99 scratch: Vec::new(),
100 #[cfg(feature = "zstd-codec")]
101 zc: None,
102 #[cfg(feature = "zstd-codec")]
103 zd: None,
104 }
105 }
106
107 pub fn compress_into(
114 &mut self,
115 cfg: &CompressionConfig,
116 hint: FileHint,
117 input: &[u8],
118 out: &mut Vec<u8>,
119 ) -> Result<Encoded> {
120 let raw_len = input.len();
121 let algo = select_algorithm(cfg, hint, input);
122
123 if algo == Algorithm::None {
124 out.extend_from_slice(input);
125 return Ok(Encoded {
126 algorithm: Algorithm::None,
127 raw_len,
128 });
129 }
130
131 let produced = self.run(algo, input, cfg.level)?;
132
133 if !worth_it(produced, raw_len, cfg.min_gain) {
136 out.extend_from_slice(input);
137 return Ok(Encoded {
138 algorithm: Algorithm::None,
139 raw_len,
140 });
141 }
142
143 out.extend_from_slice(&self.scratch[..produced]);
144 Ok(Encoded {
145 algorithm: algo,
146 raw_len,
147 })
148 }
149
150 pub fn compress_in_place(
157 &mut self,
158 cfg: &CompressionConfig,
159 hint: FileHint,
160 buf: &mut Vec<u8>,
161 prefix: usize,
162 ) -> Result<Encoded> {
163 let raw_len = buf.len() - prefix;
164
165 if let Some(fmt) = hint
171 .audio
172 .filter(|_| cfg.audio_codec && cfg.mode != CompressionMode::Off)
173 {
174 self.scratch.clear();
175 if let Some(n) =
176 super::pcm::encode(&fmt, hint.chunk_offset, &buf[prefix..], &mut self.scratch)
177 {
178 if worth_it(n, raw_len, cfg.min_gain) {
179 buf.truncate(prefix);
180 buf.extend_from_slice(&self.scratch[..n]);
181 return Ok(Encoded {
182 algorithm: Algorithm::Pcm,
183 raw_len,
184 });
185 }
186 }
187 }
188
189 let algo = select_algorithm(cfg, hint, &buf[prefix..]);
190 if algo == Algorithm::None {
191 return Ok(Encoded {
192 algorithm: Algorithm::None,
193 raw_len,
194 });
195 }
196
197 let produced = self.run_from(algo, buf, prefix, cfg.level)?;
198 if !worth_it(produced, raw_len, cfg.min_gain) {
199 return Ok(Encoded {
200 algorithm: Algorithm::None,
201 raw_len,
202 });
203 }
204
205 buf.truncate(prefix);
206 buf.extend_from_slice(&self.scratch[..produced]);
207 Ok(Encoded {
208 algorithm: algo,
209 raw_len,
210 })
211 }
212
213 pub fn decompress_into(
217 &mut self,
218 algo: Algorithm,
219 raw_len: usize,
220 input: &[u8],
221 out: &mut Vec<u8>,
222 ) -> Result<()> {
223 match algo {
224 Algorithm::None => {
225 if input.len() != raw_len {
226 return Err(Error::Compress(format!(
227 "raw chunk length {} does not match declared {}",
228 input.len(),
229 raw_len
230 )));
231 }
232 out.extend_from_slice(input);
233 Ok(())
234 }
235 Algorithm::Zstd => self.decompress_zstd(input, raw_len, out),
236 Algorithm::Lz4 => self.decompress_lz4(input, raw_len, out),
237 Algorithm::Pcm => {
238 let before = out.len();
239 super::pcm::decode(input, out)?;
240 if out.len() - before != raw_len {
241 out.truncate(before);
242 return Err(Error::Compress(format!(
243 "pcm produced {} bytes, header declared {raw_len}",
244 out.len() - before
245 )));
246 }
247 Ok(())
248 }
249 }
250 }
251
252 #[cfg_attr(not(feature = "lz4-codec"), allow(dead_code))]
255 fn scratch_at_least(&mut self, n: usize) -> &mut [u8] {
256 if self.scratch.len() < n {
257 self.scratch.resize(n, 0);
258 }
259 &mut self.scratch[..n]
260 }
261
262 fn run(&mut self, algo: Algorithm, input: &[u8], level: i32) -> Result<usize> {
263 match algo {
264 Algorithm::Zstd => self.compress_zstd(input, level),
265 Algorithm::Lz4 => self.compress_lz4(input),
266 Algorithm::Pcm | Algorithm::None => unreachable!("handled by the caller"),
269 }
270 }
271
272 fn run_from(
274 &mut self,
275 algo: Algorithm,
276 buf: &[u8],
277 prefix: usize,
278 level: i32,
279 ) -> Result<usize> {
280 self.run(algo, &buf[prefix..], level)
281 }
282}
283
284#[inline]
286fn worth_it(produced: usize, raw_len: usize, min_gain: f32) -> bool {
287 let gain = 1.0 - (produced as f32 / raw_len.max(1) as f32);
288 gain >= min_gain
289}
290
291thread_local! {
292 static TLS_CODEC: std::cell::RefCell<Codec> = std::cell::RefCell::new(Codec::new());
296}
297
298pub fn compress_into(
300 cfg: &CompressionConfig,
301 hint: FileHint,
302 input: &[u8],
303 out: &mut Vec<u8>,
304) -> Result<Encoded> {
305 TLS_CODEC.with(|c| c.borrow_mut().compress_into(cfg, hint, input, out))
306}
307
308pub fn decompress_into(
310 algo: Algorithm,
311 raw_len: usize,
312 input: &[u8],
313 out: &mut Vec<u8>,
314) -> Result<()> {
315 TLS_CODEC.with(|c| c.borrow_mut().decompress_into(algo, raw_len, input, out))
316}
317
318pub fn with_codec<R>(f: impl FnOnce(&mut Codec) -> R) -> R {
320 TLS_CODEC.with(|c| f(&mut c.borrow_mut()))
321}
322
323#[derive(Debug, Clone, Copy, Default)]
329pub struct FileHint {
330 pub known_incompressible: bool,
332 pub audio: Option<super::pcm::AudioFormat>,
335 pub chunk_offset: u64,
338}
339
340fn select_algorithm(cfg: &CompressionConfig, hint: FileHint, input: &[u8]) -> Algorithm {
341 match cfg.mode {
342 CompressionMode::Off => return Algorithm::None,
343 CompressionMode::Always => {
344 return if cfg.algorithm.available() {
345 cfg.algorithm
346 } else {
347 Algorithm::None
348 }
349 }
350 CompressionMode::Adaptive => {}
351 }
352
353 if !cfg.algorithm.available() || hint.known_incompressible {
354 return Algorithm::None;
355 }
356 if input.len() < 1024 {
358 return Algorithm::None;
359 }
360 if looks_incompressible(input, cfg.probe_bytes) {
361 return Algorithm::None;
362 }
363 cfg.algorithm
364}
365
366fn looks_incompressible(input: &[u8], probe_bytes: usize) -> bool {
374 let n = probe_bytes.min(input.len());
375 if n < 256 {
376 return false;
377 }
378 let sample = &input[..n];
380
381 let mut hist = [0u32; 256];
382 for &b in sample {
383 hist[b as usize] += 1;
384 }
385
386 let len = n as f32;
387 let mut entropy = 0.0f32;
388 for &c in hist.iter() {
389 if c != 0 {
390 let p = c as f32 / len;
391 entropy -= p * p.log2();
392 }
393 }
394
395 entropy > 7.8
398}
399
400pub fn is_incompressible_extension(cfg: &CompressionConfig, path: &str) -> bool {
402 let ext = match path.rsplit_once('.') {
403 Some((_, e)) if !e.is_empty() && e.len() <= 12 => e,
404 _ => return false,
405 };
406 let lower = ext.to_ascii_lowercase();
407 cfg.incompressible_extensions.contains(&lower)
408}
409
410impl Codec {
415 #[cfg(feature = "zstd-codec")]
416 fn compress_zstd(&mut self, input: &[u8], level: i32) -> Result<usize> {
417 if self.zc.is_none() {
418 self.zc = Some(
419 zstd::bulk::Compressor::new(level)
420 .map_err(|e| Error::Compress(format!("zstd context: {e}")))?,
421 );
422 self.level = level;
423 }
424 if self.level != level {
425 self.zc
426 .as_mut()
427 .expect("just built")
428 .set_compression_level(level)
429 .map_err(|e| Error::Compress(format!("zstd level: {e}")))?;
430 self.level = level;
431 }
432 let bound = zstd::zstd_safe::compress_bound(input.len());
433 if self.scratch.len() < bound {
434 self.scratch.resize(bound, 0);
435 }
436 let (zc, scratch) = (
437 self.zc.as_mut().expect("just built"),
438 &mut self.scratch[..bound],
439 );
440 zc.compress_to_buffer(input, scratch)
441 .map_err(|e| Error::Compress(format!("zstd: {e}")))
442 }
443
444 #[cfg(not(feature = "zstd-codec"))]
445 fn compress_zstd(&mut self, _input: &[u8], _level: i32) -> Result<usize> {
446 Err(Error::Compress("zstd support not compiled in".into()))
447 }
448
449 #[cfg(feature = "zstd-codec")]
450 fn decompress_zstd(&mut self, input: &[u8], raw_len: usize, out: &mut Vec<u8>) -> Result<()> {
451 if self.zd.is_none() {
452 self.zd = Some(
453 zstd::bulk::Decompressor::new()
454 .map_err(|e| Error::Compress(format!("zstd context: {e}")))?,
455 );
456 }
457 let before = out.len();
458 out.resize(before + raw_len, 0);
459 let written = self
460 .zd
461 .as_mut()
462 .expect("just built")
463 .decompress_to_buffer(input, &mut out[before..])
464 .map_err(|e| Error::Compress(format!("zstd decode: {e}")))?;
465 if written != raw_len {
466 out.truncate(before);
467 return Err(Error::Compress(format!(
468 "zstd produced {written} bytes, header declared {raw_len}"
469 )));
470 }
471 Ok(())
472 }
473
474 #[cfg(not(feature = "zstd-codec"))]
475 fn decompress_zstd(
476 &mut self,
477 _input: &[u8],
478 _raw_len: usize,
479 _out: &mut Vec<u8>,
480 ) -> Result<()> {
481 Err(Error::Compress(
482 "peer used zstd but zstd support is not compiled in".into(),
483 ))
484 }
485
486 #[cfg(feature = "lz4-codec")]
487 fn compress_lz4(&mut self, input: &[u8]) -> Result<usize> {
488 let bound = lz4_flex::block::get_maximum_output_size(input.len());
489 let dst = self.scratch_at_least(bound);
490 lz4_flex::block::compress_into(input, dst).map_err(|e| Error::Compress(format!("lz4: {e}")))
491 }
492
493 #[cfg(not(feature = "lz4-codec"))]
494 fn compress_lz4(&mut self, _input: &[u8]) -> Result<usize> {
495 Err(Error::Compress("lz4 support not compiled in".into()))
496 }
497
498 #[cfg(feature = "lz4-codec")]
499 fn decompress_lz4(&mut self, input: &[u8], raw_len: usize, out: &mut Vec<u8>) -> Result<()> {
500 let before = out.len();
501 out.resize(before + raw_len, 0);
502 let written = lz4_flex::block::decompress_into(input, &mut out[before..])
503 .map_err(|e| Error::Compress(format!("lz4 decode: {e}")))?;
504 if written != raw_len {
505 out.truncate(before);
506 return Err(Error::Compress(format!(
507 "lz4 produced {written} bytes, header declared {raw_len}"
508 )));
509 }
510 Ok(())
511 }
512
513 #[cfg(not(feature = "lz4-codec"))]
514 fn decompress_lz4(&mut self, _input: &[u8], _raw_len: usize, _out: &mut Vec<u8>) -> Result<()> {
515 Err(Error::Compress(
516 "peer used lz4 but lz4 support is not compiled in".into(),
517 ))
518 }
519}
520
521#[cfg(test)]
522mod tests {
523 use super::*;
524
525 fn text_chunk() -> Vec<u8> {
526 "the quick brown fox jumps over the lazy dog. "
527 .repeat(4000)
528 .into_bytes()
529 }
530
531 fn random_chunk(n: usize) -> Vec<u8> {
532 let mut s = 0x2545F4914F6CDD1Du64;
534 (0..n)
535 .map(|_| {
536 s ^= s << 13;
537 s ^= s >> 7;
538 s ^= s << 17;
539 (s >> 24) as u8
540 })
541 .collect()
542 }
543
544 #[test]
545 fn roundtrip_all_algorithms() {
546 let data = text_chunk();
547 for algo in [Algorithm::None, Algorithm::Lz4, Algorithm::Zstd] {
548 if !algo.available() {
549 continue;
550 }
551 let cfg = CompressionConfig {
552 mode: CompressionMode::Always,
553 algorithm: algo,
554 ..Default::default()
555 };
556 let mut enc = Vec::new();
557 let e = compress_into(&cfg, FileHint::default(), &data, &mut enc).unwrap();
558 let mut dec = Vec::new();
559 decompress_into(e.algorithm, e.raw_len, &enc, &mut dec).unwrap();
560 assert_eq!(dec, data, "roundtrip failed for {algo:?}");
561 }
562 }
563
564 #[test]
565 fn adaptive_skips_high_entropy_data() {
566 let cfg = CompressionConfig::default();
567 let data = random_chunk(256 * 1024);
568 let mut enc = Vec::new();
569 let e = compress_into(&cfg, FileHint::default(), &data, &mut enc).unwrap();
570 assert_eq!(e.algorithm, Algorithm::None);
571 assert_eq!(enc.len(), data.len());
573 }
574
575 #[test]
576 fn adaptive_compresses_text() {
577 let cfg = CompressionConfig::default();
578 let data = text_chunk();
579 let mut enc = Vec::new();
580 let e = compress_into(&cfg, FileHint::default(), &data, &mut enc).unwrap();
581 assert_ne!(e.algorithm, Algorithm::None);
582 assert!(enc.len() < data.len() / 2);
583 }
584
585 #[test]
586 fn extension_hint_forces_raw() {
587 let cfg = CompressionConfig::default();
588 let data = text_chunk();
589 let hint = FileHint {
590 known_incompressible: true,
591 ..Default::default()
592 };
593 let mut enc = Vec::new();
594 let e = compress_into(&cfg, hint, &data, &mut enc).unwrap();
595 assert_eq!(e.algorithm, Algorithm::None);
596 }
597
598 #[test]
599 fn extension_matching() {
600 let cfg = CompressionConfig::default();
601 assert!(is_incompressible_extension(&cfg, "song.FLAC"));
602 assert!(is_incompressible_extension(&cfg, "a/b/movie.mkv"));
603 assert!(!is_incompressible_extension(&cfg, "master.wav"));
605 assert!(!is_incompressible_extension(&cfg, "notes.txt"));
606 assert!(!is_incompressible_extension(&cfg, "no_extension"));
607 }
608
609 #[test]
610 fn decompress_rejects_length_mismatch() {
611 let cfg = CompressionConfig {
612 mode: CompressionMode::Always,
613 ..Default::default()
614 };
615 let data = text_chunk();
616 let mut enc = Vec::new();
617 let e = compress_into(&cfg, FileHint::default(), &data, &mut enc).unwrap();
618 if e.algorithm == Algorithm::None {
619 return;
620 }
621 let mut dec = Vec::new();
622 assert!(decompress_into(e.algorithm, e.raw_len / 2, &enc, &mut dec).is_err());
624 }
625}