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