Skip to main content

kacrab_protocol/compression/
zstd.rs

1//! Zstandard codec (`zstd` crate).
2//!
3//! The module path `crate::compression::zstd` shadows the external `zstd`
4//! crate inside this file — the codec is reached via a fully-qualified
5//! `::zstd::...` path.
6
7use super::{Compression, CompressionError, CompressionErrorKind, Result};
8
9const DEFAULT_ZSTD_LEVEL: i32 = 3;
10
11/// Compress `payload` at the given level (`None` -> codec default `3`, range `1..=22`).
12pub fn compress_with_level(payload: &[u8], level: Option<i32>) -> Result<Vec<u8>> {
13    let lvl = level.unwrap_or(DEFAULT_ZSTD_LEVEL);
14    ::zstd::bulk::compress(payload, lvl).map_err(|e| encode_err(e.to_string()))
15}
16
17/// Decompress `payload`, bounded by [`super::MAX_DECOMPRESSED_LEN`].
18pub fn decompress(payload: &[u8]) -> Result<Vec<u8>> {
19    decompress_bounded(payload, super::MAX_DECOMPRESSED_LEN)
20}
21
22/// Decompress `payload`, refusing to produce more than `max_len` bytes —
23/// zstd frames can claim enormous content sizes, so an unbounded decode is a
24/// decompression bomb.
25pub fn decompress_bounded(payload: &[u8], max_len: usize) -> Result<Vec<u8>> {
26    let decoder =
27        ::zstd::stream::read::Decoder::new(payload).map_err(|e| decode_err(e.to_string()))?;
28    super::read_to_end_bounded(decoder, max_len, Compression::Zstd)
29}
30
31const fn encode_err(message: String) -> CompressionError {
32    CompressionError::new(
33        Compression::Zstd,
34        CompressionErrorKind::EncodeFailed { message },
35    )
36}
37
38const fn decode_err(message: String) -> CompressionError {
39    CompressionError::new(
40        Compression::Zstd,
41        CompressionErrorKind::DecodeFailed { message },
42    )
43}
44
45#[cfg(test)]
46mod tests {
47    use super::{super::CompressionErrorKind, compress_with_level, decompress, decompress_bounded};
48
49    #[test]
50    fn decompress_bounded_rejects_a_decompression_bomb() {
51        let payload = vec![0u8; 4096];
52        let compressed = compress_with_level(&payload, None).unwrap();
53
54        let err = decompress_bounded(&compressed, 64).unwrap_err();
55        assert!(
56            matches!(
57                err.kind,
58                CompressionErrorKind::DecompressedTooLarge { limit: 64 }
59            ),
60            "expected DecompressedTooLarge, got {:?}",
61            err.kind
62        );
63    }
64
65    #[test]
66    fn decompress_bounded_allows_output_at_exactly_the_limit() {
67        let payload = vec![0u8; 4096];
68        let compressed = compress_with_level(&payload, None).unwrap();
69
70        assert_eq!(decompress_bounded(&compressed, 4096).unwrap(), payload);
71        assert_eq!(decompress(&compressed).unwrap(), payload);
72    }
73}