Skip to main content

conspire/io/deflate/
mod.rs

1#[cfg(test)]
2mod test;
3
4use std::io::{Error, ErrorKind, Result};
5
6const MAX_BITS: usize = 15;
7const WINDOW: usize = 32768;
8const MIN_MATCH: usize = 3;
9const MAX_MATCH: usize = 258;
10
11struct BitReader<'a> {
12    data: &'a [u8],
13    pos: usize,
14}
15
16impl<'a> BitReader<'a> {
17    fn new(data: &'a [u8]) -> Self {
18        Self { data, pos: 0 }
19    }
20    fn read_bit(&mut self) -> Result<u32> {
21        let byte = self.pos / 8;
22        let bit = self.pos % 8;
23        let value = *self
24            .data
25            .get(byte)
26            .ok_or_else(|| invalid("truncated deflate stream"))?;
27        self.pos += 1;
28        Ok(((value >> bit) & 1) as u32)
29    }
30    fn read_bits(&mut self, n: u32) -> Result<u32> {
31        let mut value = 0;
32        for i in 0..n {
33            value |= self.read_bit()? << i;
34        }
35        Ok(value)
36    }
37    fn align(&mut self) {
38        self.pos = self.pos.div_ceil(8) * 8;
39    }
40    fn read_bytes(&mut self, n: usize) -> Result<&'a [u8]> {
41        let start = self.pos / 8;
42        let end = start + n;
43        if end > self.data.len() {
44            return Err(invalid("truncated deflate stream"));
45        }
46        self.pos = end * 8;
47        Ok(&self.data[start..end])
48    }
49    fn peek_bits(&self, n: u32) -> u32 {
50        let mut code = 0;
51        for pos in self.pos..self.pos + n as usize {
52            let byte = self.data.get(pos / 8).copied().unwrap_or(0);
53            let bit = (byte >> (pos % 8)) & 1;
54            code = (code << 1) | bit as u32;
55        }
56        code
57    }
58    fn consume(&mut self, n: u32) {
59        self.pos += n as usize;
60    }
61}
62
63struct BitWriter {
64    bytes: Vec<u8>,
65    current: u8,
66    count: u32,
67}
68
69impl BitWriter {
70    fn new() -> Self {
71        Self {
72            bytes: Vec::new(),
73            current: 0,
74            count: 0,
75        }
76    }
77    fn write_bit(&mut self, bit: u32) {
78        self.current |= ((bit & 1) as u8) << self.count;
79        self.count += 1;
80        if self.count == 8 {
81            self.bytes.push(self.current);
82            self.current = 0;
83            self.count = 0;
84        }
85    }
86    fn write_bits(&mut self, value: u32, n: u32) {
87        for i in 0..n {
88            self.write_bit((value >> i) & 1);
89        }
90    }
91    fn write_huffman(&mut self, code: u32, len: u32) {
92        for i in (0..len).rev() {
93            self.write_bit((code >> i) & 1);
94        }
95    }
96    fn align(&mut self) {
97        if self.count > 0 {
98            self.bytes.push(self.current);
99            self.current = 0;
100            self.count = 0;
101        }
102    }
103    fn finish(mut self) -> Vec<u8> {
104        self.align();
105        self.bytes
106    }
107}
108
109struct Huffman {
110    table: Vec<(u16, u8)>,
111}
112
113impl Huffman {
114    fn build(lengths: &[u8]) -> Self {
115        let mut bl_count = [0u32; MAX_BITS + 1];
116        for &length in lengths {
117            bl_count[length as usize] += 1;
118        }
119        bl_count[0] = 0;
120        let mut next_code = [0u32; MAX_BITS + 1];
121        let mut code = 0u32;
122        for bits in 1..=MAX_BITS {
123            code = (code + bl_count[bits - 1]) << 1;
124            next_code[bits] = code;
125        }
126        let mut table = vec![(0u16, 0u8); 1 << MAX_BITS];
127        for (symbol, &length) in lengths.iter().enumerate() {
128            if length == 0 {
129                continue;
130            }
131            let length = length as usize;
132            let code = next_code[length];
133            next_code[length] += 1;
134            let shift = MAX_BITS - length;
135            let base = (code as usize) << shift;
136            table[base..base + (1 << shift)].fill((symbol as u16, length as u8));
137        }
138        Self { table }
139    }
140    fn decode(&self, reader: &mut BitReader) -> Result<u16> {
141        let window = reader.peek_bits(MAX_BITS as u32) as usize;
142        let (symbol, length) = self.table[window];
143        if length == 0 || reader.pos + length as usize > reader.data.len() * 8 {
144            return Err(invalid("invalid Huffman code in deflate stream"));
145        }
146        reader.consume(length as u32);
147        Ok(symbol)
148    }
149}
150
151fn fixed_literal_lengths() -> Vec<u8> {
152    let mut lengths = vec![0u8; 288];
153    lengths[0..144].fill(8);
154    lengths[144..256].fill(9);
155    lengths[256..280].fill(7);
156    lengths[280..288].fill(8);
157    lengths
158}
159
160fn fixed_distance_lengths() -> Vec<u8> {
161    vec![5u8; 30]
162}
163
164const LENGTH_BASE: [u16; 29] = [
165    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,
166    163, 195, 227, 258,
167];
168const LENGTH_EXTRA: [u32; 29] = [
169    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,
170];
171const DIST_BASE: [u32; 30] = [
172    1, 2, 3, 4, 5, 7, 9, 13, 17, 25, 33, 49, 65, 97, 129, 193, 257, 385, 513, 769, 1025, 1537,
173    2049, 3073, 4097, 6145, 8193, 12289, 16385, 24577,
174];
175const DIST_EXTRA: [u32; 30] = [
176    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,
177    13,
178];
179const CODE_LENGTH_ORDER: [usize; 19] = [
180    16, 17, 18, 0, 8, 7, 9, 6, 10, 5, 11, 4, 12, 3, 13, 2, 14, 1, 15,
181];
182
183fn length_code(length: usize) -> (u16, u32, u32) {
184    let index = LENGTH_BASE
185        .iter()
186        .rposition(|&base| base as usize <= length)
187        .unwrap();
188    (
189        257 + index as u16,
190        LENGTH_EXTRA[index],
191        (length - LENGTH_BASE[index] as usize) as u32,
192    )
193}
194
195fn distance_code(distance: usize) -> (u16, u32, u32) {
196    let index = DIST_BASE
197        .iter()
198        .rposition(|&base| base as usize <= distance)
199        .unwrap();
200    (
201        index as u16,
202        DIST_EXTRA[index],
203        (distance - DIST_BASE[index] as usize) as u32,
204    )
205}
206
207pub fn inflate(data: &[u8]) -> Result<Vec<u8>> {
208    let mut reader = BitReader::new(data);
209    let mut output = Vec::new();
210    loop {
211        let is_final = reader.read_bit()? == 1;
212        match reader.read_bits(2)? {
213            0 => {
214                reader.align();
215                let len_bytes = reader.read_bytes(4)?;
216                let len = u16::from_le_bytes([len_bytes[0], len_bytes[1]]) as usize;
217                let nlen = u16::from_le_bytes([len_bytes[2], len_bytes[3]]);
218                if nlen != !(len as u16) {
219                    return Err(invalid("stored block LEN/NLEN mismatch"));
220                }
221                output.extend_from_slice(reader.read_bytes(len)?);
222            }
223            1 => {
224                let literal = Huffman::build(&fixed_literal_lengths());
225                let distance = Huffman::build(&fixed_distance_lengths());
226                inflate_block(&mut reader, &literal, &distance, &mut output)?;
227            }
228            2 => {
229                let (literal, distance) = read_dynamic_tables(&mut reader)?;
230                inflate_block(&mut reader, &literal, &distance, &mut output)?;
231            }
232            _ => return Err(invalid("reserved block type in deflate stream")),
233        }
234        if is_final {
235            break;
236        }
237    }
238    Ok(output)
239}
240
241fn read_dynamic_tables(reader: &mut BitReader) -> Result<(Huffman, Huffman)> {
242    let hlit = reader.read_bits(5)? as usize + 257;
243    let hdist = reader.read_bits(5)? as usize + 1;
244    let hclen = reader.read_bits(4)? as usize + 4;
245    let mut code_length_lengths = [0u8; 19];
246    for &order in CODE_LENGTH_ORDER.iter().take(hclen) {
247        code_length_lengths[order] = reader.read_bits(3)? as u8;
248    }
249    let code_length_huffman = Huffman::build(&code_length_lengths);
250    let mut lengths = Vec::with_capacity(hlit + hdist);
251    while lengths.len() < hlit + hdist {
252        match code_length_huffman.decode(reader)? {
253            symbol @ 0..=15 => lengths.push(symbol as u8),
254            16 => {
255                let repeat = reader.read_bits(2)? + 3;
256                let previous = *lengths
257                    .last()
258                    .ok_or_else(|| invalid("repeat code 16 with no previous length"))?;
259                lengths.extend(std::iter::repeat_n(previous, repeat as usize));
260            }
261            17 => {
262                let repeat = reader.read_bits(3)? + 3;
263                lengths.extend(std::iter::repeat_n(0, repeat as usize));
264            }
265            _ => {
266                let repeat = reader.read_bits(7)? + 11;
267                lengths.extend(std::iter::repeat_n(0, repeat as usize));
268            }
269        }
270    }
271    if lengths.len() != hlit + hdist {
272        return Err(invalid("dynamic Huffman code-length overflow"));
273    }
274    let literal = Huffman::build(&lengths[..hlit]);
275    let distance = Huffman::build(&lengths[hlit..]);
276    Ok((literal, distance))
277}
278
279fn inflate_block(
280    reader: &mut BitReader,
281    literal: &Huffman,
282    distance: &Huffman,
283    output: &mut Vec<u8>,
284) -> Result<()> {
285    loop {
286        let symbol = literal.decode(reader)?;
287        match symbol {
288            0..=255 => output.push(symbol as u8),
289            256 => return Ok(()),
290            257..=285 => {
291                let index = (symbol - 257) as usize;
292                let length =
293                    LENGTH_BASE[index] as usize + reader.read_bits(LENGTH_EXTRA[index])? as usize;
294                let dist_symbol = distance.decode(reader)? as usize;
295                if dist_symbol >= DIST_BASE.len() {
296                    return Err(invalid("invalid distance symbol"));
297                }
298                let dist = DIST_BASE[dist_symbol] as usize
299                    + reader.read_bits(DIST_EXTRA[dist_symbol])? as usize;
300                if dist > output.len() {
301                    return Err(invalid("distance exceeds output so far"));
302                }
303                let start = output.len() - dist;
304                for i in 0..length {
305                    let byte = output[start + i];
306                    output.push(byte);
307                }
308            }
309            other => return Err(invalid(format!("invalid literal/length symbol {other}"))),
310        }
311    }
312}
313
314enum Token {
315    Literal(u8),
316    Match { length: u16, distance: u16 },
317}
318
319fn lz77(data: &[u8]) -> Vec<Token> {
320    const HASH_BITS: usize = 15;
321    const HASH_SIZE: usize = 1 << HASH_BITS;
322    const MAX_CHAIN: usize = 128;
323    let n = data.len();
324    let mut head = vec![u32::MAX; HASH_SIZE];
325    let mut prev = vec![u32::MAX; n];
326    let hash = |i: usize| -> usize {
327        ((data[i] as u32) ^ ((data[i + 1] as u32) << 5) ^ ((data[i + 2] as u32) << 10)) as usize
328            & (HASH_SIZE - 1)
329    };
330    let insert = |i: usize, head: &mut [u32], prev: &mut [u32]| {
331        if i + MIN_MATCH <= n {
332            let h = hash(i);
333            prev[i] = head[h];
334            head[h] = i as u32;
335        }
336    };
337    let mut tokens = Vec::new();
338    let mut i = 0;
339    while i < n {
340        let mut best_len = 0;
341        let mut best_dist = 0;
342        if i + MIN_MATCH <= n {
343            let h = hash(i);
344            let mut candidate = head[h];
345            let mut tries = 0;
346            let max_len = (n - i).min(MAX_MATCH);
347            while candidate != u32::MAX && tries < MAX_CHAIN {
348                let c = candidate as usize;
349                if i - c > WINDOW {
350                    break;
351                }
352                if best_len < max_len && data[c + best_len] == data[i + best_len] {
353                    let mut len = 0;
354                    while len < max_len && data[c + len] == data[i + len] {
355                        len += 1;
356                    }
357                    if len > best_len {
358                        best_len = len;
359                        best_dist = i - c;
360                    }
361                    if best_len >= max_len {
362                        break;
363                    }
364                }
365                candidate = prev[c];
366                tries += 1;
367            }
368        }
369        if best_len >= MIN_MATCH {
370            tokens.push(Token::Match {
371                length: best_len as u16,
372                distance: best_dist as u16,
373            });
374            let end = i + best_len;
375            while i < end {
376                insert(i, &mut head, &mut prev);
377                i += 1;
378            }
379        } else {
380            tokens.push(Token::Literal(data[i]));
381            insert(i, &mut head, &mut prev);
382            i += 1;
383        }
384    }
385    tokens
386}
387
388struct FixedCodes {
389    literal: [(u32, u32); 288],
390    distance: [(u32, u32); 30],
391}
392
393fn fixed_codes() -> FixedCodes {
394    let mut literal = [(0u32, 0u32); 288];
395    for (symbol, code) in literal.iter_mut().enumerate().take(144) {
396        *code = (0b0011_0000 + symbol as u32, 8);
397    }
398    for (symbol, code) in literal.iter_mut().enumerate().take(256).skip(144) {
399        *code = (0b1_1001_0000 + (symbol as u32 - 144), 9);
400    }
401    for (symbol, code) in literal.iter_mut().enumerate().take(280).skip(256) {
402        *code = (symbol as u32 - 256, 7);
403    }
404    for (symbol, code) in literal.iter_mut().enumerate().take(288).skip(280) {
405        *code = (0b1100_0000 + (symbol as u32 - 280), 8);
406    }
407    let mut distance = [(0u32, 0u32); 30];
408    for (symbol, code) in distance.iter_mut().enumerate() {
409        *code = (symbol as u32, 5);
410    }
411    FixedCodes { literal, distance }
412}
413
414pub fn deflate(data: &[u8]) -> Vec<u8> {
415    let mut writer = BitWriter::new();
416    if data.is_empty() {
417        writer.write_bit(1);
418        writer.write_bits(1, 2);
419        let (code, len) = fixed_codes().literal[256];
420        writer.write_huffman(code, len);
421        return writer.finish();
422    }
423    let tokens = lz77(data);
424    let codes = fixed_codes();
425    writer.write_bit(1);
426    writer.write_bits(1, 2);
427    for token in &tokens {
428        match token {
429            Token::Literal(byte) => {
430                let (code, len) = codes.literal[*byte as usize];
431                writer.write_huffman(code, len);
432            }
433            Token::Match { length, distance } => {
434                let (symbol, extra_bits, extra_val) = length_code(*length as usize);
435                let (code, len) = codes.literal[symbol as usize];
436                writer.write_huffman(code, len);
437                writer.write_bits(extra_val, extra_bits);
438                let (dsymbol, dextra_bits, dextra_val) = distance_code(*distance as usize);
439                let (dcode, dlen) = codes.distance[dsymbol as usize];
440                writer.write_huffman(dcode, dlen);
441                writer.write_bits(dextra_val, dextra_bits);
442            }
443        }
444    }
445    let (code, len) = codes.literal[256];
446    writer.write_huffman(code, len);
447    writer.finish()
448}
449
450pub fn adler32(data: &[u8]) -> u32 {
451    const MOD_ADLER: u32 = 65521;
452    let mut a: u32 = 1;
453    let mut b: u32 = 0;
454    for &byte in data {
455        a = (a + byte as u32) % MOD_ADLER;
456        b = (b + a) % MOD_ADLER;
457    }
458    (b << 16) | a
459}
460
461pub fn zlib_encode(data: &[u8]) -> Vec<u8> {
462    let mut out = vec![0x78, 0x01];
463    out.extend(deflate(data));
464    out.extend_from_slice(&adler32(data).to_be_bytes());
465    out
466}
467
468pub fn zlib_decode(data: &[u8]) -> Result<Vec<u8>> {
469    if data.len() < 6 {
470        return Err(invalid("zlib stream too short"));
471    }
472    let cmf = data[0];
473    let flg = data[1];
474    if cmf & 0x0F != 8 {
475        return Err(unsupported(
476            "only the DEFLATE zlib compression method is supported",
477        ));
478    }
479    if (flg & 0x20) != 0 {
480        return Err(unsupported("zlib preset dictionaries are not supported"));
481    }
482    if !((cmf as u32) * 256 + flg as u32).is_multiple_of(31) {
483        return Err(invalid("zlib header check bits are invalid"));
484    }
485    let payload = &data[2..data.len() - 4];
486    let output = inflate(payload)?;
487    let trailer = &data[data.len() - 4..];
488    let checksum = u32::from_be_bytes([trailer[0], trailer[1], trailer[2], trailer[3]]);
489    if adler32(&output) != checksum {
490        return Err(invalid("zlib Adler-32 checksum mismatch"));
491    }
492    Ok(output)
493}
494
495fn invalid(message: impl Into<String>) -> Error {
496    Error::new(ErrorKind::InvalidData, message.into())
497}
498
499fn unsupported(message: &str) -> Error {
500    Error::new(ErrorKind::Unsupported, message.to_string())
501}