kacrab_protocol/compression/
snappy.rs1use super::{Compression, CompressionError, CompressionErrorKind, Result};
8
9const XERIAL_HEADER: [u8; 16] = [
10 0x82, b'S', b'N', b'A', b'P', b'P', b'Y', 0x00, 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01,
11];
12
13const BLOCK_SIZE: usize = 32 * 1024;
14
15pub fn compress_with_level(payload: &[u8], level: Option<i32>) -> Result<Vec<u8>> {
18 let _ = level;
19 let capacity = XERIAL_HEADER
20 .len()
21 .checked_add(payload.len())
22 .ok_or_else(|| encode_err("snappy input too large".into()))?;
23 let mut output = Vec::with_capacity(capacity);
24 output.extend_from_slice(&XERIAL_HEADER);
25
26 if payload.is_empty() {
27 return Ok(output);
28 }
29
30 let mut encoder = snap::raw::Encoder::new();
31 for chunk in payload.chunks(BLOCK_SIZE) {
32 let compressed = encoder
33 .compress_vec(chunk)
34 .map_err(|e| encode_err(e.to_string()))?;
35 let len = u32::try_from(compressed.len())
36 .map_err(|_| encode_err("snappy block exceeds u32".into()))?;
37 output.extend_from_slice(&len.to_be_bytes());
38 output.extend_from_slice(&compressed);
39 }
40
41 Ok(output)
42}
43
44pub fn decompress(payload: &[u8]) -> Result<Vec<u8>> {
46 decompress_bounded(payload, super::MAX_DECOMPRESSED_LEN)
47}
48
49pub fn decompress_bounded(payload: &[u8], max_len: usize) -> Result<Vec<u8>> {
55 if payload.len() < XERIAL_HEADER.len()
56 || payload.get(..XERIAL_HEADER.len()) != Some(&XERIAL_HEADER)
57 {
58 check_claimed_len(payload, max_len, 0)?;
59 return snap::raw::Decoder::new()
60 .decompress_vec(payload)
61 .map_err(|e| decode_err(e.to_string()));
62 }
63
64 let mut decoder = snap::raw::Decoder::new();
65 let mut output = Vec::new();
66 let mut pos = XERIAL_HEADER.len();
67
68 while pos < payload.len() {
69 let length_end = pos
70 .checked_add(4)
71 .ok_or_else(|| decode_err("snappy: block length offset overflow".into()))?;
72 let Some(len_bytes) = payload.get(pos..length_end).and_then(|s| s.try_into().ok()) else {
73 return Err(decode_err("snappy: truncated block length".into()));
74 };
75 let block_len = u32::from_be_bytes(len_bytes) as usize;
76 pos = length_end;
77
78 let block_end = pos
79 .checked_add(block_len)
80 .ok_or_else(|| decode_err("snappy: block length offset overflow".into()))?;
81 if block_end > payload.len() {
82 let remaining = payload.len().saturating_sub(pos);
83 return Err(decode_err(format!(
84 "snappy: block length {block_len} extends past input (remaining: {remaining})"
85 )));
86 }
87
88 let Some(block) = payload.get(pos..block_end) else {
89 return Err(decode_err("snappy: block out of range".into()));
90 };
91 check_claimed_len(block, max_len, output.len())?;
92 let decompressed = decoder
93 .decompress_vec(block)
94 .map_err(|e| decode_err(e.to_string()))?;
95 output.extend_from_slice(&decompressed);
96 pos = block_end;
97 }
98
99 Ok(output)
100}
101
102fn check_claimed_len(block: &[u8], max_len: usize, already_produced: usize) -> Result<()> {
105 let claimed = snap::raw::decompress_len(block).map_err(|e| decode_err(e.to_string()))?;
106 if already_produced.saturating_add(claimed) > max_len {
107 return Err(CompressionError::new(
108 Compression::Snappy,
109 CompressionErrorKind::DecompressedTooLarge { limit: max_len },
110 ));
111 }
112 Ok(())
113}
114
115const fn encode_err(message: String) -> CompressionError {
116 CompressionError::new(
117 Compression::Snappy,
118 CompressionErrorKind::EncodeFailed { message },
119 )
120}
121
122const fn decode_err(message: String) -> CompressionError {
123 CompressionError::new(
124 Compression::Snappy,
125 CompressionErrorKind::DecodeFailed { message },
126 )
127}
128
129#[cfg(test)]
130mod tests {
131 use super::{super::CompressionErrorKind, compress_with_level, decompress, decompress_bounded};
132
133 #[test]
134 fn decompress_bounded_rejects_a_decompression_bomb() {
135 let payload = vec![0u8; 96 * 1024];
137 let compressed = compress_with_level(&payload, None).unwrap();
138
139 let err = decompress_bounded(&compressed, 64).unwrap_err();
140 assert!(
141 matches!(
142 err.kind,
143 CompressionErrorKind::DecompressedTooLarge { limit: 64 }
144 ),
145 "expected DecompressedTooLarge, got {:?}",
146 err.kind
147 );
148 }
149
150 #[test]
151 fn decompress_bounded_rejects_a_raw_block_with_a_hostile_claimed_length() {
152 let raw = snap::raw::Encoder::new()
155 .compress_vec(&[0u8; 1024])
156 .unwrap();
157
158 let err = decompress_bounded(&raw, 64).unwrap_err();
159 assert!(
160 matches!(
161 err.kind,
162 CompressionErrorKind::DecompressedTooLarge { limit: 64 }
163 ),
164 "expected DecompressedTooLarge, got {:?}",
165 err.kind
166 );
167 }
168
169 #[test]
170 fn decompress_bounded_allows_output_at_exactly_the_limit() {
171 let payload = vec![0u8; 96 * 1024];
172 let compressed = compress_with_level(&payload, None).unwrap();
173
174 assert_eq!(decompress_bounded(&compressed, 96 * 1024).unwrap(), payload);
175 assert_eq!(decompress(&compressed).unwrap(), payload);
176 }
177}