Skip to main content

compression_codecs/zstd/
decoder.rs

1use crate::{
2    zstd::{params::DParameter, OperationExt},
3    {DecodeV2, DecodedSize},
4};
5use compression_core::{
6    unshared::Unshared,
7    util::{PartialBuffer, WriteBuffer},
8};
9use libzstd::stream::raw::Decoder;
10use std::{
11    convert::TryInto,
12    io::{self, Result},
13};
14use zstd_safe::get_error_name;
15
16#[derive(Debug)]
17pub struct ZstdDecoder {
18    decoder: Unshared<Decoder<'static>>,
19    stream_ended: bool,
20    needs_input: bool,
21}
22
23impl Default for ZstdDecoder {
24    fn default() -> Self {
25        Self {
26            decoder: Unshared::new(Decoder::new().unwrap()),
27            stream_ended: false,
28            needs_input: false,
29        }
30    }
31}
32
33impl ZstdDecoder {
34    pub fn new() -> Self {
35        Self::default()
36    }
37
38    pub fn new_with_params(params: &[DParameter]) -> Self {
39        let mut decoder = Decoder::new().unwrap();
40        for param in params {
41            decoder.set_parameter(param.as_zstd()).unwrap();
42        }
43        Self {
44            decoder: Unshared::new(decoder),
45            stream_ended: false,
46            needs_input: false,
47        }
48    }
49
50    pub fn new_with_dict(dictionary: &[u8]) -> io::Result<Self> {
51        let decoder = Decoder::with_dictionary(dictionary)?;
52        Ok(Self {
53            decoder: Unshared::new(decoder),
54            stream_ended: false,
55            needs_input: false,
56        })
57    }
58}
59
60impl DecodeV2 for ZstdDecoder {
61    fn reinit(&mut self) -> Result<()> {
62        self.decoder.reinit()?;
63        self.stream_ended = false;
64        self.needs_input = false;
65        Ok(())
66    }
67
68    fn decode(
69        &mut self,
70        input: &mut PartialBuffer<&[u8]>,
71        output: &mut WriteBuffer<'_>,
72    ) -> Result<bool> {
73        let empty_input = input.unwritten().is_empty();
74        // Once buffered output is drained, wait for input before calling zstd
75        // again. Repeated empty-input calls (e.g. while a reader is Pending)
76        // otherwise trigger its no-forward-progress error and poison the context.
77        if empty_input && self.needs_input {
78            return Ok(false);
79        }
80        let before = output.written_len();
81        let finished = self.decoder.run(input, output)?;
82        self.needs_input =
83            empty_input && output.written_len() == before && !output.has_no_spare_space();
84        if finished {
85            self.stream_ended = true;
86        }
87        Ok(finished)
88    }
89
90    fn flush(&mut self, output: &mut WriteBuffer<'_>) -> Result<bool> {
91        // Note: stream_ended is not updated here because zstd's flush only flushes
92        // buffered output and doesn't indicate stream completion. Stream completion
93        // is detected in decode() when status.remaining == 0.
94        self.decoder.flush(output)
95    }
96
97    fn finish(&mut self, output: &mut WriteBuffer<'_>) -> Result<bool> {
98        self.decoder.finish(output)?;
99
100        if self.stream_ended {
101            Ok(true)
102        } else {
103            Err(io::Error::new(
104                io::ErrorKind::UnexpectedEof,
105                "zstd stream did not finish",
106            ))
107        }
108    }
109}
110
111impl DecodedSize for ZstdDecoder {
112    fn decoded_size(input: &[u8]) -> Result<u64> {
113        zstd_safe::find_frame_compressed_size(input)
114            .map_err(|error_code| io::Error::other(get_error_name(error_code)))
115            .and_then(|size| {
116                size.try_into()
117                    .map_err(|_| io::Error::from(io::ErrorKind::FileTooLarge))
118            })
119    }
120}