1use crate::{Error, Result};
21
22use crate::checksum::Adler32;
23
24#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
26pub struct Level(u8);
27
28impl Level {
29 pub const NONE: Self = Self(0);
31 pub const FAST: Self = Self(1);
33 pub const DEFAULT: Self = Self(6);
35 pub const BEST: Self = Self(9);
37
38 pub fn new(level: u8) -> Result<Self> {
44 if level > 9 {
45 return Err(Error::malformed(
46 "level",
47 format!("compression level must be 0..=9, got {level}"),
48 ));
49 }
50 Ok(Self(level))
51 }
52
53 #[must_use]
55 pub const fn get(self) -> u8 {
56 self.0
57 }
58
59 const fn search_depth(self) -> usize {
61 match self.0 {
62 0 => 0,
63 1 => 4,
64 2 => 8,
65 3 => 16,
66 4 => 32,
67 5 => 64,
68 6 => 128,
69 7 => 256,
70 8 => 512,
71 _ => 1024,
72 }
73 }
74}
75
76impl Default for Level {
77 fn default() -> Self {
78 Self::DEFAULT
79 }
80}
81
82#[derive(Debug, Default)]
84struct BitWriter {
85 out: Vec<u8>,
86 bits: u32,
87 count: u32,
88}
89
90impl BitWriter {
91 fn write(&mut self, value: u32, n: u32) {
93 self.bits |= value << self.count;
94 self.count += n;
95 while self.count >= 8 {
96 self.out.push((self.bits & 0xFF) as u8);
97 self.bits >>= 8;
98 self.count -= 8;
99 }
100 }
101
102 fn write_code(&mut self, value: u32, n: u32) {
104 for i in (0..n).rev() {
105 self.write((value >> i) & 1, 1);
106 }
107 }
108
109 fn align(&mut self) {
111 if self.count > 0 {
112 self.out.push((self.bits & 0xFF) as u8);
113 self.bits = 0;
114 self.count = 0;
115 }
116 }
117
118 fn finish(mut self) -> Vec<u8> {
119 self.align();
120 self.out
121 }
122}
123
124const fn fixed_literal_code(symbol: u16) -> (u32, u32) {
128 match symbol {
129 0..=143 => (0x30 + symbol as u32, 8),
130 144..=255 => (0x190 + (symbol as u32 - 144), 9),
131 256..=279 => (symbol as u32 - 256, 7),
132 _ => (0xC0 + (symbol as u32 - 280), 8),
133 }
134}
135
136const LENGTH_BASE: [u16; 29] = [
138 3, 4, 5, 6, 7, 8, 9, 10, 11, 13, 15, 17, 19, 23, 27, 31, 35, 43, 51, 59, 67, 83, 99, 115, 131,
139 163, 195, 227, 258,
140];
141const LENGTH_EXTRA: [u8; 29] = [
143 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 2, 2, 2, 2, 3, 3, 3, 3, 4, 4, 4, 4, 5, 5, 5, 5, 0,
144];
145const DISTANCE_BASE: [u16; 30] = [
147 1, 2, 3, 4, 5, 7, 9, 13, 17, 25, 33, 49, 65, 97, 129, 193, 257, 385, 513, 769, 1025, 1537,
148 2049, 3073, 4097, 6145, 8193, 12289, 16385, 24577,
149];
150const DISTANCE_EXTRA: [u8; 30] = [
152 0, 0, 0, 0, 1, 1, 2, 2, 3, 3, 4, 4, 5, 5, 6, 6, 7, 7, 8, 8, 9, 9, 10, 10, 11, 11, 12, 12, 13,
153 13,
154];
155
156const MAX_MATCH: usize = 258;
158const MIN_MATCH: usize = 3;
160const WINDOW: usize = 32768;
162const HASH_SIZE: usize = 1 << 15;
164
165fn length_code(length: usize) -> Option<(u16, u32, u32)> {
167 for index in (0..LENGTH_BASE.len()).rev() {
168 let base = LENGTH_BASE.get(index).copied()? as usize;
169 if length >= base {
170 let extra_bits = LENGTH_EXTRA.get(index).copied()? as u32;
171 let extra = (length - base) as u32;
172 return Some((257 + index as u16, extra, extra_bits));
173 }
174 }
175 None
176}
177
178fn distance_code(distance: usize) -> Option<(u16, u32, u32)> {
180 for index in (0..DISTANCE_BASE.len()).rev() {
181 let base = DISTANCE_BASE.get(index).copied()? as usize;
182 if distance >= base {
183 let extra_bits = DISTANCE_EXTRA.get(index).copied()? as u32;
184 let extra = (distance - base) as u32;
185 return Some((index as u16, extra, extra_bits));
186 }
187 }
188 None
189}
190
191fn hash3(data: &[u8], at: usize) -> usize {
193 let a = data.get(at).copied().unwrap_or(0) as usize;
194 let b = data.get(at + 1).copied().unwrap_or(0) as usize;
195 let c = data.get(at + 2).copied().unwrap_or(0) as usize;
196 ((a << 10) ^ (b << 5) ^ c) & (HASH_SIZE - 1)
197}
198
199pub fn deflate(data: &[u8], level: Level) -> Result<Vec<u8>> {
206 if level == Level::NONE {
207 return Ok(deflate_stored(data));
208 }
209 let mut writer = BitWriter::default();
210 writer.write(1, 1); writer.write(1, 2); let mut head = vec![usize::MAX; HASH_SIZE];
219 let mut prev = vec![usize::MAX; data.len().max(1)];
220 let depth = level.search_depth();
221
222 let mut position = 0;
223 while position < data.len() {
224 let (mut best_length, mut best_distance) = (0_usize, 0_usize);
225
226 if position + MIN_MATCH <= data.len() {
227 let bucket = hash3(data, position);
228 let mut candidate = head.get(bucket).copied().unwrap_or(usize::MAX);
229 let limit = position.saturating_sub(WINDOW);
230 let mut tries = depth;
231
232 while candidate != usize::MAX && candidate >= limit && tries > 0 {
233 tries -= 1;
234 let length = match_length(data, candidate, position);
235 if length > best_length {
236 best_length = length;
237 best_distance = position - candidate;
238 if best_length >= MAX_MATCH {
239 break;
240 }
241 }
242 let next = prev.get(candidate).copied().unwrap_or(usize::MAX);
243 if next >= candidate {
245 break;
246 }
247 candidate = next;
248 }
249 }
250
251 if best_length >= MIN_MATCH {
252 let (code, extra, extra_bits) = length_code(best_length).ok_or_else(|| {
253 Error::malformed("deflate", "no length code for a computed match")
254 })?;
255 let (literal_code, literal_bits) = fixed_literal_code(code);
256 writer.write_code(literal_code, literal_bits);
257 if extra_bits > 0 {
258 writer.write(extra, extra_bits);
259 }
260 let (dcode, dextra, dextra_bits) = distance_code(best_distance).ok_or_else(|| {
261 Error::malformed("deflate", "no distance code for a computed match")
262 })?;
263 writer.write_code(u32::from(dcode), 5);
265 if dextra_bits > 0 {
266 writer.write(dextra, dextra_bits);
267 }
268 for offset in 0..best_length {
271 insert(data, &mut head, &mut prev, position + offset);
272 }
273 position += best_length;
274 } else {
275 let byte = data.get(position).copied().unwrap_or(0);
276 let (code, bits) = fixed_literal_code(u16::from(byte));
277 writer.write_code(code, bits);
278 insert(data, &mut head, &mut prev, position);
279 position += 1;
280 }
281 }
282
283 let (code, bits) = fixed_literal_code(256);
285 writer.write_code(code, bits);
286 Ok(writer.finish())
287}
288
289fn insert(data: &[u8], head: &mut [usize], prev: &mut [usize], at: usize) {
291 if at + MIN_MATCH > data.len() {
292 return;
293 }
294 let bucket = hash3(data, at);
295 let Some(slot) = head.get_mut(bucket) else {
296 return;
297 };
298 if let Some(chain) = prev.get_mut(at) {
299 *chain = *slot;
300 }
301 *slot = at;
302}
303
304fn match_length(data: &[u8], candidate: usize, position: usize) -> usize {
306 let available = data.len() - position;
307 let max = available.min(MAX_MATCH);
308 let mut length = 0;
309 while length < max {
310 let a = data.get(candidate + length).copied();
311 let b = data.get(position + length).copied();
312 if a.is_none() || a != b {
313 break;
314 }
315 length += 1;
316 }
317 length
318}
319
320fn deflate_stored(data: &[u8]) -> Vec<u8> {
322 const MAX_STORED: usize = 65535;
324 let mut out = Vec::with_capacity(data.len() + data.len() / MAX_STORED * 5 + 5);
325 if data.is_empty() {
326 out.push(0x01);
327 out.extend_from_slice(&0_u16.to_le_bytes());
328 out.extend_from_slice(&(!0_u16).to_le_bytes());
329 return out;
330 }
331 let mut chunks = data.chunks(MAX_STORED).peekable();
332 while let Some(chunk) = chunks.next() {
333 let final_block = u8::from(chunks.peek().is_none());
334 out.push(final_block);
335 let length = chunk.len() as u16;
336 out.extend_from_slice(&length.to_le_bytes());
337 out.extend_from_slice(&(!length).to_le_bytes());
338 out.extend_from_slice(chunk);
339 }
340 out
341}
342
343pub fn zlib_compress(data: &[u8], level: Level) -> Result<Vec<u8>> {
349 let cmf = 0x78_u8;
351 let level_bits = match level.get() {
353 0..=1 => 0_u8,
354 2..=5 => 1,
355 6 => 2,
356 _ => 3,
357 };
358 let mut flg = level_bits << 6;
359 let check = (u16::from(cmf) << 8) | u16::from(flg);
360 flg += (31 - (check % 31) % 31) as u8;
361
362 let mut out = Vec::new();
363 out.push(cmf);
364 out.push(flg);
365 out.extend_from_slice(&deflate(data, level)?);
366 out.extend_from_slice(&Adler32::of(data).to_be_bytes());
367 Ok(out)
368}
369
370#[cfg(test)]
371#[allow(
372 clippy::unwrap_used,
373 clippy::expect_used,
374 clippy::indexing_slicing,
375 clippy::panic,
376 reason = "tests operate on known-good values and assert shapes directly"
377)]
378mod tests {
379 use super::*;
380 use crate::inflate::{inflate_to, zlib_decompress};
381
382 fn corpus() -> Vec<(&'static str, Vec<u8>)> {
384 vec![
385 ("empty", Vec::new()),
386 ("one byte", vec![42]),
387 ("two bytes", vec![1, 2]),
388 ("below min match", vec![7, 7]),
389 ("exactly min match", vec![7, 7, 7]),
390 ("all zeros", vec![0; 10_000]),
391 ("repeating text", b"the quick brown fox. ".repeat(300)),
392 (
393 "incompressible",
394 (0..8192).map(|i| ((i * 37 + 11) % 256) as u8).collect(),
395 ),
396 ("long run of one byte", vec![0xAB; 70_000]),
397 ("alternating", (0..5000).map(|i| (i % 2) as u8).collect()),
398 (
399 "match at max length",
400 std::iter::repeat_n(b'z', MAX_MATCH * 3).collect::<Vec<u8>>(),
401 ),
402 ("binary", (0..=255_u8).cycle().take(20_000).collect()),
403 ]
404 }
405
406 #[test]
407 fn every_payload_round_trips_at_every_level() {
408 for (name, data) in corpus() {
411 for level in 0..=9 {
412 let level = Level::new(level).unwrap();
413 let compressed = deflate(&data, level).unwrap();
414 let out = inflate_to(&compressed, data.len().max(1)).unwrap();
415 assert_eq!(out, data, "`{name}` at level {}", level.get());
416 }
417 }
418 }
419
420 #[test]
421 fn zlib_wrapping_round_trips_at_every_level() {
422 for (name, data) in corpus() {
423 for level in 0..=9 {
424 let level = Level::new(level).unwrap();
425 let compressed = zlib_compress(&data, level).unwrap();
426 let out = zlib_decompress(&compressed, data.len().max(1)).unwrap();
427 assert_eq!(out, data, "`{name}` at level {}", level.get());
428 }
429 }
430 }
431
432 #[test]
433 fn the_zlib_header_is_well_formed() {
434 for level in 0..=9 {
435 let level = Level::new(level).unwrap();
436 let stream = zlib_compress(b"hello", level).unwrap();
437 let header = (u16::from(stream[0]) << 8) | u16::from(stream[1]);
438 assert_eq!(stream[0] & 0x0F, 8, "compression method must be deflate");
439 assert_eq!(header % 31, 0, "header check bits at level {}", level.get());
440 assert_eq!(stream[1] & 0x20, 0, "no preset dictionary");
441 }
442 }
443
444 #[test]
445 fn compression_actually_compresses_compressible_data() {
446 let data = b"the quick brown fox. ".repeat(500);
449 let compressed = deflate(&data, Level::DEFAULT).unwrap();
450 assert!(
451 compressed.len() < data.len() / 10,
452 "compressed {} bytes to {}, expected under {}",
453 data.len(),
454 compressed.len(),
455 data.len() / 10
456 );
457 }
458
459 #[test]
460 fn higher_levels_do_not_compress_worse() {
461 let data = b"abcabcabd".repeat(2000);
462 let fast = deflate(&data, Level::FAST).unwrap().len();
463 let best = deflate(&data, Level::BEST).unwrap().len();
464 assert!(
465 best <= fast,
466 "level 9 produced {best} bytes, level 1 produced {fast}"
467 );
468 }
469
470 #[test]
471 fn level_zero_stores_without_compressing() {
472 let data = vec![0_u8; 1000];
473 let compressed = deflate(&data, Level::NONE).unwrap();
474 assert!(compressed.len() > data.len(), "stored blocks add framing");
475 assert_eq!(inflate_to(&compressed, 1000).unwrap(), data);
476 }
477
478 #[test]
479 fn stored_blocks_split_at_the_sixteen_bit_length_limit() {
480 let data = vec![7_u8; 200_000];
483 let compressed = deflate(&data, Level::NONE).unwrap();
484 assert_eq!(inflate_to(&compressed, 200_000).unwrap(), data);
485 }
486
487 #[test]
488 fn levels_are_validated() {
489 assert!(Level::new(0).is_ok());
490 assert!(Level::new(9).is_ok());
491 let err = Level::new(10).unwrap_err();
492 assert!(err.detail().contains("0..=9"), "{err}");
493 assert_eq!(Level::default(), Level::DEFAULT);
494 assert_eq!(Level::DEFAULT.get(), 6);
495 }
496
497 #[test]
498 fn matches_at_the_window_boundary_round_trip() {
499 let mut data = vec![0_u8; WINDOW + 64];
502 for (i, slot) in data.iter_mut().enumerate() {
503 *slot = ((i * 7) % 251) as u8;
504 }
505 let head: Vec<u8> = data[..32].to_vec();
507 data.extend_from_slice(&head);
508 let compressed = deflate(&data, Level::BEST).unwrap();
509 assert_eq!(inflate_to(&compressed, data.len()).unwrap(), data);
510 }
511
512 #[test]
513 fn maximum_length_matches_round_trip() {
514 let data = vec![0x5A_u8; MAX_MATCH * 5 + 7];
516 let compressed = deflate(&data, Level::BEST).unwrap();
517 assert_eq!(inflate_to(&compressed, data.len()).unwrap(), data);
518 }
519
520 #[test]
521 fn our_output_is_decodable_after_a_round_trip_through_our_decoder() {
522 for (name, data) in corpus() {
525 let once = zlib_compress(&data, Level::DEFAULT).unwrap();
526 let back = zlib_decompress(&once, data.len().max(1)).unwrap();
527 let twice = zlib_compress(&back, Level::DEFAULT).unwrap();
528 assert_eq!(once, twice, "`{name}` is not deterministic");
529 }
530 }
531
532 #[test]
533 fn compression_is_deterministic() {
534 let data = b"determinism matters. ".repeat(100);
536 let first = zlib_compress(&data, Level::BEST).unwrap();
537 for _ in 0..5 {
538 assert_eq!(zlib_compress(&data, Level::BEST).unwrap(), first);
539 }
540 }
541}