use crate::compress::{FrameCompressor, MtStreamState, ZSTDMT_JOBSIZE_MIN, compress_mt};
use crate::error::Error;
#[derive(PartialEq, Eq, Clone, Copy)]
enum EndOp {
Continue,
Flush,
End,
}
pub struct StreamEncoder {
level: i32,
requested_pledged: Option<u64>,
checksum: bool,
dict: Option<Vec<u8>>,
nb_workers: u32,
job_size: u64,
overlap_log: i32,
state: Option<StreamState>,
mt_state: Option<MtStreamState>,
frame_ended: bool,
}
struct StreamState {
fc: FrameCompressor,
in_buff: Vec<u8>,
in_to_compress: usize,
in_buff_pos: usize,
in_buff_target: usize,
}
impl StreamEncoder {
pub fn new(level: i32) -> Self {
StreamEncoder {
level,
requested_pledged: None,
checksum: false,
dict: None,
nb_workers: 0,
job_size: 0,
overlap_log: 0,
state: None,
mt_state: None,
frame_ended: false,
}
}
pub fn with_dictionary(level: i32, dict: &[u8]) -> Self {
StreamEncoder {
level,
requested_pledged: None,
checksum: false,
dict: if dict.is_empty() {
None
} else {
Some(dict.to_vec())
},
nb_workers: 0,
job_size: 0,
overlap_log: 0,
state: None,
mt_state: None,
frame_ended: false,
}
}
pub fn with_pledged_src_size(level: i32, size: u64) -> Self {
StreamEncoder {
level,
requested_pledged: Some(size),
checksum: false,
dict: None,
nb_workers: 0,
job_size: 0,
overlap_log: 0,
state: None,
mt_state: None,
frame_ended: false,
}
}
pub fn with_checksum(mut self, on: bool) -> Self {
assert!(
self.state.is_none() && self.mt_state.is_none(),
"checksum flag must be set before streaming starts"
);
self.checksum = on;
self
}
pub fn with_workers(mut self, nb_workers: u32, job_size: u64, overlap_log: i32) -> Self {
assert!(
self.state.is_none() && self.mt_state.is_none(),
"workers must be set before streaming starts"
);
self.nb_workers = nb_workers;
self.job_size = job_size;
self.overlap_log = overlap_log;
self
}
pub fn compress(&mut self, input: &[u8], out: &mut Vec<u8>) -> Result<(), Error> {
self.stream_op(input, EndOp::Continue, out)
}
pub fn flush(&mut self, out: &mut Vec<u8>) -> Result<(), Error> {
self.stream_op(&[], EndOp::Flush, out)
}
pub fn finish(mut self, input: &[u8], out: &mut Vec<u8>) -> Result<(), Error> {
self.stream_op(input, EndOp::End, out)
}
fn init(&mut self, end_op: EndOp, in_size: usize) -> Result<(), Error> {
if self.nb_workers > 0 {
let pledged = if end_op == EndOp::End {
Some(in_size as u64)
} else {
self.requested_pledged
};
let engages = match pledged {
Some(p) => p > ZSTDMT_JOBSIZE_MIN,
None => true,
};
if engages {
self.mt_state = Some(MtStreamState::new(
self.level,
self.job_size,
self.overlap_log,
self.checksum,
pledged,
self.dict.as_deref(),
)?);
return Ok(());
}
}
let pledged = if end_op == EndOp::End {
Some(in_size as u64)
} else {
self.requested_pledged
};
let one_block_pledge = |bs: usize| (pledged == Some(bs as u64)) as usize;
self.state = Some(if let Some(dict) = &self.dict {
let init =
crate::compress::streaming_cdict_init(dict, self.level, pledged, self.checksum)?;
let block_size = init.fc.block_size_max();
let in_buff_target = init.content_len + block_size + one_block_pledge(block_size);
StreamState {
fc: init.fc,
in_buff: init.in_buff,
in_to_compress: init.content_len,
in_buff_pos: init.content_len,
in_buff_target,
}
} else {
let fc = FrameCompressor::new(self.level, pledged, self.checksum);
let block_size = fc.block_size_max();
let in_buff_size = fc.window_size() + block_size;
StreamState {
fc,
in_buff: vec![0u8; in_buff_size],
in_to_compress: 0,
in_buff_pos: 0,
in_buff_target: block_size + one_block_pledge(block_size),
}
});
Ok(())
}
fn mt_drive(&mut self, input: &[u8], op: EndOp, out: &mut Vec<u8>) -> Result<(), Error> {
match op {
EndOp::Continue => self.mt_state.as_mut().unwrap().push(input, out),
EndOp::Flush => self.mt_state.as_mut().unwrap().flush(out),
EndOp::End => {
self.mt_state.as_mut().unwrap().end(input, out)?;
self.frame_ended = true;
Ok(())
}
}
}
fn stream_op(&mut self, mut input: &[u8], op: EndOp, out: &mut Vec<u8>) -> Result<(), Error> {
if self.frame_ended {
return Err(Error::Encode("frame already finished"));
}
if self.state.is_none() && self.mt_state.is_none() {
if op == EndOp::End
&& self.nb_workers > 0
&& self.dict.is_none()
&& input.len() as u64 > ZSTDMT_JOBSIZE_MIN
{
out.extend_from_slice(&compress_mt(
input,
self.level,
self.nb_workers,
self.job_size,
self.overlap_log,
self.checksum,
)?);
self.frame_ended = true;
return Ok(());
}
self.init(op, input.len())?;
}
if self.mt_state.is_some() {
return self.mt_drive(input, op, out);
}
let st = self.state.as_mut().expect("initialized above");
let block_size = st.fc.block_size_max();
loop {
let to_load = st.in_buff_target - st.in_buff_pos;
let loaded = to_load.min(input.len());
st.in_buff[st.in_buff_pos..st.in_buff_pos + loaded].copy_from_slice(&input[..loaded]);
st.in_buff_pos += loaded;
input = &input[loaded..];
if op == EndOp::Continue && st.in_buff_pos < st.in_buff_target {
break;
}
if op == EndOp::Flush && st.in_buff_pos == st.in_to_compress {
break;
}
if st.fc.cdict_attach_overflow(st.in_buff_pos) {
return Err(Error::Encode(
"streaming source outgrew the window with an attached dictionary \
(large dict streams are not supported yet)",
));
}
let last_block = op == EndOp::End && input.is_empty();
if last_block {
st.fc
.compress_end(out, &st.in_buff, st.in_to_compress, st.in_buff_pos)?;
self.frame_ended = true;
} else {
st.fc.compress_continue(
out,
&st.in_buff,
st.in_to_compress,
st.in_buff_pos,
false,
)?;
}
st.in_buff_target = st.in_buff_pos + block_size;
if st.in_buff_target > st.in_buff.len() {
st.in_buff_pos = 0;
st.in_buff_target = block_size;
}
st.in_to_compress = st.in_buff_pos;
if self.frame_ended {
break;
}
}
Ok(())
}
}