use std::io::{self, Read, Write};
use crate::frame::compress::{
lz4f_compress_begin, lz4f_compress_bound, lz4f_compress_end, lz4f_compress_update,
lz4f_create_compression_context,
};
use crate::frame::decompress::{
lz4f_create_decompression_context, lz4f_decompress, lz4f_get_frame_info, Lz4FDCtx,
};
use crate::frame::types::{
BlockSizeId, Lz4FCCtx, Lz4FError, Preferences, LZ4F_VERSION, MAX_FH_SIZE,
};
fn block_size_from_id(id: BlockSizeId) -> usize {
match id {
BlockSizeId::Default | BlockSizeId::Max64Kb => 64 * 1024,
BlockSizeId::Max256Kb => 256 * 1024,
BlockSizeId::Max1Mb => 1024 * 1024,
BlockSizeId::Max4Mb => 4 * 1024 * 1024,
}
}
pub struct Lz4ReadFile<R: Read> {
dctx: Box<Lz4FDCtx>,
inner: R,
src_buf: Vec<u8>,
src_buf_size: usize,
src_buf_next: usize,
}
impl<R: Read> Lz4ReadFile<R> {
pub fn open(mut reader: R) -> Result<Self, Lz4FError> {
let mut dctx = lz4f_create_decompression_context(LZ4F_VERSION)?;
let mut header_buf = [0u8; MAX_FH_SIZE];
let mut total_read = 0usize;
while total_read < MAX_FH_SIZE {
let n = reader
.read(&mut header_buf[total_read..])
.map_err(|_| Lz4FError::IoRead)?;
if n == 0 {
break; }
total_read += n;
}
if total_read == 0 {
return Err(Lz4FError::IoRead);
}
let (frame_info, consumed, _hint) =
lz4f_get_frame_info(&mut dctx, &header_buf[..total_read])?;
let src_buf_max_size = block_size_from_id(frame_info.block_size_id);
let leftover = total_read - consumed;
let mut src_buf = vec![0u8; src_buf_max_size];
src_buf[..leftover].copy_from_slice(&header_buf[consumed..total_read]);
Ok(Lz4ReadFile {
dctx,
inner: reader,
src_buf,
src_buf_size: leftover,
src_buf_next: 0,
})
}
}
impl<R: Read> Read for Lz4ReadFile<R> {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
let size = buf.len();
let mut next: usize = 0;
while next < size {
let src_avail = self.src_buf_size - self.src_buf_next;
if src_avail == 0 {
let n = self.inner.read(&mut self.src_buf)?;
if n == 0 {
break; }
self.src_buf_size = n;
self.src_buf_next = 0;
}
let src_avail = self.src_buf_size - self.src_buf_next;
let src_copy = self.src_buf[self.src_buf_next..self.src_buf_next + src_avail].to_vec();
let (src_consumed, dst_written, _hint) =
lz4f_decompress(&mut self.dctx, Some(&mut buf[next..]), &src_copy, None)
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e.to_string()))?;
self.src_buf_next += src_consumed;
next += dst_written;
}
Ok(next)
}
}
impl<R: Read> Drop for Lz4ReadFile<R> {
fn drop(&mut self) {}
}
pub struct Lz4WriteFile<W: Write> {
cctx: Box<Lz4FCCtx>,
inner: Option<W>,
dst_buf: Vec<u8>,
max_write_size: usize,
errored: bool,
}
impl<W: Write> Lz4WriteFile<W> {
pub fn open(mut writer: W, prefs: Option<&Preferences>) -> Result<Self, Lz4FError> {
let max_write_size = prefs
.map(|p| block_size_from_id(p.frame_info.block_size_id))
.unwrap_or(64 * 1024);
let dst_buf_max_size = lz4f_compress_bound(max_write_size, prefs);
let mut dst_buf = vec![0u8; dst_buf_max_size];
let mut cctx = lz4f_create_compression_context(LZ4F_VERSION)?;
let header_size = lz4f_compress_begin(&mut cctx, &mut dst_buf, prefs)?;
writer
.write_all(&dst_buf[..header_size])
.map_err(|_| Lz4FError::IoWrite)?;
Ok(Lz4WriteFile {
cctx,
inner: Some(writer),
dst_buf,
max_write_size,
errored: false,
})
}
pub fn finish(mut self) -> Result<W, Lz4FError> {
if !self.errored {
let writer = self.inner.as_mut().expect("inner writer already taken");
let end_size = lz4f_compress_end(&mut self.cctx, &mut self.dst_buf, None)?;
writer
.write_all(&self.dst_buf[..end_size])
.map_err(|_| Lz4FError::IoWrite)?;
}
Ok(self.inner.take().expect("inner writer already taken"))
}
}
impl<W: Write> Write for Lz4WriteFile<W> {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
let mut remain = buf.len();
let mut p = 0usize;
while remain > 0 {
let chunk = remain.min(self.max_write_size);
let compressed =
lz4f_compress_update(&mut self.cctx, &mut self.dst_buf, &buf[p..p + chunk], None)
.map_err(|e| {
self.errored = true;
io::Error::other(e.to_string())
})?;
self.inner
.as_mut()
.expect("inner writer already taken")
.write_all(&self.dst_buf[..compressed])
.inspect_err(|e| {
self.errored = true;
})?;
p += chunk;
remain -= chunk;
}
Ok(buf.len())
}
fn flush(&mut self) -> io::Result<()> {
self.inner
.as_mut()
.expect("inner writer already taken")
.flush()
}
}
impl<W: Write> Drop for Lz4WriteFile<W> {
fn drop(&mut self) {
if self.inner.is_none() || self.errored {
return; }
let _ = lz4f_compress_end(&mut self.cctx, &mut self.dst_buf, None).and_then(|end_size| {
self.inner
.as_mut()
.unwrap()
.write_all(&self.dst_buf[..end_size])
.map_err(|_| Lz4FError::IoWrite)
});
}
}
pub fn lz4_write_frame<W: Write>(data: &[u8], writer: W) -> Result<W, Lz4FError> {
let mut lz4w = Lz4WriteFile::open(writer, None)?;
lz4w.write_all(data).map_err(|_| Lz4FError::IoWrite)?;
lz4w.finish()
}
pub fn lz4_read_frame<R: Read>(reader: R, writer: &mut impl Write) -> Result<(), Lz4FError> {
let mut lz4r = Lz4ReadFile::open(reader)?;
let mut tmp = [0u8; 64 * 1024];
loop {
let n = lz4r.read(&mut tmp).map_err(|_| Lz4FError::IoRead)?;
if n == 0 {
break;
}
writer
.write_all(&tmp[..n])
.map_err(|_| Lz4FError::IoWrite)?;
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Cursor;
#[test]
fn round_trip_small() {
let original = b"Hello, LZ4 world! This is a test of the lz4 file streaming API.";
let compressed = lz4_write_frame(original, Vec::new()).unwrap();
let mut recovered = Vec::new();
lz4_read_frame(Cursor::new(&compressed), &mut recovered).unwrap();
assert_eq!(recovered, original);
}
#[test]
fn round_trip_multi_block() {
let original: Vec<u8> = (0u8..=255).cycle().take(200 * 1024).collect();
let compressed = lz4_write_frame(&original, Vec::new()).unwrap();
let mut recovered = Vec::new();
lz4_read_frame(Cursor::new(&compressed), &mut recovered).unwrap();
assert_eq!(recovered, original);
}
#[test]
fn streaming_write_read() {
let original: Vec<u8> = b"streaming test data"
.iter()
.cycle()
.take(4096)
.cloned()
.collect();
let mut lz4w = Lz4WriteFile::open(Vec::new(), None).unwrap();
for chunk in original.chunks(256) {
lz4w.write_all(chunk).unwrap();
}
let compressed = lz4w.finish().unwrap();
let mut lz4r = Lz4ReadFile::open(Cursor::new(&compressed)).unwrap();
let mut recovered = Vec::new();
let mut tmp = [0u8; 512];
loop {
let n = lz4r.read(&mut tmp).unwrap();
if n == 0 {
break;
}
recovered.extend_from_slice(&tmp[..n]);
}
assert_eq!(recovered, original);
}
#[test]
fn round_trip_one_byte() {
let original = b"x";
let compressed = lz4_write_frame(original, Vec::new()).unwrap();
let mut recovered = Vec::new();
lz4_read_frame(Cursor::new(&compressed), &mut recovered).unwrap();
assert_eq!(recovered.as_slice(), original.as_ref());
}
#[test]
fn round_trip_empty() {
let original: &[u8] = b"";
let compressed = lz4_write_frame(original, Vec::new()).unwrap();
let mut recovered = Vec::new();
lz4_read_frame(Cursor::new(&compressed), &mut recovered).unwrap();
assert_eq!(recovered.as_slice(), original);
}
}