#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum Wrapper {
Zlib,
Raw,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum InflateError {
OutputExhausted,
Incomplete,
Malformed,
}
impl core::fmt::Display for InflateError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
let s = match self {
InflateError::OutputExhausted => "decoded length exceeds the bound",
InflateError::Incomplete => "truncated deflate stream",
InflateError::Malformed => "malformed deflate stream",
};
f.write_str(s)
}
}
const DIRECT_PREALLOC_MAX: usize = 256 * 1024 * 1024;
fn initial_window(encoded: &[u8], cap: usize) -> usize {
encoded.len().saturating_mul(2).max(1).min(cap.max(1))
}
pub(crate) fn inflate_bounded(
encoded: &[u8],
cap: usize,
wrapper: Wrapper,
) -> Result<Vec<u8>, InflateError> {
if cap <= DIRECT_PREALLOC_MAX {
inflate_direct(encoded, cap, wrapper)
} else {
inflate_growing(encoded, cap, wrapper)
}
}
pub(crate) fn inflate_len(
encoded: &[u8],
cap: usize,
wrapper: Wrapper,
) -> Result<u64, InflateError> {
let mut inflate = zlib_rs::Inflate::new(wrapper == Wrapper::Zlib, 15);
let mut buf: Vec<u8> = Vec::new();
let mut chunk = initial_window(encoded, cap);
let mut consumed: u64 = 0;
let mut produced: u64 = 0;
loop {
if buf.len() < chunk {
buf.resize(chunk, 0u8);
}
let start = usize::try_from(consumed).map_err(|_| InflateError::Malformed)?;
let remaining = encoded.get(start..).ok_or(InflateError::Malformed)?;
let status =
inflate.decompress(remaining, &mut buf[..chunk], zlib_rs::InflateFlush::Finish);
let total_out = inflate.total_out();
let total_in = inflate.total_in();
match status {
Ok(zlib_rs::Status::StreamEnd) => {
if total_out > cap as u64 {
return Err(InflateError::OutputExhausted);
}
return Ok(total_out);
}
Ok(zlib_rs::Status::BufError) | Ok(zlib_rs::Status::Ok) => {
if total_out >= cap as u64 {
return Err(InflateError::OutputExhausted);
}
if total_out == produced && total_in == consumed {
return Err(InflateError::Incomplete);
}
consumed = total_in;
produced = total_out;
chunk = chunk
.saturating_mul(2)
.min(cap.max(1))
.min(u32::MAX as usize);
}
Err(_) => return Err(InflateError::Malformed),
}
}
}
fn inflate_direct(encoded: &[u8], cap: usize, wrapper: Wrapper) -> Result<Vec<u8>, InflateError> {
let mut out = vec![0u8; cap.max(1)];
let mut inflate = zlib_rs::Inflate::new(wrapper == Wrapper::Zlib, 15);
match inflate.decompress(encoded, &mut out, zlib_rs::InflateFlush::Finish) {
Ok(zlib_rs::Status::StreamEnd) => {
let n = usize::try_from(inflate.total_out()).map_err(|_| InflateError::Malformed)?;
if n > cap {
return Err(InflateError::OutputExhausted);
}
out.truncate(n);
Ok(out)
}
Ok(zlib_rs::Status::BufError) => Err(InflateError::OutputExhausted),
Ok(zlib_rs::Status::Ok) => Err(InflateError::Incomplete),
Err(_) => Err(InflateError::Malformed),
}
}
fn inflate_growing(encoded: &[u8], cap: usize, wrapper: Wrapper) -> Result<Vec<u8>, InflateError> {
let mut inflate = zlib_rs::Inflate::new(wrapper == Wrapper::Zlib, 15);
let mut out: Vec<u8> = Vec::new();
let mut consumed: u64 = 0;
let mut produced: u64 = 0;
let mut chunk = initial_window(encoded, cap);
loop {
let mut buf = vec![0u8; chunk];
let start = usize::try_from(consumed).map_err(|_| InflateError::Malformed)?;
let remaining = encoded.get(start..).ok_or(InflateError::Malformed)?;
let status = inflate.decompress(remaining, &mut buf, zlib_rs::InflateFlush::Finish);
let total_out = inflate.total_out();
let total_in = inflate.total_in();
let written = usize::try_from(total_out - produced).map_err(|_| InflateError::Malformed)?;
if written > buf.len() {
return Err(InflateError::Malformed);
}
match status {
Ok(zlib_rs::Status::StreamEnd) => {
out.extend_from_slice(&buf[..written]);
if out.len() > cap {
return Err(InflateError::OutputExhausted);
}
return Ok(out);
}
Ok(zlib_rs::Status::BufError) | Ok(zlib_rs::Status::Ok) => {
out.extend_from_slice(&buf[..written]);
if out.len() >= cap {
return Err(InflateError::OutputExhausted);
}
if written == 0 && total_in == consumed {
return Err(InflateError::Incomplete);
}
consumed = total_in;
produced = total_out;
chunk = chunk
.saturating_mul(2)
.min(cap.max(1))
.min(u32::MAX as usize);
}
Err(_) => return Err(InflateError::Malformed),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Write;
use flate2::Compression;
use flate2::write::{DeflateEncoder, ZlibEncoder};
fn zlib_encode(data: &[u8]) -> Vec<u8> {
let mut e = ZlibEncoder::new(Vec::new(), Compression::default());
e.write_all(data).expect("zlib encode");
e.finish().expect("zlib finish")
}
fn raw_deflate_encode(data: &[u8]) -> Vec<u8> {
let mut e = DeflateEncoder::new(Vec::new(), Compression::default());
e.write_all(data).expect("deflate encode");
e.finish().expect("deflate finish")
}
fn battery() -> Vec<Vec<u8>> {
vec![
Vec::new(),
b"hello, world".to_vec(),
vec![0u8; 200_000],
b"abcdefghijklmnop".repeat(20_000),
(0..=255u8).cycle().take(100_000).collect(),
]
}
#[test]
fn direct_path_roundtrips_both_wrappers() {
for (i, data) in battery().into_iter().enumerate() {
let z = zlib_encode(&data);
assert_eq!(
inflate_bounded(&z, data.len(), Wrapper::Zlib).unwrap(),
data,
"case {i}: zlib"
);
let r = raw_deflate_encode(&data);
assert_eq!(
inflate_bounded(&r, data.len(), Wrapper::Raw).unwrap(),
data,
"case {i}: raw"
);
}
}
#[test]
fn growth_path_decodes_a_small_stream() {
let data = b"growth path witness".to_vec();
let z = zlib_encode(&data);
let got = inflate_bounded(&z, DIRECT_PREALLOC_MAX + 1, Wrapper::Zlib).unwrap();
assert_eq!(got, data);
}
#[test]
fn len_path_agrees_with_bounded_and_never_overallocates() {
for (i, data) in battery().into_iter().enumerate() {
let z = zlib_encode(&data);
let n = inflate_len(&z, DIRECT_PREALLOC_MAX + 1, Wrapper::Zlib).unwrap();
assert_eq!(n as usize, data.len(), "case {i}: zlib len");
let r = raw_deflate_encode(&data);
let n = inflate_len(&r, 32 * 1024 * 1024 + 1, Wrapper::Raw).unwrap();
assert_eq!(n as usize, data.len(), "case {i}: raw len");
}
}
#[test]
fn wrong_wrapper_is_a_typed_error() {
let z = zlib_encode(b"wrapped");
assert!(inflate_bounded(&z, 64, Wrapper::Raw).is_err());
}
#[test]
fn truncated_and_corrupt_streams_are_typed_errors() {
let data = b"abcdefghij".repeat(100);
let z = zlib_encode(&data);
assert!(inflate_bounded(&z[..z.len() / 2], data.len(), Wrapper::Zlib).is_err());
assert!(inflate_len(&z[..z.len() / 2], 1 << 20, Wrapper::Zlib).is_err());
let mut corrupt = z.clone();
corrupt[0] ^= 0xFF; assert!(inflate_bounded(&corrupt, data.len(), Wrapper::Zlib).is_err());
assert!(inflate_len(&corrupt, 1 << 20, Wrapper::Zlib).is_err());
assert_eq!(
inflate_bounded(&z, data.len() - 1, Wrapper::Zlib),
Err(InflateError::OutputExhausted)
);
assert_eq!(
inflate_len(&z, data.len() - 1, Wrapper::Zlib),
Err(InflateError::OutputExhausted)
);
}
}