use crate::{Error, Result};
use crate::checksum::Adler32;
const WINDOW: usize = 32 * 1024;
const COMPACT_AT: usize = 64 * 1024;
#[derive(Debug)]
enum Halt {
NeedInput,
Fatal(Error),
}
impl From<Error> for Halt {
fn from(error: Error) -> Self {
Self::Fatal(error)
}
}
type Step<T> = std::result::Result<T, Halt>;
#[derive(Debug, Default)]
struct BitReader {
data: Vec<u8>,
position: usize,
bits: u64,
count: u32,
ended: bool,
}
#[derive(Debug, Clone, Copy)]
struct Checkpoint {
position: usize,
bits: u64,
count: u32,
}
impl BitReader {
fn feed(&mut self, more: &[u8]) {
self.data.extend_from_slice(more);
}
const fn end(&mut self) {
self.ended = true;
}
fn checkpoint(&self) -> Checkpoint {
Checkpoint {
position: self.position,
bits: self.bits,
count: self.count,
}
}
fn restore(&mut self, at: Checkpoint) {
self.position = at.position;
self.bits = at.bits;
self.count = at.count;
}
fn compact(&mut self) {
if self.position >= COMPACT_AT && self.position * 2 >= self.data.len() {
self.data.drain(..self.position);
self.position = 0;
}
}
fn drain_unconsumed(&mut self) -> Vec<u8> {
self.align();
let mut out = Vec::new();
while self.count >= 8 {
out.push((self.bits & 0xFF) as u8);
self.bits >>= 8;
self.count -= 8;
}
if let Some(rest) = self.data.get(self.position..) {
out.extend_from_slice(rest);
}
self.position = self.data.len();
out
}
fn fill(&mut self, want: u32) {
while self.count < want {
let Some(&byte) = self.data.get(self.position) else {
break;
};
self.bits |= u64::from(byte) << self.count;
self.position += 1;
self.count += 8;
}
}
fn short(&self) -> Halt {
if self.ended {
Halt::Fatal(truncated())
} else {
Halt::NeedInput
}
}
fn take(&mut self, n: u32) -> Step<u32> {
if n == 0 {
return Ok(0);
}
self.fill(n);
if self.count < n {
return Err(self.short());
}
let mask = (1_u64 << n) - 1;
let value = (self.bits & mask) as u32;
self.bits >>= n;
self.count -= n;
Ok(value)
}
fn peek(&mut self, n: u32) -> u32 {
self.fill(n);
let mask = (1_u64 << n) - 1;
(self.bits & mask) as u32
}
fn skip(&mut self, n: u32) -> Step<()> {
if self.count < n {
return Err(self.short());
}
self.bits >>= n;
self.count -= n;
Ok(())
}
fn align(&mut self) {
let extra = self.count % 8;
self.bits >>= extra;
self.count -= extra;
}
fn take_bytes_upto(&mut self, n: usize, out: &mut Vec<u8>) -> usize {
let mut taken = 0;
while taken < n && self.count >= 8 {
out.push((self.bits & 0xFF) as u8);
self.bits >>= 8;
self.count -= 8;
taken += 1;
}
let rest = self.data.get(self.position..).unwrap_or(&[]);
let run = rest.get(..(n - taken).min(rest.len())).unwrap_or(&[]);
out.extend_from_slice(run);
self.position += run.len();
taken + run.len()
}
}
fn truncated() -> Error {
Error::malformed("deflate", "stream ended in the middle of a symbol")
}
const MAX_BITS: usize = 15;
const FAST_BITS: u32 = 10;
#[derive(Debug, Clone)]
struct Huffman {
counts: [u16; MAX_BITS + 1],
symbols: Vec<u16>,
fast: Vec<u16>,
}
impl Huffman {
fn new(lengths: &[u8]) -> Result<Self> {
let mut counts = [0_u16; MAX_BITS + 1];
for &length in lengths {
let length = length as usize;
if length > MAX_BITS {
return Err(Error::malformed(
"deflate",
format!("code length {length} exceeds the {MAX_BITS}-bit maximum"),
));
}
if let Some(slot) = counts.get_mut(length) {
*slot += 1;
}
}
if let Some(slot) = counts.get_mut(0) {
*slot = 0;
}
let mut left = 1_i32;
for length in 1..=MAX_BITS {
left <<= 1;
left -= i32::from(counts.get(length).copied().unwrap_or(0));
if left < 0 {
return Err(Error::malformed(
"deflate",
"Huffman table is over-subscribed",
));
}
}
let mut offsets = [0_u16; MAX_BITS + 2];
for length in 1..=MAX_BITS {
let next = offsets.get(length).copied().unwrap_or(0)
+ counts.get(length).copied().unwrap_or(0);
if let Some(slot) = offsets.get_mut(length + 1) {
*slot = next;
}
}
let total: usize = counts.iter().map(|&c| c as usize).sum();
let mut symbols = vec![0_u16; total];
let mut cursor = offsets;
for (symbol, &length) in lengths.iter().enumerate() {
if length == 0 {
continue;
}
let length = length as usize;
let Some(at) = cursor.get_mut(length) else {
continue;
};
let index = *at as usize;
*at += 1;
if let Some(slot) = symbols.get_mut(index) {
*slot = symbol as u16;
}
}
let mut next_code = [0_u32; MAX_BITS + 1];
let mut code = 0_u32;
for length in 1..=MAX_BITS {
code = (code + u32::from(counts.get(length - 1).copied().unwrap_or(0))) << 1;
if let Some(slot) = next_code.get_mut(length) {
*slot = code;
}
}
let mut fast = vec![0_u16; 1 << FAST_BITS];
for (symbol, &length) in lengths.iter().enumerate() {
let length = u32::from(length);
if length == 0 || length > FAST_BITS {
continue;
}
let Some(slot) = next_code.get_mut(length as usize) else {
continue;
};
let code = *slot;
*slot += 1;
let reversed = code.reverse_bits() >> (32 - length);
let entry = (symbol as u16) << 4 | length as u16;
for index in (reversed as usize..fast.len()).step_by(1 << length) {
if let Some(cell) = fast.get_mut(index) {
*cell = entry;
}
}
}
Ok(Self {
counts,
symbols,
fast,
})
}
fn decode(&self, reader: &mut BitReader) -> Step<u16> {
reader.fill(MAX_BITS as u32);
if reader.count < MAX_BITS as u32 && !reader.ended {
return Err(Halt::NeedInput);
}
let mut code = 0_i32;
let mut first = 0_i32;
let mut index = 0_i32;
let peeked = reader.peek(MAX_BITS as u32);
let entry = self
.fast
.get((peeked & ((1 << FAST_BITS) - 1)) as usize)
.copied()
.unwrap_or(0);
if entry != 0 {
reader.skip(u32::from(entry & 0xF))?;
return Ok(entry >> 4);
}
for length in 1..=MAX_BITS {
code |= ((peeked >> (length - 1)) & 1) as i32;
let count = i32::from(self.counts.get(length).copied().unwrap_or(0));
if code - first < count {
reader.skip(length as u32)?;
let position = (index + (code - first)) as usize;
return self.symbols.get(position).copied().ok_or_else(|| {
Halt::Fatal(Error::malformed("deflate", "invalid Huffman symbol"))
});
}
index += count;
first = (first + count) << 1;
code <<= 1;
}
Err(Halt::Fatal(Error::malformed(
"deflate",
"no Huffman code matched within 15 bits",
)))
}
}
const LENGTH_BASE: [u16; 29] = [
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,
163, 195, 227, 258,
];
const LENGTH_EXTRA: [u8; 29] = [
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,
];
const DISTANCE_BASE: [u16; 30] = [
1, 2, 3, 4, 5, 7, 9, 13, 17, 25, 33, 49, 65, 97, 129, 193, 257, 385, 513, 769, 1025, 1537,
2049, 3073, 4097, 6145, 8193, 12289, 16385, 24577,
];
const DISTANCE_EXTRA: [u8; 30] = [
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,
13,
];
const CODE_LENGTH_ORDER: [usize; 19] = [
16, 17, 18, 0, 8, 7, 9, 6, 10, 5, 11, 4, 12, 3, 13, 2, 14, 1, 15,
];
fn fixed_literal_table() -> Result<Huffman> {
let mut lengths = [0_u8; 288];
for (symbol, slot) in lengths.iter_mut().enumerate() {
*slot = match symbol {
0..=143 => 8,
144..=255 => 9,
256..=279 => 7,
_ => 8,
};
}
Huffman::new(&lengths)
}
fn fixed_distance_table() -> Result<Huffman> {
Huffman::new(&[5_u8; 30])
}
#[derive(Debug)]
enum State {
BlockHeader,
Stored { remaining: usize, last: bool },
Coded {
literals: Box<Huffman>,
distances: Box<Huffman>,
last: bool,
},
Done,
}
#[derive(Debug)]
pub struct Inflater {
reader: BitReader,
state: State,
window: Vec<u8>,
pending: usize,
produced: usize,
limit: usize,
}
impl Inflater {
#[must_use]
pub fn new(limit: usize) -> Self {
Self {
reader: BitReader::default(),
state: State::BlockHeader,
window: Vec::new(),
pending: 0,
produced: 0,
limit,
}
}
pub fn feed(&mut self, data: &[u8]) {
self.reader.feed(data);
}
pub const fn end_of_input(&mut self) {
self.reader.end();
}
#[must_use]
pub const fn is_finished(&self) -> bool {
matches!(self.state, State::Done)
}
#[must_use]
pub const fn produced(&self) -> usize {
self.produced
}
#[must_use]
pub fn retained(&self) -> usize {
self.window.len()
}
fn drain_unconsumed_input(&mut self) -> Vec<u8> {
self.reader.drain_unconsumed()
}
pub fn take_output(&mut self) -> Vec<u8> {
let out = self.window.get(self.pending..).unwrap_or(&[]).to_vec();
if self.window.len() > WINDOW {
self.window.drain(..self.window.len() - WINDOW);
}
self.pending = self.window.len();
out
}
pub fn decode(&mut self) -> Result<()> {
loop {
if matches!(self.state, State::Done) {
return Ok(());
}
let at = self.reader.checkpoint();
match self.step() {
Ok(()) => {
self.reader.compact();
}
Err(Halt::NeedInput) => {
self.reader.restore(at);
return Ok(());
}
Err(Halt::Fatal(error)) => return Err(error),
}
}
}
fn step(&mut self) -> Step<()> {
match &self.state {
State::Done => Ok(()),
State::BlockHeader => {
let last = self.reader.take(1)? == 1;
let kind = self.reader.take(2)?;
self.state = match kind {
0 => {
self.reader.align();
let length = self.reader.take(16)? as usize;
let complement = self.reader.take(16)? as usize;
if length ^ 0xFFFF != complement {
return Err(Halt::Fatal(Error::malformed(
"deflate",
"stored block length does not match its complement",
)));
}
self.check_limit(length)?;
State::Stored {
remaining: length,
last,
}
}
1 => State::Coded {
literals: Box::new(fixed_literal_table()?),
distances: Box::new(fixed_distance_table()?),
last,
},
2 => {
let (literals, distances) = read_dynamic_tables(&mut self.reader)?;
State::Coded {
literals: Box::new(literals),
distances: Box::new(distances),
last,
}
}
_ => {
return Err(Halt::Fatal(Error::malformed(
"deflate",
"reserved block type 3",
)));
}
};
Ok(())
}
State::Stored { remaining, last } => {
let (remaining, last) = (*remaining, *last);
if remaining == 0 {
self.state = if last {
State::Done
} else {
State::BlockHeader
};
return Ok(());
}
let taken = self.reader.take_bytes_upto(remaining, &mut self.window);
self.produced += taken;
if taken == 0 {
return Err(self.reader.short());
}
self.state = State::Stored {
remaining: remaining - taken,
last,
};
Ok(())
}
State::Coded { .. } => self.step_coded(),
}
}
fn step_coded(&mut self) -> Step<()> {
let Self {
reader,
state,
window,
produced,
limit,
..
} = self;
let State::Coded {
literals,
distances,
last,
} = state
else {
return Ok(());
};
let mut progressed = false;
loop {
let at = reader.checkpoint();
match decode_symbol(reader, literals, distances, window, produced, *limit) {
Ok(true) => progressed = true,
Ok(false) => break,
Err(Halt::NeedInput) => {
reader.restore(at);
return if progressed {
Ok(())
} else {
Err(Halt::NeedInput)
};
}
Err(fatal) => return Err(fatal),
}
}
*state = if *last {
State::Done
} else {
State::BlockHeader
};
Ok(())
}
fn check_limit(&self, adding: usize) -> Step<()> {
check_limit(self.produced, adding, self.limit)
}
}
fn check_limit(produced: usize, adding: usize, limit: usize) -> Step<()> {
if produced.saturating_add(adding) > limit {
return Err(Halt::Fatal(Error::malformed(
"deflate",
format!("stream expands beyond the {limit} byte limit implied by the image header"),
)));
}
Ok(())
}
fn decode_symbol(
reader: &mut BitReader,
literals: &Huffman,
distances: &Huffman,
window: &mut Vec<u8>,
produced: &mut usize,
limit: usize,
) -> Step<bool> {
let symbol = literals.decode(reader)?;
match symbol {
0..=255 => {
check_limit(*produced, 1, limit)?;
window.push(symbol as u8);
*produced += 1;
Ok(true)
}
256 => Ok(false),
257..=285 => {
let index = symbol as usize - 257;
let base = LENGTH_BASE
.get(index)
.copied()
.ok_or_else(|| Halt::Fatal(Error::malformed("deflate", "invalid length code")))?;
let extra = LENGTH_EXTRA.get(index).copied().unwrap_or(0);
let length = base as usize + reader.take(u32::from(extra))? as usize;
let distance_symbol = distances.decode(reader)? as usize;
let distance_base = DISTANCE_BASE
.get(distance_symbol)
.copied()
.ok_or_else(|| Halt::Fatal(Error::malformed("deflate", "invalid distance code")))?;
let distance_extra = DISTANCE_EXTRA.get(distance_symbol).copied().unwrap_or(0);
let distance =
distance_base as usize + reader.take(u32::from(distance_extra))? as usize;
if distance == 0 || distance > window.len() {
return Err(Halt::Fatal(Error::malformed(
"deflate",
format!(
"back-reference of distance {distance} points before the start of \
the {produced} bytes decoded so far"
),
)));
}
check_limit(*produced, length, limit)?;
let start = window.len() - distance;
let mut remaining = length;
while remaining > 0 {
let piece = remaining.min(window.len() - start);
window.extend_from_within(start..start + piece);
remaining -= piece;
}
*produced += length;
Ok(true)
}
_ => Err(Halt::Fatal(Error::malformed(
"deflate",
format!("literal/length symbol {symbol} is out of range"),
))),
}
}
fn read_dynamic_tables(reader: &mut BitReader) -> Step<(Huffman, Huffman)> {
let literal_count = reader.take(5)? as usize + 257;
let distance_count = reader.take(5)? as usize + 1;
let code_length_count = reader.take(4)? as usize + 4;
if literal_count > 288 || distance_count > 30 {
return Err(Halt::Fatal(Error::malformed(
"deflate",
"dynamic block declares too many codes",
)));
}
let mut code_lengths = [0_u8; 19];
for index in 0..code_length_count {
let bits = reader.take(3)? as u8;
let Some(&position) = CODE_LENGTH_ORDER.get(index) else {
break;
};
if let Some(slot) = code_lengths.get_mut(position) {
*slot = bits;
}
}
let code_length_table = Huffman::new(&code_lengths)?;
let total = literal_count + distance_count;
let mut lengths = vec![0_u8; total];
let mut index = 0;
while index < total {
let symbol = code_length_table.decode(reader)?;
match symbol {
0..=15 => {
if let Some(slot) = lengths.get_mut(index) {
*slot = symbol as u8;
}
index += 1;
}
16 => {
let previous = index
.checked_sub(1)
.and_then(|i| lengths.get(i).copied())
.ok_or_else(|| {
Halt::Fatal(Error::malformed(
"deflate",
"repeat code with no previous length",
))
})?;
let repeat = reader.take(2)? as usize + 3;
fill(&mut lengths, &mut index, previous, repeat, total)?;
}
17 => {
let repeat = reader.take(3)? as usize + 3;
fill(&mut lengths, &mut index, 0, repeat, total)?;
}
18 => {
let repeat = reader.take(7)? as usize + 11;
fill(&mut lengths, &mut index, 0, repeat, total)?;
}
_ => {
return Err(Halt::Fatal(Error::malformed(
"deflate",
"invalid code length symbol",
)));
}
}
}
let (literal_lengths, distance_lengths) = lengths.split_at(literal_count);
let literals = Huffman::new(literal_lengths)?;
let distances = Huffman::new(distance_lengths)?;
Ok((literals, distances))
}
fn fill(lengths: &mut [u8], index: &mut usize, value: u8, repeat: usize, total: usize) -> Step<()> {
if *index + repeat > total {
return Err(Halt::Fatal(Error::malformed(
"deflate",
"code length repeat runs past the end of the table",
)));
}
for _ in 0..repeat {
if let Some(slot) = lengths.get_mut(*index) {
*slot = value;
}
*index += 1;
}
Ok(())
}
pub fn inflate_to(data: &[u8], limit: usize) -> Result<Vec<u8>> {
let mut inflater = Inflater::new(limit);
inflater.feed(data);
inflater.end_of_input();
inflater.decode()?;
if !inflater.is_finished() {
return Err(truncated());
}
Ok(inflater.take_output())
}
#[derive(Debug)]
pub struct ZlibStream {
header: Vec<u8>,
inflater: Inflater,
adler: Adler32,
trailer: Vec<u8>,
ended: bool,
}
impl ZlibStream {
#[must_use]
pub fn new(limit: usize) -> Self {
Self {
header: Vec::with_capacity(2),
inflater: Inflater::new(limit),
adler: Adler32::new(),
trailer: Vec::with_capacity(4),
ended: false,
}
}
pub fn push(&mut self, mut data: &[u8]) -> Result<Vec<u8>> {
while self.header.len() < 2 {
let Some((&byte, rest)) = data.split_first() else {
return Ok(Vec::new());
};
self.header.push(byte);
data = rest;
if self.header.len() == 2 {
validate_zlib_header(&self.header)?;
}
}
if self.inflater.is_finished() {
self.collect_trailer(data);
return Ok(Vec::new());
}
self.inflater.feed(data);
self.inflater.decode()?;
let out = self.inflater.take_output();
self.adler.update(&out);
if self.inflater.is_finished() {
let leftover = self.inflater.drain_unconsumed_input();
self.collect_trailer(&leftover);
}
Ok(out)
}
fn collect_trailer(&mut self, data: &[u8]) {
for &byte in data {
if self.trailer.len() < 4 {
self.trailer.push(byte);
}
}
}
pub fn finish(&mut self) -> Result<Vec<u8>> {
if self.ended {
return Ok(Vec::new());
}
self.ended = true;
if self.header.len() < 2 {
return Err(Error::malformed(
"zlib",
"stream is shorter than its 2-byte header",
));
}
self.inflater.end_of_input();
self.inflater.decode()?;
let out = self.inflater.take_output();
self.adler.update(&out);
if !self.inflater.is_finished() {
return Err(truncated());
}
let leftover = self.inflater.drain_unconsumed_input();
self.collect_trailer(&leftover);
if self.trailer.len() < 4 {
return Err(Error::malformed(
"zlib",
"stream is missing its Adler-32 trailer",
));
}
let expected = u32::from_be_bytes([
self.trailer.first().copied().unwrap_or(0),
self.trailer.get(1).copied().unwrap_or(0),
self.trailer.get(2).copied().unwrap_or(0),
self.trailer.get(3).copied().unwrap_or(0),
]);
let actual = self.adler.finish();
if actual != expected {
return Err(Error::malformed(
"zlib",
format!(
"Adler-32 mismatch: stream declares {expected:#010x}, data is {actual:#010x}"
),
));
}
Ok(out)
}
}
fn validate_zlib_header(header: &[u8]) -> Result<()> {
let (&cmf, &flg) = match (header.first(), header.get(1)) {
(Some(cmf), Some(flg)) => (cmf, flg),
_ => {
return Err(Error::malformed(
"zlib",
"stream is shorter than its 2-byte header",
));
}
};
if cmf & 0x0F != 8 {
return Err(Error::malformed(
"zlib",
format!("compression method {} is not deflate", cmf & 0x0F),
));
}
if (u16::from(cmf) << 8 | u16::from(flg)) % 31 != 0 {
return Err(Error::malformed("zlib", "header check bits are wrong"));
}
if flg & 0x20 != 0 {
return Err(Error::malformed(
"zlib",
"preset dictionaries are not supported",
));
}
Ok(())
}
pub fn zlib_decompress(data: &[u8], limit: usize) -> Result<Vec<u8>> {
let mut stream = ZlibStream::new(limit);
let mut out = stream.push(data)?;
out.extend_from_slice(&stream.finish()?);
Ok(out)
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::indexing_slicing,
clippy::panic,
reason = "tests operate on known-good values and assert shapes directly"
)]
mod tests {
use super::*;
fn stored_stream(payload: &[u8]) -> Vec<u8> {
let mut out = vec![0x01];
let length = payload.len() as u16;
out.extend_from_slice(&length.to_le_bytes());
out.extend_from_slice(&(!length).to_le_bytes());
out.extend_from_slice(payload);
out
}
#[test]
fn stored_blocks_round_trip() {
let payload = b"the quick brown fox";
let out = inflate_to(&stored_stream(payload), 1024).unwrap();
assert_eq!(out, payload);
}
#[test]
fn an_empty_stored_block_yields_nothing() {
assert_eq!(
inflate_to(&stored_stream(b""), 16).unwrap(),
Vec::<u8>::new()
);
}
#[test]
fn a_stored_block_with_a_bad_complement_is_rejected() {
let mut stream = stored_stream(b"abc");
stream[3] ^= 0xFF;
let err = inflate_to(&stream, 1024).unwrap_err();
assert!(err.to_string().contains("complement"), "{err}");
}
#[test]
fn fixed_huffman_decodes_a_known_stream() {
let stream = [0xCB, 0x48, 0xCD, 0xC9, 0xC9, 0x07, 0x00];
assert_eq!(inflate_to(&stream, 64).unwrap(), b"hello");
}
#[test]
fn zlib_wrapped_streams_verify_their_checksum() {
let stream = [
0x78, 0xDA, 0xCB, 0x48, 0xCD, 0xC9, 0xC9, 0x57, 0x28, 0xCF, 0x2F, 0xCA, 0x49, 0x01,
0x00, 0x1A, 0x0B, 0x04, 0x5D,
];
assert_eq!(zlib_decompress(&stream, 64).unwrap(), b"hello world");
}
#[test]
fn a_corrupted_adler_is_reported() {
let mut stream = vec![
0x78, 0xDA, 0xCB, 0x48, 0xCD, 0xC9, 0xC9, 0x57, 0x28, 0xCF, 0x2F, 0xCA, 0x49, 0x01,
0x00, 0x1A, 0x0B, 0x04, 0x5D,
];
let last = stream.len() - 1;
stream[last] ^= 0xFF;
let err = zlib_decompress(&stream, 64).unwrap_err();
assert!(err.to_string().contains("Adler-32"), "{err}");
}
#[test]
fn zlib_headers_are_validated() {
assert!(zlib_decompress(&[], 16).is_err(), "empty");
assert!(zlib_decompress(&[0x78], 16).is_err(), "one byte");
assert!(zlib_decompress(&[0x77, 0x00, 0x00], 16).is_err());
assert!(zlib_decompress(&[0x78, 0x00, 0x00], 16).is_err());
let err = zlib_decompress(&[0x78, 0x3F, 0x00], 16).unwrap_err();
assert!(err.to_string().contains("dictionar"), "{err}");
}
#[test]
fn reserved_block_type_three_is_rejected() {
let err = inflate_to(&[0x07], 16).unwrap_err();
assert!(err.to_string().contains("reserved"), "{err}");
}
#[test]
fn a_back_reference_before_the_start_is_rejected() {
let err = inflate_to(&[0x03, 0x02], 1024).unwrap_err();
assert_eq!(err.format(), "deflate", "{err}");
assert!(err.to_string().contains("back-reference"), "{err}");
}
#[test]
fn output_beyond_the_limit_is_malformed_not_an_allocation() {
let bomb = stored_stream(&vec![0_u8; 65535]);
let err = inflate_to(&bomb, 1024).unwrap_err();
assert_eq!(err.format(), "deflate", "{err}");
assert!(err.to_string().contains("limit"), "{err}");
assert_eq!(inflate_to(&bomb, 65535).unwrap().len(), 65535);
}
#[test]
fn every_truncation_of_a_valid_stream_is_an_error_not_a_panic() {
let full = [
0x78, 0xDA, 0xCB, 0x48, 0xCD, 0xC9, 0xC9, 0x57, 0x28, 0xCF, 0x2F, 0xCA, 0x49, 0x01,
0x00, 0x1A, 0x0B, 0x04, 0x5D,
];
for len in 0..full.len() {
let _ = zlib_decompress(&full[..len], 4096);
}
assert!(
zlib_decompress(&full, 4096).is_ok(),
"the untruncated stream still works"
);
}
#[test]
fn arbitrary_bytes_never_panic() {
let mut state = 0x1234_5678_u32;
for _ in 0..2000 {
let mut bytes = Vec::new();
for _ in 0..32 {
state = state.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
bytes.push((state >> 24) as u8);
}
let _ = inflate_to(&bytes, 4096);
let _ = zlib_decompress(&bytes, 4096);
}
}
#[test]
fn over_subscribed_huffman_tables_are_rejected() {
assert!(Huffman::new(&[1, 1, 1]).is_err());
assert!(Huffman::new(&[1, 1]).is_ok());
assert!(Huffman::new(&[16]).is_err());
}
#[test]
fn overlapping_back_references_encode_runs() {
let stream = [0x4B, 0x4C, 0x84, 0x00, 0x00];
assert_eq!(inflate_to(&stream, 64).unwrap(), b"aaaaaaaa");
}
fn compress(data: &[u8], level: u8) -> Vec<u8> {
crate::deflate::zlib_compress(data, crate::deflate::Level::new(level).unwrap()).unwrap()
}
fn deflate_body(data: &[u8], level: u8) -> Vec<u8> {
compress(data, level).split_off(2)
}
fn encode_canonical(lengths: &[u8], symbols: &[usize]) -> Vec<u8> {
let mut codes = vec![0_u32; lengths.len()];
let mut code = 0_u32;
for length in 1..=MAX_BITS as u8 {
for (symbol, &l) in lengths.iter().enumerate() {
if l == length {
codes[symbol] = code;
code += 1;
}
}
code <<= 1;
}
let (mut out, mut bits, mut count) = (Vec::new(), 0_u64, 0_u32);
for &symbol in symbols {
let length = u32::from(lengths[symbol]);
let reversed = codes[symbol].reverse_bits() >> (32 - length);
bits |= u64::from(reversed) << count;
count += length;
while count >= 8 {
out.push(bits as u8);
bits >>= 8;
count -= 8;
}
}
if count > 0 {
out.push(bits as u8);
}
out
}
#[test]
fn codes_of_every_length_decode_through_table_and_walk() {
let mut lengths: Vec<u8> = (1..=15).collect();
lengths.push(15);
let table = Huffman::new(&lengths).unwrap();
let symbols: Vec<usize> = (0..lengths.len()).chain((0..lengths.len()).rev()).collect();
let mut reader = BitReader::default();
reader.feed(&encode_canonical(&lengths, &symbols));
reader.end();
for &expected in &symbols {
assert_eq!(usize::from(table.decode(&mut reader).unwrap()), expected);
}
}
#[test]
fn overlapping_references_of_every_short_distance_decode() {
let mut original = Vec::new();
for period in (1..=9).chain([31, 258]) {
let pattern: Vec<u8> = (0..period).map(|i| (i * 37 + period) as u8).collect();
for _ in 0..600 / period + 3 {
original.extend_from_slice(&pattern);
}
}
for level in [1_u8, 6, 9] {
assert_eq!(
zlib_decompress(&compress(&original, level), 1 << 20).unwrap(),
original
);
}
}
#[test]
fn a_stream_past_the_compaction_threshold_decodes_whole_and_in_pieces() {
let mut state = 0x2545_f491_u32;
let original: Vec<u8> = (0..400_000)
.map(|i| {
state ^= state << 13;
state ^= state >> 17;
state ^= state << 5;
if i % 5 == 0 { b'a' } else { state as u8 }
})
.collect();
let stream = compress(&original, 6);
assert!(stream.len() > 2 * COMPACT_AT);
assert_eq!(zlib_decompress(&stream, 1 << 20).unwrap(), original);
let mut zlib = ZlibStream::new(1 << 20);
let mut out = Vec::new();
for piece in stream.chunks(997) {
out.extend_from_slice(&zlib.push(piece).unwrap());
}
out.extend_from_slice(&zlib.finish().unwrap());
assert_eq!(out, original);
}
#[test]
fn feeding_one_byte_at_a_time_decodes_identically() {
for level in [0_u8, 1, 6, 9] {
let original = b"the quick brown fox jumps over the lazy dog. ".repeat(120);
let stream = compress(&original, level);
let mut zlib = ZlibStream::new(1 << 20);
let mut out = Vec::new();
for byte in &stream {
out.extend_from_slice(&zlib.push(std::slice::from_ref(byte)).unwrap());
}
out.extend_from_slice(&zlib.finish().unwrap());
assert_eq!(
out, original,
"level {level} differed when fed byte by byte"
);
}
}
#[test]
fn every_chunk_size_decodes_identically() {
let original: Vec<u8> = (0..40_000).map(|i| ((i * 7) % 251) as u8).collect();
let stream = compress(&original, 6);
for chunk in [1, 2, 3, 7, 64, 1024, 65_536] {
let mut zlib = ZlibStream::new(1 << 20);
let mut out = Vec::new();
for piece in stream.chunks(chunk) {
out.extend_from_slice(&zlib.push(piece).unwrap());
}
out.extend_from_slice(&zlib.finish().unwrap());
assert_eq!(out, original, "chunk size {chunk} differed");
}
}
#[test]
fn a_drained_inflater_retains_only_its_window() {
let original = vec![0_u8; 8 * 1024 * 1024];
let stream = deflate_body(&original, 9);
let mut inflater = Inflater::new(16 * 1024 * 1024);
let mut total = 0_usize;
for piece in stream.chunks(4096) {
inflater.feed(piece);
inflater.decode().unwrap();
total += inflater.take_output().len();
assert!(
inflater.retained() <= WINDOW + 4096,
"retained {} bytes after {total} of output",
inflater.retained()
);
}
inflater.end_of_input();
inflater.decode().unwrap();
total += inflater.take_output().len();
assert_eq!(total, original.len());
assert_eq!(inflater.produced(), original.len());
}
#[test]
fn a_back_reference_reaching_across_a_drain_still_resolves() {
let original = b"abcdefgh".repeat(200_000);
let stream = deflate_body(&original, 9);
let mut inflater = Inflater::new(4 * 1024 * 1024);
let mut out = Vec::new();
for piece in stream.chunks(777) {
inflater.feed(piece);
inflater.decode().unwrap();
out.extend_from_slice(&inflater.take_output());
}
inflater.end_of_input();
inflater.decode().unwrap();
out.extend_from_slice(&inflater.take_output());
assert_eq!(out, original);
}
#[test]
fn an_unfinished_stream_is_not_reported_as_complete() {
let stream = deflate_body(&vec![7_u8; 100_000], 6);
let mut inflater = Inflater::new(1 << 20);
inflater.feed(&stream[..stream.len() / 2]);
inflater.decode().unwrap();
assert!(!inflater.is_finished());
inflater.end_of_input();
assert!(inflater.decode().is_err() || !inflater.is_finished());
}
#[test]
fn the_limit_is_enforced_incrementally_not_at_the_end() {
let stream = deflate_body(&vec![0_u8; 4 * 1024 * 1024], 9);
let mut inflater = Inflater::new(1024);
inflater.feed(&stream);
let error = inflater.decode().unwrap_err();
assert_eq!(error.format(), "deflate", "{error}");
assert!(
inflater.produced() <= 1024,
"produced {} bytes",
inflater.produced()
);
}
}