1use alloc::vec::Vec;
10
11use crate::err::{Error, Result};
12
13#[derive(Debug)]
15pub struct Inflated {
16 pub data: Vec<u8>,
17 pub consumed: usize,
18}
19
20pub fn inflate_zlib(input: &[u8], size_hint: Option<usize>) -> Result<Inflated> {
22 if input.len() < 2 {
23 return Err(Error::UnexpectedEof);
24 }
25 let cmf = input[0];
26 let flg = input[1];
27 if cmf & 0x0f != 8 {
28 return Err(Error::Unsupported("zlib compression method"));
29 }
30 if (u16::from(cmf) << 8 | u16::from(flg)) % 31 != 0 {
31 return Err(Error::Corrupt("zlib header check"));
32 }
33 if flg & 0x20 != 0 {
34 return Err(Error::Unsupported("zlib preset dictionary"));
35 }
36
37 let mut br = BitReader::new(&input[2..]);
38 let data = inflate_raw(&mut br, size_hint)?;
39
40 br.align_to_byte();
42 let pos = 2 + br.byte_pos();
43 let sum = input
44 .get(pos..pos + 4)
45 .ok_or(Error::UnexpectedEof)
46 .map(|b| u32::from_be_bytes(b.try_into().unwrap()))?;
47 if sum != adler32(&data) {
48 return Err(Error::Checksum("zlib adler32"));
49 }
50
51 Ok(Inflated {
52 data,
53 consumed: pos + 4,
54 })
55}
56
57#[cfg(feature = "write")]
67pub fn deflate_zlib(data: &[u8]) -> Vec<u8> {
68 let stored_len = 2 + data.len() + data.len().div_ceil(STORED_BLOCK).max(1) * 5 + 4;
69 let mut out = Vec::with_capacity(stored_len.min(2 + data.len() / 2 + 64));
70 out.extend_from_slice(&[0x78, 0x01]); let mut bw = BitWriter::new(&mut out);
73 bw.bits(1, 1); bw.bits(1, 2); deflate_fixed_block(data, &mut bw);
76 bw.code(FIXED_END_CODE, FIXED_END_LEN);
77 bw.flush();
78
79 if out.len() + 4 >= stored_len {
80 return deflate_zlib_stored(data);
81 }
82 out.extend_from_slice(&adler32(data).to_be_bytes());
83 out
84}
85
86#[cfg(feature = "write")]
92pub fn deflate_zlib_stored(data: &[u8]) -> Vec<u8> {
93 let blocks = data.len().div_ceil(STORED_BLOCK).max(1);
95 let mut out = Vec::with_capacity(2 + data.len() + blocks * 5 + 4);
96 out.extend_from_slice(&[0x78, 0x01]);
97
98 if data.is_empty() {
99 out.extend_from_slice(&[0x01, 0x00, 0x00, 0xff, 0xff]);
100 } else {
101 let mut chunks = data.chunks(STORED_BLOCK).peekable();
102 while let Some(chunk) = chunks.next() {
103 out.push(if chunks.peek().is_none() { 0x01 } else { 0x00 });
104 let len = chunk.len() as u16;
105 out.extend_from_slice(&len.to_le_bytes());
106 out.extend_from_slice(&(!len).to_le_bytes());
107 out.extend_from_slice(chunk);
108 }
109 }
110
111 out.extend_from_slice(&adler32(data).to_be_bytes());
112 out
113}
114
115#[cfg(feature = "write")]
117const STORED_BLOCK: usize = 65_535;
118
119#[cfg(feature = "write")]
123const FIXED_END_CODE: u32 = 0;
124#[cfg(feature = "write")]
125const FIXED_END_LEN: u32 = 7;
126
127#[cfg(feature = "write")]
129const HASH_BITS: u32 = 12;
130#[cfg(feature = "write")]
131const HASH_SIZE: usize = 1 << HASH_BITS;
132#[cfg(feature = "write")]
134const MAX_DIST: usize = 32_768;
135#[cfg(feature = "write")]
137const MIN_MATCH: usize = 3;
138#[cfg(feature = "write")]
139const MAX_MATCH: usize = 258;
140
141#[cfg(feature = "write")]
142struct BitWriter<'a> {
143 out: &'a mut Vec<u8>,
144 bit_buf: u64,
145 bit_cnt: u32,
146}
147
148#[cfg(feature = "write")]
149impl<'a> BitWriter<'a> {
150 fn new(out: &'a mut Vec<u8>) -> Self {
151 Self {
152 out,
153 bit_buf: 0,
154 bit_cnt: 0,
155 }
156 }
157
158 fn bits(&mut self, value: u32, n: u32) {
160 self.bit_buf |= u64::from(value) << self.bit_cnt;
161 self.bit_cnt += n;
162 while self.bit_cnt >= 8 {
163 self.out.push(self.bit_buf as u8);
164 self.bit_buf >>= 8;
165 self.bit_cnt -= 8;
166 }
167 }
168
169 fn code(&mut self, code: u32, len: u32) {
171 let reversed = code.reverse_bits() >> (32 - len);
172 self.bits(reversed, len);
173 }
174
175 fn flush(&mut self) {
177 if self.bit_cnt > 0 {
178 self.out.push(self.bit_buf as u8);
179 self.bit_buf = 0;
180 self.bit_cnt = 0;
181 }
182 }
183
184 fn fixed_lit(&mut self, sym: u16) {
186 let (code, len) = match sym {
187 0..=143 => (0x30 + u32::from(sym), 8),
188 144..=255 => (0x190 + u32::from(sym - 144), 9),
189 256..=279 => (u32::from(sym - 256), 7),
190 _ => (0xc0 + u32::from(sym - 280), 8),
191 };
192 self.code(code, len);
193 }
194
195 fn fixed_match(&mut self, len: usize, dist: usize) {
197 let li = LEN_BASE
198 .iter()
199 .rposition(|&b| usize::from(b) <= len)
200 .unwrap();
201 self.fixed_lit(257 + li as u16);
202 self.bits(
203 (len - usize::from(LEN_BASE[li])) as u32,
204 u32::from(LEN_EXTRA[li]),
205 );
206
207 let di = DIST_BASE
208 .iter()
209 .rposition(|&b| usize::from(b) <= dist)
210 .unwrap();
211 self.code(di as u32, 5);
213 self.bits(
214 (dist - usize::from(DIST_BASE[di])) as u32,
215 u32::from(DIST_EXTRA[di]),
216 );
217 }
218}
219
220#[cfg(feature = "write")]
222fn hash3(data: &[u8], pos: usize) -> usize {
223 let v = u32::from(data[pos]) | u32::from(data[pos + 1]) << 8 | u32::from(data[pos + 2]) << 16;
224 (v.wrapping_mul(0x9E37_79B1) >> (32 - HASH_BITS)) as usize
225}
226
227#[cfg(feature = "write")]
230fn deflate_fixed_block(data: &[u8], bw: &mut BitWriter<'_>) {
231 let mut head = alloc::vec![0u32; HASH_SIZE];
233 let mut pos = 0;
234 while pos < data.len() {
235 let remaining = data.len() - pos;
236 let mut best = 0;
237 let mut best_dist = 0;
238 if remaining >= MIN_MATCH {
239 let h = hash3(data, pos);
240 let cand = head[h] as usize;
241 head[h] = (pos + 1) as u32;
242 if cand > 0 {
243 let cand = cand - 1;
244 let dist = pos - cand;
245 if dist <= MAX_DIST {
246 let limit = remaining.min(MAX_MATCH);
247 let len = (0..limit)
248 .take_while(|&k| data[cand + k] == data[pos + k])
249 .count();
250 if len >= MIN_MATCH {
251 best = len;
252 best_dist = dist;
253 }
254 }
255 }
256 }
257
258 if best == 0 {
259 bw.fixed_lit(u16::from(data[pos]));
260 pos += 1;
261 } else {
262 bw.fixed_match(best, best_dist);
263 for p in pos + 1..pos + best {
265 if data.len() - p >= MIN_MATCH {
266 head[hash3(data, p)] = (p + 1) as u32;
267 }
268 }
269 pos += best;
270 }
271 }
272}
273
274pub fn adler32(data: &[u8]) -> u32 {
276 const MOD: u32 = 65_521;
277 let mut a: u32 = 1;
278 let mut b: u32 = 0;
279 for chunk in data.chunks(5552) {
281 for &byte in chunk {
282 a += u32::from(byte);
283 b += a;
284 }
285 a %= MOD;
286 b %= MOD;
287 }
288 (b << 16) | a
289}
290
291struct BitReader<'a> {
294 data: &'a [u8],
295 pos: usize,
297 bit_buf: u32,
298 bit_cnt: u32,
299}
300
301impl<'a> BitReader<'a> {
302 fn new(data: &'a [u8]) -> Self {
303 Self {
304 data,
305 pos: 0,
306 bit_buf: 0,
307 bit_cnt: 0,
308 }
309 }
310
311 fn bits(&mut self, n: u32) -> Result<u32> {
313 while self.bit_cnt < n {
314 let byte = *self.data.get(self.pos).ok_or(Error::UnexpectedEof)?;
315 self.bit_buf |= u32::from(byte) << self.bit_cnt;
316 self.bit_cnt += 8;
317 self.pos += 1;
318 }
319 let out = self.bit_buf & ((1 << n) - 1);
320 self.bit_buf >>= n;
321 self.bit_cnt -= n;
322 Ok(out)
323 }
324
325 fn align_to_byte(&mut self) {
327 self.bit_buf >>= self.bit_cnt % 8;
328 self.bit_cnt -= self.bit_cnt % 8;
329 self.pos -= (self.bit_cnt / 8) as usize;
331 self.bit_buf = 0;
332 self.bit_cnt = 0;
333 }
334
335 fn byte_pos(&self) -> usize {
337 debug_assert_eq!(self.bit_cnt, 0);
338 self.pos
339 }
340
341 fn read_bytes(&mut self, n: usize, out: &mut Vec<u8>) -> Result<()> {
342 let end = self.pos.checked_add(n).ok_or(Error::UnexpectedEof)?;
343 let src = self.data.get(self.pos..end).ok_or(Error::UnexpectedEof)?;
344 out.extend_from_slice(src);
345 self.pos = end;
346 Ok(())
347 }
348}
349
350struct Huffman {
355 count: [u16; 16],
357 symbols: Vec<u16>,
359}
360
361impl Huffman {
362 fn new(lengths: &[u8]) -> Result<Self> {
364 let mut count = [0u16; 16];
365 for &len in lengths {
366 if len > 15 {
367 return Err(Error::Corrupt("huffman code length"));
368 }
369 count[usize::from(len)] += 1;
370 }
371
372 let mut left: i32 = 1;
374 for &len_count in &count[1..] {
375 left = (left << 1) - i32::from(len_count);
376 if left < 0 {
377 return Err(Error::Corrupt("oversubscribed huffman table"));
378 }
379 }
380
381 let mut offsets = [0u16; 16];
382 for n in 1..15 {
383 offsets[n + 1] = offsets[n] + count[n];
384 }
385 let mut symbols = alloc::vec![0u16; lengths.iter().filter(|&&l| l != 0).count()];
386 for (sym, &len) in lengths.iter().enumerate() {
387 if len != 0 {
388 symbols[usize::from(offsets[usize::from(len)])] = sym as u16;
389 offsets[usize::from(len)] += 1;
390 }
391 }
392 Ok(Self { count, symbols })
393 }
394
395 fn decode(&self, br: &mut BitReader<'_>) -> Result<u16> {
396 let mut code: u32 = 0;
397 let mut first: u32 = 0;
398 let mut index: u32 = 0;
399 for len in 1..16 {
400 code |= br.bits(1)?;
401 let count = u32::from(self.count[len]);
402 if code < first + count {
403 return Ok(self.symbols[(index + code - first) as usize]);
404 }
405 index += count;
406 first = (first + count) << 1;
407 code <<= 1;
408 }
409 Err(Error::Corrupt("invalid huffman code"))
410 }
411}
412
413const LEN_BASE: [u16; 29] = [
414 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,
415 163, 195, 227, 258,
416];
417const LEN_EXTRA: [u8; 29] = [
418 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,
419];
420const DIST_BASE: [u16; 30] = [
421 1, 2, 3, 4, 5, 7, 9, 13, 17, 25, 33, 49, 65, 97, 129, 193, 257, 385, 513, 769, 1025, 1537,
422 2049, 3073, 4097, 6145, 8193, 12289, 16385, 24577,
423];
424const DIST_EXTRA: [u8; 30] = [
425 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,
426 13,
427];
428
429const CLEN_ORDER: [usize; 19] = [
431 16, 17, 18, 0, 8, 7, 9, 6, 10, 5, 11, 4, 12, 3, 13, 2, 14, 1, 15,
432];
433
434fn inflate_raw(br: &mut BitReader<'_>, size_hint: Option<usize>) -> Result<Vec<u8>> {
435 let mut out = Vec::with_capacity(size_hint.unwrap_or(0));
436 loop {
437 let bfinal = br.bits(1)?;
438 match br.bits(2)? {
439 0 => inflate_stored(br, &mut out)?,
440 1 => {
441 let (lit, dist) = fixed_tables()?;
442 inflate_block(br, &lit, &dist, &mut out)?;
443 }
444 2 => {
445 let (lit, dist) = dynamic_tables(br)?;
446 inflate_block(br, &lit, &dist, &mut out)?;
447 }
448 _ => return Err(Error::Corrupt("deflate block type")),
449 }
450 if bfinal == 1 {
451 return Ok(out);
452 }
453 }
454}
455
456fn inflate_stored(br: &mut BitReader<'_>, out: &mut Vec<u8>) -> Result<()> {
457 br.align_to_byte();
458 let pos = br.byte_pos();
459 let header = br.data.get(pos..pos + 4).ok_or(Error::UnexpectedEof)?;
460 let len = u16::from_le_bytes(header[0..2].try_into().unwrap());
461 let nlen = u16::from_le_bytes(header[2..4].try_into().unwrap());
462 if len != !nlen {
463 return Err(Error::Corrupt("stored block length"));
464 }
465 br.pos = pos + 4;
466 br.read_bytes(usize::from(len), out)
467}
468
469fn fixed_tables() -> Result<(Huffman, Huffman)> {
470 let mut lit_lengths = [0u8; 288];
471 lit_lengths[0..144].fill(8);
472 lit_lengths[144..256].fill(9);
473 lit_lengths[256..280].fill(7);
474 lit_lengths[280..288].fill(8);
475 Ok((Huffman::new(&lit_lengths)?, Huffman::new(&[5u8; 30])?))
476}
477
478fn dynamic_tables(br: &mut BitReader<'_>) -> Result<(Huffman, Huffman)> {
479 let hlit = br.bits(5)? as usize + 257;
480 let hdist = br.bits(5)? as usize + 1;
481 let hclen = br.bits(4)? as usize + 4;
482 if hlit > 286 || hdist > 30 {
483 return Err(Error::Corrupt("dynamic table size"));
484 }
485
486 let mut clen_lengths = [0u8; 19];
487 for &idx in CLEN_ORDER.iter().take(hclen) {
488 clen_lengths[idx] = br.bits(3)? as u8;
489 }
490 let clen = Huffman::new(&clen_lengths)?;
491
492 let mut lengths = [0u8; 286 + 30];
494 let mut i = 0;
495 while i < hlit + hdist {
496 let sym = clen.decode(br)?;
497 match sym {
498 0..=15 => {
499 lengths[i] = sym as u8;
500 i += 1;
501 }
502 16 => {
503 if i == 0 {
504 return Err(Error::Corrupt("length repeat without previous"));
505 }
506 let prev = lengths[i - 1];
507 let n = br.bits(2)? as usize + 3;
508 repeat(&mut lengths, &mut i, hlit + hdist, prev, n)?;
509 }
510 17 => {
511 let n = br.bits(3)? as usize + 3;
512 repeat(&mut lengths, &mut i, hlit + hdist, 0, n)?;
513 }
514 18 => {
515 let n = br.bits(7)? as usize + 11;
516 repeat(&mut lengths, &mut i, hlit + hdist, 0, n)?;
517 }
518 _ => return Err(Error::Corrupt("code length symbol")),
519 }
520 }
521
522 if lengths[256] == 0 {
523 return Err(Error::Corrupt("missing end-of-block code"));
524 }
525 Ok((
526 Huffman::new(&lengths[..hlit])?,
527 Huffman::new(&lengths[hlit..hlit + hdist])?,
528 ))
529}
530
531fn repeat(lengths: &mut [u8], i: &mut usize, limit: usize, value: u8, n: usize) -> Result<()> {
532 if *i + n > limit {
533 return Err(Error::Corrupt("length repeat overflow"));
534 }
535 lengths[*i..*i + n].fill(value);
536 *i += n;
537 Ok(())
538}
539
540fn inflate_block(
541 br: &mut BitReader<'_>,
542 lit: &Huffman,
543 dist: &Huffman,
544 out: &mut Vec<u8>,
545) -> Result<()> {
546 loop {
547 let sym = lit.decode(br)?;
548 match sym {
549 0..=255 => out.push(sym as u8),
550 256 => return Ok(()),
551 257..=285 => {
552 let idx = usize::from(sym - 257);
553 let len = usize::from(LEN_BASE[idx]) + br.bits(u32::from(LEN_EXTRA[idx]))? as usize;
554
555 let dsym = usize::from(dist.decode(br)?);
556 if dsym >= 30 {
557 return Err(Error::Corrupt("distance symbol"));
558 }
559 let distance =
560 usize::from(DIST_BASE[dsym]) + br.bits(u32::from(DIST_EXTRA[dsym]))? as usize;
561 if distance > out.len() {
562 return Err(Error::Corrupt("distance beyond output"));
563 }
564
565 let start = out.len() - distance;
567 for k in 0..len {
568 let byte = out[start + k];
569 out.push(byte);
570 }
571 }
572 _ => return Err(Error::Corrupt("literal/length symbol")),
573 }
574 }
575}
576
577#[cfg(test)]
578mod tests {
579 use super::*;
580
581 const HELLO_FIXED: &[u8] = &[120, 156, 203, 72, 205, 201, 201, 7, 0, 6, 44, 2, 21];
584 const HELLO_STORED: &[u8] = &[
586 120, 1, 1, 5, 0, 250, 255, 104, 101, 108, 108, 111, 6, 44, 2, 21,
587 ];
588
589 #[test]
590 fn fixed_block() {
591 let r = inflate_zlib(HELLO_FIXED, None).unwrap();
592 assert_eq!(r.data, b"hello");
593 assert_eq!(r.consumed, HELLO_FIXED.len());
594 }
595
596 #[test]
597 fn stored_block() {
598 let r = inflate_zlib(HELLO_STORED, Some(5)).unwrap();
599 assert_eq!(r.data, b"hello");
600 assert_eq!(r.consumed, HELLO_STORED.len());
601 }
602
603 #[test]
605 fn trailing_data_ignored() {
606 let mut input = HELLO_FIXED.to_vec();
607 input.extend_from_slice(&[0xde, 0xad, 0xbe, 0xef]);
608 let r = inflate_zlib(&input, None).unwrap();
609 assert_eq!(r.data, b"hello");
610 assert_eq!(r.consumed, HELLO_FIXED.len());
611 }
612
613 #[test]
614 fn corrupt_adler32_rejected() {
615 let mut input = HELLO_FIXED.to_vec();
616 let last = input.len() - 1;
617 input[last] ^= 0xff;
618 assert_eq!(
619 inflate_zlib(&input, None).unwrap_err(),
620 Error::Checksum("zlib adler32")
621 );
622 }
623
624 #[test]
625 fn truncated_input_rejected() {
626 for n in 0..HELLO_FIXED.len() {
627 assert!(inflate_zlib(&HELLO_FIXED[..n], None).is_err(), "n={n}");
628 }
629 }
630
631 #[test]
632 fn adler32_vectors() {
633 assert_eq!(adler32(b""), 1);
634 assert_eq!(adler32(b"Wikipedia"), 0x11e6_0398);
635 }
636
637 #[cfg(feature = "write")]
640 #[test]
641 fn fixed_deflate_roundtrip() {
642 let mut seed: u32 = 12345;
644 let mut rand = || {
645 seed = seed.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
646 (seed >> 24) as u8
647 };
648 let cases: Vec<Vec<u8>> = vec![
649 Vec::new(),
650 vec![0x41],
651 b"hello".to_vec(),
652 vec![0u8; 1000], (0..=255u8).collect(), (0..40_000).map(|i| (i % 7) as u8).collect(), (0..100_000).map(|i| (i % 251) as u8).collect(), (0..70_000).map(|_| rand()).collect(), b"tree 1234\nparent abcd\nauthor A <a@example.com> 1 +0000\n\nmsg\n".repeat(50),
658 ];
659 for (i, data) in cases.iter().enumerate() {
660 let compressed = deflate_zlib(data);
661 let inflated = inflate_zlib(&compressed, Some(data.len())).unwrap();
662 assert_eq!(&inflated.data, data, "case {i}");
663 assert_eq!(inflated.consumed, compressed.len(), "case {i}");
664 assert!(
665 compressed.len() <= deflate_zlib_stored(data).len(),
666 "case {i}: 圧縮結果が stored より大きい"
667 );
668 }
669 let text = cases.last().unwrap();
671 assert!(deflate_zlib(text).len() < text.len() / 10);
672 }
673
674 #[cfg(feature = "write")]
675 #[test]
676 fn fixed_deflate_boundary_distances() {
677 for gap in [MAX_DIST - 1, MAX_DIST, MAX_DIST + 1] {
679 let mut data = b"ABCDEFGHIJKLMNOP".to_vec();
680 data.extend((0..gap - 16).map(|i| (i % 13) as u8 + b'a'));
681 data.extend_from_slice(b"ABCDEFGHIJKLMNOP");
682 let compressed = deflate_zlib(&data);
683 let inflated = inflate_zlib(&compressed, None).unwrap();
684 assert_eq!(inflated.data, data, "gap={gap}");
685 }
686 let data = vec![b'x'; 258 * 3 + 7];
688 let inflated = inflate_zlib(&deflate_zlib(&data), None).unwrap();
689 assert_eq!(inflated.data, data);
690 }
691
692 #[cfg(feature = "write")]
693 #[test]
694 fn fixed_deflate_matches_reference_stream() {
695 assert_eq!(deflate_zlib(b"hello")[2..], HELLO_FIXED[2..]);
699 }
700
701 #[cfg(feature = "write")]
702 #[test]
703 fn stored_deflate_roundtrip() {
704 for len in [0usize, 1, 100, 65_534, 65_535, 65_536, 200_000] {
705 let data: Vec<u8> = (0..len).map(|i| (i % 251) as u8).collect();
706 let compressed = deflate_zlib_stored(&data);
707 let inflated = inflate_zlib(&compressed, Some(len)).unwrap();
708 assert_eq!(inflated.data, data, "len={len}");
709 assert_eq!(inflated.consumed, compressed.len(), "len={len}");
710 }
711 }
712}