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}