1use crate::{Error, Result};
38
39use crate::checksum::Adler32;
40
41const WINDOW: usize = 32 * 1024;
44
45const COMPACT_AT: usize = 64 * 1024;
48
49#[derive(Debug)]
56enum Halt {
57 NeedInput,
59 Fatal(Error),
61}
62
63impl From<Error> for Halt {
64 fn from(error: Error) -> Self {
65 Self::Fatal(error)
66 }
67}
68
69type Step<T> = std::result::Result<T, Halt>;
71
72#[derive(Debug, Default)]
78struct BitReader {
79 data: Vec<u8>,
80 position: usize,
82 bits: u64,
84 count: u32,
86 ended: bool,
88}
89
90#[derive(Debug, Clone, Copy)]
92struct Checkpoint {
93 position: usize,
94 bits: u64,
95 count: u32,
96}
97
98impl BitReader {
99 fn feed(&mut self, more: &[u8]) {
101 self.data.extend_from_slice(more);
102 }
103
104 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 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 fn drain_unconsumed(&mut self) -> Vec<u8> {
143 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 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 fn short(&self) -> Halt {
175 if self.ended {
176 Halt::Fatal(truncated())
177 } else {
178 Halt::NeedInput
179 }
180 }
181
182 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 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 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 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 fn align(&mut self) {
218 let extra = self.count % 8;
219 self.bits >>= extra;
220 self.count -= extra;
221 }
222
223 fn take_bytes_upto(&mut self, n: usize, out: &mut Vec<u8>) -> usize {
228 let mut taken = 0;
229 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
244fn truncated() -> Error {
246 Error::malformed("deflate", "stream ended in the middle of a symbol")
247}
248
249const MAX_BITS: usize = 15;
251
252const FAST_BITS: u32 = 10;
254
255#[derive(Debug, Clone)]
263struct Huffman {
264 counts: [u16; MAX_BITS + 1],
266 symbols: Vec<u16>,
268 fast: Vec<u16>,
271}
272
273impl Huffman {
274 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 if let Some(slot) = counts.get_mut(0) {
291 *slot = 0;
292 }
293
294 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 *slot = symbol as u16;
332 }
333 }
334
335 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 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 fn decode(&self, reader: &mut BitReader) -> Step<u16> {
378 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 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 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
422const 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];
427const 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];
431const 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];
436const 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];
441const 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
446fn 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
460fn fixed_distance_table() -> Result<Huffman> {
462 Huffman::new(&[5_u8; 30])
463}
464
465#[derive(Debug)]
467enum State {
468 BlockHeader,
470 Stored { remaining: usize, last: bool },
472 Coded {
474 literals: Box<Huffman>,
475 distances: Box<Huffman>,
476 last: bool,
477 },
478 Done,
480}
481
482#[derive(Debug)]
489pub struct Inflater {
490 reader: BitReader,
491 state: State,
492 window: Vec<u8>,
494 pending: usize,
496 produced: usize,
498 limit: usize,
499}
500
501impl Inflater {
502 #[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 pub fn feed(&mut self, data: &[u8]) {
517 self.reader.feed(data);
518 }
519
520 pub const fn end_of_input(&mut self) {
522 self.reader.end();
523 }
524
525 #[must_use]
527 pub const fn is_finished(&self) -> bool {
528 matches!(self.state, State::Done)
529 }
530
531 #[must_use]
533 pub const fn produced(&self) -> usize {
534 self.produced
535 }
536
537 #[must_use]
542 pub fn retained(&self) -> usize {
543 self.window.len()
544 }
545
546 fn drain_unconsumed_input(&mut self) -> Vec<u8> {
549 self.reader.drain_unconsumed()
550 }
551
552 pub fn take_output(&mut self) -> Vec<u8> {
558 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 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 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 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 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 fn check_limit(&self, adding: usize) -> Step<()> {
718 check_limit(self.produced, adding, self.limit)
719 }
720}
721
722fn 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
733fn 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 0..=255 => {
749 check_limit(*produced, 1, limit)?;
750 window.push(symbol as u8);
751 *produced += 1;
752 Ok(true)
753 }
754 256 => Ok(false),
756 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 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 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
815fn 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 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 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
889fn 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
906pub 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#[derive(Debug)]
933pub struct ZlibStream {
934 header: Vec<u8>,
935 inflater: Inflater,
936 adler: Adler32,
937 trailer: Vec<u8>,
939 ended: bool,
940}
941
942impl ZlibStream {
943 #[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 pub fn push(&mut self, mut data: &[u8]) -> Result<Vec<u8>> {
962 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 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 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 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
1056fn 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 return Err(Error::malformed(
1080 "zlib",
1081 "preset dictionaries are not supported",
1082 ));
1083 }
1084 Ok(())
1085}
1086
1087pub 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 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 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 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 assert!(zlib_decompress(&[0x77, 0x00, 0x00], 16).is_err());
1179 assert!(zlib_decompress(&[0x78, 0x00, 0x00], 16).is_err());
1181 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 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 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 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 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 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 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 assert!(Huffman::new(&[1, 1, 1]).is_err());
1255 assert!(Huffman::new(&[1, 1]).is_ok());
1257 assert!(Huffman::new(&[16]).is_err());
1259 }
1260
1261 #[test]
1262 fn overlapping_back_references_encode_runs() {
1263 let stream = [0x4B, 0x4C, 0x84, 0x00, 0x00];
1267 assert_eq!(inflate_to(&stream, 64).unwrap(), b"aaaaaaaa");
1268 }
1269
1270 fn compress(data: &[u8], level: u8) -> Vec<u8> {
1272 crate::deflate::zlib_compress(data, crate::deflate::Level::new(level).unwrap()).unwrap()
1273 }
1274
1275 fn deflate_body(data: &[u8], level: u8) -> Vec<u8> {
1278 compress(data, level).split_off(2)
1279 }
1280
1281 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 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 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 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 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 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 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 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 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}