Skip to main content

async_compression/generic/bufread/
decoder.rs

1use crate::{
2    codecs::DecodeV2,
3    core::util::{PartialBuffer, WriteBuffer},
4};
5use std::{io::Result, ops::ControlFlow, panic::AssertUnwindSafe};
6
7#[derive(Debug)]
8enum State {
9    Decoding,
10    Flushing,
11    Done,
12    Next,
13    Error(AssertUnwindSafe<std::io::Error>),
14}
15
16#[derive(Debug)]
17pub struct Decoder {
18    state: State,
19    multiple_members: bool,
20}
21
22impl Default for Decoder {
23    fn default() -> Self {
24        Self {
25            state: State::Decoding,
26            multiple_members: false,
27        }
28    }
29}
30
31impl Decoder {
32    pub fn multiple_members(&mut self, enabled: bool) {
33        self.multiple_members = enabled;
34    }
35
36    pub fn do_poll_read(
37        &mut self,
38        output: &mut WriteBuffer<'_>,
39        decoder: &mut dyn DecodeV2,
40        input: &mut PartialBuffer<&[u8]>,
41        mut first: bool,
42    ) -> ControlFlow<Result<()>> {
43        loop {
44            self.state = match self.state {
45                State::Decoding => {
46                    if input.unwritten().is_empty() && !first {
47                        // Avoid attempting to reinitialise the decoder if the
48                        // reader has returned EOF.
49                        self.multiple_members = false;
50
51                        State::Flushing
52                    } else {
53                        match decoder.decode(input, output) {
54                            Ok(true) => State::Flushing,
55                            // ignore the first error, occurs when input is empty
56                            // but we need to run decode to flush
57                            Err(err) if !first => {
58                                self.state = State::Error(AssertUnwindSafe(err));
59                                if output.written_len() > 0 {
60                                    return ControlFlow::Break(Ok(()));
61                                } else {
62                                    continue;
63                                }
64                            }
65                            // poll for more data for the next decode
66                            _ => break,
67                        }
68                    }
69                }
70
71                State::Flushing => {
72                    let before = output.written_len();
73                    let result = decoder.finish(output).and_then(|finished| {
74                        if !finished
75                            && output.written_len() == before
76                            && !output.has_no_spare_space()
77                        {
78                            Err(std::io::ErrorKind::UnexpectedEof.into())
79                        } else {
80                            Ok(finished)
81                        }
82                    });
83                    match result {
84                        Ok(true) => {
85                            if self.multiple_members {
86                                if let Err(err) = decoder.reinit() {
87                                    self.state = State::Error(AssertUnwindSafe(err));
88                                    if output.written_len() > 0 {
89                                        return ControlFlow::Break(Ok(()));
90                                    } else {
91                                        continue;
92                                    }
93                                }
94
95                                // Poll again only if the decode stage consumed all the input.
96                                first = input.unwritten().is_empty();
97                                State::Next
98                            } else {
99                                State::Done
100                            }
101                        }
102                        Ok(false) => State::Flushing,
103                        Err(err) => {
104                            self.state = State::Error(AssertUnwindSafe(err));
105                            if output.written_len() > 0 {
106                                return ControlFlow::Break(Ok(()));
107                            } else {
108                                continue;
109                            }
110                        }
111                    }
112                }
113
114                State::Done => return ControlFlow::Break(Ok(())),
115
116                State::Next => {
117                    if input.unwritten().is_empty() {
118                        if first {
119                            // poll for more data to check if there's another stream
120                            break;
121                        }
122                        State::Done
123                    } else {
124                        State::Decoding
125                    }
126                }
127
128                State::Error(_) => {
129                    let State::Error(err) = std::mem::replace(&mut self.state, State::Done) else {
130                        unreachable!()
131                    };
132                    return ControlFlow::Break(Err(err.0));
133                }
134            };
135
136            if output.has_no_spare_space() {
137                return ControlFlow::Break(Ok(()));
138            }
139        }
140
141        if output.has_no_spare_space() {
142            ControlFlow::Break(Ok(()))
143        } else {
144            ControlFlow::Continue(())
145        }
146    }
147}
148
149macro_rules! impl_decoder {
150    () => {
151        use crate::generic::bufread::Decoder as GenericDecoder;
152
153        use std::{ops::ControlFlow, task::ready};
154
155        use pin_project_lite::pin_project;
156
157        pin_project! {
158            #[derive(Debug)]
159            pub struct Decoder<R, D> {
160                #[pin]
161                reader: R,
162                decoder: D,
163                inner: GenericDecoder,
164            }
165        }
166
167        impl<R: AsyncBufRead, D: DecodeV2> Decoder<R, D> {
168            pub fn new(reader: R, decoder: D) -> Self {
169                Self {
170                    reader,
171                    decoder,
172                    inner: GenericDecoder::default(),
173                }
174            }
175        }
176
177        impl<R, D> Decoder<R, D> {
178            pub fn get_ref(&self) -> &R {
179                &self.reader
180            }
181
182            pub fn get_mut(&mut self) -> &mut R {
183                &mut self.reader
184            }
185
186            pub fn get_pin_mut(self: Pin<&mut Self>) -> Pin<&mut R> {
187                self.project().reader
188            }
189
190            pub fn into_inner(self) -> R {
191                self.reader
192            }
193
194            pub fn multiple_members(&mut self, enabled: bool) {
195                self.inner.multiple_members(enabled);
196            }
197        }
198
199        fn do_poll_read(
200            inner: &mut GenericDecoder,
201            decoder: &mut dyn DecodeV2,
202            mut reader: Pin<&mut dyn AsyncBufRead>,
203            cx: &mut Context<'_>,
204            output: &mut WriteBuffer<'_>,
205        ) -> Poll<Result<()>> {
206            if let ControlFlow::Break(res) =
207                inner.do_poll_read(output, decoder, &mut PartialBuffer::new(&[][..]), true)
208            {
209                return Poll::Ready(res);
210            }
211
212            loop {
213                let mut input = PartialBuffer::new(match reader.as_mut().poll_fill_buf(cx)? {
214                    Poll::Ready(input) => input,
215                    Poll::Pending if output.written().is_empty() => return Poll::Pending,
216                    _ => return Poll::Ready(Ok(())),
217                });
218
219                let control_flow = inner.do_poll_read(output, decoder, &mut input, false);
220
221                let bytes_read = input.written().len();
222                reader.as_mut().consume(bytes_read);
223
224                if let ControlFlow::Break(res) = control_flow {
225                    break Poll::Ready(res);
226                }
227            }
228        }
229
230        impl<R: AsyncBufRead, D: DecodeV2> Decoder<R, D> {
231            fn do_poll_read(
232                self: Pin<&mut Self>,
233                cx: &mut Context<'_>,
234                output: &mut WriteBuffer<'_>,
235            ) -> Poll<Result<()>> {
236                let this = self.project();
237
238                do_poll_read(this.inner, this.decoder, this.reader, cx, output)
239            }
240        }
241    };
242}
243pub(crate) use impl_decoder;