1#![forbid(unsafe_code)]
2
3use std::io::{self, Read};
4
5use crate::BlockDecodeWorkspace;
6use crate::exec::SequenceOutputScope;
7use crate::literals::decode_literals_ws;
8use crate::sequences::{SequenceDecodeTables, parse_sequence_count, parse_sequence_tables_ws};
9
10use crate::decode_sequences_dispatch;
11use zrip_core::block::{BlockType, parse_block_header};
12use zrip_core::dict::Dictionary;
13use zrip_core::error::DecompressError;
14use zrip_core::frame::header::parse_frame_header;
15use zrip_core::frame::{MAX_BLOCK_SIZE, MAX_WINDOW_SIZE};
16use zrip_core::xxhash::Xxh64State;
17
18enum State {
19 FrameHeader,
20 BlockHeader,
21 BlockData {
22 block_type: BlockType,
23 block_size: usize,
24 last: bool,
25 },
26 Checksum,
27 Done,
28}
29
30pub struct FrameDecoder<R: Read> {
47 inner: R,
48 state: State,
49 read_buf: Vec<u8>,
50 output_buf: Vec<u8>,
51 output_pos: usize,
52 ws: Box<BlockDecodeWorkspace>,
53 seq_tables: SequenceDecodeTables,
54 rep_offsets: [u32; 3],
55 hasher: Option<Xxh64State>,
56 content_checksum: bool,
57 max_output: usize,
58 bytes_output: usize,
59 frame_content_size: Option<u64>,
60 frame_bytes: usize,
61 dict: Option<Dictionary>,
62 decode_history: Vec<u8>,
63 window_size: usize,
64}
65
66impl<R: Read> FrameDecoder<R> {
67 pub fn new(reader: R) -> Self {
69 Self::with_limit(reader, zrip_core::DEFAULT_DECOMPRESS_LIMIT)
70 }
71
72 pub fn with_limit(reader: R, max_output: usize) -> Self {
74 Self {
75 inner: reader,
76 state: State::FrameHeader,
77 read_buf: Vec::new(),
78 output_buf: Vec::new(),
79 output_pos: 0,
80 ws: Box::new(BlockDecodeWorkspace::new()),
81 seq_tables: SequenceDecodeTables::new_default(),
82 rep_offsets: [1, 4, 8],
83 hasher: None,
84 content_checksum: false,
85 max_output,
86 bytes_output: 0,
87 frame_content_size: None,
88 frame_bytes: 0,
89 dict: None,
90 decode_history: Vec::new(),
91 window_size: 0,
92 }
93 }
94
95 pub fn with_dict(reader: R, dict: Dictionary) -> Self {
97 Self::with_dict_and_limit(reader, dict, zrip_core::DEFAULT_DECOMPRESS_LIMIT)
98 }
99
100 pub fn with_dict_and_limit(reader: R, dict: Dictionary, max_output: usize) -> Self {
102 Self {
103 inner: reader,
104 state: State::FrameHeader,
105 read_buf: Vec::new(),
106 output_buf: Vec::new(),
107 output_pos: 0,
108 ws: Box::new(BlockDecodeWorkspace::new()),
109 seq_tables: SequenceDecodeTables::new_default(),
110 rep_offsets: [1, 4, 8],
111 hasher: None,
112 content_checksum: false,
113 max_output,
114 bytes_output: 0,
115 frame_content_size: None,
116 frame_bytes: 0,
117 dict: Some(dict),
118 decode_history: Vec::new(),
119 window_size: 0,
120 }
121 }
122
123 pub fn into_inner(self) -> R {
125 self.inner
126 }
127
128 pub fn reset(&mut self, new_reader: R) -> R {
131 let old = core::mem::replace(&mut self.inner, new_reader);
132 self.state = State::FrameHeader;
133 self.output_buf.clear();
134 self.output_pos = 0;
135 self.rep_offsets = [1, 4, 8];
136 self.seq_tables = SequenceDecodeTables::new_default();
137 self.ws.reset_huffman_state();
138 self.hasher = None;
139 self.content_checksum = false;
140 self.bytes_output = 0;
141 self.frame_content_size = None;
142 self.frame_bytes = 0;
143 self.decode_history.clear();
144 self.window_size = 0;
145 old
146 }
147
148 fn fill_output(&mut self) -> io::Result<()> {
149 loop {
150 match self.state {
151 State::Done => return Ok(()),
152 State::FrameHeader => self.read_frame_header()?,
153 State::BlockHeader => self.read_block_header()?,
154 State::BlockData {
155 block_type,
156 block_size,
157 last,
158 } => {
159 self.read_block_data(block_type, block_size, last)?;
160 if self.output_pos < self.output_buf.len() {
161 return Ok(());
162 }
163 }
164 State::Checksum => self.read_checksum()?,
165 }
166 }
167 }
168
169 fn read_frame_header(&mut self) -> io::Result<()> {
170 self.read_buf.resize(18, 0);
171 self.inner.read_exact(&mut self.read_buf[..5])?;
172
173 let magic = u32::from_le_bytes([
174 self.read_buf[0],
175 self.read_buf[1],
176 self.read_buf[2],
177 self.read_buf[3],
178 ]);
179
180 if (magic & 0xFFFF_FFF0) == 0x184D_2A50 {
181 self.inner.read_exact(&mut self.read_buf[5..9])?;
182 let skip_size = u32::from_le_bytes([
183 self.read_buf[5],
184 self.read_buf[6],
185 self.read_buf[7],
186 self.read_buf[8],
187 ]) as usize;
188 io::copy(
189 &mut self.inner.by_ref().take(skip_size as u64),
190 &mut io::sink(),
191 )?;
192 return Ok(());
193 }
194
195 let descriptor = self.read_buf[4];
196 let single_segment = (descriptor & 0x20) != 0;
197 let dict_id_flag = descriptor & 0x03;
198 let fcs_flag = (descriptor >> 6) & 0x03;
199
200 let mut hdr_len = 5usize;
201 if !single_segment {
202 hdr_len += 1;
203 }
204 hdr_len += match dict_id_flag {
205 0 => 0,
206 1 => 1,
207 2 => 2,
208 3 => 4,
209 _ => unreachable!(),
210 };
211 hdr_len += match fcs_flag {
212 0 if single_segment => 1,
213 0 => 0,
214 1 => 2,
215 2 => 4,
216 3 => 8,
217 _ => unreachable!(),
218 };
219
220 if hdr_len > 5 {
221 self.inner.read_exact(&mut self.read_buf[5..hdr_len])?;
222 }
223
224 let header = parse_frame_header(&self.read_buf[..hdr_len])
225 .map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
226
227 if let Some(frame_dict_id) = header.dict_id {
228 match &self.dict {
229 Some(d) if d.id() == frame_dict_id => {}
230 Some(d) => {
231 return Err(io::Error::new(
232 io::ErrorKind::InvalidData,
233 DecompressError::DictMismatch {
234 expected: frame_dict_id,
235 got: d.id(),
236 },
237 ));
238 }
239 None => {
240 return Err(io::Error::new(
241 io::ErrorKind::InvalidData,
242 DecompressError::DictRequired,
243 ));
244 }
245 }
246 }
247
248 let window_size = if header.window_size > MAX_WINDOW_SIZE {
249 if header.single_segment {
250 MAX_WINDOW_SIZE as usize
251 } else {
252 return Err(io::Error::new(
253 io::ErrorKind::InvalidData,
254 DecompressError::WindowTooLarge {
255 requested: header.window_size,
256 max: MAX_WINDOW_SIZE,
257 },
258 ));
259 }
260 } else {
261 header.window_size as usize
262 };
263
264 if let Some(fcs) = header.frame_content_size
265 && fcs as usize > self.max_output
266 {
267 return Err(io::Error::new(
268 io::ErrorKind::InvalidData,
269 DecompressError::OutputTooSmall,
270 ));
271 }
272
273 self.window_size = window_size;
274 self.decode_history.clear();
275 self.frame_content_size = header.frame_content_size;
276 self.frame_bytes = 0;
277 self.content_checksum = header.content_checksum;
278 self.hasher = if header.content_checksum {
279 Some(Xxh64State::new(0))
280 } else {
281 None
282 };
283
284 if let Some(ref d) = self.dict {
285 self.rep_offsets = *d.rep_offsets();
286 self.decode_history.extend_from_slice(d.content());
287 let mut st = SequenceDecodeTables::new_default();
288 if let Some((t, l)) = d.of_table() {
289 st.of_table = crate::seq_table::SeqTable::promote_of(t);
290 st.of_accuracy = l;
291 st.of_set = true;
292 }
293 if let Some((t, l)) = d.ml_table() {
294 st.ml_table = crate::seq_table::SeqTable::promote_ml(t);
295 st.ml_accuracy = l;
296 st.ml_set = true;
297 }
298 if let Some((t, l)) = d.ll_table() {
299 st.ll_table = crate::seq_table::SeqTable::promote_ll(t);
300 st.ll_accuracy = l;
301 st.ll_set = true;
302 }
303 self.seq_tables = st;
304 self.ws.reset_huffman_state();
305 if let Some((t, l)) = d.huf_table() {
306 self.ws.huf_table.clear();
307 self.ws.huf_table.extend_from_slice(t);
308 self.ws.huf_table_log = l;
309 self.ws.huf_valid = true;
310 }
311 } else {
312 self.rep_offsets = [1, 4, 8];
313 self.seq_tables = SequenceDecodeTables::new_default();
314 self.ws.reset_huffman_state();
315 }
316
317 self.state = State::BlockHeader;
318 Ok(())
319 }
320
321 fn read_block_header(&mut self) -> io::Result<()> {
322 let mut hdr = [0u8; 3];
323 self.inner.read_exact(&mut hdr)?;
324 let block_header =
325 parse_block_header(&hdr).map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
326
327 let block_size = block_header.block_size as usize;
328
329 match block_header.block_type {
330 BlockType::Raw | BlockType::Rle if block_size > MAX_BLOCK_SIZE => {
331 return Err(io::Error::new(
332 io::ErrorKind::InvalidData,
333 DecompressError::BlockTooLarge,
334 ));
335 }
336 _ => {}
337 }
338
339 self.state = State::BlockData {
340 block_type: block_header.block_type,
341 block_size,
342 last: block_header.last_block,
343 };
344 Ok(())
345 }
346
347 fn read_block_data(
348 &mut self,
349 block_type: BlockType,
350 block_size: usize,
351 last: bool,
352 ) -> io::Result<()> {
353 self.output_buf.clear();
354 self.output_pos = 0;
355
356 match block_type {
357 BlockType::Raw => {
358 self.output_buf.resize(block_size, 0);
359 self.inner.read_exact(&mut self.output_buf)?;
360 }
361 BlockType::Rle => {
362 let mut byte = [0u8; 1];
363 self.inner.read_exact(&mut byte)?;
364 self.output_buf.resize(block_size, byte[0]);
365 }
366 BlockType::Compressed => {
367 self.read_buf.resize(block_size, 0);
368 self.inner.read_exact(&mut self.read_buf[..block_size])?;
369 self.decode_compressed_block(block_size)?;
370 }
371 }
372
373 if let Some(ref mut hasher) = self.hasher {
374 hasher.update(&self.output_buf);
375 }
376 self.bytes_output += self.output_buf.len();
377 self.frame_bytes += self.output_buf.len();
378 if self.bytes_output > self.max_output {
379 return Err(io::Error::new(
380 io::ErrorKind::InvalidData,
381 DecompressError::OutputTooSmall,
382 ));
383 }
384
385 if self.window_size > 0 {
386 self.decode_history.extend_from_slice(&self.output_buf);
387 if self.decode_history.len() > self.window_size {
388 let start = self.decode_history.len() - self.window_size;
389 self.decode_history.copy_within(start.., 0);
390 self.decode_history.truncate(self.window_size);
391 }
392 }
393
394 self.state = if last {
395 if let Some(fcs) = self.frame_content_size
396 && self.frame_bytes as u64 != fcs
397 {
398 return Err(io::Error::new(
399 io::ErrorKind::InvalidData,
400 DecompressError::FrameSizeMismatch,
401 ));
402 }
403 if self.content_checksum {
404 State::Checksum
405 } else {
406 State::FrameHeader
407 }
408 } else {
409 State::BlockHeader
410 };
411
412 Ok(())
413 }
414
415 fn decode_compressed_block(&mut self, block_size: usize) -> io::Result<()> {
416 let history: &[u8] = &self.decode_history;
417 let block_data = &self.read_buf[..block_size];
418
419 let lit_consumed = decode_literals_ws(block_data, &mut self.ws)
420 .map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
421
422 let remaining = &block_data[lit_consumed..];
423
424 if remaining.is_empty() {
425 self.output_buf.extend_from_slice(&self.ws.literal_buf);
426 return Ok(());
427 }
428
429 let (num_sequences, seq_count_size) = parse_sequence_count(remaining)
430 .map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
431
432 if num_sequences == 0 {
433 self.output_buf.extend_from_slice(&self.ws.literal_buf);
434 return Ok(());
435 }
436
437 let table_data = &remaining[seq_count_size..];
438 let tables_consumed =
439 parse_sequence_tables_ws(table_data, &mut self.seq_tables, &mut self.ws)
440 .map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
441
442 let seq_data = &table_data[tables_consumed..];
443
444 let before = self.output_buf.len();
445 let max_block_output = self
446 .max_output
447 .saturating_sub(self.bytes_output)
448 .min(MAX_BLOCK_SIZE);
449 let scope = SequenceOutputScope {
450 output_base: 0,
451 max_block_output,
452 history,
453 };
454 decode_sequences_dispatch(
455 seq_data,
456 num_sequences,
457 &mut self.seq_tables,
458 &mut self.rep_offsets,
459 &self.ws.literal_buf,
460 &mut self.output_buf,
461 scope,
462 )
463 .map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
464 if self.output_buf.len() - before > MAX_BLOCK_SIZE {
465 return Err(io::Error::new(
466 io::ErrorKind::InvalidData,
467 DecompressError::BlockTooLarge,
468 ));
469 }
470 Ok(())
471 }
472
473 fn read_checksum(&mut self) -> io::Result<()> {
474 let mut buf = [0u8; 4];
475 self.inner.read_exact(&mut buf)?;
476 let stored = u32::from_le_bytes(buf);
477
478 if let Some(ref hasher) = self.hasher {
479 let hash = hasher.finish();
480 let expected = (hash & 0xFFFF_FFFF) as u32;
481 if expected != stored {
482 return Err(io::Error::new(
483 io::ErrorKind::InvalidData,
484 DecompressError::ChecksumMismatch {
485 expected: stored,
486 got: expected,
487 },
488 ));
489 }
490 }
491
492 self.state = State::FrameHeader;
493 Ok(())
494 }
495}
496
497impl<R: Read> Read for FrameDecoder<R> {
498 fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
499 if self.output_pos >= self.output_buf.len() {
500 if let State::Done = &self.state {
501 return Ok(0);
502 }
503
504 self.output_buf.clear();
505 self.output_pos = 0;
506
507 match self.fill_output() {
508 Ok(()) => {}
509 Err(e) if e.kind() == io::ErrorKind::UnexpectedEof => match &self.state {
510 State::FrameHeader => {
511 self.state = State::Done;
512 return Ok(0);
513 }
514 _ => return Err(e),
515 },
516 Err(e) => return Err(e),
517 }
518 }
519
520 let available = &self.output_buf[self.output_pos..];
521 let n = buf.len().min(available.len());
522 buf[..n].copy_from_slice(&available[..n]);
523 self.output_pos += n;
524 Ok(n)
525 }
526}