Skip to main content

zrip_decode/
lib.rs

1#![cfg_attr(not(feature = "std"), no_std)]
2#![cfg_attr(feature = "nightly", feature(optimize_attribute))]
3#![cfg_attr(feature = "paranoid", forbid(unsafe_code))]
4
5#[cfg(feature = "alloc")]
6extern crate alloc;
7
8pub(crate) mod block_decoder;
9#[cfg(feature = "std")]
10pub mod context;
11pub(crate) mod exec;
12pub(crate) mod fast_vec;
13pub(crate) mod literals;
14pub(crate) mod ring_buffer;
15pub(crate) mod sequences;
16#[cfg(feature = "std")]
17pub mod streaming;
18
19#[cfg(feature = "alloc")]
20use alloc::boxed::Box;
21#[cfg(feature = "alloc")]
22use alloc::vec::Vec;
23
24use crate::exec::decode_execute_sequences;
25use crate::literals::decode_literals_ws;
26use crate::sequences::{SequenceDecodeTables, parse_sequence_count, parse_sequence_tables_ws};
27use zrip_core::block::{BlockType, parse_block_header};
28use zrip_core::error::DecompressError;
29use zrip_core::frame::MAX_WINDOW_SIZE;
30use zrip_core::frame::header::parse_frame_header;
31use zrip_core::huffman::HuffmanDecodeEntry;
32use zrip_core::xxhash::Xxh64State;
33
34pub(crate) struct BlockDecodeWorkspace {
35    pub literal_buf: Vec<u8>,
36    pub huf_table: Vec<HuffmanDecodeEntry>,
37    pub huf_table_log: u8,
38    pub huf_valid: bool,
39    pub huf_all_weights: Vec<u8>,
40    pub huf_rank_count: Vec<u32>,
41    pub huf_rank_start: Vec<u32>,
42    pub fse_dist: Vec<i16>,
43    pub fse_symbol_next: Vec<u16>,
44    pub fse_build_buf: Vec<zrip_core::fse::FseDecodeEntry>,
45}
46
47impl BlockDecodeWorkspace {
48    pub(crate) fn new() -> Self {
49        Self {
50            literal_buf: Vec::new(),
51            huf_table: Vec::new(),
52            huf_table_log: 0,
53            huf_valid: false,
54            huf_all_weights: Vec::new(),
55            huf_rank_count: Vec::new(),
56            huf_rank_start: Vec::new(),
57            fse_dist: Vec::new(),
58            fse_symbol_next: Vec::new(),
59            fse_build_buf: Vec::new(),
60        }
61    }
62}
63
64pub(crate) fn skip_skippable_frame(data: &[u8]) -> Option<usize> {
65    if data.len() < 8 {
66        return None;
67    }
68    let magic = u32::from_le_bytes([data[0], data[1], data[2], data[3]]);
69    if (magic & 0xFFFF_FFF0) != 0x184D_2A50 {
70        return None;
71    }
72    let frame_size = u32::from_le_bytes([data[4], data[5], data[6], data[7]]) as usize;
73    let total = 8 + frame_size;
74    if total > data.len() {
75        return None;
76    }
77    Some(total)
78}
79
80pub fn decompress(input: &[u8]) -> Result<Vec<u8>, DecompressError> {
81    decompress_with_dict(input, None)
82}
83
84/// Decompress with an explicit output size limit.
85///
86/// Returns [`DecompressError::OutputTooSmall`] if the decompressed output would
87/// exceed `max_output_size` bytes. Use [`SAFE_DECOMPRESS_LIMIT`](zrip_core::SAFE_DECOMPRESS_LIMIT)
88/// when processing untrusted input to prevent memory exhaustion attacks.
89pub fn decompress_with_limit(
90    input: &[u8],
91    max_output_size: usize,
92) -> Result<Vec<u8>, DecompressError> {
93    let mut output = Vec::new();
94    let mut ws = Box::new(BlockDecodeWorkspace::new());
95    let mut offset = 0;
96    while offset < input.len() {
97        let remaining = &input[offset..];
98        if let Some(skip_len) = skip_skippable_frame(remaining) {
99            offset += skip_len;
100            continue;
101        }
102        let consumed = decompress_frame(remaining, &mut output, max_output_size, None, &mut ws)?;
103        offset += consumed;
104    }
105    Ok(output)
106}
107
108pub fn decompress_into(input: &[u8], output: &mut Vec<u8>) -> Result<usize, DecompressError> {
109    let max_output = zrip_core::DEFAULT_DECOMPRESS_LIMIT;
110    let mut ws = Box::new(BlockDecodeWorkspace::new());
111    let start = output.len();
112    let mut offset = 0;
113    while offset < input.len() {
114        let remaining = &input[offset..];
115        if let Some(skip_len) = skip_skippable_frame(remaining) {
116            offset += skip_len;
117            continue;
118        }
119        let consumed = decompress_frame(remaining, output, max_output, None, &mut ws)?;
120        offset += consumed;
121    }
122    Ok(output.len() - start)
123}
124
125pub fn decompress_with_dict(
126    input: &[u8],
127    dict: Option<&zrip_core::dict::Dictionary>,
128) -> Result<Vec<u8>, DecompressError> {
129    let max_output = zrip_core::DEFAULT_DECOMPRESS_LIMIT;
130    let mut output = Vec::new();
131    let mut ws = Box::new(BlockDecodeWorkspace::new());
132    let mut offset = 0;
133
134    while offset < input.len() {
135        let remaining = &input[offset..];
136        if let Some(skip_len) = skip_skippable_frame(remaining) {
137            offset += skip_len;
138            continue;
139        }
140        let consumed = decompress_frame(remaining, &mut output, max_output, dict, &mut ws)?;
141        offset += consumed;
142    }
143
144    Ok(output)
145}
146
147pub(crate) fn decompress_frame(
148    input: &[u8],
149    output: &mut Vec<u8>,
150    max_output: usize,
151    dict: Option<&zrip_core::dict::Dictionary>,
152    ws: &mut BlockDecodeWorkspace,
153) -> Result<usize, DecompressError> {
154    let header = parse_frame_header(input)?;
155
156    if header.window_size > MAX_WINDOW_SIZE && !header.single_segment {
157        return Err(DecompressError::WindowTooLarge {
158            requested: header.window_size,
159            max: MAX_WINDOW_SIZE,
160        });
161    }
162
163    if let Some(frame_dict_id) = header.dict_id {
164        match dict {
165            Some(d) if d.id() == frame_dict_id => {}
166            Some(d) => {
167                return Err(DecompressError::DictMismatch {
168                    expected: frame_dict_id,
169                    got: d.id(),
170                });
171            }
172            None => return Err(DecompressError::DictRequired),
173        }
174    }
175
176    if let Some(fcs) = header.frame_content_size {
177        if max_output < usize::MAX && fcs as usize > max_output {
178            return Err(DecompressError::OutputTooSmall);
179        }
180        let hint = (fcs as usize).min(MAX_WINDOW_SIZE as usize);
181        output.reserve(hint + 32);
182    }
183
184    let mut offset = header.header_size;
185    let output_start = output.len();
186
187    let dict_history: &[u8] = if let Some(d) = dict { d.content() } else { &[] };
188
189    let mut seq_tables = if let Some(d) = dict {
190        let mut st = SequenceDecodeTables::new_default();
191        if let Some((t, l)) = d.of_table() {
192            st.of_table = crate::sequences::into_table(&zrip_core::fse::promote_of_table(t));
193            st.of_accuracy = l;
194            st.of_set = true;
195        }
196        if let Some((t, l)) = d.ml_table() {
197            st.ml_table = crate::sequences::into_table(&zrip_core::fse::promote_ml_table(t));
198            st.ml_accuracy = l;
199            st.ml_set = true;
200        }
201        if let Some((t, l)) = d.ll_table() {
202            st.ll_table = crate::sequences::into_table(&zrip_core::fse::promote_ll_table(t));
203            st.ll_accuracy = l;
204            st.ll_set = true;
205        }
206        st
207    } else {
208        SequenceDecodeTables::new_default()
209    };
210    let mut rep_offsets: [u32; 3] = if let Some(d) = dict {
211        *d.rep_offsets()
212    } else {
213        [1, 4, 8]
214    };
215    ws.huf_valid = false;
216    if let Some(d) = dict
217        && let Some((t, l)) = d.huf_table()
218    {
219        ws.huf_table.clear();
220        ws.huf_table.extend_from_slice(t);
221        ws.huf_table_log = l;
222        ws.huf_valid = true;
223    }
224
225    let mut hasher = if header.content_checksum {
226        Some(Xxh64State::new(0))
227    } else {
228        None
229    };
230
231    loop {
232        if offset + 3 > input.len() {
233            return Err(DecompressError::InputExhausted);
234        }
235        let block_header = parse_block_header(&input[offset..])?;
236        offset += 3;
237
238        let block_size = block_header.block_size as usize;
239
240        if block_size > zrip_core::frame::MAX_BLOCK_SIZE {
241            match block_header.block_type {
242                BlockType::Raw | BlockType::Rle => {
243                    return Err(DecompressError::BlockTooLarge);
244                }
245                BlockType::Compressed => {}
246            }
247        }
248
249        match block_header.block_type {
250            BlockType::Raw => {
251                if offset + block_size > input.len() {
252                    return Err(DecompressError::InputExhausted);
253                }
254                if output.len() - output_start + block_size > max_output {
255                    return Err(DecompressError::OutputTooSmall);
256                }
257                output.extend_from_slice(&input[offset..offset + block_size]);
258                offset += block_size;
259            }
260            BlockType::Rle => {
261                if offset >= input.len() {
262                    return Err(DecompressError::InputExhausted);
263                }
264                if output.len() - output_start + block_size > max_output {
265                    return Err(DecompressError::OutputTooSmall);
266                }
267                let byte = input[offset];
268                output.resize(output.len() + block_size, byte);
269                offset += 1;
270            }
271            BlockType::Compressed => {
272                if offset + block_size > input.len() {
273                    return Err(DecompressError::InputExhausted);
274                }
275                let block_data = &input[offset..offset + block_size];
276                decode_compressed_block(
277                    block_data,
278                    output,
279                    output_start,
280                    max_output,
281                    &mut seq_tables,
282                    &mut rep_offsets,
283                    ws,
284                    dict_history,
285                )?;
286                offset += block_size;
287            }
288        }
289
290        if block_header.last_block {
291            break;
292        }
293    }
294
295    if let Some(ref mut hasher) = hasher {
296        hasher.update(&output[output_start..]);
297        let hash = hasher.finish();
298        let expected_checksum = (hash & 0xFFFF_FFFF) as u32;
299
300        if offset + 4 > input.len() {
301            return Err(DecompressError::InputExhausted);
302        }
303        let stored_checksum = u32::from_le_bytes([
304            input[offset],
305            input[offset + 1],
306            input[offset + 2],
307            input[offset + 3],
308        ]);
309        offset += 4;
310
311        if expected_checksum != stored_checksum {
312            return Err(DecompressError::ChecksumMismatch {
313                expected: stored_checksum,
314                got: expected_checksum,
315            });
316        }
317    }
318
319    if let Some(fcs) = header.frame_content_size
320        && (output.len() - output_start) as u64 != fcs
321    {
322        return Err(DecompressError::FrameSizeMismatch);
323    }
324
325    Ok(offset)
326}
327
328#[allow(clippy::too_many_arguments)]
329fn decode_compressed_block(
330    data: &[u8],
331    output: &mut Vec<u8>,
332    output_start: usize,
333    max_output: usize,
334    seq_tables: &mut SequenceDecodeTables,
335    rep_offsets: &mut [u32; 3],
336    ws: &mut BlockDecodeWorkspace,
337    dict_history: &[u8],
338) -> Result<(), DecompressError> {
339    let lit_consumed = decode_literals_ws(data, ws)?;
340
341    let remaining = &data[lit_consumed..];
342
343    if remaining.is_empty() {
344        if output.len() - output_start + ws.literal_buf.len() > max_output {
345            return Err(DecompressError::OutputTooSmall);
346        }
347        output.extend_from_slice(&ws.literal_buf);
348        return Ok(());
349    }
350
351    let (num_sequences, seq_count_size) = parse_sequence_count(remaining)?;
352
353    if num_sequences == 0 {
354        if output.len() - output_start + ws.literal_buf.len() > max_output {
355            return Err(DecompressError::OutputTooSmall);
356        }
357        output.extend_from_slice(&ws.literal_buf);
358        return Ok(());
359    }
360
361    let table_data = &remaining[seq_count_size..];
362    let tables_consumed = parse_sequence_tables_ws(table_data, seq_tables, ws)?;
363
364    let seq_data = &table_data[tables_consumed..];
365
366    let before = output.len();
367
368    let result = decode_sequences_dispatch(
369        seq_data,
370        num_sequences,
371        seq_tables,
372        rep_offsets,
373        &ws.literal_buf,
374        output,
375        dict_history,
376    );
377    result?;
378    if output.len() - before > zrip_core::frame::MAX_BLOCK_SIZE {
379        return Err(DecompressError::BlockTooLarge);
380    }
381
382    Ok(())
383}
384
385#[inline(always)]
386pub(crate) fn decode_sequences_dispatch(
387    seq_data: &[u8],
388    num_sequences: u32,
389    seq_tables: &mut SequenceDecodeTables,
390    rep_offsets: &mut [u32; 3],
391    literals: &[u8],
392    output: &mut Vec<u8>,
393    history: &[u8],
394) -> Result<(), DecompressError> {
395    #[cfg(all(feature = "std", feature = "simd"))]
396    {
397        use std::sync::OnceLock;
398        static LEVEL: OnceLock<fearless_simd::Level> = OnceLock::new();
399        let level = *LEVEL.get_or_init(fearless_simd::Level::new);
400        return fearless_simd::dispatch!(level, _simd => {
401            decode_execute_sequences(
402                seq_data,
403                num_sequences,
404                seq_tables,
405                rep_offsets,
406                literals,
407                output,
408                history,
409            )
410        });
411    }
412
413    #[allow(unreachable_code)]
414    decode_execute_sequences(
415        seq_data,
416        num_sequences,
417        seq_tables,
418        rep_offsets,
419        literals,
420        output,
421        history,
422    )
423}