1use std::error::Error;
20use std::fmt;
21use std::io::{self};
22use std::pin::Pin;
23
24use async_compression::tokio::bufread::{BrotliDecoder, GzipDecoder, ZlibDecoder, ZstdDecoder};
25use bytes::Bytes;
26use futures::stream::Peekable;
27use futures::task::{Context, Poll};
28use futures::{Future, Stream};
29use futures_util::StreamExt;
30use headers::{ContentLength, HeaderMapExt};
31use http_body_util::BodyExt;
32use hyper::Response;
33use hyper::body::Body;
34use hyper::header::{CONTENT_ENCODING, HeaderValue, TRANSFER_ENCODING};
35use tokio_util::codec::{BytesCodec, FramedRead};
36use tokio_util::io::StreamReader;
37
38use crate::connector::BoxedBody;
39
40pub const DECODER_BUFFER_SIZE: usize = 8192;
41
42#[derive(Debug)]
44pub struct BodyStreamError(pub Box<dyn Error + Send + Sync>);
45
46impl fmt::Display for BodyStreamError {
47 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
48 self.0.fmt(f)
49 }
50}
51
52impl Error for BodyStreamError {
53 fn source(&self) -> Option<&(dyn Error + 'static)> {
54 Some(self.0.as_ref())
55 }
56}
57
58pub fn map_decode_error(err: io::Error) -> io::Error {
61 if err.kind() == io::ErrorKind::InvalidData {
62 return err;
63 }
64
65 let mut source: Option<&(dyn Error + 'static)> = err.get_ref().map(|e| e as _);
66 while let Some(e) = source {
67 if e.is::<BodyStreamError>() {
68 return err;
69 }
70
71 source = match e.downcast_ref::<io::Error>() {
72 Some(io_error) => io_error.get_ref().map(|e| e as _),
73 None => e.source(),
74 };
75 }
76 io::Error::new(io::ErrorKind::InvalidData, err)
77}
78
79pub struct Decoder {
83 inner: Inner,
84}
85
86#[derive(PartialEq)]
87enum DecoderType {
88 Gzip,
89 Brotli,
90 Deflate,
91 Zstd,
92}
93
94enum Inner {
95 PlainText(BodyStream),
97 Gzip(FramedRead<GzipDecoder<StreamReader<Peekable<BodyStream>, Bytes>>, BytesCodec>),
99 Deflate(FramedRead<ZlibDecoder<StreamReader<Peekable<BodyStream>, Bytes>>, BytesCodec>),
101 Brotli(FramedRead<BrotliDecoder<StreamReader<Peekable<BodyStream>, Bytes>>, BytesCodec>),
103 Zstd(FramedRead<ZstdDecoder<StreamReader<Peekable<BodyStream>, Bytes>>, BytesCodec>),
105 Pending(Pending),
107}
108
109struct Pending {
111 body: Peekable<BodyStream>,
112 type_: DecoderType,
113}
114
115impl fmt::Debug for Decoder {
116 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
117 f.debug_struct("Decoder").finish()
118 }
119}
120
121impl Decoder {
122 #[inline]
126 fn plain_text(
127 body: BoxedBody,
128 is_secure_scheme: bool,
129 content_length: Option<ContentLength>,
130 ) -> Decoder {
131 Decoder {
132 inner: Inner::PlainText(BodyStream::new(body, is_secure_scheme, content_length)),
133 }
134 }
135
136 #[inline]
140 fn pending(
141 body: BoxedBody,
142 type_: DecoderType,
143 is_secure_scheme: bool,
144 content_length: Option<ContentLength>,
145 ) -> Decoder {
146 Decoder {
147 inner: Inner::Pending(Pending {
148 body: BodyStream::new(body, is_secure_scheme, content_length).peekable(),
149 type_,
150 }),
151 }
152 }
153
154 pub fn is_encoded(&self) -> bool {
156 !matches!(self.inner, Inner::PlainText(_))
157 }
158
159 pub fn detect(response: Response<BoxedBody>, is_secure_scheme: bool) -> Response<Decoder> {
166 let values = response
167 .headers()
168 .get_all(CONTENT_ENCODING)
169 .iter()
170 .chain(response.headers().get_all(TRANSFER_ENCODING).iter());
171 let decoder = values.fold(None, |acc, enc| {
172 acc.or_else(|| {
173 if enc == HeaderValue::from_static("gzip") {
174 Some(DecoderType::Gzip)
175 } else if enc == HeaderValue::from_static("br") {
176 Some(DecoderType::Brotli)
177 } else if enc == HeaderValue::from_static("deflate") {
178 Some(DecoderType::Deflate)
179 } else if enc == HeaderValue::from_static("zstd") {
180 Some(DecoderType::Zstd)
181 } else {
182 None
183 }
184 })
185 });
186 let content_length = response.headers().typed_get::<ContentLength>();
187 match decoder {
188 Some(type_) => {
189 response.map(|r| Decoder::pending(r, type_, is_secure_scheme, content_length))
190 },
191 None => response.map(|r| Decoder::plain_text(r, is_secure_scheme, content_length)),
192 }
193 }
194}
195
196impl Stream for Decoder {
197 type Item = Result<Bytes, io::Error>;
198
199 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
200 match self.inner {
202 Inner::Pending(ref mut future) => match futures_core::ready!(Pin::new(future).poll(cx))
203 {
204 Ok(inner) => {
205 self.inner = inner;
206 self.poll_next(cx)
207 },
208 Err(e) => Poll::Ready(Some(Err(e))),
209 },
210 Inner::PlainText(ref mut body) => Pin::new(body).poll_next(cx),
211 Inner::Gzip(ref mut decoder) => {
212 match futures_core::ready!(Pin::new(decoder).poll_next(cx)) {
213 Some(Ok(bytes)) => Poll::Ready(Some(Ok(bytes.freeze()))),
214 Some(Err(err)) => Poll::Ready(Some(Err(map_decode_error(err)))),
215 None => Poll::Ready(None),
216 }
217 },
218 Inner::Brotli(ref mut decoder) => {
219 match futures_core::ready!(Pin::new(decoder).poll_next(cx)) {
220 Some(Ok(bytes)) => Poll::Ready(Some(Ok(bytes.freeze()))),
221 Some(Err(err)) => Poll::Ready(Some(Err(map_decode_error(err)))),
222 None => Poll::Ready(None),
223 }
224 },
225 Inner::Deflate(ref mut decoder) => {
226 match futures_core::ready!(Pin::new(decoder).poll_next(cx)) {
227 Some(Ok(bytes)) => Poll::Ready(Some(Ok(bytes.freeze()))),
228 Some(Err(err)) => Poll::Ready(Some(Err(map_decode_error(err)))),
229 None => Poll::Ready(None),
230 }
231 },
232 Inner::Zstd(ref mut decoder) => {
233 match futures_core::ready!(Pin::new(decoder).poll_next(cx)) {
234 Some(Ok(bytes)) => Poll::Ready(Some(Ok(bytes.freeze()))),
235 Some(Err(err)) => Poll::Ready(Some(Err(map_decode_error(err)))),
236 None => Poll::Ready(None),
237 }
238 },
239 }
240 }
241}
242
243impl Future for Pending {
244 type Output = Result<Inner, io::Error>;
245
246 fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
247 match futures_core::ready!(Pin::new(&mut self.body).poll_peek(cx)) {
248 Some(Ok(_)) => {
249 },
251 Some(Err(_e)) => {
252 return Poll::Ready(Err(futures_core::ready!(
254 Pin::new(&mut self.body).poll_next(cx)
255 )
256 .expect("just peeked Some")
257 .unwrap_err()));
258 },
259 None => return Poll::Ready(Ok(Inner::PlainText(BodyStream::empty()))),
260 };
261
262 let body = std::mem::replace(&mut self.body, BodyStream::empty().peekable());
263
264 match self.type_ {
265 DecoderType::Brotli => Poll::Ready(Ok(Inner::Brotli(FramedRead::with_capacity(
266 BrotliDecoder::new(StreamReader::new(body)),
267 BytesCodec::new(),
268 DECODER_BUFFER_SIZE,
269 )))),
270 DecoderType::Gzip => Poll::Ready(Ok(Inner::Gzip(FramedRead::with_capacity(
271 GzipDecoder::new(StreamReader::new(body)),
272 BytesCodec::new(),
273 DECODER_BUFFER_SIZE,
274 )))),
275 DecoderType::Deflate => Poll::Ready(Ok(Inner::Deflate(FramedRead::with_capacity(
276 ZlibDecoder::new(StreamReader::new(body)),
277 BytesCodec::new(),
278 DECODER_BUFFER_SIZE,
279 )))),
280 DecoderType::Zstd => Poll::Ready(Ok(Inner::Zstd(FramedRead::with_capacity(
281 ZstdDecoder::new(StreamReader::new(body)),
282 BytesCodec::new(),
283 DECODER_BUFFER_SIZE,
284 )))),
285 }
286 }
287}
288
289struct BodyStream {
290 body: BoxedBody,
291 is_secure_scheme: bool,
292 content_length: Option<ContentLength>,
293 total_read: u64,
294}
295
296impl BodyStream {
297 fn empty() -> Self {
298 BodyStream {
299 body: http_body_util::Empty::new()
300 .map_err(|_| unreachable!())
301 .boxed(),
302 is_secure_scheme: false,
303 content_length: None,
304 total_read: 0,
305 }
306 }
307
308 fn new(body: BoxedBody, is_secure_scheme: bool, content_length: Option<ContentLength>) -> Self {
309 BodyStream {
310 body,
311 is_secure_scheme,
312 content_length,
313 total_read: 0,
314 }
315 }
316}
317
318impl Stream for BodyStream {
319 type Item = Result<Bytes, io::Error>;
320
321 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Option<Self::Item>> {
322 match futures_core::ready!(Pin::new(&mut self.body).poll_frame(cx)) {
323 Some(Ok(bytes)) => {
324 let Ok(bytes) = bytes.into_data() else {
325 return Poll::Ready(None);
326 };
327 self.total_read += bytes.len() as u64;
328 Poll::Ready(Some(Ok(bytes)))
329 },
330 Some(Err(err)) => {
331 let all_content_read = self.content_length.is_some_and(|c| c.0 == self.total_read);
338 if self.is_secure_scheme && all_content_read {
339 let source = err.source();
340 let is_unexpected_eof = source
341 .and_then(|e| e.downcast_ref::<io::Error>())
342 .is_some_and(|e| e.kind() == io::ErrorKind::UnexpectedEof);
343 if is_unexpected_eof {
344 return Poll::Ready(None);
345 }
346 }
347 Poll::Ready(Some(Err(io::Error::other(BodyStreamError(err.into())))))
348 },
349 None => Poll::Ready(None),
350 }
351 }
352}