Skip to main content

otf_pixels_compress/
deflate.rs

1//! DEFLATE compression (RFC 1951) and the zlib wrapper (RFC 1950).
2//!
3//! Written from scratch per ADR-0010. Correctness here means "a conforming
4//! decoder reproduces the input exactly"; compression ratio is a tuning
5//! question, not a correctness one, and this implementation deliberately
6//! favours being obviously right.
7//!
8//! # Strategy
9//!
10//! Level 0 emits stored blocks. Levels 1–9 run LZ77 over a hash chain of
11//! three-byte prefixes and emit fixed-Huffman blocks. The level controls how
12//! far back the matcher searches, trading time for ratio.
13//!
14//! Fixed Huffman rather than dynamic is a deliberate simplification: dynamic
15//! tables would compress better, but they are a second encoder to get right,
16//! and every conforming decoder accepts fixed blocks. The `Encoder` trait
17//! boundary makes a later dynamic-table implementation a drop-in, exactly as
18//! ADR-0004 makes whole codecs swappable.
19
20use crate::{Error, Result};
21
22use crate::checksum::Adler32;
23
24/// Compression effort, 0 (stored) to 9 (most search).
25#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
26pub struct Level(u8);
27
28impl Level {
29    /// No compression: stored blocks only. Fastest, always expands slightly.
30    pub const NONE: Self = Self(0);
31    /// Fastest compression.
32    pub const FAST: Self = Self(1);
33    /// The default balance.
34    pub const DEFAULT: Self = Self(6);
35    /// Most search effort.
36    pub const BEST: Self = Self(9);
37
38    /// A level from 0..=9.
39    ///
40    /// # Errors
41    ///
42    /// Returns [`Error`] outside 0..=9.
43    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    /// The raw level.
54    #[must_use]
55    pub const fn get(self) -> u8 {
56        self.0
57    }
58
59    /// How many chain positions the matcher inspects at this level.
60    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/// Writes bits least-significant-first, as DEFLATE specifies.
83#[derive(Debug, Default)]
84struct BitWriter {
85    out: Vec<u8>,
86    bits: u32,
87    count: u32,
88}
89
90impl BitWriter {
91    /// Append `n` bits of `value`, least-significant first.
92    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    /// Append `n` bits of `value`, most-significant first (Huffman codes).
103    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    /// Pad to a byte boundary with zero bits.
110    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
124/// The fixed literal/length code for `symbol`, as (code, bit length).
125///
126/// From RFC 1951 §3.2.6.
127const 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
136/// Base lengths for length codes 257..=285.
137const 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];
141/// Extra bits for length codes 257..=285.
142const 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];
145/// Base distances for distance codes 0..=29.
146const 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];
150/// Extra bits for distance codes 0..=29.
151const 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
156/// The largest back-reference DEFLATE can express.
157const MAX_MATCH: usize = 258;
158/// The shortest back-reference worth emitting.
159const MIN_MATCH: usize = 3;
160/// The sliding window size.
161const WINDOW: usize = 32768;
162/// Hash table size; a power of two so the mask is cheap.
163const HASH_SIZE: usize = 1 << 15;
164
165/// Find the length code and extra bits for a match length.
166fn 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
178/// Find the distance code and extra bits for a match distance.
179fn 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
191/// Hash three bytes into a chain bucket.
192fn 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
199/// Compress `data` into a raw DEFLATE stream.
200///
201/// # Errors
202///
203/// Returns [`Error`] only if an internal invariant is violated;
204/// compression itself cannot fail on valid input.
205pub 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    // One fixed-Huffman block for the whole input. Block splitting would
211    // improve ratio on heterogeneous data; it does not affect correctness.
212    writer.write(1, 1); // BFINAL
213    writer.write(1, 2); // BTYPE = 01, fixed Huffman
214
215    // `head[hash]` is the most recent position with that hash; `prev[pos]` is
216    // the previous position in the same chain. Together they form the standard
217    // hash-chain matcher.
218    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                // Chains must strictly decrease, or a corrupt chain could loop.
244                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            // Distance codes use a fixed 5-bit code in fixed-Huffman blocks.
264            writer.write_code(u32::from(dcode), 5);
265            if dextra_bits > 0 {
266                writer.write(dextra, dextra_bits);
267            }
268            // Insert every position the match covers, so later matches can
269            // start inside it.
270            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    // End-of-block.
284    let (code, bits) = fixed_literal_code(256);
285    writer.write_code(code, bits);
286    Ok(writer.finish())
287}
288
289/// Record `at` in the hash chains.
290fn 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
304/// How many bytes match between `candidate` and `position`.
305fn 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
320/// Emit `data` as stored (uncompressed) blocks.
321fn deflate_stored(data: &[u8]) -> Vec<u8> {
322    // A stored block's length field is 16 bits, so long inputs are split.
323    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
343/// Compress `data` into a zlib stream (RFC 1950).
344///
345/// # Errors
346///
347/// As [`deflate`].
348pub fn zlib_compress(data: &[u8], level: Level) -> Result<Vec<u8>> {
349    // CMF: deflate (8) with a 32 KiB window (7 << 4).
350    let cmf = 0x78_u8;
351    // FLG carries the level hint and makes the 16-bit header divisible by 31.
352    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    /// Payloads that exercise different compressor behaviour.
383    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        // The correctness property that matters: our decoder reproduces the
409        // input exactly, at every level, for every shape of data.
410        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        // Not a correctness property, but a ratio this bad would mean the
447        // matcher is not finding anything at all.
448        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        // A stored block's length field is 16 bits, so >65535 bytes must be
481        // split across blocks or the length silently wraps.
482        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        // A match exactly 32768 bytes back is the furthest DEFLATE can express;
500        // one byte further must not be emitted as a match.
501        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        // Repeat the opening bytes at the very edge of the window.
506        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        // 258 is the longest expressible match; longer runs must be split.
515        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        // Guards against the encoder and decoder drifting together: the data
523        // is compressed, decompressed, recompressed, and compared.
524        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        // SPEC §Guarantees 2 reaches the encoder too: same input, same bytes.
535        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}