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
17pub 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 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 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 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(¶ms);
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 pub fn finish(mut self) -> Result<W, io::Error> {
122 self.finish_frame()?;
123 Ok(self.inner)
124 }
125
126 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}