Skip to main content

zrip_encode/
streaming.rs

1#![forbid(unsafe_code)]
2
3use std::io::{self, Write};
4
5use crate::block_encoder::{self, BlockEncodeWorkspace};
6use crate::dfast;
7use crate::fast;
8#[cfg(feature = "ldm")]
9use crate::ldm::LdmState;
10use crate::strategy::{self, LevelParams, Strategy};
11use crate::write_frame_header_without_content_size;
12use zrip_core::Sequence;
13use zrip_core::dict::Dictionary;
14use zrip_core::error::CompressError;
15use zrip_core::frame::MAX_BLOCK_SIZE;
16use zrip_core::xxhash::Xxh64State;
17
18#[derive(Clone, Copy)]
19enum EncoderState {
20    Ready,
21    Finished,
22    Failed,
23}
24
25/// Streaming zstd compressor implementing [`Write`].
26///
27/// Buffers input until a full block (128 KiB) is ready, then compresses
28/// and writes it to the underlying writer. Call [`finish`](Self::finish)
29/// to flush the final block, write the content checksum, and recover the
30/// writer.
31///
32/// After a write, flush, or finalization error, the encoder cannot be reused.
33/// Later writes, flushes, [`finish`](Self::finish), and [`reset`](Self::reset)
34/// return errors. Discard the encoder and its incomplete output.
35///
36/// Internal buffers (hash tables, sequence scratch, block encoder workspace)
37/// are allocated once and reused across blocks. To reuse them across
38/// multiple frames, call [`reset`](Self::reset) instead of `finish`:
39///
40/// ```
41/// use std::io::Write;
42///
43/// let mut encoder = zrip::FrameEncoder::new(Vec::new(), 1).unwrap();
44/// encoder.write_all(b"first frame").unwrap();
45/// let first = encoder.reset(Vec::new()).unwrap();   // reuses buffers
46/// encoder.write_all(b"second frame").unwrap();
47/// let second = encoder.finish().unwrap();
48/// ```
49pub struct FrameEncoder<W: Write> {
50    inner: W,
51    params: LevelParams,
52    buffer: Vec<u8>,
53    rep_offsets: [u32; 3],
54    hasher: Xxh64State,
55    header_written: bool,
56    state: EncoderState,
57    workspace: BlockEncodeWorkspace,
58    dict: Option<Dictionary>,
59    first_block: bool,
60    hash_table: Vec<u32>,
61    hash_long: Vec<u32>,
62    sequences: Vec<Sequence>,
63    block_out: Vec<u8>,
64    window_buf: Vec<u8>,
65    #[cfg(feature = "ldm")]
66    ldm_state: Option<LdmState>,
67}
68
69impl<W: Write> FrameEncoder<W> {
70    /// Creates a new streaming encoder at the given level (-8..=4).
71    pub fn new(writer: W, level: i32) -> Result<Self, CompressError> {
72        let params = strategy::level_params(level).ok_or(CompressError::InvalidLevel(level))?;
73        Self::from_params(writer, params, None)
74    }
75
76    /// Creates a new streaming encoder with [`Options`](strategy::Options)
77    /// for large windows and/or LDM.
78    pub fn with_options(
79        writer: W,
80        level: i32,
81        opts: &strategy::Options,
82    ) -> Result<Self, CompressError> {
83        let mut params = strategy::level_params(level).ok_or(CompressError::InvalidLevel(level))?;
84        strategy::apply_options(&mut params, opts);
85        Self::from_params(writer, params, None)
86    }
87
88    /// Creates a new streaming encoder with a dictionary at the given level (-8..=4).
89    pub fn with_dict(writer: W, level: i32, dict: Dictionary) -> Result<Self, CompressError> {
90        let params = strategy::level_params(level).ok_or(CompressError::InvalidLevel(level))?;
91        Self::from_params(writer, params, Some(dict))
92    }
93
94    #[allow(clippy::unnecessary_wraps)]
95    fn from_params(
96        writer: W,
97        params: LevelParams,
98        dict: Option<Dictionary>,
99    ) -> Result<Self, CompressError> {
100        let (hash_table, hash_long) = alloc_hash_tables(&params);
101        let (rep_offsets, first_block) = match &dict {
102            Some(d) => (*d.rep_offsets(), true),
103            None => ([1, 4, 8], false),
104        };
105        let mut window_buf = Vec::new();
106        if let Some(ref d) = dict {
107            window_buf.extend_from_slice(d.content());
108        }
109        #[cfg(feature = "ldm")]
110        let ldm_state = params.ldm_params.as_ref().map(LdmState::new);
111        Ok(Self {
112            inner: writer,
113            params,
114            buffer: Vec::new(),
115            rep_offsets,
116            hasher: Xxh64State::new(0),
117            header_written: false,
118            state: EncoderState::Ready,
119            workspace: BlockEncodeWorkspace::new(),
120            dict,
121            first_block,
122            hash_table,
123            hash_long,
124            sequences: Vec::new(),
125            block_out: Vec::new(),
126            window_buf,
127            #[cfg(feature = "ldm")]
128            ldm_state,
129        })
130    }
131
132    /// Flushes remaining data, writes the content checksum, and returns the inner writer.
133    pub fn finish(mut self) -> Result<W, io::Error> {
134        self.finish_frame()?;
135        Ok(self.inner)
136    }
137
138    /// Finishes the current frame and installs `new_writer` for the next one.
139    ///
140    /// Returns the previous writer containing the completed frame. All
141    /// internal buffers (hash tables, workspace, block scratch) stay
142    /// allocated and are reused for the next frame.
143    pub fn reset(&mut self, new_writer: W) -> Result<W, io::Error> {
144        self.finish_frame()?;
145        let old = core::mem::replace(&mut self.inner, new_writer);
146        self.header_written = false;
147        self.state = EncoderState::Ready;
148        self.first_block = self.dict.is_some();
149        self.rep_offsets = match &self.dict {
150            Some(d) => *d.rep_offsets(),
151            None => [1, 4, 8],
152        };
153        self.hasher = Xxh64State::new(0);
154        self.workspace.prev_huffman = None;
155        self.window_buf.clear();
156        if let Some(ref d) = self.dict {
157            self.window_buf.extend_from_slice(d.content());
158        }
159        #[cfg(feature = "ldm")]
160        if let Some(ref mut ldm) = self.ldm_state {
161            ldm.reset();
162        }
163        Ok(old)
164    }
165
166    fn finish_frame(&mut self) -> io::Result<()> {
167        if matches!(self.state, EncoderState::Failed) {
168            return Err(io::Error::other("encoder cannot be reused after an error"));
169        }
170        if matches!(self.state, EncoderState::Finished) {
171            return Ok(());
172        }
173        self.state = EncoderState::Failed;
174
175        if !self.header_written {
176            self.write_header()?;
177        }
178
179        self.flush_block(true)?;
180
181        let hash = self.hasher.finish();
182        let checksum = (hash & 0xFFFF_FFFF) as u32;
183        self.inner.write_all(&checksum.to_le_bytes())?;
184        self.state = EncoderState::Finished;
185        Ok(())
186    }
187
188    fn write_header(&mut self) -> io::Result<()> {
189        self.block_out.clear();
190        write_frame_header_without_content_size(
191            &mut self.block_out,
192            self.dict.as_ref().map(Dictionary::id),
193            self.params.window_log,
194        )
195        .map_err(io::Error::other)?;
196        self.inner.write_all(&self.block_out)?;
197        self.header_written = true;
198        Ok(())
199    }
200
201    fn flush_block(&mut self, last: bool) -> io::Result<()> {
202        if self.buffer.is_empty() && last {
203            self.block_out.clear();
204            block_encoder::encode_raw_block(&[], true, &mut self.block_out)
205                .map_err(io::Error::other)?;
206            self.inner.write_all(&self.block_out)?;
207            return Ok(());
208        }
209
210        if self.buffer.is_empty() {
211            return Ok(());
212        }
213
214        let chunk = core::mem::take(&mut self.buffer);
215        let seed_dict = self.first_block && self.dict.is_some();
216
217        let plen = self.window_buf.len();
218        self.window_buf.extend_from_slice(&chunk);
219
220        self.block_out.clear();
221        self.block_out.reserve(chunk.len() + 32);
222        if crate::block_looks_incompressible(&chunk) {
223            block_encoder::encode_raw_block(&chunk, last, &mut self.block_out)
224                .map_err(io::Error::other)?;
225        } else {
226            if seed_dict {
227                match self.params.strategy {
228                    Strategy::Fast => {
229                        fast::prefill_hash_table(
230                            &self.window_buf,
231                            plen,
232                            self.params.hash_log,
233                            &mut self.hash_table,
234                        );
235                    }
236                    Strategy::DFast => {
237                        dfast::prefill_hash_tables(
238                            &self.window_buf,
239                            plen,
240                            self.params.hash_log,
241                            self.params.chain_log,
242                            self.params.min_match,
243                            &mut self.hash_table,
244                            &mut self.hash_long,
245                        );
246                    }
247                }
248            } else if plen == 0 {
249                self.hash_table.fill(0);
250                if !self.hash_long.is_empty() {
251                    self.hash_long.fill(0);
252                }
253            }
254
255            #[cfg(feature = "ldm")]
256            let used_ldm = if let Some(ref mut ldm) = self.ldm_state {
257                ldm.compress_block(
258                    &self.window_buf,
259                    plen,
260                    self.window_buf.len(),
261                    &self.params,
262                    &self.rep_offsets,
263                    &mut self.hash_table,
264                    &mut self.hash_long,
265                    &mut self.sequences,
266                );
267                true
268            } else {
269                false
270            };
271            #[cfg(not(feature = "ldm"))]
272            let used_ldm = false;
273
274            if !used_ldm {
275                match self.params.strategy {
276                    Strategy::Fast => {
277                        fast::compress_fast_block(
278                            &self.window_buf,
279                            plen,
280                            self.window_buf.len(),
281                            &self.params,
282                            &self.rep_offsets,
283                            &mut self.hash_table,
284                            &mut self.sequences,
285                        );
286                    }
287                    Strategy::DFast => {
288                        dfast::compress_dfast_block(
289                            &self.window_buf,
290                            plen,
291                            self.window_buf.len(),
292                            &self.params,
293                            &self.rep_offsets,
294                            &mut self.hash_table,
295                            &mut self.hash_long,
296                            &mut self.sequences,
297                        );
298                    }
299                }
300            }
301
302            if self.params.force_raw_literals {
303                block_encoder::encode_compressed_block_raw(
304                    &chunk,
305                    &self.sequences,
306                    &mut self.rep_offsets,
307                    last,
308                    &mut self.block_out,
309                    &mut self.workspace,
310                )
311                .map_err(io::Error::other)?;
312            } else {
313                block_encoder::encode_compressed_block(
314                    &chunk,
315                    &self.sequences,
316                    &mut self.rep_offsets,
317                    last,
318                    &mut self.block_out,
319                    &mut self.workspace,
320                    strategy::use_custom_sequence_tables(&self.params, chunk.len()),
321                )
322                .map_err(io::Error::other)?;
323            }
324        }
325
326        let window_size = 1usize << self.params.window_log;
327        if self.window_buf.len() > window_size * 2 {
328            let shift = self.window_buf.len() - window_size;
329            reduce_hash_table(&mut self.hash_table, shift as u32);
330            if !self.hash_long.is_empty() {
331                reduce_hash_table(&mut self.hash_long, shift as u32);
332            }
333            #[cfg(feature = "ldm")]
334            if let Some(ref mut ldm) = self.ldm_state {
335                ldm.reduce_positions(shift as u32);
336            }
337            self.window_buf.copy_within(shift.., 0);
338            self.window_buf.truncate(window_size);
339        }
340
341        self.first_block = false;
342        self.inner.write_all(&self.block_out)?;
343        Ok(())
344    }
345}
346
347fn alloc_hash_tables(params: &LevelParams) -> (Vec<u32>, Vec<u32>) {
348    match params.strategy {
349        Strategy::Fast => (vec![0u32; 1usize << params.hash_log], Vec::new()),
350        Strategy::DFast => (
351            vec![0u32; 1usize << params.chain_log],
352            vec![0u32; 1usize << params.hash_log],
353        ),
354    }
355}
356
357fn reduce_hash_table(table: &mut [u32], shift: u32) {
358    for entry in table.iter_mut() {
359        if *entry < shift {
360            *entry = 0;
361        } else {
362            *entry -= shift;
363        }
364    }
365}
366
367impl<W: Write> Write for FrameEncoder<W> {
368    fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
369        if matches!(self.state, EncoderState::Failed) {
370            return Err(io::Error::other("encoder cannot be reused after an error"));
371        }
372        if matches!(self.state, EncoderState::Finished) {
373            return Err(io::Error::other("encoder already finished"));
374        }
375        // Any early return leaves the encoder failed, including partial writes.
376        self.state = EncoderState::Failed;
377
378        if !self.header_written {
379            self.write_header()?;
380        }
381
382        self.hasher.update(buf);
383
384        let mut consumed = 0;
385        while consumed < buf.len() {
386            let space = MAX_BLOCK_SIZE - self.buffer.len();
387            let n = space.min(buf.len() - consumed);
388            self.buffer.extend_from_slice(&buf[consumed..consumed + n]);
389            consumed += n;
390
391            if self.buffer.len() >= MAX_BLOCK_SIZE {
392                self.flush_block(false)?;
393            }
394        }
395
396        self.state = EncoderState::Ready;
397        Ok(consumed)
398    }
399
400    fn flush(&mut self) -> io::Result<()> {
401        if matches!(self.state, EncoderState::Failed) {
402            return Err(io::Error::other("encoder cannot be reused after an error"));
403        }
404        let state = self.state;
405        self.state = EncoderState::Failed;
406        if !self.buffer.is_empty() {
407            self.flush_block(false)?;
408        }
409        self.inner.flush()?;
410        self.state = state;
411        Ok(())
412    }
413}