async_compression/generic/write/
decoder.rs1use 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;