use std::io::Read;
use crate::chunk;
use crate::PARALLEL_THRESHOLD;
#[derive(Debug)]
pub struct GzError(pub &'static str);
pub const OUTPUT_TOO_LARGE: &str = "decompressed output exceeds max_out cap";
#[inline]
pub fn output_too_large() -> GzError {
GzError(OUTPUT_TOO_LARGE)
}
impl std::fmt::Display for GzError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "lgz: {}", self.0)
}
}
impl std::error::Error for GzError {}
struct GzipHeader {
deflate_start: usize,
}
pub fn deflate_offset(data: &[u8]) -> Result<usize, GzError> {
Ok(parse_gzip_header(data)?.deflate_start)
}
fn parse_gzip_header(data: &[u8]) -> Result<GzipHeader, GzError> {
if data.len() < 10 {
return Err(GzError("input too short for gzip header"));
}
if data[0] != 0x1f || data[1] != 0x8b {
return Err(GzError("not a gzip file (bad magic)"));
}
if data[2] != 0x08 {
return Err(GzError("unsupported compression method (not DEFLATE)"));
}
let flags = data[3];
let mut pos = 10;
if flags & 0x04 != 0 {
if pos + 2 > data.len() { return Err(GzError("truncated FEXTRA")); }
let xlen = u16::from_le_bytes([data[pos], data[pos + 1]]) as usize;
pos += 2 + xlen;
}
if flags & 0x08 != 0 {
while pos < data.len() && data[pos] != 0 { pos += 1; }
pos += 1; }
if flags & 0x10 != 0 {
while pos < data.len() && data[pos] != 0 { pos += 1; }
pos += 1;
}
if flags & 0x02 != 0 {
pos += 2;
}
if pos >= data.len() {
return Err(GzError("gzip header extends past input"));
}
Ok(GzipHeader { deflate_start: pos })
}
pub fn decompress_gz(data: &[u8]) -> Result<Vec<u8>, GzError> {
let header = parse_gzip_header(data)?;
let deflate_data = &data[header.deflate_start..];
let deflate_len = if deflate_data.len() > 8 {
deflate_data.len() - 8
} else {
deflate_data.len()
};
let raw_deflate = &deflate_data[..deflate_len];
if raw_deflate.len() >= PARALLEL_THRESHOLD {
let n_workers = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(4);
match chunk::decode_chunk(raw_deflate, n_workers, true) {
Ok((output, _consumed)) => {
if deflate_data.len() >= 8 {
let trailer = &deflate_data[deflate_data.len() - 8..];
let expected_crc =
u32::from_le_bytes([trailer[0], trailer[1], trailer[2], trailer[3]]);
let expected_isize =
u32::from_le_bytes([trailer[4], trailer[5], trailer[6], trailer[7]]);
let mut crc = flate2::Crc::new();
crc.update(&output);
if crc.sum() != expected_crc || (output.len() as u32) != expected_isize {
return Err(GzError("gzip checksum mismatch"));
}
}
return Ok(output);
}
Err(_) => {
}
}
}
decompress_gz_flate2(data)
}
pub fn decompress_gz_capped(data: &[u8], max_out: usize) -> Result<Vec<u8>, GzError> {
let _ = parse_gzip_header(data)?;
let mut decoder = flate2::read::GzDecoder::new(data);
read_capped(&mut decoder, max_out)
}
pub(crate) fn read_capped<R: Read>(reader: &mut R, max_out: usize) -> Result<Vec<u8>, GzError> {
const CHUNK: usize = 64 * 1024;
let mut out: Vec<u8> = Vec::new();
let mut buf = [0u8; CHUNK];
loop {
let n = reader
.read(&mut buf)
.map_err(|_| GzError("gzip decompression failed"))?;
if n == 0 {
break;
}
if out.len() + n > max_out {
return Err(output_too_large());
}
out.extend_from_slice(&buf[..n]);
}
Ok(out)
}
pub fn decompress_gz_into(data: &[u8], out: &mut [u8]) -> Result<usize, GzError> {
let decompressed = decompress_gz(data)?;
if decompressed.len() > out.len() {
return Err(GzError("output buffer too small"));
}
out[..decompressed.len()].copy_from_slice(&decompressed);
Ok(decompressed.len())
}
pub fn decompress_gz_stream<R: Read>(mut reader: R) -> Result<Vec<u8>, GzError> {
let mut data = Vec::new();
reader.read_to_end(&mut data)
.map_err(|_| GzError("failed to read input"))?;
decompress_gz(&data)
}
fn decompress_gz_flate2(data: &[u8]) -> Result<Vec<u8>, GzError> {
let mut decoder = flate2::read::GzDecoder::new(data);
let mut output = Vec::new();
decoder.read_to_end(&mut output)
.map_err(|_| GzError("gzip decompression failed"))?;
Ok(output)
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Write;
fn make_gz(data: &[u8]) -> Vec<u8> {
use flate2::write::GzEncoder;
use flate2::Compression;
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
encoder.write_all(data).unwrap();
encoder.finish().unwrap()
}
fn make_gz_with_flushes(chunks: &[&[u8]]) -> Vec<u8> {
use flate2::write::GzEncoder;
use flate2::Compression;
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
for (i, chunk) in chunks.iter().enumerate() {
encoder.write_all(chunk).unwrap();
if i < chunks.len() - 1 {
encoder.flush().unwrap();
}
}
encoder.finish().unwrap()
}
#[test]
fn roundtrip_basic() {
let original = b"Hello, lgz! This is a test of parallel gzip decompression.";
let compressed = make_gz(original);
let decompressed = decompress_gz(&compressed).unwrap();
assert_eq!(&decompressed, original);
}
#[test]
fn roundtrip_large() {
let original: Vec<u8> = (0..100_000).map(|i| (i % 256) as u8).collect();
let compressed = make_gz(&original);
let decompressed = decompress_gz(&compressed).unwrap();
assert_eq!(decompressed, original);
}
#[test]
fn decompress_into_buffer() {
let original = b"fixed size test";
let compressed = make_gz(original);
let mut buf = vec![0u8; 1024];
let n = decompress_gz_into(&compressed, &mut buf).unwrap();
assert_eq!(&buf[..n], original.as_slice());
}
#[test]
fn invalid_data_returns_error() {
let bad = b"not gzip at all";
assert!(decompress_gz(bad).is_err());
}
#[test]
fn stream_decompress() {
let original = b"streaming test data for lgz";
let compressed = make_gz(original);
let decompressed = decompress_gz_stream(compressed.as_slice()).unwrap();
assert_eq!(&decompressed, original);
}
#[test]
fn parse_header_basic() {
let compressed = make_gz(b"test");
let header = parse_gzip_header(&compressed).unwrap();
assert!(header.deflate_start >= 10);
assert!(header.deflate_start < compressed.len());
}
#[test]
fn capped_small_input_roundtrips() {
let original = b"a normal, small payload that fits well under the cap";
let compressed = make_gz(original);
let out = decompress_gz_capped(&compressed, 1024).unwrap();
assert_eq!(&out, original);
}
#[test]
fn capped_bomb_errors_not_ooms() {
let bomb = make_gz(&vec![0u8; 256 * 1024]);
assert!(bomb.len() < 1024, "zeros must compress tiny: {} bytes", bomb.len());
let err = decompress_gz_capped(&bomb, 4096).unwrap_err();
assert_eq!(err.0, OUTPUT_TOO_LARGE, "expected the output-too-large sentinel");
}
#[test]
fn capped_exact_boundary_ok() {
let original = vec![7u8; 4096];
let compressed = make_gz(&original);
let out = decompress_gz_capped(&compressed, 4096).unwrap();
assert_eq!(out, original);
}
#[test]
fn capped_invalid_data_errors() {
assert!(decompress_gz_capped(b"not gzip at all", 1024).is_err());
}
#[test]
fn roundtrip_with_flushes() {
let chunk_data: Vec<u8> = (0..200_000).map(|i| (i % 251) as u8).collect();
let chunks: Vec<&[u8]> = chunk_data.chunks(50_000).collect();
let compressed = make_gz_with_flushes(&chunks);
let decompressed = decompress_gz(&compressed).unwrap();
assert_eq!(decompressed, chunk_data);
}
}