use std::io;
use crate::frame::types::LZ4F_VERSION;
use crate::frame::{
lz4f_create_decompression_context, lz4f_decompress, DecompressOptions, Lz4FDCtx,
};
pub type DecFunctionF = fn(
decompressor: &mut FrameDecompressor,
src: &[u8],
dst: &mut Vec<u8>,
dst_capacity: usize,
skip_checksums: bool,
) -> io::Result<usize>;
#[derive(Debug, Default)]
pub struct FrameDecompressor;
impl FrameDecompressor {
pub fn new() -> Self {
FrameDecompressor
}
}
pub fn decompress_frame_block(
_decompressor: &mut FrameDecompressor,
src: &[u8],
dst: &mut Vec<u8>,
dst_capacity: usize,
skip_checksums: bool,
) -> io::Result<usize> {
let mut dctx: Box<Lz4FDCtx> = lz4f_create_decompression_context(LZ4F_VERSION)
.map_err(|e| io::Error::other(e.to_string()))?;
let opts = DecompressOptions {
stable_dst: true,
skip_checksums,
};
let mut tmp = vec![0u8; 64 * 1024];
let before = dst.len();
let mut src_pos: usize = 0;
let mut total_written: usize = 0;
loop {
let (src_consumed, dst_written, next_src_hint) =
lz4f_decompress(&mut dctx, Some(&mut tmp), &src[src_pos..], Some(&opts))
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e.to_string()))?;
if dst_written > 0 {
total_written += dst_written;
if total_written > dst_capacity {
dst.truncate(before);
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"decompressed output exceeds dst_capacity",
));
}
dst.extend_from_slice(&tmp[..dst_written]);
}
src_pos += src_consumed;
if next_src_hint == 0 {
break;
}
if src_consumed == 0 && dst_written == 0 {
dst.truncate(before);
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"decompressor stalled — no progress on source or destination",
));
}
}
if src_pos != src.len() {
dst.truncate(before);
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"decompressor did not consume all input bytes",
));
}
Ok(total_written)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::frame::{lz4f_compress_frame, lz4f_compress_frame_bound};
fn compress_frame(data: &[u8]) -> Vec<u8> {
let bound = lz4f_compress_frame_bound(data.len(), None);
let mut buf = vec![0u8; bound];
let n = lz4f_compress_frame(&mut buf, data, None).unwrap();
buf.truncate(n);
buf
}
#[test]
fn round_trip_basic() {
let original = b"hello, lz4 frame decompressor!";
let frame = compress_frame(original);
let mut decompressor = FrameDecompressor::new();
let mut dst = Vec::new();
let n = decompress_frame_block(&mut decompressor, &frame, &mut dst, original.len(), false)
.unwrap();
assert_eq!(n, original.len());
assert_eq!(dst, original);
}
#[test]
fn round_trip_1mb() {
let original: Vec<u8> = (0u8..=255).cycle().take(1024 * 1024).collect();
let frame = compress_frame(&original);
let mut decompressor = FrameDecompressor::new();
let mut dst = Vec::new();
let n = decompress_frame_block(&mut decompressor, &frame, &mut dst, original.len(), false)
.unwrap();
assert_eq!(n, original.len());
assert_eq!(dst, original);
}
#[test]
fn skip_checksums_flag_accepted() {
let data = b"test data";
let frame = compress_frame(data);
let mut dec = FrameDecompressor::new();
let mut dst = Vec::new();
let result = decompress_frame_block(&mut dec, &frame, &mut dst, data.len(), true);
assert!(result.is_ok());
}
#[test]
fn dec_function_f_callable_via_type_alias() {
let f: DecFunctionF = decompress_frame_block;
let data = b"type alias test";
let frame = compress_frame(data);
let mut dec = FrameDecompressor::new();
let mut dst = Vec::new();
let n = f(&mut dec, &frame, &mut dst, data.len(), false).unwrap();
assert_eq!(n, data.len());
}
#[test]
fn invalid_frame_returns_error() {
let mut dec = FrameDecompressor::new();
let mut dst = Vec::new();
let result = decompress_frame_block(&mut dec, b"not valid lz4 data", &mut dst, 1024, false);
assert!(result.is_err());
}
#[test]
fn dst_capacity_exceeded_returns_error() {
let data = b"hello, capacity check!";
let frame = compress_frame(data);
let mut dec = FrameDecompressor::new();
let mut dst = Vec::new();
let result = decompress_frame_block(&mut dec, &frame, &mut dst, data.len() - 1, false);
assert!(result.is_err());
assert!(dst.is_empty());
}
}