Skip to main content

otf_pixels_compress/
inflate.rs

1//! DEFLATE decompression (RFC 1951) and the zlib wrapper (RFC 1950).
2//!
3//! Written from scratch per ADR-0010. Every path returns an error rather than
4//! panicking: this code parses attacker-controlled bytes, and
5//! `unsafe_code = "forbid"` plus explicit bounds checks mean the classic
6//! decompressor failure — an out-of-bounds write through a back-reference —
7//! is unrepresentable rather than merely avoided.
8//!
9//! # Shape
10//!
11//! A DEFLATE stream is a sequence of blocks, each either stored (raw),
12//! fixed-Huffman (a table the spec hardcodes) or dynamic-Huffman (a table the
13//! block carries). Literal/length and distance codes decode into either a
14//! literal byte or a back-reference into the last 32 KiB of output.
15//!
16//! # Incremental by construction
17//!
18//! [`Inflater`] is the real decompressor: bytes are fed in as they arrive and
19//! output is drained as it is produced, so a caller never needs the whole
20//! stream in memory. That is what lets PNG decode row by row rather than
21//! image by image (SPEC §Guarantees 1).
22//!
23//! Resuming mid-symbol is handled by *checkpointing* rather than by a symbol
24//! -level state machine: before each step the bit reader's position is saved,
25//! and if the input runs out the step is rewound and retried when more
26//! arrives. The decode logic is therefore written once, in its natural
27//! straight-line form, and is identical whether the input arrives in one piece
28//! or in a thousand.
29//!
30//! # Bounded output
31//!
32//! Every entry point takes a byte limit. A short input can expand enormously —
33//! a decompression bomb — so the caller states how much output it is prepared
34//! to accept, which for PNG is exactly the filtered raster size derived from
35//! the header (SPEC §Safety).
36
37use crate::{Error, Result};
38
39use crate::checksum::Adler32;
40
41/// The largest back-reference distance DEFLATE permits, and therefore the
42/// amount of already-emitted output that must stay reachable.
43const WINDOW: usize = 32 * 1024;
44
45/// How much consumed input [`BitReader::compact`] lets accumulate before
46/// dropping it.
47const COMPACT_AT: usize = 64 * 1024;
48
49/// Why a decode step stopped early.
50///
51/// `NeedInput` is not an error: it means the step was rewound and should be
52/// retried once more bytes arrive. Keeping it distinct from a real failure is
53/// what makes "the stream is truncated" and "the stream is not finished yet"
54/// different answers rather than the same one.
55#[derive(Debug)]
56enum Halt {
57    /// The input ran out mid-item; the reader has been rewound.
58    NeedInput,
59    /// The stream is invalid and no amount of further input will help.
60    Fatal(Error),
61}
62
63impl From<Error> for Halt {
64    fn from(error: Error) -> Self {
65        Self::Fatal(error)
66    }
67}
68
69/// The result of a step that may pause for more input.
70type Step<T> = std::result::Result<T, Halt>;
71
72/// Reads bits least-significant-first, as DEFLATE specifies.
73///
74/// Owns its input so that more can be appended mid-stream. Consumed bytes are
75/// dropped by [`BitReader::compact`], which is what keeps a long stream from
76/// accumulating in memory.
77#[derive(Debug, Default)]
78struct BitReader {
79    data: Vec<u8>,
80    /// Index of the next byte to load.
81    position: usize,
82    /// Bits not yet consumed, right-aligned.
83    bits: u64,
84    /// How many bits in `bits` are valid.
85    count: u32,
86    /// Whether the caller has promised there is no more input.
87    ended: bool,
88}
89
90/// A saved reader position, so a step that runs out of input can be rewound.
91#[derive(Debug, Clone, Copy)]
92struct Checkpoint {
93    position: usize,
94    bits: u64,
95    count: u32,
96}
97
98impl BitReader {
99    /// Append more input.
100    fn feed(&mut self, more: &[u8]) {
101        self.data.extend_from_slice(more);
102    }
103
104    /// Declare that no further input will arrive.
105    const fn end(&mut self) {
106        self.ended = true;
107    }
108
109    fn checkpoint(&self) -> Checkpoint {
110        Checkpoint {
111            position: self.position,
112            bits: self.bits,
113            count: self.count,
114        }
115    }
116
117    fn restore(&mut self, at: Checkpoint) {
118        self.position = at.position;
119        self.bits = at.bits;
120        self.count = at.count;
121    }
122
123    /// Drop input that has been consumed and can never be rewound to.
124    ///
125    /// Safe only at a checkpoint boundary, which is where the caller calls it.
126    /// Dropping shifts every unread byte down, so it waits until the consumed
127    /// prefix is both large and at least half the buffer: done after every
128    /// symbol, that shift made decoding quadratic in the input size.
129    fn compact(&mut self) {
130        if self.position >= COMPACT_AT && self.position * 2 >= self.data.len() {
131            self.data.drain(..self.position);
132            self.position = 0;
133        }
134    }
135
136    /// Take back every unconsumed byte, including whole bytes that were
137    /// pulled into the bit buffer but never used.
138    ///
139    /// Whole bytes in `bits` are real input; a partial byte is the padding
140    /// that ends a DEFLATE stream. Forgetting the former is how a trailer
141    /// comes to look one byte short.
142    fn drain_unconsumed(&mut self) -> Vec<u8> {
143        // Realign first. Mid-byte, `bits` holds the tail of a partially read
144        // byte in its low positions, so popping eight bits from the low end
145        // would yield a byte straddling two real ones. That padding is exactly
146        // what ends a DEFLATE stream, so discarding it is also correct.
147        self.align();
148        let mut out = Vec::new();
149        while self.count >= 8 {
150            out.push((self.bits & 0xFF) as u8);
151            self.bits >>= 8;
152            self.count -= 8;
153        }
154        if let Some(rest) = self.data.get(self.position..) {
155            out.extend_from_slice(rest);
156        }
157        self.position = self.data.len();
158        out
159    }
160
161    /// Ensure at least `want` bits are buffered, if the input has them.
162    fn fill(&mut self, want: u32) {
163        while self.count < want {
164            let Some(&byte) = self.data.get(self.position) else {
165                break;
166            };
167            self.bits |= u64::from(byte) << self.count;
168            self.position += 1;
169            self.count += 8;
170        }
171    }
172
173    /// Report a shortage as either "wait" or "the stream is truncated".
174    fn short(&self) -> Halt {
175        if self.ended {
176            Halt::Fatal(truncated())
177        } else {
178            Halt::NeedInput
179        }
180    }
181
182    /// Consume `n` bits (`n <= 32`).
183    fn take(&mut self, n: u32) -> Step<u32> {
184        if n == 0 {
185            return Ok(0);
186        }
187        self.fill(n);
188        if self.count < n {
189            return Err(self.short());
190        }
191        // `n <= 32`, so the mask fits and the cast cannot lose bits.
192        let mask = (1_u64 << n) - 1;
193        let value = (self.bits & mask) as u32;
194        self.bits >>= n;
195        self.count -= n;
196        Ok(value)
197    }
198
199    /// Look at up to `n` buffered bits without consuming them.
200    fn peek(&mut self, n: u32) -> u32 {
201        self.fill(n);
202        let mask = (1_u64 << n) - 1;
203        (self.bits & mask) as u32
204    }
205
206    /// Drop `n` already-peeked bits.
207    fn skip(&mut self, n: u32) -> Step<()> {
208        if self.count < n {
209            return Err(self.short());
210        }
211        self.bits >>= n;
212        self.count -= n;
213        Ok(())
214    }
215
216    /// Discard buffered bits back to a byte boundary.
217    fn align(&mut self) {
218        let extra = self.count % 8;
219        self.bits >>= extra;
220        self.count -= extra;
221    }
222
223    /// Take up to `n` whole bytes, which must be byte-aligned already.
224    ///
225    /// Returns fewer than `n` when the input is exhausted, so a stored block
226    /// can be copied across as many feeds as it takes.
227    fn take_bytes_upto(&mut self, n: usize, out: &mut Vec<u8>) -> usize {
228        let mut taken = 0;
229        // Buffered bits are consumed first, then raw input in one run.
230        while taken < n && self.count >= 8 {
231            out.push((self.bits & 0xFF) as u8);
232            self.bits >>= 8;
233            self.count -= 8;
234            taken += 1;
235        }
236        let rest = self.data.get(self.position..).unwrap_or(&[]);
237        let run = rest.get(..(n - taken).min(rest.len())).unwrap_or(&[]);
238        out.extend_from_slice(run);
239        self.position += run.len();
240        taken + run.len()
241    }
242}
243
244/// The error for a stream that ended mid-symbol.
245fn truncated() -> Error {
246    Error::malformed("deflate", "stream ended in the middle of a symbol")
247}
248
249/// The maximum code length DEFLATE permits.
250const MAX_BITS: usize = 15;
251
252/// Code lengths up to this many bits decode with one table lookup.
253const FAST_BITS: u32 = 10;
254
255/// A canonical Huffman decoding table.
256///
257/// Codes up to [`FAST_BITS`] long, which are nearly all of them in practice,
258/// decode with a single lookup in `fast`. Longer codes fall back to walking
259/// the code length by length, which needs no table-size arithmetic at all;
260/// `fast` is indexed by a masked peek, so every lookup is in bounds by
261/// construction.
262#[derive(Debug, Clone)]
263struct Huffman {
264    /// `counts[n]` is how many codes have length `n`.
265    counts: [u16; MAX_BITS + 1],
266    /// Symbols ordered by code length, then by symbol value.
267    symbols: Vec<u16>,
268    /// Indexed by the next [`FAST_BITS`] input bits: `symbol << 4 | length`
269    /// for a code of at most that length, or 0 for none.
270    fast: Vec<u16>,
271}
272
273impl Huffman {
274    /// Build a table from per-symbol code lengths (zero meaning "unused").
275    fn new(lengths: &[u8]) -> Result<Self> {
276        let mut counts = [0_u16; MAX_BITS + 1];
277        for &length in lengths {
278            let length = length as usize;
279            if length > MAX_BITS {
280                return Err(Error::malformed(
281                    "deflate",
282                    format!("code length {length} exceeds the {MAX_BITS}-bit maximum"),
283                ));
284            }
285            if let Some(slot) = counts.get_mut(length) {
286                *slot += 1;
287            }
288        }
289        // Length zero means "no code", so it never participates.
290        if let Some(slot) = counts.get_mut(0) {
291            *slot = 0;
292        }
293
294        // Reject an over-subscribed table: more codes than the tree can hold.
295        let mut left = 1_i32;
296        for length in 1..=MAX_BITS {
297            left <<= 1;
298            left -= i32::from(counts.get(length).copied().unwrap_or(0));
299            if left < 0 {
300                return Err(Error::malformed(
301                    "deflate",
302                    "Huffman table is over-subscribed",
303                ));
304            }
305        }
306
307        let mut offsets = [0_u16; MAX_BITS + 2];
308        for length in 1..=MAX_BITS {
309            let next = offsets.get(length).copied().unwrap_or(0)
310                + counts.get(length).copied().unwrap_or(0);
311            if let Some(slot) = offsets.get_mut(length + 1) {
312                *slot = next;
313            }
314        }
315
316        let total: usize = counts.iter().map(|&c| c as usize).sum();
317        let mut symbols = vec![0_u16; total];
318        let mut cursor = offsets;
319        for (symbol, &length) in lengths.iter().enumerate() {
320            if length == 0 {
321                continue;
322            }
323            let length = length as usize;
324            let Some(at) = cursor.get_mut(length) else {
325                continue;
326            };
327            let index = *at as usize;
328            *at += 1;
329            if let Some(slot) = symbols.get_mut(index) {
330                // Symbols are `u16`; a DEFLATE alphabet never exceeds 288.
331                *slot = symbol as u16;
332            }
333        }
334
335        // Canonical codes (RFC 1951 §3.2.2): within a length, consecutive
336        // in symbol order; each length starts where the previous one ended.
337        let mut next_code = [0_u32; MAX_BITS + 1];
338        let mut code = 0_u32;
339        for length in 1..=MAX_BITS {
340            code = (code + u32::from(counts.get(length - 1).copied().unwrap_or(0))) << 1;
341            if let Some(slot) = next_code.get_mut(length) {
342                *slot = code;
343            }
344        }
345        let mut fast = vec![0_u16; 1 << FAST_BITS];
346        for (symbol, &length) in lengths.iter().enumerate() {
347            let length = u32::from(length);
348            if length == 0 || length > FAST_BITS {
349                continue;
350            }
351            let Some(slot) = next_code.get_mut(length as usize) else {
352                continue;
353            };
354            let code = *slot;
355            *slot += 1;
356            // Codes are sent most-significant bit first, but the reader
357            // peeks least-significant first, so the table is indexed by the
358            // code reversed. Every index whose low `length` bits are that
359            // code decodes to this symbol, whatever the bits above them.
360            let reversed = code.reverse_bits() >> (32 - length);
361            let entry = (symbol as u16) << 4 | length as u16;
362            for index in (reversed as usize..fast.len()).step_by(1 << length) {
363                if let Some(cell) = fast.get_mut(index) {
364                    *cell = entry;
365                }
366            }
367        }
368
369        Ok(Self {
370            counts,
371            symbols,
372            fast,
373        })
374    }
375
376    /// Decode one symbol from `reader`.
377    fn decode(&self, reader: &mut BitReader) -> Step<u16> {
378        // A short buffer would let zero padding masquerade as real bits, so
379        // ask for the maximum first and pause if the stream cannot supply it.
380        reader.fill(MAX_BITS as u32);
381        if reader.count < MAX_BITS as u32 && !reader.ended {
382            return Err(Halt::NeedInput);
383        }
384
385        let mut code = 0_i32;
386        let mut first = 0_i32;
387        let mut index = 0_i32;
388        // Peek the maximum, then consume exactly the bits actually used.
389        let peeked = reader.peek(MAX_BITS as u32);
390        let entry = self
391            .fast
392            .get((peeked & ((1 << FAST_BITS) - 1)) as usize)
393            .copied()
394            .unwrap_or(0);
395        if entry != 0 {
396            reader.skip(u32::from(entry & 0xF))?;
397            return Ok(entry >> 4);
398        }
399        for length in 1..=MAX_BITS {
400            // DEFLATE codes are stored most-significant-bit first within the
401            // LSB-first bit stream, so the code is rebuilt bit by bit.
402            code |= ((peeked >> (length - 1)) & 1) as i32;
403            let count = i32::from(self.counts.get(length).copied().unwrap_or(0));
404            if code - first < count {
405                reader.skip(length as u32)?;
406                let position = (index + (code - first)) as usize;
407                return self.symbols.get(position).copied().ok_or_else(|| {
408                    Halt::Fatal(Error::malformed("deflate", "invalid Huffman symbol"))
409                });
410            }
411            index += count;
412            first = (first + count) << 1;
413            code <<= 1;
414        }
415        Err(Halt::Fatal(Error::malformed(
416            "deflate",
417            "no Huffman code matched within 15 bits",
418        )))
419    }
420}
421
422/// Base lengths for length codes 257..=285.
423const LENGTH_BASE: [u16; 29] = [
424    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,
425    163, 195, 227, 258,
426];
427/// Extra bits for length codes 257..=285.
428const LENGTH_EXTRA: [u8; 29] = [
429    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,
430];
431/// Base distances for distance codes 0..=29.
432const DISTANCE_BASE: [u16; 30] = [
433    1, 2, 3, 4, 5, 7, 9, 13, 17, 25, 33, 49, 65, 97, 129, 193, 257, 385, 513, 769, 1025, 1537,
434    2049, 3073, 4097, 6145, 8193, 12289, 16385, 24577,
435];
436/// Extra bits for distance codes 0..=29.
437const DISTANCE_EXTRA: [u8; 30] = [
438    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,
439    13,
440];
441/// The order code-length codes appear in a dynamic block header.
442const CODE_LENGTH_ORDER: [usize; 19] = [
443    16, 17, 18, 0, 8, 7, 9, 6, 10, 5, 11, 4, 12, 3, 13, 2, 14, 1, 15,
444];
445
446/// The fixed literal/length table from RFC 1951 §3.2.6.
447fn fixed_literal_table() -> Result<Huffman> {
448    let mut lengths = [0_u8; 288];
449    for (symbol, slot) in lengths.iter_mut().enumerate() {
450        *slot = match symbol {
451            0..=143 => 8,
452            144..=255 => 9,
453            256..=279 => 7,
454            _ => 8,
455        };
456    }
457    Huffman::new(&lengths)
458}
459
460/// The fixed distance table: 30 codes of 5 bits each.
461fn fixed_distance_table() -> Result<Huffman> {
462    Huffman::new(&[5_u8; 30])
463}
464
465/// Where the decoder is in the block structure.
466#[derive(Debug)]
467enum State {
468    /// Between blocks, about to read a three-bit block header.
469    BlockHeader,
470    /// Inside a stored block with `remaining` bytes still to copy.
471    Stored { remaining: usize, last: bool },
472    /// Inside a Huffman-coded block.
473    Coded {
474        literals: Box<Huffman>,
475        distances: Box<Huffman>,
476        last: bool,
477    },
478    /// The final block has ended.
479    Done,
480}
481
482/// An incremental DEFLATE decompressor.
483///
484/// Feed bytes with [`Inflater::feed`], drain output with
485/// [`Inflater::take_output`], and call [`Inflater::end_of_input`] when the
486/// stream is over. Output beyond the configured limit is a malformed-input
487/// error, never an allocation.
488#[derive(Debug)]
489pub struct Inflater {
490    reader: BitReader,
491    state: State,
492    /// Retained history plus not-yet-drained output.
493    window: Vec<u8>,
494    /// Index into `window` of the first byte not yet handed to the caller.
495    pending: usize,
496    /// Total bytes ever produced, which is what `limit` bounds.
497    produced: usize,
498    limit: usize,
499}
500
501impl Inflater {
502    /// A decompressor that will produce at most `limit` bytes.
503    #[must_use]
504    pub fn new(limit: usize) -> Self {
505        Self {
506            reader: BitReader::default(),
507            state: State::BlockHeader,
508            window: Vec::new(),
509            pending: 0,
510            produced: 0,
511            limit,
512        }
513    }
514
515    /// Supply more compressed bytes.
516    pub fn feed(&mut self, data: &[u8]) {
517        self.reader.feed(data);
518    }
519
520    /// Declare that no further compressed bytes will arrive.
521    pub const fn end_of_input(&mut self) {
522        self.reader.end();
523    }
524
525    /// Whether the final block has been decoded.
526    #[must_use]
527    pub const fn is_finished(&self) -> bool {
528        matches!(self.state, State::Done)
529    }
530
531    /// How many bytes have been produced in total.
532    #[must_use]
533    pub const fn produced(&self) -> usize {
534        self.produced
535    }
536
537    /// How many output bytes are currently held — history plus undrained.
538    ///
539    /// Exposed so callers can assert the streaming guarantee rather than
540    /// trust it: a drained inflater retains the 32 KiB window and no more.
541    #[must_use]
542    pub fn retained(&self) -> usize {
543        self.window.len()
544    }
545
546    /// Take back compressed bytes that were fed in but are not part of the
547    /// DEFLATE stream — for zlib, that is the Adler-32 trailer.
548    fn drain_unconsumed_input(&mut self) -> Vec<u8> {
549        self.reader.drain_unconsumed()
550    }
551
552    /// Take everything decoded since the last call.
553    ///
554    /// Draining is what bounds memory: the decompressor keeps only the 32 KiB
555    /// of history that back-references can still reach, so a caller that
556    /// drains regularly holds a fixed amount regardless of stream length.
557    pub fn take_output(&mut self) -> Vec<u8> {
558        // Copied, not moved: delivered bytes are still back-reference history
559        // until they fall out of the window. Moving them out is how a
560        // long-range reference comes to point at nothing.
561        let out = self.window.get(self.pending..).unwrap_or(&[]).to_vec();
562        if self.window.len() > WINDOW {
563            self.window.drain(..self.window.len() - WINDOW);
564        }
565        self.pending = self.window.len();
566        out
567    }
568
569    /// Decode as far as the buffered input allows.
570    ///
571    /// # Errors
572    ///
573    /// Returns [`Error`] for an invalid stream, for output
574    /// exceeding the limit, or for a stream that ended mid-symbol after
575    /// [`Inflater::end_of_input`].
576    pub fn decode(&mut self) -> Result<()> {
577        loop {
578            if matches!(self.state, State::Done) {
579                return Ok(());
580            }
581            let at = self.reader.checkpoint();
582            match self.step() {
583                Ok(()) => {
584                    // Only past a completed step is consumed input unreachable.
585                    self.reader.compact();
586                }
587                Err(Halt::NeedInput) => {
588                    self.reader.restore(at);
589                    return Ok(());
590                }
591                Err(Halt::Fatal(error)) => return Err(error),
592            }
593        }
594    }
595
596    /// Perform one unit of work: a block header, a stored run, or a symbol.
597    fn step(&mut self) -> Step<()> {
598        match &self.state {
599            State::Done => Ok(()),
600            State::BlockHeader => {
601                let last = self.reader.take(1)? == 1;
602                let kind = self.reader.take(2)?;
603                self.state = match kind {
604                    0 => {
605                        self.reader.align();
606                        let length = self.reader.take(16)? as usize;
607                        let complement = self.reader.take(16)? as usize;
608                        if length ^ 0xFFFF != complement {
609                            return Err(Halt::Fatal(Error::malformed(
610                                "deflate",
611                                "stored block length does not match its complement",
612                            )));
613                        }
614                        self.check_limit(length)?;
615                        State::Stored {
616                            remaining: length,
617                            last,
618                        }
619                    }
620                    1 => State::Coded {
621                        literals: Box::new(fixed_literal_table()?),
622                        distances: Box::new(fixed_distance_table()?),
623                        last,
624                    },
625                    2 => {
626                        let (literals, distances) = read_dynamic_tables(&mut self.reader)?;
627                        State::Coded {
628                            literals: Box::new(literals),
629                            distances: Box::new(distances),
630                            last,
631                        }
632                    }
633                    _ => {
634                        return Err(Halt::Fatal(Error::malformed(
635                            "deflate",
636                            "reserved block type 3",
637                        )));
638                    }
639                };
640                Ok(())
641            }
642            State::Stored { remaining, last } => {
643                let (remaining, last) = (*remaining, *last);
644                if remaining == 0 {
645                    self.state = if last {
646                        State::Done
647                    } else {
648                        State::BlockHeader
649                    };
650                    return Ok(());
651                }
652                let taken = self.reader.take_bytes_upto(remaining, &mut self.window);
653                self.produced += taken;
654                if taken == 0 {
655                    return Err(self.reader.short());
656                }
657                self.state = State::Stored {
658                    remaining: remaining - taken,
659                    last,
660                };
661                Ok(())
662            }
663            State::Coded { .. } => self.step_coded(),
664        }
665    }
666
667    /// Decode literals and back-references from the current coded block until
668    /// it ends or the input runs out.
669    ///
670    /// Each symbol is its own checkpoint, so a pause rewinds only the symbol
671    /// in flight: if any symbol completed this is progress and returns `Ok`,
672    /// and the next step pauses at once. Decoding a run per step, rather than
673    /// one symbol, keeps the per-step bookkeeping off the hot path.
674    fn step_coded(&mut self) -> Step<()> {
675        let Self {
676            reader,
677            state,
678            window,
679            produced,
680            limit,
681            ..
682        } = self;
683        let State::Coded {
684            literals,
685            distances,
686            last,
687        } = state
688        else {
689            return Ok(());
690        };
691        let mut progressed = false;
692        loop {
693            let at = reader.checkpoint();
694            match decode_symbol(reader, literals, distances, window, produced, *limit) {
695                Ok(true) => progressed = true,
696                Ok(false) => break,
697                Err(Halt::NeedInput) => {
698                    reader.restore(at);
699                    return if progressed {
700                        Ok(())
701                    } else {
702                        Err(Halt::NeedInput)
703                    };
704                }
705                Err(fatal) => return Err(fatal),
706            }
707        }
708        *state = if *last {
709            State::Done
710        } else {
711            State::BlockHeader
712        };
713        Ok(())
714    }
715
716    /// Reject output that would exceed the limit.
717    fn check_limit(&self, adding: usize) -> Step<()> {
718        check_limit(self.produced, adding, self.limit)
719    }
720}
721
722/// Reject output that would take `produced` past `limit`.
723fn check_limit(produced: usize, adding: usize, limit: usize) -> Step<()> {
724    if produced.saturating_add(adding) > limit {
725        return Err(Halt::Fatal(Error::malformed(
726            "deflate",
727            format!("stream expands beyond the {limit} byte limit implied by the image header"),
728        )));
729    }
730    Ok(())
731}
732
733/// Decode one literal or back-reference into `window`.
734///
735/// Returns `false` at the end of the block. Nothing is written until every
736/// bit of the symbol has been read, so a pause leaves `window` untouched.
737fn decode_symbol(
738    reader: &mut BitReader,
739    literals: &Huffman,
740    distances: &Huffman,
741    window: &mut Vec<u8>,
742    produced: &mut usize,
743    limit: usize,
744) -> Step<bool> {
745    let symbol = literals.decode(reader)?;
746    match symbol {
747        // A literal byte.
748        0..=255 => {
749            check_limit(*produced, 1, limit)?;
750            window.push(symbol as u8);
751            *produced += 1;
752            Ok(true)
753        }
754        // End of block.
755        256 => Ok(false),
756        // A back-reference.
757        257..=285 => {
758            let index = symbol as usize - 257;
759            let base = LENGTH_BASE
760                .get(index)
761                .copied()
762                .ok_or_else(|| Halt::Fatal(Error::malformed("deflate", "invalid length code")))?;
763            let extra = LENGTH_EXTRA.get(index).copied().unwrap_or(0);
764            let length = base as usize + reader.take(u32::from(extra))? as usize;
765
766            let distance_symbol = distances.decode(reader)? as usize;
767            let distance_base = DISTANCE_BASE
768                .get(distance_symbol)
769                .copied()
770                .ok_or_else(|| Halt::Fatal(Error::malformed("deflate", "invalid distance code")))?;
771            let distance_extra = DISTANCE_EXTRA.get(distance_symbol).copied().unwrap_or(0);
772            let distance =
773                distance_base as usize + reader.take(u32::from(distance_extra))? as usize;
774
775            // The reference must land inside what has already been emitted.
776            // This is the check whose absence is the classic decompressor
777            // out-of-bounds read. Comparing against the retained window
778            // rather than total output is what makes it still correct once
779            // old output has been drained away.
780            if distance == 0 || distance > window.len() {
781                return Err(Halt::Fatal(Error::malformed(
782                    "deflate",
783                    format!(
784                        "back-reference of distance {distance} points before the start of \
785                         the {produced} bytes decoded so far"
786                    ),
787                )));
788            }
789            check_limit(*produced, length, limit)?;
790
791            // Overlapping references are legal and are how DEFLATE encodes
792            // runs: output byte `i` is byte `i - distance`, which may itself
793            // have been written by this copy. Everything from `start` on is
794            // periodic with period `distance`, so copying it from `start`
795            // again continues the run correctly — and each piece doubles, so
796            // a long run takes a handful of copies, and a reference that does
797            // not overlap takes one.
798            let start = window.len() - distance;
799            let mut remaining = length;
800            while remaining > 0 {
801                let piece = remaining.min(window.len() - start);
802                window.extend_from_within(start..start + piece);
803                remaining -= piece;
804            }
805            *produced += length;
806            Ok(true)
807        }
808        _ => Err(Halt::Fatal(Error::malformed(
809            "deflate",
810            format!("literal/length symbol {symbol} is out of range"),
811        ))),
812    }
813}
814
815/// Read the code-length-coded tables of a dynamic block.
816fn read_dynamic_tables(reader: &mut BitReader) -> Step<(Huffman, Huffman)> {
817    let literal_count = reader.take(5)? as usize + 257;
818    let distance_count = reader.take(5)? as usize + 1;
819    let code_length_count = reader.take(4)? as usize + 4;
820    if literal_count > 288 || distance_count > 30 {
821        return Err(Halt::Fatal(Error::malformed(
822            "deflate",
823            "dynamic block declares too many codes",
824        )));
825    }
826
827    let mut code_lengths = [0_u8; 19];
828    for index in 0..code_length_count {
829        let bits = reader.take(3)? as u8;
830        let Some(&position) = CODE_LENGTH_ORDER.get(index) else {
831            break;
832        };
833        if let Some(slot) = code_lengths.get_mut(position) {
834            *slot = bits;
835        }
836    }
837    let code_length_table = Huffman::new(&code_lengths)?;
838
839    // The two tables are coded as one run, so repeats may straddle the seam.
840    let total = literal_count + distance_count;
841    let mut lengths = vec![0_u8; total];
842    let mut index = 0;
843    while index < total {
844        let symbol = code_length_table.decode(reader)?;
845        match symbol {
846            0..=15 => {
847                if let Some(slot) = lengths.get_mut(index) {
848                    *slot = symbol as u8;
849                }
850                index += 1;
851            }
852            16 => {
853                // Repeat the previous length 3..=6 times.
854                let previous = index
855                    .checked_sub(1)
856                    .and_then(|i| lengths.get(i).copied())
857                    .ok_or_else(|| {
858                        Halt::Fatal(Error::malformed(
859                            "deflate",
860                            "repeat code with no previous length",
861                        ))
862                    })?;
863                let repeat = reader.take(2)? as usize + 3;
864                fill(&mut lengths, &mut index, previous, repeat, total)?;
865            }
866            17 => {
867                let repeat = reader.take(3)? as usize + 3;
868                fill(&mut lengths, &mut index, 0, repeat, total)?;
869            }
870            18 => {
871                let repeat = reader.take(7)? as usize + 11;
872                fill(&mut lengths, &mut index, 0, repeat, total)?;
873            }
874            _ => {
875                return Err(Halt::Fatal(Error::malformed(
876                    "deflate",
877                    "invalid code length symbol",
878                )));
879            }
880        }
881    }
882
883    let (literal_lengths, distance_lengths) = lengths.split_at(literal_count);
884    let literals = Huffman::new(literal_lengths)?;
885    let distances = Huffman::new(distance_lengths)?;
886    Ok((literals, distances))
887}
888
889/// Write `value` into `lengths` `repeat` times, refusing to overrun.
890fn fill(lengths: &mut [u8], index: &mut usize, value: u8, repeat: usize, total: usize) -> Step<()> {
891    if *index + repeat > total {
892        return Err(Halt::Fatal(Error::malformed(
893            "deflate",
894            "code length repeat runs past the end of the table",
895        )));
896    }
897    for _ in 0..repeat {
898        if let Some(slot) = lengths.get_mut(*index) {
899            *slot = value;
900        }
901        *index += 1;
902    }
903    Ok(())
904}
905
906/// Decompress a raw DEFLATE stream, refusing to exceed `limit` output bytes.
907///
908/// This is [`Inflater`] used all at once, which is what a caller that already
909/// holds the whole stream wants.
910///
911/// # Errors
912///
913/// Returns [`Error`] for any invalid or truncated stream, and
914/// for output exceeding `limit` — a decompression bomb is malformed input, not
915/// a resource the caller must survive.
916pub fn inflate_to(data: &[u8], limit: usize) -> Result<Vec<u8>> {
917    let mut inflater = Inflater::new(limit);
918    inflater.feed(data);
919    inflater.end_of_input();
920    inflater.decode()?;
921    if !inflater.is_finished() {
922        return Err(truncated());
923    }
924    Ok(inflater.take_output())
925}
926
927/// An incremental zlib (RFC 1950) decompressor.
928///
929/// Wraps [`Inflater`] with the two-byte header, the running Adler-32 and the
930/// four-byte trailer. The checksum is computed as output is produced, so
931/// verifying it costs nothing extra and does not require retaining the output.
932#[derive(Debug)]
933pub struct ZlibStream {
934    header: Vec<u8>,
935    inflater: Inflater,
936    adler: Adler32,
937    /// The trailer, accumulated once the deflate stream is finished.
938    trailer: Vec<u8>,
939    ended: bool,
940}
941
942impl ZlibStream {
943    /// A decompressor that will produce at most `limit` bytes.
944    #[must_use]
945    pub fn new(limit: usize) -> Self {
946        Self {
947            header: Vec::with_capacity(2),
948            inflater: Inflater::new(limit),
949            adler: Adler32::new(),
950            trailer: Vec::with_capacity(4),
951            ended: false,
952        }
953    }
954
955    /// Supply more compressed bytes and return whatever they decoded to.
956    ///
957    /// # Errors
958    ///
959    /// Returns [`Error`] for a bad header, an unsupported
960    /// compression method, or any DEFLATE error.
961    pub fn push(&mut self, mut data: &[u8]) -> Result<Vec<u8>> {
962        // The two-byte header must be complete before anything else happens,
963        // and it may itself be split across feeds.
964        while self.header.len() < 2 {
965            let Some((&byte, rest)) = data.split_first() else {
966                return Ok(Vec::new());
967            };
968            self.header.push(byte);
969            data = rest;
970            if self.header.len() == 2 {
971                validate_zlib_header(&self.header)?;
972            }
973        }
974
975        if self.inflater.is_finished() {
976            self.collect_trailer(data);
977            return Ok(Vec::new());
978        }
979
980        self.inflater.feed(data);
981        self.inflater.decode()?;
982        let out = self.inflater.take_output();
983        self.adler.update(&out);
984
985        // Anything the deflate stream did not consume is the trailer. Taking
986        // it from the reader rather than slicing `data` keeps this correct
987        // when the trailer straddles two feeds.
988        if self.inflater.is_finished() {
989            let leftover = self.inflater.drain_unconsumed_input();
990            self.collect_trailer(&leftover);
991        }
992        Ok(out)
993    }
994
995    /// Keep up to four trailing bytes, which carry the Adler-32.
996    fn collect_trailer(&mut self, data: &[u8]) {
997        for &byte in data {
998            if self.trailer.len() < 4 {
999                self.trailer.push(byte);
1000            }
1001        }
1002    }
1003
1004    /// Declare the input over and verify the checksum.
1005    ///
1006    /// # Errors
1007    ///
1008    /// Returns [`Error`] if the stream ended mid-symbol, if
1009    /// the trailer is missing, or if the Adler-32 does not match.
1010    pub fn finish(&mut self) -> Result<Vec<u8>> {
1011        if self.ended {
1012            return Ok(Vec::new());
1013        }
1014        self.ended = true;
1015        if self.header.len() < 2 {
1016            return Err(Error::malformed(
1017                "zlib",
1018                "stream is shorter than its 2-byte header",
1019            ));
1020        }
1021        self.inflater.end_of_input();
1022        self.inflater.decode()?;
1023        let out = self.inflater.take_output();
1024        self.adler.update(&out);
1025        if !self.inflater.is_finished() {
1026            return Err(truncated());
1027        }
1028        let leftover = self.inflater.drain_unconsumed_input();
1029        self.collect_trailer(&leftover);
1030
1031        if self.trailer.len() < 4 {
1032            return Err(Error::malformed(
1033                "zlib",
1034                "stream is missing its Adler-32 trailer",
1035            ));
1036        }
1037        let expected = u32::from_be_bytes([
1038            self.trailer.first().copied().unwrap_or(0),
1039            self.trailer.get(1).copied().unwrap_or(0),
1040            self.trailer.get(2).copied().unwrap_or(0),
1041            self.trailer.get(3).copied().unwrap_or(0),
1042        ]);
1043        let actual = self.adler.finish();
1044        if actual != expected {
1045            return Err(Error::malformed(
1046                "zlib",
1047                format!(
1048                    "Adler-32 mismatch: stream declares {expected:#010x}, data is {actual:#010x}"
1049                ),
1050            ));
1051        }
1052        Ok(out)
1053    }
1054}
1055
1056/// Validate the two-byte zlib header (RFC 1950 §2.2).
1057fn validate_zlib_header(header: &[u8]) -> Result<()> {
1058    let (&cmf, &flg) = match (header.first(), header.get(1)) {
1059        (Some(cmf), Some(flg)) => (cmf, flg),
1060        _ => {
1061            return Err(Error::malformed(
1062                "zlib",
1063                "stream is shorter than its 2-byte header",
1064            ));
1065        }
1066    };
1067    if cmf & 0x0F != 8 {
1068        return Err(Error::malformed(
1069            "zlib",
1070            format!("compression method {} is not deflate", cmf & 0x0F),
1071        ));
1072    }
1073    if (u16::from(cmf) << 8 | u16::from(flg)) % 31 != 0 {
1074        return Err(Error::malformed("zlib", "header check bits are wrong"));
1075    }
1076    if flg & 0x20 != 0 {
1077        // A preset dictionary would change what the back-references mean, and
1078        // PNG forbids it (PNG spec §10.3).
1079        return Err(Error::malformed(
1080            "zlib",
1081            "preset dictionaries are not supported",
1082        ));
1083    }
1084    Ok(())
1085}
1086
1087/// Decompress a zlib stream (RFC 1950), verifying its Adler-32.
1088///
1089/// # Errors
1090///
1091/// Returns [`Error`] for a bad header, an unsupported
1092/// compression method, a checksum mismatch, or any DEFLATE error.
1093pub fn zlib_decompress(data: &[u8], limit: usize) -> Result<Vec<u8>> {
1094    let mut stream = ZlibStream::new(limit);
1095    let mut out = stream.push(data)?;
1096    out.extend_from_slice(&stream.finish()?);
1097    Ok(out)
1098}
1099#[cfg(test)]
1100#[allow(
1101    clippy::unwrap_used,
1102    clippy::expect_used,
1103    clippy::indexing_slicing,
1104    clippy::panic,
1105    reason = "tests operate on known-good values and assert shapes directly"
1106)]
1107mod tests {
1108    use super::*;
1109
1110    /// A stored-block DEFLATE stream wrapping `payload`.
1111    fn stored_stream(payload: &[u8]) -> Vec<u8> {
1112        let mut out = vec![0x01];
1113        let length = payload.len() as u16;
1114        out.extend_from_slice(&length.to_le_bytes());
1115        out.extend_from_slice(&(!length).to_le_bytes());
1116        out.extend_from_slice(payload);
1117        out
1118    }
1119
1120    #[test]
1121    fn stored_blocks_round_trip() {
1122        let payload = b"the quick brown fox";
1123        let out = inflate_to(&stored_stream(payload), 1024).unwrap();
1124        assert_eq!(out, payload);
1125    }
1126
1127    #[test]
1128    fn an_empty_stored_block_yields_nothing() {
1129        assert_eq!(
1130            inflate_to(&stored_stream(b""), 16).unwrap(),
1131            Vec::<u8>::new()
1132        );
1133    }
1134
1135    #[test]
1136    fn a_stored_block_with_a_bad_complement_is_rejected() {
1137        let mut stream = stored_stream(b"abc");
1138        stream[3] ^= 0xFF;
1139        let err = inflate_to(&stream, 1024).unwrap_err();
1140        assert!(err.to_string().contains("complement"), "{err}");
1141    }
1142
1143    #[test]
1144    fn fixed_huffman_decodes_a_known_stream() {
1145        // zlib's output for "hello" at the default level, minus the wrapper:
1146        // a single fixed-Huffman block.
1147        let stream = [0xCB, 0x48, 0xCD, 0xC9, 0xC9, 0x07, 0x00];
1148        assert_eq!(inflate_to(&stream, 64).unwrap(), b"hello");
1149    }
1150
1151    #[test]
1152    fn zlib_wrapped_streams_verify_their_checksum() {
1153        // zlib -9 output for "hello world".
1154        let stream = [
1155            0x78, 0xDA, 0xCB, 0x48, 0xCD, 0xC9, 0xC9, 0x57, 0x28, 0xCF, 0x2F, 0xCA, 0x49, 0x01,
1156            0x00, 0x1A, 0x0B, 0x04, 0x5D,
1157        ];
1158        assert_eq!(zlib_decompress(&stream, 64).unwrap(), b"hello world");
1159    }
1160
1161    #[test]
1162    fn a_corrupted_adler_is_reported() {
1163        let mut stream = vec![
1164            0x78, 0xDA, 0xCB, 0x48, 0xCD, 0xC9, 0xC9, 0x57, 0x28, 0xCF, 0x2F, 0xCA, 0x49, 0x01,
1165            0x00, 0x1A, 0x0B, 0x04, 0x5D,
1166        ];
1167        let last = stream.len() - 1;
1168        stream[last] ^= 0xFF;
1169        let err = zlib_decompress(&stream, 64).unwrap_err();
1170        assert!(err.to_string().contains("Adler-32"), "{err}");
1171    }
1172
1173    #[test]
1174    fn zlib_headers_are_validated() {
1175        assert!(zlib_decompress(&[], 16).is_err(), "empty");
1176        assert!(zlib_decompress(&[0x78], 16).is_err(), "one byte");
1177        // Compression method 7 is not deflate.
1178        assert!(zlib_decompress(&[0x77, 0x00, 0x00], 16).is_err());
1179        // Header check bits that do not divide by 31.
1180        assert!(zlib_decompress(&[0x78, 0x00, 0x00], 16).is_err());
1181        // Preset dictionary flag set (0x78, 0x3F is divisible by 31).
1182        let err = zlib_decompress(&[0x78, 0x3F, 0x00], 16).unwrap_err();
1183        assert!(err.to_string().contains("dictionar"), "{err}");
1184    }
1185
1186    #[test]
1187    fn reserved_block_type_three_is_rejected() {
1188        // BFINAL=1, BTYPE=3 packs to 0b111 in the first byte.
1189        let err = inflate_to(&[0x07], 16).unwrap_err();
1190        assert!(err.to_string().contains("reserved"), "{err}");
1191    }
1192
1193    #[test]
1194    fn a_back_reference_before_the_start_is_rejected() {
1195        // A fixed-Huffman block whose very first symbol is a length code:
1196        // nothing has been emitted, so the distance points before the start of
1197        // the output. This is the classic decompressor out-of-bounds read, and
1198        // the bytes are hand-assembled because no real encoder emits it.
1199        //
1200        // BFINAL=1, BTYPE=01, symbol 257 (7-bit code 0000001), distance code 0.
1201        let err = inflate_to(&[0x03, 0x02], 1024).unwrap_err();
1202        assert_eq!(err.format(), "deflate", "{err}");
1203        assert!(err.to_string().contains("back-reference"), "{err}");
1204    }
1205
1206    #[test]
1207    fn output_beyond_the_limit_is_malformed_not_an_allocation() {
1208        // A decompression bomb: 65535 zero bytes from a tiny stored block.
1209        let bomb = stored_stream(&vec![0_u8; 65535]);
1210        let err = inflate_to(&bomb, 1024).unwrap_err();
1211        assert_eq!(err.format(), "deflate", "{err}");
1212        assert!(err.to_string().contains("limit"), "{err}");
1213        // The same stream within a generous limit is fine.
1214        assert_eq!(inflate_to(&bomb, 65535).unwrap().len(), 65535);
1215    }
1216
1217    #[test]
1218    fn every_truncation_of_a_valid_stream_is_an_error_not_a_panic() {
1219        let full = [
1220            0x78, 0xDA, 0xCB, 0x48, 0xCD, 0xC9, 0xC9, 0x57, 0x28, 0xCF, 0x2F, 0xCA, 0x49, 0x01,
1221            0x00, 0x1A, 0x0B, 0x04, 0x5D,
1222        ];
1223        for len in 0..full.len() {
1224            // Must not panic. Some prefixes decode a valid shorter stream, so
1225            // success is acceptable; a crash is not.
1226            let _ = zlib_decompress(&full[..len], 4096);
1227        }
1228        assert!(
1229            zlib_decompress(&full, 4096).is_ok(),
1230            "the untruncated stream still works"
1231        );
1232    }
1233
1234    #[test]
1235    fn arbitrary_bytes_never_panic() {
1236        // A cheap deterministic sweep; the real corpus fuzzing is in the
1237        // crate's fuzz tests. This is here so the module is never committed in
1238        // a state that crashes on trivial garbage.
1239        let mut state = 0x1234_5678_u32;
1240        for _ in 0..2000 {
1241            let mut bytes = Vec::new();
1242            for _ in 0..32 {
1243                state = state.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
1244                bytes.push((state >> 24) as u8);
1245            }
1246            let _ = inflate_to(&bytes, 4096);
1247            let _ = zlib_decompress(&bytes, 4096);
1248        }
1249    }
1250
1251    #[test]
1252    fn over_subscribed_huffman_tables_are_rejected() {
1253        // Three one-bit codes cannot form a prefix code.
1254        assert!(Huffman::new(&[1, 1, 1]).is_err());
1255        // Two one-bit codes can.
1256        assert!(Huffman::new(&[1, 1]).is_ok());
1257        // A length beyond 15 bits is out of spec.
1258        assert!(Huffman::new(&[16]).is_err());
1259    }
1260
1261    #[test]
1262    fn overlapping_back_references_encode_runs() {
1263        // A run is encoded as a copy from distance 1, so the source overlaps
1264        // the destination and includes bytes the copy loop is itself writing.
1265        // Produced by zlib, not by hand.
1266        let stream = [0x4B, 0x4C, 0x84, 0x00, 0x00];
1267        assert_eq!(inflate_to(&stream, 64).unwrap(), b"aaaaaaaa");
1268    }
1269
1270    /// A zlib stream of `data` at the given level, via our own compressor.
1271    fn compress(data: &[u8], level: u8) -> Vec<u8> {
1272        crate::deflate::zlib_compress(data, crate::deflate::Level::new(level).unwrap()).unwrap()
1273    }
1274
1275    /// The raw DEFLATE body of that stream, for tests that drive `Inflater`
1276    /// directly rather than through the zlib wrapper.
1277    fn deflate_body(data: &[u8], level: u8) -> Vec<u8> {
1278        compress(data, level).split_off(2)
1279    }
1280
1281    /// Write `symbols` with the canonical code `lengths` defines, most
1282    /// significant code bit first, into DEFLATE's LSB-first bit order.
1283    fn encode_canonical(lengths: &[u8], symbols: &[usize]) -> Vec<u8> {
1284        let mut codes = vec![0_u32; lengths.len()];
1285        let mut code = 0_u32;
1286        for length in 1..=MAX_BITS as u8 {
1287            for (symbol, &l) in lengths.iter().enumerate() {
1288                if l == length {
1289                    codes[symbol] = code;
1290                    code += 1;
1291                }
1292            }
1293            code <<= 1;
1294        }
1295        let (mut out, mut bits, mut count) = (Vec::new(), 0_u64, 0_u32);
1296        for &symbol in symbols {
1297            let length = u32::from(lengths[symbol]);
1298            let reversed = codes[symbol].reverse_bits() >> (32 - length);
1299            bits |= u64::from(reversed) << count;
1300            count += length;
1301            while count >= 8 {
1302                out.push(bits as u8);
1303                bits >>= 8;
1304                count -= 8;
1305            }
1306        }
1307        if count > 0 {
1308            out.push(bits as u8);
1309        }
1310        out
1311    }
1312
1313    #[test]
1314    fn codes_of_every_length_decode_through_table_and_walk() {
1315        // Lengths 1..=15 plus a second 15: a complete code whose short codes
1316        // hit the lookup table and whose long ones fall back to the walk.
1317        let mut lengths: Vec<u8> = (1..=15).collect();
1318        lengths.push(15);
1319        let table = Huffman::new(&lengths).unwrap();
1320        let symbols: Vec<usize> = (0..lengths.len()).chain((0..lengths.len()).rev()).collect();
1321        let mut reader = BitReader::default();
1322        reader.feed(&encode_canonical(&lengths, &symbols));
1323        reader.end();
1324        for &expected in &symbols {
1325            assert_eq!(usize::from(table.decode(&mut reader).unwrap()), expected);
1326        }
1327    }
1328
1329    #[test]
1330    fn overlapping_references_of_every_short_distance_decode() {
1331        // A run is a back-reference shorter than its length; the copy must
1332        // read bytes it is itself writing. Periods 1..=9 and a long one.
1333        let mut original = Vec::new();
1334        for period in (1..=9).chain([31, 258]) {
1335            let pattern: Vec<u8> = (0..period).map(|i| (i * 37 + period) as u8).collect();
1336            for _ in 0..600 / period + 3 {
1337                original.extend_from_slice(&pattern);
1338            }
1339        }
1340        for level in [1_u8, 6, 9] {
1341            assert_eq!(
1342                zlib_decompress(&compress(&original, level), 1 << 20).unwrap(),
1343                original
1344            );
1345        }
1346    }
1347
1348    #[test]
1349    fn a_stream_past_the_compaction_threshold_decodes_whole_and_in_pieces() {
1350        // Incompressible-ish data keeps the compressed stream well past
1351        // `COMPACT_AT`, so consumed input is dropped mid-stream.
1352        let mut state = 0x2545_f491_u32;
1353        let original: Vec<u8> = (0..400_000)
1354            .map(|i| {
1355                state ^= state << 13;
1356                state ^= state >> 17;
1357                state ^= state << 5;
1358                if i % 5 == 0 { b'a' } else { state as u8 }
1359            })
1360            .collect();
1361        let stream = compress(&original, 6);
1362        assert!(stream.len() > 2 * COMPACT_AT);
1363        assert_eq!(zlib_decompress(&stream, 1 << 20).unwrap(), original);
1364
1365        let mut zlib = ZlibStream::new(1 << 20);
1366        let mut out = Vec::new();
1367        for piece in stream.chunks(997) {
1368            out.extend_from_slice(&zlib.push(piece).unwrap());
1369        }
1370        out.extend_from_slice(&zlib.finish().unwrap());
1371        assert_eq!(out, original);
1372    }
1373
1374    #[test]
1375    fn feeding_one_byte_at_a_time_decodes_identically() {
1376        // The property that makes streaming real: how the input is chunked
1377        // must not change the output. One byte at a time is the worst case,
1378        // because every symbol is interrupted.
1379        for level in [0_u8, 1, 6, 9] {
1380            let original = b"the quick brown fox jumps over the lazy dog. ".repeat(120);
1381            let stream = compress(&original, level);
1382
1383            let mut zlib = ZlibStream::new(1 << 20);
1384            let mut out = Vec::new();
1385            for byte in &stream {
1386                out.extend_from_slice(&zlib.push(std::slice::from_ref(byte)).unwrap());
1387            }
1388            out.extend_from_slice(&zlib.finish().unwrap());
1389            assert_eq!(
1390                out, original,
1391                "level {level} differed when fed byte by byte"
1392            );
1393        }
1394    }
1395
1396    #[test]
1397    fn every_chunk_size_decodes_identically() {
1398        let original: Vec<u8> = (0..40_000).map(|i| ((i * 7) % 251) as u8).collect();
1399        let stream = compress(&original, 6);
1400        for chunk in [1, 2, 3, 7, 64, 1024, 65_536] {
1401            let mut zlib = ZlibStream::new(1 << 20);
1402            let mut out = Vec::new();
1403            for piece in stream.chunks(chunk) {
1404                out.extend_from_slice(&zlib.push(piece).unwrap());
1405            }
1406            out.extend_from_slice(&zlib.finish().unwrap());
1407            assert_eq!(out, original, "chunk size {chunk} differed");
1408        }
1409    }
1410
1411    #[test]
1412    fn a_drained_inflater_retains_only_its_window() {
1413        // The whole point of the restructure: decoding a stream far larger
1414        // than the window must not accumulate it. A caller that drains holds
1415        // the window and nothing more.
1416        let original = vec![0_u8; 8 * 1024 * 1024];
1417        let stream = deflate_body(&original, 9);
1418
1419        let mut inflater = Inflater::new(16 * 1024 * 1024);
1420        let mut total = 0_usize;
1421        for piece in stream.chunks(4096) {
1422            inflater.feed(piece);
1423            inflater.decode().unwrap();
1424            total += inflater.take_output().len();
1425            assert!(
1426                inflater.retained() <= WINDOW + 4096,
1427                "retained {} bytes after {total} of output",
1428                inflater.retained()
1429            );
1430        }
1431        inflater.end_of_input();
1432        inflater.decode().unwrap();
1433        total += inflater.take_output().len();
1434        assert_eq!(total, original.len());
1435        assert_eq!(inflater.produced(), original.len());
1436    }
1437
1438    #[test]
1439    fn a_back_reference_reaching_across_a_drain_still_resolves() {
1440        // Draining discards delivered output, so a back-reference pointing
1441        // into it must still find its bytes in the retained window. A run
1442        // longer than the window is where that breaks if it is going to.
1443        let original = b"abcdefgh".repeat(200_000);
1444        let stream = deflate_body(&original, 9);
1445
1446        let mut inflater = Inflater::new(4 * 1024 * 1024);
1447        let mut out = Vec::new();
1448        for piece in stream.chunks(777) {
1449            inflater.feed(piece);
1450            inflater.decode().unwrap();
1451            out.extend_from_slice(&inflater.take_output());
1452        }
1453        inflater.end_of_input();
1454        inflater.decode().unwrap();
1455        out.extend_from_slice(&inflater.take_output());
1456        assert_eq!(out, original);
1457    }
1458
1459    #[test]
1460    fn an_unfinished_stream_is_not_reported_as_complete() {
1461        // Half a stream must read as "not done", never as a short success —
1462        // that distinction is what stops a truncated PNG decoding to a
1463        // partial image.
1464        let stream = deflate_body(&vec![7_u8; 100_000], 6);
1465        let mut inflater = Inflater::new(1 << 20);
1466        inflater.feed(&stream[..stream.len() / 2]);
1467        inflater.decode().unwrap();
1468        assert!(!inflater.is_finished());
1469
1470        inflater.end_of_input();
1471        assert!(inflater.decode().is_err() || !inflater.is_finished());
1472    }
1473
1474    #[test]
1475    fn the_limit_is_enforced_incrementally_not_at_the_end() {
1476        // A bomb must be refused while decoding, not after the allocation it
1477        // was trying to provoke.
1478        let stream = deflate_body(&vec![0_u8; 4 * 1024 * 1024], 9);
1479        let mut inflater = Inflater::new(1024);
1480        inflater.feed(&stream);
1481        let error = inflater.decode().unwrap_err();
1482        assert_eq!(error.format(), "deflate", "{error}");
1483        assert!(
1484            inflater.produced() <= 1024,
1485            "produced {} bytes",
1486            inflater.produced()
1487        );
1488    }
1489}