Skip to main content

zrip_decode/
streaming.rs

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
30/// Streaming zstd decompressor implementing [`Read`].
31///
32/// Wraps a reader of compressed data and yields decompressed bytes.
33/// Supports multi-frame streams and skippable frames.
34///
35/// ```no_run
36/// use std::io::Read;
37///
38/// let data = b"hello, streaming world!".repeat(100);
39/// let compressed = zrip::compress(&data, 1).unwrap();
40///
41/// let mut decoder = zrip::FrameDecoder::new(&compressed[..]);
42/// let mut output = Vec::new();
43/// decoder.read_to_end(&mut output).unwrap();
44/// assert_eq!(output, data);
45/// ```
46pub 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    /// Creates a decoder with [`DEFAULT_DECOMPRESS_LIMIT`](zrip_core::DEFAULT_DECOMPRESS_LIMIT).
68    pub fn new(reader: R) -> Self {
69        Self::with_limit(reader, zrip_core::DEFAULT_DECOMPRESS_LIMIT)
70    }
71
72    /// Creates a decoder with an explicit output size limit.
73    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    /// Creates a decoder with a dictionary and default output limit.
96    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    /// Creates a decoder with a dictionary and explicit output limit.
101    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    /// Consumes the decoder and returns the underlying reader.
124    pub fn into_inner(self) -> R {
125        self.inner
126    }
127
128    /// Installs a new reader for the next frame, keeping all internal
129    /// buffers allocated. Returns the previous reader.
130    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}