compression_codecs/zstd/
decoder.rs1use 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 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 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}