use std::io::{self, Read, Write};
use crate::frame::types::LZ4F_VERSION;
use crate::frame::{
lz4f_create_decompression_context, lz4f_decompress, lz4f_decompress_using_dict, Lz4FDCtx,
};
use crate::io::decompress_resources::DecompressResources;
use crate::io::prefs::{display_level, Prefs, DISPLAY_LEVEL, LZ4IO_MAGICNUMBER};
const DECOMP_BUF_SIZE: usize = 64 * 1024;
fn lz4f_err_to_io(e: crate::frame::Lz4FError) -> io::Error {
io::Error::new(io::ErrorKind::InvalidData, format!("LZ4F error: {e}"))
}
pub fn decompress_lz4f(
src: &mut impl Read,
dst: &mut impl Write,
prefs: &Prefs,
resources: &mut DecompressResources,
) -> io::Result<u64> {
if let Some(dict) = &resources.dict_buffer {
let dict = dict.clone(); return decompress_lz4f_st_dict(src, dst, prefs, &dict);
}
if prefs.nb_workers > 1 {
decompress_lz4f_st(src, dst, prefs)
} else {
decompress_lz4f_st(src, dst, prefs)
}
}
fn feed_to_decompressor(
dctx: &mut Lz4FDCtx,
input: &[u8],
dst_buf: &mut [u8],
dst: &mut impl Write,
prefs: &Prefs,
filesize: &mut u64,
) -> io::Result<usize> {
let mut pos = 0usize;
let mut next_hint: usize = 1;
while pos < input.len() {
let (src_consumed, dst_written, hint) =
lz4f_decompress(dctx, Some(dst_buf), &input[pos..], None).map_err(lz4f_err_to_io)?;
pos += src_consumed;
next_hint = hint;
if dst_written > 0 {
*filesize += dst_written as u64;
if !prefs.test_mode {
dst.write_all(&dst_buf[..dst_written])
.map_err(|e| io::Error::new(e.kind(), format!("Write error: {e}")))?;
}
if DISPLAY_LEVEL.load(std::sync::atomic::Ordering::Relaxed) >= 2 {
display_level(2, &format!("\rDecompressed : {} MiB ", *filesize >> 20));
}
}
if next_hint == 0 {
break;
}
if src_consumed == 0 && dst_written == 0 {
break;
}
}
Ok(next_hint)
}
fn feed_to_decompressor_dict(
dctx: &mut Lz4FDCtx,
input: &[u8],
dict: &[u8],
dst_buf: &mut [u8],
dst: &mut impl Write,
prefs: &Prefs,
filesize: &mut u64,
) -> io::Result<usize> {
let mut pos = 0usize;
let mut next_hint: usize = 1;
while pos < input.len() {
let (src_consumed, dst_written, hint) =
lz4f_decompress_using_dict(dctx, Some(dst_buf), &input[pos..], dict, None)
.map_err(lz4f_err_to_io)?;
pos += src_consumed;
next_hint = hint;
if dst_written > 0 {
*filesize += dst_written as u64;
if !prefs.test_mode {
dst.write_all(&dst_buf[..dst_written])
.map_err(|e| io::Error::new(e.kind(), format!("Write error: {e}")))?;
}
if DISPLAY_LEVEL.load(std::sync::atomic::Ordering::Relaxed) >= 2 {
display_level(2, &format!("\rDecompressed : {} MiB ", *filesize >> 20));
}
}
if next_hint == 0 {
break;
}
if src_consumed == 0 && dst_written == 0 {
break;
}
}
Ok(next_hint)
}
fn decompress_lz4f_st(src: &mut impl Read, dst: &mut impl Write, prefs: &Prefs) -> io::Result<u64> {
let mut dctx = lz4f_create_decompression_context(LZ4F_VERSION).map_err(lz4f_err_to_io)?;
let mut src_buf = vec![0u8; DECOMP_BUF_SIZE];
let mut dst_buf = vec![0u8; DECOMP_BUF_SIZE];
let mut filesize: u64 = 0;
let magic_bytes = LZ4IO_MAGICNUMBER.to_le_bytes();
let mut next_hint = feed_to_decompressor(
&mut dctx,
&magic_bytes,
&mut dst_buf,
dst,
prefs,
&mut filesize,
)?;
while next_hint != 0 {
let to_read = next_hint.min(src_buf.len());
let read_n = src
.read(&mut src_buf[..to_read])
.map_err(|e| io::Error::new(e.kind(), format!("Read error: {e}")))?;
if read_n == 0 {
break; }
next_hint = feed_to_decompressor(
&mut dctx,
&src_buf[..read_n],
&mut dst_buf,
dst,
prefs,
&mut filesize,
)?;
}
if next_hint != 0 {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"Truncated LZ4 frame",
));
}
Ok(filesize)
}
fn decompress_lz4f_st_dict(
src: &mut impl Read,
dst: &mut impl Write,
prefs: &Prefs,
dict: &[u8],
) -> io::Result<u64> {
let mut dctx = lz4f_create_decompression_context(LZ4F_VERSION).map_err(lz4f_err_to_io)?;
let mut src_buf = vec![0u8; DECOMP_BUF_SIZE];
let mut dst_buf = vec![0u8; DECOMP_BUF_SIZE];
let mut filesize: u64 = 0;
let magic_bytes = LZ4IO_MAGICNUMBER.to_le_bytes();
let mut next_hint = feed_to_decompressor_dict(
&mut dctx,
&magic_bytes,
dict,
&mut dst_buf,
dst,
prefs,
&mut filesize,
)?;
while next_hint != 0 {
let to_read = next_hint.min(src_buf.len());
let read_n = src
.read(&mut src_buf[..to_read])
.map_err(|e| io::Error::new(e.kind(), format!("Read error: {e}")))?;
if read_n == 0 {
break; }
next_hint = feed_to_decompressor_dict(
&mut dctx,
&src_buf[..read_n],
dict,
&mut dst_buf,
dst,
prefs,
&mut filesize,
)?;
}
if next_hint != 0 {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"Truncated LZ4 frame (dictionary decompression)",
));
}
Ok(filesize)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::io::decompress_resources::DecompressResources;
use crate::io::prefs::Prefs;
use std::io::Write;
fn compress_frame(data: &[u8]) -> Vec<u8> {
use crate::frame::{lz4f_compress_frame, lz4f_compress_frame_bound};
let bound = lz4f_compress_frame_bound(data.len(), None);
let mut dst = vec![0u8; bound];
let n = lz4f_compress_frame(&mut dst, data, None).unwrap();
dst.truncate(n);
dst
}
#[test]
fn round_trip_st_no_dict() {
let original: Vec<u8> = (0u8..=255).cycle().take(4096).collect();
let compressed = compress_frame(&original);
let mut compressed_body = &compressed[4..];
let prefs = Prefs::default();
let mut res = DecompressResources::new(&prefs).unwrap();
let mut output = Vec::new();
let n = decompress_lz4f(&mut compressed_body, &mut output, &prefs, &mut res).unwrap();
assert_eq!(n as usize, original.len(), "byte count mismatch");
assert_eq!(output, original, "decompressed content mismatch");
}
#[test]
fn test_mode_discards_output() {
let original: Vec<u8> = b"hello, test mode!".to_vec();
let compressed = compress_frame(&original);
let mut compressed_body = &compressed[4..];
let mut prefs = Prefs::default();
prefs.test_mode = true;
let mut res = DecompressResources::new(&prefs).unwrap();
let mut output = Vec::new();
let n = decompress_lz4f(&mut compressed_body, &mut output, &prefs, &mut res).unwrap();
assert_eq!(
n as usize,
original.len(),
"byte count should match even in test mode"
);
assert!(output.is_empty(), "test_mode must not write anything");
}
#[test]
fn empty_frame_returns_zero() {
let compressed = compress_frame(&[]);
let mut compressed_body = &compressed[4..];
let prefs = Prefs::default();
let mut res = DecompressResources::new(&prefs).unwrap();
let mut output = Vec::new();
let n = decompress_lz4f(&mut compressed_body, &mut output, &prefs, &mut res).unwrap();
assert_eq!(n, 0);
assert!(output.is_empty());
}
#[test]
fn dict_path_round_trip_no_dict_buffer() {
let original: Vec<u8> = b"hello dict path".to_vec();
let compressed = compress_frame(&original);
let mut compressed_body = &compressed[4..];
let prefs = Prefs::default();
let mut res = DecompressResources::new(&prefs).unwrap();
res.dict_buffer = Some(Vec::new());
let mut output = Vec::new();
let n = decompress_lz4f(&mut compressed_body, &mut output, &prefs, &mut res).unwrap();
assert_eq!(
n as usize,
original.len(),
"byte count mismatch (dict path)"
);
assert_eq!(output, original, "output mismatch (dict path)");
}
#[test]
fn corrupt_input_returns_error() {
let garbage: &[u8] = b"\x00\x01\x02\x03\xFF\xFE\xFD";
let mut src = &garbage[..];
let prefs = Prefs::default();
let mut res = DecompressResources::new(&prefs).unwrap();
let mut output = Vec::new();
let result = decompress_lz4f(&mut src, &mut output, &prefs, &mut res);
assert!(result.is_err(), "corrupt input must return Err");
}
#[test]
fn large_frame_round_trip() {
let original: Vec<u8> = (0u8..=255)
.cycle()
.enumerate()
.map(|(i, b)| b.wrapping_add((i >> 8) as u8))
.take(256 * 1024)
.collect();
let compressed = compress_frame(&original);
let mut compressed_body = &compressed[4..];
let prefs = Prefs::default();
let mut res = DecompressResources::new(&prefs).unwrap();
let mut output = Vec::new();
let n = decompress_lz4f(&mut compressed_body, &mut output, &prefs, &mut res).unwrap();
assert_eq!(n as usize, original.len());
assert_eq!(output, original);
}
}