Skip to main content

async_compression/generic/write/
decoder.rs

1use crate::{
2    codecs::DecodeV2,
3    core::util::{PartialBuffer, WriteBuffer},
4    generic::write::AsyncBufWrite,
5};
6use std::{
7    io,
8    pin::Pin,
9    task::{ready, Context, Poll},
10};
11
12#[derive(Debug)]
13enum State {
14    Decoding,
15    Finishing,
16    Done,
17}
18
19#[derive(Debug)]
20pub struct Decoder {
21    state: State,
22}
23
24impl Default for Decoder {
25    fn default() -> Self {
26        Self {
27            state: State::Decoding,
28        }
29    }
30}
31
32impl Decoder {
33    fn do_poll_write(
34        &mut self,
35        cx: &mut Context<'_>,
36        input: &mut PartialBuffer<&[u8]>,
37        mut writer: Pin<&mut dyn AsyncBufWrite>,
38        decoder: &mut dyn DecodeV2,
39    ) -> Poll<io::Result<()>> {
40        loop {
41            let mut output = ready!(writer.as_mut().poll_partial_flush_buf(cx))?;
42            let output = &mut output.write_buffer;
43
44            self.state = match self.state {
45                State::Decoding => {
46                    if decoder.decode(input, output)? {
47                        State::Finishing
48                    } else {
49                        State::Decoding
50                    }
51                }
52
53                State::Finishing => {
54                    if decoder.finish(output)? {
55                        State::Done
56                    } else {
57                        State::Finishing
58                    }
59                }
60
61                State::Done => {
62                    return Poll::Ready(Err(io::Error::other("Write after end of stream")));
63                }
64            };
65
66            if let State::Done = self.state {
67                return Poll::Ready(Ok(()));
68            }
69
70            if input.unwritten().is_empty() {
71                return Poll::Ready(Ok(()));
72            }
73        }
74    }
75
76    pub fn poll_write(
77        &mut self,
78        cx: &mut Context<'_>,
79        buf: &[u8],
80        writer: Pin<&mut dyn AsyncBufWrite>,
81        decoder: &mut dyn DecodeV2,
82    ) -> Poll<io::Result<usize>> {
83        if buf.is_empty() {
84            return Poll::Ready(Ok(0));
85        }
86
87        let mut input = PartialBuffer::new(buf);
88
89        match self.do_poll_write(cx, &mut input, writer, decoder)? {
90            Poll::Pending if input.written().is_empty() => Poll::Pending,
91            _ => Poll::Ready(Ok(input.written().len())),
92        }
93    }
94
95    pub fn do_poll_flush(
96        &mut self,
97        cx: &mut Context<'_>,
98        mut writer: Pin<&mut dyn AsyncBufWrite>,
99        decoder: &mut dyn DecodeV2,
100    ) -> Poll<io::Result<()>> {
101        loop {
102            let mut output = ready!(writer.as_mut().poll_partial_flush_buf(cx))?;
103            let output = &mut output.write_buffer;
104
105            let (state, done) = match self.state {
106                State::Decoding => {
107                    let done = decoder.flush(output)?;
108                    (State::Decoding, done)
109                }
110
111                State::Finishing => {
112                    let before = output.written_len();
113                    if decoder.finish(output)? {
114                        (State::Done, false)
115                    } else if output.written_len() == before && !output.has_no_spare_space() {
116                        return Poll::Ready(Err(io::ErrorKind::UnexpectedEof.into()));
117                    } else {
118                        (State::Finishing, false)
119                    }
120                }
121
122                State::Done => (State::Done, true),
123            };
124
125            self.state = state;
126
127            if done {
128                break Poll::Ready(Ok(()));
129            }
130        }
131    }
132
133    pub fn do_close(&mut self) {
134        if let State::Decoding = self.state {
135            self.state = State::Finishing;
136        }
137    }
138
139    pub fn is_done(&self) -> bool {
140        matches!(self.state, State::Done)
141    }
142}
143
144macro_rules! impl_decoder {
145    ($poll_close: tt) => {
146        use crate::{
147            codecs::DecodeV2, core::util::PartialBuffer, generic::write::Decoder as GenericDecoder,
148        };
149        use pin_project_lite::pin_project;
150        use std::task::ready;
151
152        pin_project! {
153            #[derive(Debug)]
154            pub struct Decoder<W, D> {
155                #[pin]
156                writer: BufWriter<W>,
157                decoder: D,
158                inner: GenericDecoder,
159            }
160        }
161
162        impl<W: AsyncWrite, D: DecodeV2> Decoder<W, D> {
163            pub fn new(writer: W, decoder: D) -> Self {
164                Self {
165                    writer: BufWriter::new(writer),
166                    decoder,
167                    inner: Default::default(),
168                }
169            }
170        }
171
172        impl<W, D> Decoder<W, D> {
173            pub fn get_ref(&self) -> &W {
174                self.writer.get_ref()
175            }
176
177            pub fn get_mut(&mut self) -> &mut W {
178                self.writer.get_mut()
179            }
180
181            pub fn get_pin_mut(self: Pin<&mut Self>) -> Pin<&mut W> {
182                self.project().writer.get_pin_mut()
183            }
184
185            pub fn into_inner(self) -> W {
186                self.writer.into_inner()
187            }
188        }
189
190        impl<W: AsyncWrite, D: DecodeV2> Decoder<W, D> {
191            fn do_poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
192                let mut this = self.project();
193
194                this.inner.do_poll_flush(cx, this.writer, this.decoder)
195            }
196        }
197
198        impl<W: AsyncWrite, D: DecodeV2> AsyncWrite for Decoder<W, D> {
199            fn poll_write(
200                self: Pin<&mut Self>,
201                cx: &mut Context<'_>,
202                buf: &[u8],
203            ) -> Poll<io::Result<usize>> {
204                let mut this = self.project();
205
206                this.inner.poll_write(cx, buf, this.writer, this.decoder)
207            }
208
209            fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
210                ready!(self.as_mut().do_poll_flush(cx))?;
211                self.project().writer.poll_flush(cx)
212            }
213
214            fn $poll_close(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
215                self.as_mut().project().inner.do_close();
216
217                ready!(self.as_mut().do_poll_flush(cx))?;
218
219                let this = self.project();
220                if this.inner.is_done() {
221                    this.writer.$poll_close(cx)
222                } else {
223                    Poll::Ready(Err(io::Error::other(
224                        "Attempt to close before finishing input",
225                    )))
226                }
227            }
228        }
229
230        impl<W: AsyncBufRead, D> AsyncBufRead for Decoder<W, D> {
231            fn poll_fill_buf(
232                self: Pin<&mut Self>,
233                cx: &mut Context<'_>,
234            ) -> Poll<io::Result<&[u8]>> {
235                self.get_pin_mut().poll_fill_buf(cx)
236            }
237
238            fn consume(self: Pin<&mut Self>, amt: usize) {
239                self.get_pin_mut().consume(amt)
240            }
241        }
242    };
243}
244pub(crate) use impl_decoder;