Skip to main content

urna_format/encoding/
zstd_codec.rs

1//! zstd compression for non-embedding sections. Embeddings are never
2//! zstd-compressed - they live in mmap and the runtime reads them via
3//! SIMD straight from disk.
4
5use crate::error::UrnaError;
6
7/// Default zstd compression level. 19 is in the "high" tier - slow to
8/// encode but a one-time cost and yields ~30% smaller text payloads
9/// than the default level 3.
10pub const DEFAULT_ZSTD_LEVEL: i32 = 19;
11
12/// Compress with zstd at `DEFAULT_ZSTD_LEVEL`. Returns the compressed
13/// bytes ready to write as the section payload.
14pub fn zstd_encode(bytes: &[u8]) -> crate::Result<Vec<u8>> {
15    zstd::encode_all(bytes, DEFAULT_ZSTD_LEVEL)
16        .map_err(|e| UrnaError::InvalidInput(format!("zstd compression failed: {}", e)))
17}
18
19/// hostile-input guard: an attacker can ship a tiny zstd frame that inflates
20/// to many GiB and OOM-kills the process at file `open()` (a classic
21/// decompression bomb). The decompressed size is bounded to
22/// `max(MIN_DECOMPRESS_CAP, MAX_DECOMPRESS_RATIO * compressed_len)`: the floor
23/// keeps small legitimate sections working, the ratio bounds amplification on
24/// larger inputs. The floor matches the per-stream `STREAM_CAP` the dict codec
25/// already uses (`zstd_dict.rs`).
26const MIN_DECOMPRESS_CAP: usize = 64 * 1024 * 1024;
27const MAX_DECOMPRESS_RATIO: usize = 128;
28
29fn decompress_cap(compressed_len: usize) -> usize {
30    compressed_len
31        .saturating_mul(MAX_DECOMPRESS_RATIO)
32        .max(MIN_DECOMPRESS_CAP)
33}
34
35/// Decompress a zstd payload. Internal helper used by `decode_payload`.
36/// Bounded by [`decompress_cap`] so a decompression bomb cannot exhaust memory
37/// at open time: a frame that declares (or expands past) more than the cap is
38/// rejected before the large allocation, never a panic.
39pub(super) fn zstd_decode(bytes: &[u8]) -> crate::Result<Vec<u8>> {
40    let cap = decompress_cap(bytes.len());
41    let bomb = |reason: String| UrnaError::MalformedSectionPayload {
42        section_id: 0,
43        reason,
44    };
45    let fail = |e: std::io::Error| bomb(format!("zstd decompression failed: {}", e));
46    match zstd::bulk::Decompressor::upper_bound(bytes) {
47        // declared frame size already exceeds the cap: refuse before allocating.
48        Some(declared) if declared > cap => Err(bomb(format!(
49            "zstd frame declares {} bytes, exceeds cap {}",
50            declared, cap
51        ))),
52        // declared size known and within cap: allocate exactly it.
53        Some(declared) => zstd::bulk::decompress(bytes, declared).map_err(fail),
54        // no size hint: allow up to the cap, erroring if the frame exceeds it.
55        None => zstd::bulk::decompress(bytes, cap).map_err(fail),
56    }
57}
58
59#[cfg(test)]
60mod tests {
61    use super::*;
62
63    #[test]
64    #[cfg_attr(miri, ignore)] // zstd is c code, miri cannot call it
65    fn roundtrip_within_cap() {
66        let original = b"the file is the database ".repeat(256);
67        let compressed = zstd_encode(&original).unwrap();
68        let decoded = zstd_decode(&compressed).unwrap();
69        assert_eq!(decoded, original);
70    }
71
72    #[test]
73    #[cfg_attr(miri, ignore)] // zstd is c code, miri cannot call it
74    fn decompression_bomb_is_rejected_not_ooming() {
75        // a tiny compressed frame declaring more than the cap must error,
76        // never inflate. zeros compress to a few bytes but declare their full
77        // size in the frame header, which `upper_bound` reads.
78        let bomb = vec![0u8; MIN_DECOMPRESS_CAP + 1];
79        let compressed = zstd_encode(&bomb).unwrap();
80        assert!(compressed.len() < 4096, "bomb should compress tiny");
81        let err = zstd_decode(&compressed).unwrap_err();
82        assert!(matches!(err, UrnaError::MalformedSectionPayload { .. }));
83    }
84
85    #[test]
86    fn cap_has_floor_and_ratio() {
87        assert_eq!(decompress_cap(0), MIN_DECOMPRESS_CAP);
88        assert_eq!(decompress_cap(usize::MAX), usize::MAX); // saturating, no overflow
89        assert_eq!(
90            decompress_cap(MIN_DECOMPRESS_CAP),
91            MIN_DECOMPRESS_CAP * MAX_DECOMPRESS_RATIO
92        );
93    }
94}