urna_format/encoding/
zstd_codec.rs1use crate::error::UrnaError;
6
7pub const DEFAULT_ZSTD_LEVEL: i32 = 19;
11
12pub 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
19const 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
35pub(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 Some(declared) if declared > cap => Err(bomb(format!(
49 "zstd frame declares {} bytes, exceeds cap {}",
50 declared, cap
51 ))),
52 Some(declared) => zstd::bulk::decompress(bytes, declared).map_err(fail),
54 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)] 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)] fn decompression_bomb_is_rejected_not_ooming() {
75 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); assert_eq!(
90 decompress_cap(MIN_DECOMPRESS_CAP),
91 MIN_DECOMPRESS_CAP * MAX_DECOMPRESS_RATIO
92 );
93 }
94}