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 .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}