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