1use bytes::Bytes;
34use flate2::{Compress, Decompress, FlushCompress, FlushDecompress, Status};
35
36pub const DEFAULT_LEVEL: u32 = 6;
39
40pub const DEFAULT_MAX_FRAME_SIZE: u64 = 64 * 1024 * 1024;
42
43const SYNC_FLUSH_TAIL: [u8; 4] = [0x00, 0x00, 0xff, 0xff];
45
46const CHUNK: usize = 8 * 1024;
48
49#[derive(thiserror::Error, Debug, Clone, PartialEq, Eq)]
51#[non_exhaustive]
52pub enum Error {
53 #[error("decompression failed")]
55 Decompress,
56
57 #[error("decompressed frame exceeded {0} bytes")]
59 TooLarge(u64),
60}
61
62pub type Result<T> = std::result::Result<T, Error>;
64
65pub struct Encoder(Box<Compress>);
70
71impl Encoder {
72 pub fn new() -> Self {
74 Self::with_level(DEFAULT_LEVEL)
75 }
76
77 pub fn with_level(level: u32) -> Self {
80 Self(Box::new(Compress::new(flate2::Compression::new(level.min(9)), false)))
82 }
83
84 pub fn frame(&mut self, payload: &[u8]) -> Bytes {
88 if payload.is_empty() {
89 return Bytes::new();
90 }
91
92 let mut out = Vec::with_capacity(payload.len() / 2 + 16);
93 let mut tmp = [0u8; CHUNK];
94 let mut input = payload;
95
96 loop {
99 let before_in = self.0.total_in();
100 let before_out = self.0.total_out();
101 self.0.compress(input, &mut tmp, FlushCompress::Sync).expect("deflate");
102 let consumed = (self.0.total_in() - before_in) as usize;
103 let produced = (self.0.total_out() - before_out) as usize;
104 out.extend_from_slice(&tmp[..produced]);
105 input = &input[consumed..];
106 if produced < tmp.len() {
107 break;
108 }
109 }
110
111 assert!(
115 out.ends_with(&SYNC_FLUSH_TAIL),
116 "a sync flush must end in the deflate marker"
117 );
118 out.truncate(out.len() - SYNC_FLUSH_TAIL.len());
119 Bytes::from(out)
120 }
121}
122
123impl Default for Encoder {
124 fn default() -> Self {
125 Self::new()
126 }
127}
128
129pub struct Decoder {
132 inner: Box<Decompress>,
133 max_frame_size: u64,
134}
135
136impl Decoder {
137 pub fn new() -> Self {
139 Self::with_max_frame_size(DEFAULT_MAX_FRAME_SIZE)
140 }
141
142 pub fn with_max_frame_size(max_frame_size: u64) -> Self {
147 Self {
149 inner: Box::new(Decompress::new(false)),
150 max_frame_size,
151 }
152 }
153
154 pub fn frame(&mut self, slice: &[u8]) -> Result<Bytes> {
160 let mut out = Vec::new();
161 self.frame_into(slice, &mut out)?;
162 Ok(Bytes::from(out))
163 }
164
165 pub fn frame_into(&mut self, slice: &[u8], out: &mut Vec<u8>) -> Result<()> {
169 out.clear();
170 if slice.is_empty() {
171 return Ok(());
172 }
173 let mut tmp = [0u8; CHUNK];
174
175 for segment in [slice, &SYNC_FLUSH_TAIL] {
178 let mut input = segment;
179 loop {
180 let before_in = self.inner.total_in();
181 let before_out = self.inner.total_out();
182 let status = self
183 .inner
184 .decompress(input, &mut tmp, FlushDecompress::Sync)
185 .map_err(|_| Error::Decompress)?;
186 let consumed = (self.inner.total_in() - before_in) as usize;
187 let produced = (self.inner.total_out() - before_out) as usize;
188 if out.len() as u64 + produced as u64 > self.max_frame_size {
190 return Err(Error::TooLarge(self.max_frame_size));
191 }
192 out.extend_from_slice(&tmp[..produced]);
193 input = &input[consumed..];
194
195 if matches!(status, Status::StreamEnd) || (input.is_empty() && produced < tmp.len()) {
198 break;
199 }
200 if consumed == 0 && produced == 0 {
201 break;
202 }
203 }
204 }
205
206 Ok(())
207 }
208}
209
210impl Default for Decoder {
211 fn default() -> Self {
212 Self::new()
213 }
214}
215
216#[cfg(test)]
217mod test {
218 use super::*;
219
220 fn roundtrip(frames: &[&[u8]]) -> Vec<Vec<u8>> {
222 let mut enc = Encoder::new();
223 let slices: Vec<Bytes> = frames.iter().map(|f| enc.frame(f)).collect();
224
225 let mut dec = Decoder::new();
226 slices.iter().map(|s| dec.frame(s).unwrap().to_vec()).collect()
227 }
228
229 #[test]
230 fn stream_roundtrip() {
231 let frames: &[&[u8]] = &[b"the quick brown fox", b"the quick brown dog", b"the lazy fox"];
232 let got = roundtrip(frames);
233 for (a, b) in frames.iter().zip(&got) {
234 assert_eq!(*a, b.as_slice());
235 }
236 }
237
238 #[test]
239 fn frame_into_reuses_its_output_buffer() {
240 let mut encoder = Encoder::new();
241 let frames = [b"first payload".as_slice(), b"second payload".as_slice()];
242 let mut decoder = Decoder::new();
243 let mut out = Vec::with_capacity(64);
244 let ptr = out.as_ptr();
245 for frame in frames {
246 decoder.frame_into(&encoder.frame(frame), &mut out).unwrap();
247 assert_eq!(out, frame);
248 assert_eq!(out.as_ptr(), ptr);
249 }
250 decoder.frame_into(b"", &mut out).unwrap();
251 assert!(out.is_empty());
252 }
253
254 #[test]
255 fn empty_frames_roundtrip() {
256 assert!(Encoder::new().frame(b"").is_empty());
257 assert!(Decoder::new().frame(b"").unwrap().is_empty());
258 }
259
260 #[test]
261 fn cross_frame_context_shrinks() {
262 let payload = b"Media over QUIC delivers real-time latency at massive scale.".repeat(6);
265 let mut enc = Encoder::new();
266 let first = enc.frame(&payload);
267 let second = enc.frame(&payload);
268 assert!(
269 second.len() < first.len(),
270 "repeat frame {} should be smaller than first {}",
271 second.len(),
272 first.len()
273 );
274 }
275
276 fn noise(len: usize) -> Vec<u8> {
278 let mut state: u64 = 0x9E37_79B9_7F4A_7C15;
279 (0..len)
280 .map(|_| {
281 state ^= state << 13;
282 state ^= state >> 7;
283 state ^= state << 17;
284 (state >> 56) as u8
285 })
286 .collect()
287 }
288
289 #[test]
290 fn frame_larger_than_chunk_roundtrips() {
291 let payload = noise(64 * 1024);
294
295 let mut enc = Encoder::new();
296 let slice = enc.frame(&payload);
297 assert!(slice.len() > CHUNK, "slice {} should exceed CHUNK {CHUNK}", slice.len());
298
299 let mut dec = Decoder::new();
300 assert_eq!(dec.frame(&slice).unwrap(), Bytes::from(payload));
301 }
302
303 #[test]
304 fn block_boundary_at_frame_end_roundtrips() {
305 let lens: Vec<usize> = (31 * 1024..32 * 1024 + 256).step_by(16).collect();
310 let noise = noise(lens.iter().sum());
311
312 let mut enc = Encoder::new();
313 let mut dec = Decoder::new();
314 let mut rest = noise.as_slice();
315 for len in lens {
316 let (frame, next) = rest.split_at(len);
317 rest = next;
318 let got = dec
319 .frame(&enc.frame(frame))
320 .unwrap_or_else(|err| panic!("{len} byte frame: {err}"));
321 assert!(got == frame, "{len} byte frame corrupted");
322 }
323 }
324
325 #[test]
326 fn decompress_rejects_garbage() {
327 let mut dec = Decoder::new();
328 assert_eq!(dec.frame(b"not a deflate stream at all"), Err(Error::Decompress));
329 }
330
331 #[test]
332 fn enforces_max_frame_size() {
333 let payload = vec![0u8; 1024];
335 let slice = Encoder::new().frame(&payload);
336
337 let mut dec = Decoder::with_max_frame_size(512);
338 assert_eq!(dec.frame(&slice), Err(Error::TooLarge(512)));
339 }
340
341 #[test]
342 fn custom_level_roundtrips() {
343 let payload = b"compress me at maximum effort".repeat(8);
344 let mut enc = Encoder::with_level(9);
345 let slice = enc.frame(&payload);
346 let mut dec = Decoder::new();
347 assert_eq!(dec.frame(&slice).unwrap(), Bytes::from(payload));
348 }
349}