Skip to main content

net/
decoder.rs

1/* This Source Code Form is subject to the terms of the Mozilla Public
2 * License, v. 2.0. If a copy of the MPL was not distributed with this
3 * file, You can obtain one at https://mozilla.org/MPL/2.0/. */
4
5//! Adapted from an implementation in reqwest.
6
7/*!
8A non-blocking response decoder.
9
10The decoder wraps a stream of bytes and produces a new stream of decompressed bytes.
11The decompressed bytes aren't guaranteed to align to the compressed ones.
12
13If the response is plaintext then no additional work is carried out.
14Bytes are just passed along.
15
16If the response is gzip, deflate or brotli then the bytes are decompressed.
17*/
18
19use 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/// Marker wrapper for errors that originate from the network body stream
43#[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
58/// Normalize errors produced by the decompressors to `ErrorKind::InvalidData`
59/// so that `http_loader` reports them as `NetworkError::DecompressionError`.
60pub 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
79/// A response decompressor over a non-blocking stream of bytes.
80///
81/// The inner decoder may be constructed asynchronously.
82pub 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    /// A `PlainText` decoder just returns the response content as is.
96    PlainText(BodyStream),
97    /// A `Gzip` decoder will uncompress the gzipped response content before returning it.
98    Gzip(FramedRead<GzipDecoder<StreamReader<Peekable<BodyStream>, Bytes>>, BytesCodec>),
99    /// A `Delfate` decoder will uncompress the inflated response content before returning it.
100    Deflate(FramedRead<ZlibDecoder<StreamReader<Peekable<BodyStream>, Bytes>>, BytesCodec>),
101    /// A `Brotli` decoder will uncompress the brotli-encoded response content before returning it.
102    Brotli(FramedRead<BrotliDecoder<StreamReader<Peekable<BodyStream>, Bytes>>, BytesCodec>),
103    /// A `Zstd` decoder will uncompress the zstd-encoded response content before returning it.
104    Zstd(FramedRead<ZstdDecoder<StreamReader<Peekable<BodyStream>, Bytes>>, BytesCodec>),
105    /// A decoder that doesn't have a value yet.
106    Pending(Pending),
107}
108
109/// A future attempt to poll the response body for EOF so we know whether to use gzip or not.
110struct 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    /// A plain text decoder.
123    ///
124    /// This decoder will emit the underlying bytes as-is.
125    #[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    /// A pending decoder.
137    ///
138    /// This decoder will buffer and decompress bytes that are encoded in the expected format.
139    #[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    /// Returns true if the content is encoded
155    pub fn is_encoded(&self) -> bool {
156        !matches!(self.inner, Inner::PlainText(_))
157    }
158
159    /// Constructs a Decoder from a hyper response.
160    ///
161    /// A decoder is just a wrapper around the hyper response that knows
162    /// how to decode the content body of the response.
163    ///
164    /// Uses the correct variant by inspecting the Content-Encoding header.
165    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        // Do a read or poll for a pending decoder value.
201        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                // fallthrough
250            },
251            Some(Err(_e)) => {
252                // error was just a ref, so we need to really poll to move it
253                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                // To prevent truncation attacks rustls treats close connection without a close_notify as
332                // an error of type std::io::Error with ErrorKind::UnexpectedEof.
333                // https://docs.rs/rustls/latest/rustls/manual/_03_howto/index.html#unexpected-eof
334                //
335                // The error can be safely ignored if we known that all content was received or is explicitly
336                // set in preferences.
337                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}