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 (-7..=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 (-7..=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            self.inner.write_all(&self.block_out)?;
218            return Ok(());
219        }
220
221        if self.buffer.is_empty() {
222            return Ok(());
223        }
224
225        let chunk = core::mem::take(&mut self.buffer);
226        let seed_dict = self.first_block && self.dict.is_some();
227
228        let plen = self.window_buf.len();
229        self.window_buf.extend_from_slice(&chunk);
230
231        self.block_out.clear();
232        self.block_out.reserve(chunk.len() + 32);
233        if crate::block_looks_incompressible(&chunk) {
234            block_encoder::encode_raw_block(&chunk, last, &mut self.block_out);
235        } else {
236            if seed_dict {
237                match self.params.strategy {
238                    Strategy::Fast => {
239                        fast::prefill_hash_table(
240                            &self.window_buf,
241                            plen,
242                            self.params.hash_log,
243                            &mut self.hash_table,
244                        );
245                    }
246                    Strategy::DFast => {
247                        dfast::prefill_hash_tables(
248                            &self.window_buf,
249                            plen,
250                            self.params.hash_log,
251                            self.params.chain_log,
252                            self.params.min_match,
253                            &mut self.hash_table,
254                            &mut self.hash_long,
255                        );
256                    }
257                }
258            } else if plen == 0 {
259                self.hash_table.fill(0);
260                if !self.hash_long.is_empty() {
261                    self.hash_long.fill(0);
262                }
263            }
264
265            #[cfg(feature = "ldm")]
266            let used_ldm = if let Some(ref mut ldm) = self.ldm_state {
267                ldm.compress_block(
268                    &self.window_buf,
269                    plen,
270                    self.window_buf.len(),
271                    &self.params,
272                    &self.rep_offsets,
273                    &mut self.hash_table,
274                    &mut self.hash_long,
275                    &mut self.sequences,
276                );
277                true
278            } else {
279                false
280            };
281            #[cfg(not(feature = "ldm"))]
282            let used_ldm = false;
283
284            if !used_ldm {
285                match self.params.strategy {
286                    Strategy::Fast => {
287                        fast::compress_fast_block(
288                            &self.window_buf,
289                            plen,
290                            self.window_buf.len(),
291                            &self.params,
292                            &self.rep_offsets,
293                            &mut self.hash_table,
294                            &mut self.sequences,
295                        );
296                    }
297                    Strategy::DFast => {
298                        dfast::compress_dfast_block(
299                            &self.window_buf,
300                            plen,
301                            self.window_buf.len(),
302                            &self.params,
303                            &self.rep_offsets,
304                            &mut self.hash_table,
305                            &mut self.hash_long,
306                            &mut self.sequences,
307                        );
308                    }
309                }
310            }
311
312            if self.params.force_raw_literals {
313                block_encoder::encode_compressed_block_raw(
314                    &chunk,
315                    &self.sequences,
316                    &mut self.rep_offsets,
317                    last,
318                    &mut self.block_out,
319                    &mut self.workspace,
320                );
321            } else {
322                block_encoder::encode_compressed_block(
323                    &chunk,
324                    &self.sequences,
325                    &mut self.rep_offsets,
326                    last,
327                    &mut self.block_out,
328                    &mut self.workspace,
329                );
330            }
331        }
332
333        let window_size = 1usize << self.params.window_log;
334        if self.window_buf.len() > window_size * 2 {
335            let shift = self.window_buf.len() - window_size;
336            reduce_hash_table(&mut self.hash_table, shift as u32);
337            if !self.hash_long.is_empty() {
338                reduce_hash_table(&mut self.hash_long, shift as u32);
339            }
340            #[cfg(feature = "ldm")]
341            if let Some(ref mut ldm) = self.ldm_state {
342                ldm.reduce_positions(shift as u32);
343            }
344            self.window_buf.copy_within(shift.., 0);
345            self.window_buf.truncate(window_size);
346        }
347
348        self.first_block = false;
349        self.inner.write_all(&self.block_out)?;
350        Ok(())
351    }
352}
353
354fn alloc_hash_tables(params: &LevelParams) -> (Vec<u32>, Vec<u32>) {
355    match params.strategy {
356        Strategy::Fast => (vec![0u32; 1usize << params.hash_log], Vec::new()),
357        Strategy::DFast => (
358            vec![0u32; 1usize << params.chain_log],
359            vec![0u32; 1usize << params.hash_log],
360        ),
361    }
362}
363
364fn reduce_hash_table(table: &mut [u32], shift: u32) {
365    for entry in table.iter_mut() {
366        if *entry < shift {
367            *entry = 0;
368        } else {
369            *entry -= shift;
370        }
371    }
372}
373
374impl<W: Write> Write for FrameEncoder<W> {
375    fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
376        if self.finished {
377            return Err(io::Error::other("encoder already finished"));
378        }
379
380        if !self.header_written {
381            self.write_header()?;
382        }
383
384        self.hasher.update(buf);
385
386        let mut consumed = 0;
387        while consumed < buf.len() {
388            let space = MAX_BLOCK_SIZE - self.buffer.len();
389            let n = space.min(buf.len() - consumed);
390            self.buffer.extend_from_slice(&buf[consumed..consumed + n]);
391            consumed += n;
392
393            if self.buffer.len() >= MAX_BLOCK_SIZE {
394                self.flush_block(false)?;
395            }
396        }
397
398        Ok(consumed)
399    }
400
401    fn flush(&mut self) -> io::Result<()> {
402        if !self.buffer.is_empty() {
403            self.flush_block(false)?;
404        }
405        self.inner.flush()
406    }
407}