Skip to main content

ntex/http/encoding/
decoder.rs

1use std::sync::atomic::{AtomicUsize, Ordering};
2use std::{collections::VecDeque, future::Future, io, pin::Pin, task::Context, task::Poll};
3
4use flate2::{Crc, Decompress, FlushDecompress, Status};
5use zstd::zstd_safe::{DCtx, DParameter, InBuffer, OutBuffer};
6
7use super::{Spare, offload, zstd_error};
8use crate::http::error::PayloadError;
9use crate::http::header::{CONTENT_ENCODING, ContentEncoding, HeaderMap};
10use crate::rt::BlockingResult;
11use crate::util::{BufMut, BytePageSize, Bytes, BytesMut, Stream};
12
13/// Decoders write into pages of this size, so decoded chunks are never larger.
14#[cfg(test)]
15const MAX_CHUNK_SIZE: usize = 32 * 1024;
16
17/// A blocking task stops decoding once its output reaches this size.
18///
19/// The output is kept in chunks of at most 32KiB, a single page.
20const MAX_TASK_OUTPUT: usize = 256 * 1024;
21
22/// The largest `zstd` window a decoder accepts, 8MiB as required by RFC 9659.
23const ZSTD_WINDOW_LOG_MAX: u32 = 23;
24
25/// Memory all `zstd` decoders of the process may use beyond `ZSTD_FREE_MEMORY` each.
26///
27/// RFC 9659 requires 8MiB windows, so a request can make its decoder allocate
28/// about 9MiB. A frame that would take the total over this limit is rejected.
29const ZSTD_MEMORY_LIMIT: usize = 512 * 1024 * 1024;
30
31/// Memory a `zstd` decoder uses without being counted, enough for a 512KiB window.
32const ZSTD_FREE_MEMORY: usize = 1024 * 1024;
33
34static ZSTD_MEMORY: ZstdMemory = ZstdMemory::new(ZSTD_MEMORY_LIMIT);
35
36/// Payload stream decoder.
37///
38/// Decompresses a stream of payload chunks. `gzip`, `deflate` and `zstd` are
39/// decoded; other encodings pass the stream through unchanged.
40#[derive(derive_more::Debug)]
41pub struct Decoder<S> {
42    #[debug(skip)]
43    inner: Option<ContentDecoder>,
44    stream: S,
45    eof: bool,
46    /// The stream is decoded
47    decode: bool,
48    /// Input that is not decoded yet
49    #[debug(skip)]
50    pending: Option<Bytes>,
51    /// Output of a blocking task that is not returned yet
52    #[debug(skip)]
53    ready: VecDeque<Bytes>,
54    #[debug(skip)]
55    fut: Option<BlockingResult<DecodeResult>>,
56}
57
58type DecodeResult = Result<(VecDeque<Bytes>, ContentDecoder, Bytes), io::Error>;
59
60impl<S> Decoder<S>
61where
62    S: Stream<Item = Result<Bytes, PayloadError>>,
63{
64    /// Construct a decoder for the given content encoding.
65    #[inline]
66    pub fn new(stream: S, encoding: ContentEncoding) -> Decoder<S> {
67        let inner = match encoding {
68            ContentEncoding::Deflate => {
69                Some(ContentDecoder::Flate(Box::new(FlateDecoder::new(false))))
70            }
71            ContentEncoding::Gzip => Some(ContentDecoder::Flate(Box::new(FlateDecoder::new(true)))),
72            ContentEncoding::Zstd => ZstdDecoder::new()
73                .inspect_err(|err| log::error!("Cannot create zstd decoder: {err}"))
74                .ok()
75                .map(|d| ContentDecoder::Zstd(Box::new(d))),
76            _ => None,
77        };
78        Decoder {
79            decode: inner.is_some(),
80            inner,
81            stream,
82            fut: None,
83            eof: false,
84            pending: None,
85            ready: VecDeque::new(),
86        }
87    }
88
89    /// Returns `true` if the stream is decoded.
90    pub(crate) fn is_decoding(&self) -> bool {
91        self.decode
92    }
93
94    /// Construct decoder based on the `Content-Encoding` header.
95    ///
96    /// A missing or invalid header selects the `Identity` encoding.
97    #[inline]
98    pub fn from_headers(stream: S, headers: &HeaderMap) -> Decoder<S> {
99        let encoding = headers
100            .get(&CONTENT_ENCODING)
101            .and_then(|enc| enc.to_str().ok())
102            .map_or(ContentEncoding::Identity, ContentEncoding::from);
103        Self::new(stream, encoding)
104    }
105}
106
107impl<S> Stream for Decoder<S>
108where
109    S: Stream<Item = Result<Bytes, PayloadError>> + Unpin,
110{
111    type Item = Result<Bytes, PayloadError>;
112
113    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
114        let result = self.poll_decoded(cx);
115        if let Poll::Ready(Some(Err(_))) = result
116            && self.decode
117            && self.inner.is_none()
118        {
119            // the decoder state is lost, the stream must not continue with raw data
120            self.eof = true;
121            self.fut = None;
122            self.pending = None;
123        }
124        result
125    }
126}
127
128impl<S> Decoder<S>
129where
130    S: Stream<Item = Result<Bytes, PayloadError>> + Unpin,
131{
132    /// Puts the decoder back after a feed, with the input it left.
133    fn restore(&mut self, decoder: ContentDecoder, rest: Bytes) {
134        if !rest.is_empty() || decoder.has_more() {
135            self.pending = Some(rest);
136        }
137        self.inner = Some(decoder);
138    }
139
140    fn poll_decoded(&mut self, cx: &mut Context<'_>) -> Poll<Option<Result<Bytes, PayloadError>>> {
141        loop {
142            if let Some(chunk) = self.ready.pop_front() {
143                return Poll::Ready(Some(Ok(chunk)));
144            }
145
146            if let Some(ref mut fut) = self.fut {
147                let (chunks, decoder, rest) = match Pin::new(fut).poll(cx) {
148                    Poll::Ready(Ok(Ok(item))) => item,
149                    Poll::Ready(Ok(Err(e))) => return Poll::Ready(Some(Err(e.into()))),
150                    Poll::Ready(Err(e)) => return Poll::Ready(Some(Err(e.into()))),
151                    Poll::Pending => return Poll::Pending,
152                };
153                self.fut = None;
154                self.restore(decoder, rest);
155                self.ready = chunks;
156                continue;
157            }
158
159            if let Some(mut data) = self.pending.take() {
160                let mut decoder = self.inner.take().unwrap();
161                if data.len() < decoder.limit() {
162                    let chunk = decoder.feed_data(&mut data)?;
163                    self.restore(decoder, data);
164                    if let Some(chunk) = chunk {
165                        return Poll::Ready(Some(Ok(chunk)));
166                    }
167                } else {
168                    self.fut = Some(offload(move || {
169                        let chunks = decoder.feed_task(&mut data)?;
170                        Ok((chunks, decoder, data))
171                    }));
172                }
173                continue;
174            }
175
176            if self.eof {
177                return Poll::Ready(None);
178            }
179
180            match Pin::new(&mut self.stream).poll_next(cx) {
181                Poll::Ready(Some(Err(err))) => return Poll::Ready(Some(Err(err))),
182                Poll::Ready(Some(Ok(chunk))) => {
183                    if self.inner.is_some() {
184                        if !chunk.is_empty() {
185                            self.pending = Some(chunk);
186                        }
187                        continue;
188                    }
189                    return Poll::Ready(Some(Ok(chunk)));
190                }
191                Poll::Ready(None) => {
192                    self.eof = true;
193                    if let Some(decoder) = self.inner.take() {
194                        decoder.finish()?;
195                    }
196                    return Poll::Ready(None);
197                }
198                Poll::Pending => return Poll::Pending,
199            }
200        }
201    }
202}
203
204enum ContentDecoder {
205    Flate(Box<FlateDecoder>),
206    Zstd(Box<ZstdDecoder>),
207}
208
209impl ContentDecoder {
210    /// Input of this size and larger is decoded on the blocking thread pool.
211    ///
212    /// A feed on the current thread decodes a single page of output, a task
213    /// on the pool decodes several pages, up to `MAX_TASK_OUTPUT`.
214    const fn limit(&self) -> usize {
215        match self {
216            ContentDecoder::Flate(_) => 128 * 1024,
217            ContentDecoder::Zstd(_) => 512 * 1024,
218        }
219    }
220
221    /// Returns `true` if decoded output is left without new input.
222    fn has_more(&self) -> bool {
223        match self {
224            ContentDecoder::Flate(decoder) => decoder.more,
225            ContentDecoder::Zstd(decoder) => decoder.more,
226        }
227    }
228
229    /// Checks that the compressed stream is complete.
230    fn finish(&self) -> io::Result<()> {
231        match self {
232            ContentDecoder::Flate(decoder) => decoder.finish(),
233            ContentDecoder::Zstd(decoder) => decoder.finish(),
234        }
235    }
236
237    /// Decodes `data` on a blocking task until the output reaches `MAX_TASK_OUTPUT`.
238    fn feed_task(&mut self, data: &mut Bytes) -> io::Result<VecDeque<Bytes>> {
239        let mut chunks = VecDeque::new();
240        let mut size = 0;
241        while size < MAX_TASK_OUTPUT && !data.is_empty() {
242            let Some(chunk) = self.feed_data(data)? else {
243                break;
244            };
245            size += chunk.len();
246            chunks.push_back(chunk);
247        }
248        Ok(chunks)
249    }
250
251    /// Decodes `data` until a page of output is full.
252    ///
253    /// Decoded input is removed from `data`.
254    fn feed_data(&mut self, data: &mut Bytes) -> io::Result<Option<Bytes>> {
255        match self {
256            ContentDecoder::Flate(decoder) => decoder.feed(data),
257            ContentDecoder::Zstd(decoder) => decoder.feed(data),
258        }
259    }
260}
261
262/// `deflate` (zlib) or `gzip` decoder.
263struct FlateDecoder {
264    inner: Decompress,
265    buf: BytesMut,
266    /// The `gzip` header and trailer, `None` for `deflate`
267    gzip: Option<Gzip>,
268    /// The compressed data is complete
269    done: bool,
270    /// The output buffer was filled, the decoder may hold more output
271    more: bool,
272}
273
274impl FlateDecoder {
275    fn new(gzip: bool) -> Self {
276        FlateDecoder {
277            inner: Decompress::new(!gzip),
278            buf: BytesMut::with_page_size(BytePageSize::Size32),
279            gzip: gzip.then(Gzip::new),
280            done: false,
281            more: false,
282        }
283    }
284
285    /// Decodes `data` until the output buffer is full.
286    fn feed(&mut self, data: &mut Bytes) -> io::Result<Option<Bytes>> {
287        if let Some(gzip) = &mut self.gzip
288            && !gzip.header(data)?
289        {
290            return Ok(None);
291        }
292
293        self.buf.reserve_more();
294        let start = self.buf.len();
295        while !self.done && self.buf.remaining_mut() > 0 && (!data.is_empty() || self.more) {
296            let (total_in, total_out) = (self.inner.total_in(), self.inner.total_out());
297            // SAFETY: the decoder only writes to the slice, and `advance_mut`
298            // covers just the bytes it has written
299            let status = unsafe {
300                let spare = self.buf.chunk_mut().as_uninit_slice_mut();
301                self.inner
302                    .decompress_uninit(data, spare, FlushDecompress::None)
303                    .map_err(invalid_data)?
304            };
305            let read = usize::try_from(self.inner.total_in() - total_in).unwrap();
306            let written = usize::try_from(self.inner.total_out() - total_out).unwrap();
307            unsafe { self.buf.advance_mut(written) };
308            data.advance_to(read);
309
310            self.done = status == Status::StreamEnd;
311            self.more = !self.done && self.buf.remaining_mut() == 0;
312            if read == 0 && written == 0 {
313                break;
314            }
315        }
316
317        if let Some(gzip) = &mut self.gzip {
318            gzip.crc.update(&self.buf[start..]);
319            if self.done {
320                gzip.trailer(data)?;
321            }
322        }
323        if self.done && !data.is_empty() {
324            return Err(invalid_data("data after the end of the stream"));
325        }
326
327        Ok((!self.buf.is_empty()).then(|| self.buf.take()))
328    }
329
330    fn finish(&self) -> io::Result<()> {
331        if self.done
332            && self
333                .gzip
334                .as_ref()
335                .is_none_or(|gzip| gzip.state == GzState::Done)
336        {
337            Ok(())
338        } else {
339            Err(io::Error::new(
340                io::ErrorKind::UnexpectedEof,
341                "compressed stream is incomplete",
342            ))
343        }
344    }
345}
346
347const FHCRC: u8 = 0x02;
348const FEXTRA: u8 = 0x04;
349const FNAME: u8 = 0x08;
350const FCOMMENT: u8 = 0x10;
351const FRESERVED: u8 = 0xe0;
352
353#[derive(Copy, Clone, Debug, PartialEq, Eq)]
354enum GzState {
355    Header,
356    ExtraLen,
357    Extra(usize),
358    Name,
359    Comment,
360    HeaderCrc,
361    Body,
362    Trailer,
363    Done,
364}
365
366/// Parser of the `gzip` member header and trailer, RFC 1952.
367///
368/// The optional header fields are skipped. Like most decoders, only a single
369/// member is accepted.
370struct Gzip {
371    state: GzState,
372    flags: u8,
373    /// Checksum of the header
374    header_crc: Crc,
375    /// Checksum and size of the decoded data
376    crc: Crc,
377    buf: [u8; 10],
378    len: usize,
379}
380
381impl Gzip {
382    fn new() -> Self {
383        Gzip {
384            state: GzState::Header,
385            flags: 0,
386            header_crc: Crc::new(),
387            crc: Crc::new(),
388            buf: [0; 10],
389            len: 0,
390        }
391    }
392
393    /// Moves up to `n` bytes of `data` to `buf`, returns `true` once it holds `n` bytes.
394    fn fill(&mut self, data: &mut Bytes, n: usize) -> bool {
395        let size = (n - self.len).min(data.len());
396        self.buf[self.len..self.len + size].copy_from_slice(&data[..size]);
397        self.len += size;
398        data.advance_to(size);
399        self.len == n
400    }
401
402    /// Moves to `state`, the bytes of the previous field are dropped.
403    fn next(&mut self, state: GzState) {
404        self.state = state;
405        self.len = 0;
406    }
407
408    /// Parses the header, returns `true` once it is complete.
409    fn header(&mut self, data: &mut Bytes) -> io::Result<bool> {
410        loop {
411            match self.state {
412                GzState::Header => {
413                    if !self.fill(data, 10) {
414                        return Ok(false);
415                    }
416                    let [id1, id2, method, flags, ..] = self.buf;
417                    if id1 != 0x1f || id2 != 0x8b {
418                        return Err(invalid_data("invalid gzip header"));
419                    }
420                    if method != 8 || flags & FRESERVED != 0 {
421                        return Err(invalid_data("unsupported gzip header"));
422                    }
423                    self.flags = flags;
424                    self.header_crc.update(&self.buf);
425                    self.next(GzState::ExtraLen);
426                }
427                GzState::ExtraLen if self.flags & FEXTRA == 0 => self.next(GzState::Name),
428                GzState::ExtraLen => {
429                    if !self.fill(data, 2) {
430                        return Ok(false);
431                    }
432                    let len = [self.buf[0], self.buf[1]];
433                    self.header_crc.update(&len);
434                    self.next(GzState::Extra(u16::from_le_bytes(len).into()));
435                }
436                GzState::Extra(len) => {
437                    let size = len.min(data.len());
438                    self.header_crc.update(&data[..size]);
439                    data.advance_to(size);
440                    if size < len {
441                        self.state = GzState::Extra(len - size);
442                        return Ok(false);
443                    }
444                    self.next(GzState::Name);
445                }
446                GzState::Name if self.flags & FNAME == 0 => self.next(GzState::Comment),
447                GzState::Comment if self.flags & FCOMMENT == 0 => self.next(GzState::HeaderCrc),
448                GzState::Name | GzState::Comment => {
449                    // a zero-terminated string
450                    let Some(end) = data.iter().position(|b| *b == 0) else {
451                        self.header_crc.update(data);
452                        data.advance_to(data.len());
453                        return Ok(false);
454                    };
455                    self.header_crc.update(&data[..=end]);
456                    data.advance_to(end + 1);
457                    self.next(if self.state == GzState::Name {
458                        GzState::Comment
459                    } else {
460                        GzState::HeaderCrc
461                    });
462                }
463                GzState::HeaderCrc if self.flags & FHCRC == 0 => self.next(GzState::Body),
464                GzState::HeaderCrc => {
465                    if !self.fill(data, 2) {
466                        return Ok(false);
467                    }
468                    // the low 16 bits of the crc32
469                    let crc = u16::from_le_bytes([self.buf[0], self.buf[1]]);
470                    if u32::from(crc) != self.header_crc.sum() & 0xffff {
471                        return Err(invalid_data("gzip header checksum mismatch"));
472                    }
473                    self.next(GzState::Body);
474                }
475                GzState::Body | GzState::Trailer | GzState::Done => return Ok(true),
476            }
477        }
478    }
479
480    /// Checks the trailer after the compressed data, the crc32 and size of the output.
481    fn trailer(&mut self, data: &mut Bytes) -> io::Result<()> {
482        if self.state == GzState::Body {
483            self.next(GzState::Trailer);
484        }
485        if self.state == GzState::Trailer && self.fill(data, 8) {
486            let [c0, c1, c2, c3, s0, s1, s2, s3, ..] = self.buf;
487            if u32::from_le_bytes([c0, c1, c2, c3]) != self.crc.sum() {
488                return Err(invalid_data("gzip checksum mismatch"));
489            }
490            // the size modulo 2^32
491            if u32::from_le_bytes([s0, s1, s2, s3]) != self.crc.amount() {
492                return Err(invalid_data("gzip size mismatch"));
493            }
494            self.next(GzState::Done);
495        }
496        Ok(())
497    }
498}
499
500fn invalid_data<E>(err: E) -> io::Error
501where
502    E: Into<Box<dyn std::error::Error + Send + Sync>>,
503{
504    io::Error::new(io::ErrorKind::InvalidData, err)
505}
506
507/// Memory used by `zstd` decoders.
508struct ZstdMemory {
509    used: AtomicUsize,
510    limit: usize,
511}
512
513impl ZstdMemory {
514    const fn new(limit: usize) -> Self {
515        ZstdMemory {
516            used: AtomicUsize::new(0),
517            limit,
518        }
519    }
520
521    fn acquire(&self, size: usize) -> io::Result<()> {
522        self.used
523            .try_update(Ordering::AcqRel, Ordering::Acquire, |used| {
524                used.checked_add(size).filter(|total| *total <= self.limit)
525            })
526            .map(|_| ())
527            .map_err(|_| {
528                io::Error::new(
529                    io::ErrorKind::OutOfMemory,
530                    "zstd decoders memory limit is reached",
531                )
532            })
533    }
534
535    fn release(&self, size: usize) {
536        self.used.fetch_sub(size, Ordering::AcqRel);
537    }
538}
539
540struct ZstdDecoder {
541    ctx: DCtx<'static>,
542    buf: BytesMut,
543    memory: &'static ZstdMemory,
544    /// Memory counted in `memory`
545    charged: usize,
546    /// The last frame is complete
547    done: bool,
548    /// The output buffer was filled, the decoder may hold more output
549    more: bool,
550}
551
552impl ZstdDecoder {
553    fn new() -> io::Result<Self> {
554        Self::with_memory(&ZSTD_MEMORY)
555    }
556
557    fn with_memory(memory: &'static ZstdMemory) -> io::Result<Self> {
558        let mut ctx = DCtx::try_create().ok_or(io::ErrorKind::OutOfMemory)?;
559        ctx.set_parameter(DParameter::WindowLogMax(ZSTD_WINDOW_LOG_MAX))
560            .map_err(zstd_error)?;
561        Ok(ZstdDecoder {
562            ctx,
563            memory,
564            charged: 0,
565            buf: BytesMut::with_page_size(BytePageSize::Size32),
566            done: false,
567            more: false,
568        })
569    }
570
571    /// Decodes `data` until the output buffer is full.
572    fn feed(&mut self, data: &mut Bytes) -> io::Result<Option<Bytes>> {
573        self.buf.reserve_more();
574        while self.buf.remaining_mut() > 0 && (!data.is_empty() || self.more) {
575            let mut src = InBuffer::around(data);
576            let mut spare = Spare::new(&mut self.buf);
577            let mut dst = OutBuffer::around(&mut spare);
578            // a new frame starts automatically after the previous one is complete
579            let hint = self
580                .ctx
581                .decompress_stream(&mut dst, &mut src)
582                .map_err(zstd_error)?;
583            let (read, written) = (src.pos(), dst.pos());
584            self.charge()?;
585            self.more = self.buf.remaining_mut() == 0;
586            if read == 0 && written == 0 {
587                break;
588            }
589            self.done = hint == 0;
590            data.advance_to(read);
591        }
592        Ok((!self.buf.is_empty()).then(|| self.buf.take()))
593    }
594
595    /// Counts the memory of the context, it grows with the window of a frame.
596    fn charge(&mut self) -> io::Result<()> {
597        let size = self.ctx.sizeof().saturating_sub(ZSTD_FREE_MEMORY);
598        if size > self.charged {
599            self.memory.acquire(size - self.charged)?;
600        } else {
601            self.memory.release(self.charged - size);
602        }
603        self.charged = size;
604        Ok(())
605    }
606
607    fn finish(&self) -> io::Result<()> {
608        if self.done {
609            Ok(())
610        } else {
611            Err(io::Error::new(
612                io::ErrorKind::UnexpectedEof,
613                "zstd stream is incomplete",
614            ))
615        }
616    }
617}
618
619impl Drop for ZstdDecoder {
620    fn drop(&mut self) {
621        self.memory.release(self.charged);
622    }
623}
624
625#[cfg(test)]
626mod tests {
627    use std::io::Write;
628
629    use flate2::{Compression, write::GzEncoder, write::ZlibEncoder};
630    use futures_util::stream::{self, StreamExt};
631
632    use super::*;
633
634    const BOMB_SIZE: usize = 4 * 1024 * 1024;
635
636    fn bomb(encoding: ContentEncoding) -> Vec<u8> {
637        let data = vec![0u8; BOMB_SIZE];
638        if encoding == ContentEncoding::Gzip {
639            let mut e = GzEncoder::new(Vec::new(), Compression::best());
640            e.write_all(&data).unwrap();
641            e.finish().unwrap()
642        } else {
643            let mut e = ZlibEncoder::new(Vec::new(), Compression::best());
644            e.write_all(&data).unwrap();
645            e.finish().unwrap()
646        }
647    }
648
649    #[crate::rt_test]
650    async fn decoded_chunks_are_bounded() {
651        for encoding in [ContentEncoding::Gzip, ContentEncoding::Deflate] {
652            let compressed = bomb(encoding);
653            assert!(compressed.len() < 32 * 1024);
654
655            for size in [compressed.len(), 2048] {
656                let chunks: Vec<_> = compressed
657                    .chunks(size)
658                    .map(|c| Ok::<_, PayloadError>(Bytes::copy_from_slice(c)))
659                    .collect();
660                let mut decoder = Decoder::new(stream::iter(chunks), encoding);
661
662                let mut total = 0;
663                let mut max = 0;
664                while let Some(chunk) = decoder.next().await {
665                    let chunk = chunk.unwrap();
666                    assert!(chunk.iter().all(|b| *b == 0));
667                    max = max.max(chunk.len());
668                    total += chunk.len();
669                }
670                assert_eq!(total, BOMB_SIZE);
671                assert!(max <= MAX_CHUNK_SIZE, "{encoding:?} chunk of {max} bytes");
672            }
673        }
674    }
675
676    #[crate::rt_test]
677    async fn decoder_is_fused_after_error() {
678        let chunks = vec![
679            Ok::<_, PayloadError>(Bytes::from_static(b"not gzip data")),
680            Ok(Bytes::from_static(b"raw")),
681        ];
682        let mut decoder = Decoder::new(stream::iter(chunks), ContentEncoding::Gzip);
683        assert!(matches!(decoder.next().await, Some(Err(_))));
684        assert!(decoder.next().await.is_none());
685    }
686
687    #[crate::rt_test]
688    async fn decoder_from_headers() {
689        use crate::http::header::HeaderValue;
690
691        let mut headers = HeaderMap::new();
692        headers.insert(
693            CONTENT_ENCODING,
694            HeaderValue::from_bytes(b"gzip\xff").unwrap(),
695        );
696        let chunks = vec![Ok::<_, PayloadError>(Bytes::from_static(b"raw"))];
697        let mut decoder = Decoder::from_headers(stream::iter(chunks), &headers);
698        assert!(!decoder.is_decoding());
699        assert_eq!(decoder.next().await.unwrap().unwrap(), "raw");
700        assert!(decoder.next().await.is_none());
701    }
702
703    #[crate::rt_test]
704    async fn decoder_error_on_blocking_pool() {
705        let before = super::super::offloaded();
706        let chunks = vec![Ok::<_, PayloadError>(Bytes::from(vec![b'x'; 256 * 1024]))];
707        let mut decoder = Decoder::new(stream::iter(chunks), ContentEncoding::Deflate);
708        assert!(matches!(decoder.next().await, Some(Err(_))));
709        assert!(decoder.next().await.is_none());
710        assert_eq!(super::super::offloaded(), before + 1);
711    }
712
713    fn compress(encoding: ContentEncoding, data: &[u8]) -> Vec<u8> {
714        match encoding {
715            ContentEncoding::Gzip => {
716                let mut e = GzEncoder::new(Vec::new(), Compression::fast());
717                e.write_all(data).unwrap();
718                e.finish().unwrap()
719            }
720            ContentEncoding::Deflate => {
721                let mut e = ZlibEncoder::new(Vec::new(), Compression::fast());
722                e.write_all(data).unwrap();
723                e.finish().unwrap()
724            }
725            ContentEncoding::Zstd => zstd::encode_all(data, 0).unwrap(),
726            _ => unreachable!(),
727        }
728    }
729
730    async fn decode(
731        encoding: ContentEncoding,
732        chunks: Vec<Bytes>,
733    ) -> Result<Vec<u8>, PayloadError> {
734        let chunks = chunks.into_iter().map(Ok::<_, PayloadError>);
735        let mut decoder = Decoder::new(stream::iter(chunks), encoding);
736        let mut out = Vec::new();
737        while let Some(chunk) = decoder.next().await {
738            let chunk = chunk?;
739            assert!(!chunk.is_empty());
740            assert!(chunk.len() <= MAX_CHUNK_SIZE);
741            out.extend_from_slice(&chunk);
742        }
743        Ok(out)
744    }
745
746    /// Incompressible data, so the compressed size is close to `len`.
747    fn random(len: usize) -> Vec<u8> {
748        let mut x = 0x2545_f491_4f6c_dd1d_u64;
749        (0..len)
750            .map(|_| {
751                x ^= x << 13;
752                x ^= x >> 7;
753                x ^= x << 17;
754                x as u8
755            })
756            .collect()
757    }
758
759    #[crate::rt_test]
760    async fn decoder_offloads_large_input() {
761        for (encoding, limit) in [
762            (ContentEncoding::Gzip, 128 * 1024),
763            (ContentEncoding::Deflate, 128 * 1024),
764            (ContentEncoding::Zstd, 512 * 1024),
765        ] {
766            let data = random(limit + 1024);
767            let compressed = compress(encoding, &data);
768            assert!(compressed.len() > limit);
769
770            // input below the limit is decoded in place
771            let before = super::super::offloaded();
772            let chunks = compressed
773                .chunks(limit - 1)
774                .map(Bytes::copy_from_slice)
775                .collect();
776            assert_eq!(decode(encoding, chunks).await.unwrap(), data);
777            assert_eq!(super::super::offloaded(), before, "{encoding:?}");
778
779            // the first feed of a larger chunk is offloaded
780            let chunks = compressed[..limit]
781                .chunks(limit)
782                .chain(compressed[limit..].chunks(limit))
783                .map(Bytes::copy_from_slice)
784                .collect();
785            assert_eq!(decode(encoding, chunks).await.unwrap(), data);
786            assert_eq!(super::super::offloaded(), before + 1, "{encoding:?}");
787        }
788    }
789
790    #[crate::rt_test]
791    async fn zstd_decoder() {
792        let data = random(300 * 1024);
793        let compressed = zstd::encode_all(&data[..], 0).unwrap();
794        for size in [1, 100, 64 * 1024, compressed.len()] {
795            let chunks = compressed
796                .chunks(size)
797                .map(Bytes::copy_from_slice)
798                .collect();
799            assert_eq!(decode(ContentEncoding::Zstd, chunks).await.unwrap(), data);
800        }
801
802        // concatenated frames
803        let mut frames = zstd::encode_all(&b"hello "[..], 0).unwrap();
804        frames.extend(zstd::encode_all(&b"world"[..], 0).unwrap());
805        let out = decode(ContentEncoding::Zstd, vec![Bytes::from(frames)]).await;
806        assert_eq!(out.unwrap(), b"hello world");
807    }
808
809    #[crate::rt_test]
810    async fn zstd_decoder_drains_output_without_input() {
811        use zstd::stream::raw::Operation;
812
813        let data = b"hello world ".repeat(50_000);
814        let compressed = zstd::encode_all(&data[..], 0).unwrap();
815
816        // find the input byte that completes the first block
817        let mut raw = zstd::stream::raw::Decoder::new().unwrap();
818        let mut out = vec![0; data.len()];
819        let mut pos = 0;
820        let mut split = 0;
821        for (i, b) in compressed.iter().enumerate() {
822            let mut src = InBuffer::around(std::slice::from_ref(b));
823            let mut dst = OutBuffer::around_pos(&mut out[..], pos);
824            raw.run(&mut src, &mut dst).unwrap();
825            if dst.pos() - pos > MAX_CHUNK_SIZE {
826                pos = dst.pos();
827                split = i + 1;
828                break;
829            }
830            pos = dst.pos();
831        }
832        assert!(split > 0 && split < compressed.len());
833
834        // the stream stalls after the block, its output must not wait for more input
835        let chunks = compressed[..split]
836            .chunks(1)
837            .map(|c| Ok::<_, PayloadError>(Bytes::copy_from_slice(c)))
838            .collect::<Vec<_>>();
839        let stream = stream::iter(chunks).chain(stream::pending());
840        let mut decoder = Decoder::new(stream, ContentEncoding::Zstd);
841
842        let mut decoded = Vec::new();
843        while decoded.len() < pos {
844            let chunk = crate::time::timeout(crate::time::Millis(5_000), decoder.next())
845                .await
846                .expect("decoder is stuck")
847                .unwrap()
848                .unwrap();
849            assert!(chunk.len() <= MAX_CHUNK_SIZE);
850            decoded.extend_from_slice(&chunk);
851        }
852        assert_eq!(decoded, data[..pos]);
853    }
854
855    #[crate::rt_test]
856    async fn zstd_decoded_chunks_are_bounded() {
857        let compressed = zstd::encode_all(&vec![0u8; BOMB_SIZE][..], 19).unwrap();
858        assert!(compressed.len() < 1024);
859
860        for size in [1, compressed.len()] {
861            let chunks: Vec<_> = compressed
862                .chunks(size)
863                .map(|c| Ok::<_, PayloadError>(Bytes::copy_from_slice(c)))
864                .collect();
865            let mut decoder = Decoder::new(stream::iter(chunks), ContentEncoding::Zstd);
866            let mut total = 0;
867            while let Some(chunk) = decoder.next().await {
868                let chunk = chunk.unwrap();
869                assert!(chunk.iter().all(|b| *b == 0));
870                assert!(chunk.len() <= BytePageSize::Size32.capacity());
871                total += chunk.len();
872            }
873            assert_eq!(total, BOMB_SIZE);
874        }
875    }
876
877    #[crate::rt_test]
878    async fn zstd_decoder_errors() {
879        let compressed = zstd::encode_all(&random(1024)[..], 0).unwrap();
880
881        // truncated
882        let chunks = vec![Bytes::copy_from_slice(&compressed[..compressed.len() - 4])];
883        assert!(decode(ContentEncoding::Zstd, chunks).await.is_err());
884
885        // empty body
886        assert!(decode(ContentEncoding::Zstd, vec![]).await.is_err());
887
888        // corrupted
889        let out = decode(ContentEncoding::Zstd, vec![Bytes::from_static(b"not zstd")]).await;
890        assert!(out.is_err());
891
892        // a window above 8MiB
893        let mut e = zstd::stream::write::Encoder::new(Vec::new(), 0).unwrap();
894        e.set_parameter(zstd::stream::raw::CParameter::WindowLog(24))
895            .unwrap();
896        e.write_all(b"hello").unwrap();
897        let large = e.finish().unwrap();
898        let out = decode(ContentEncoding::Zstd, vec![Bytes::from(large)]).await;
899        assert!(out.is_err());
900
901        // the same window size is accepted at 8MiB
902        let mut e = zstd::stream::write::Encoder::new(Vec::new(), 0).unwrap();
903        e.set_parameter(zstd::stream::raw::CParameter::WindowLog(23))
904            .unwrap();
905        e.write_all(b"hello").unwrap();
906        let ok = e.finish().unwrap();
907        let out = decode(ContentEncoding::Zstd, vec![Bytes::from(ok)]).await;
908        assert_eq!(out.unwrap(), b"hello");
909    }
910
911    /// A frame with an unknown content size, its decoder allocates the whole window.
912    fn zstd_frame(window_log: u32, data: &[u8]) -> Vec<u8> {
913        let mut e = zstd::stream::write::Encoder::new(Vec::new(), 0).unwrap();
914        e.set_parameter(zstd::stream::raw::CParameter::WindowLog(window_log))
915            .unwrap();
916        e.write_all(data).unwrap();
917        e.finish().unwrap()
918    }
919
920    fn zstd_feed(decoder: &mut ZstdDecoder, frame: &[u8]) -> io::Result<Vec<u8>> {
921        let mut data = Bytes::copy_from_slice(frame);
922        let mut out = Vec::new();
923        while !data.is_empty() || decoder.more {
924            if let Some(chunk) = decoder.feed(&mut data)? {
925                out.extend_from_slice(&chunk);
926            }
927        }
928        Ok(out)
929    }
930
931    #[test]
932    fn zstd_decoder_memory_is_limited() {
933        static MEMORY: ZstdMemory = ZstdMemory::new(20 * 1024 * 1024);
934
935        let large = zstd_frame(23, b"hello");
936        let mut first = ZstdDecoder::with_memory(&MEMORY).unwrap();
937        assert_eq!(zstd_feed(&mut first, &large).unwrap(), b"hello");
938        let charged = first.charged;
939        assert!(charged > 7 * 1024 * 1024, "{charged}");
940        assert_eq!(MEMORY.used.load(Ordering::Acquire), charged);
941
942        let mut second = ZstdDecoder::with_memory(&MEMORY).unwrap();
943        assert_eq!(zstd_feed(&mut second, &large).unwrap(), b"hello");
944        assert_eq!(MEMORY.used.load(Ordering::Acquire), 2 * charged);
945
946        // the third decoder is over the limit
947        let mut third = ZstdDecoder::with_memory(&MEMORY).unwrap();
948        let err = zstd_feed(&mut third, &large).unwrap_err();
949        assert_eq!(err.kind(), io::ErrorKind::OutOfMemory);
950        drop(third);
951        assert_eq!(MEMORY.used.load(Ordering::Acquire), 2 * charged);
952
953        // small windows are not counted
954        let mut small = ZstdDecoder::with_memory(&MEMORY).unwrap();
955        assert_eq!(
956            zstd_feed(&mut small, &zstd_frame(19, b"hi")).unwrap(),
957            b"hi"
958        );
959        assert_eq!(small.charged, 0);
960
961        // memory of a dropped decoder is available again
962        drop(first);
963        assert_eq!(MEMORY.used.load(Ordering::Acquire), charged);
964        let mut third = ZstdDecoder::with_memory(&MEMORY).unwrap();
965        assert_eq!(zstd_feed(&mut third, &large).unwrap(), b"hello");
966        drop((second, third));
967        assert_eq!(MEMORY.used.load(Ordering::Acquire), 0);
968    }
969
970    #[test]
971    fn zstd_decoder_memory_shrinks() {
972        static MEMORY: ZstdMemory = ZstdMemory::new(20 * 1024 * 1024);
973
974        // zstd frees an oversized window after it was unused for a number of frames
975        let mut decoder = ZstdDecoder::with_memory(&MEMORY).unwrap();
976        zstd_feed(&mut decoder, &zstd_frame(23, b"hello")).unwrap();
977        assert!(decoder.charged > 0);
978        let small = zstd_frame(10, b"hi");
979        for _ in 0..256 {
980            zstd_feed(&mut decoder, &small).unwrap();
981        }
982        assert_eq!(decoder.charged, 0);
983        assert_eq!(MEMORY.used.load(Ordering::Acquire), 0);
984    }
985
986    #[crate::rt_test]
987    async fn zstd_decoder_memory_error() {
988        static MEMORY: ZstdMemory = ZstdMemory::new(1024 * 1024);
989
990        let chunks = vec![Ok::<_, PayloadError>(Bytes::from(zstd_frame(23, b"hello")))];
991        let mut decoder = Decoder {
992            inner: Some(ContentDecoder::Zstd(Box::new(
993                ZstdDecoder::with_memory(&MEMORY).unwrap(),
994            ))),
995            stream: stream::iter(chunks),
996            eof: false,
997            decode: true,
998            pending: None,
999            ready: VecDeque::new(),
1000            fut: None,
1001        };
1002        assert!(matches!(decoder.next().await, Some(Err(_))));
1003        assert!(decoder.next().await.is_none());
1004        assert_eq!(MEMORY.used.load(Ordering::Acquire), 0);
1005    }
1006
1007    #[crate::rt_test]
1008    async fn decoder_task_output() {
1009        for (encoding, limit) in [
1010            (ContentEncoding::Gzip, 128 * 1024),
1011            (ContentEncoding::Zstd, 512 * 1024),
1012        ] {
1013            let data = random(2 * 1024 * 1024);
1014            let compressed = compress(encoding, &data);
1015            let before = super::super::offloaded();
1016            let chunks = vec![Bytes::from(compressed)];
1017            assert_eq!(decode(encoding, chunks).await.unwrap(), data);
1018
1019            // each task decodes MAX_TASK_OUTPUT, the rest below the limit in place
1020            let tasks = super::super::offloaded() - before;
1021            let expected = (data.len() - limit).div_ceil(MAX_TASK_OUTPUT);
1022            assert!(
1023                (expected - 1..=expected).contains(&tasks),
1024                "{encoding:?} {tasks}"
1025            );
1026        }
1027    }
1028
1029    #[crate::rt_test]
1030    async fn decoder_truncated_stream() {
1031        let mut e = GzEncoder::new(Vec::new(), Compression::fast());
1032        e.write_all(b"hello world").unwrap();
1033        let data = e.finish().unwrap();
1034
1035        let chunks = vec![Ok::<_, PayloadError>(Bytes::copy_from_slice(
1036            &data[..data.len() - 4],
1037        ))];
1038        let mut decoder = Decoder::new(stream::iter(chunks), ContentEncoding::Gzip);
1039        let mut result = Ok(());
1040        while let Some(chunk) = decoder.next().await {
1041            if let Err(e) = chunk {
1042                result = Err(e);
1043            }
1044        }
1045        assert!(result.is_err());
1046    }
1047
1048    fn bytewise(data: &[u8]) -> Vec<Bytes> {
1049        data.chunks(1).map(Bytes::copy_from_slice).collect()
1050    }
1051
1052    /// A gzip member with all optional header fields, the header crc is at `len - 2`.
1053    fn gzip_header(flags: u8) -> Vec<u8> {
1054        let mut header = vec![0x1f, 0x8b, 8, flags, 1, 2, 3, 4, 0, 255];
1055        header.extend([3, 0, b'a', b'b', b'c']);
1056        header.extend(b"name\0comment\0");
1057        let mut crc = flate2::Crc::new();
1058        crc.update(&header);
1059        header.extend((crc.sum() as u16).to_le_bytes());
1060        header
1061    }
1062
1063    #[crate::rt_test]
1064    async fn gzip_decoder_header() {
1065        let data = random(4096);
1066        let body = &compress(ContentEncoding::Gzip, &data)[10..];
1067        let mut full = gzip_header(0x1e);
1068        full.extend(body);
1069        for chunks in [vec![Bytes::from(full.clone())], bytewise(&full)] {
1070            assert_eq!(decode(ContentEncoding::Gzip, chunks).await.unwrap(), data);
1071        }
1072
1073        // header crc, magic, method and reserved flags
1074        let crc = gzip_header(0x1e).len() - 2;
1075        for (pos, flip) in [(crc, 1), (0, 1), (1, 1), (2, 1), (3, 0x20), (3, 0x80)] {
1076            let mut bad = full.clone();
1077            bad[pos] ^= flip;
1078            let out = decode(ContentEncoding::Gzip, vec![Bytes::from(bad)]).await;
1079            assert!(out.is_err(), "{pos} {flip}");
1080        }
1081    }
1082
1083    #[crate::rt_test]
1084    async fn flate_decoder_errors() {
1085        let data = random(4096);
1086        for encoding in [ContentEncoding::Gzip, ContentEncoding::Deflate] {
1087            let compressed = compress(encoding, &data);
1088            let len = compressed.len();
1089            assert_eq!(decode(encoding, bytewise(&compressed)).await.unwrap(), data);
1090
1091            let mut trailing = compressed.clone();
1092            trailing.push(0);
1093            let out = decode(encoding, vec![Bytes::from(trailing)]).await;
1094            assert!(out.is_err(), "{encoding:?} trailing data");
1095
1096            // checksums
1097            for pos in [len - 1, len - 5] {
1098                let mut bad = compressed.clone();
1099                bad[pos] ^= 1;
1100                let out = decode(encoding, vec![Bytes::from(bad)]).await;
1101                assert!(out.is_err(), "{encoding:?} {pos}");
1102            }
1103
1104            for n in [0, 1, 11, len - 1] {
1105                let chunks = vec![Bytes::copy_from_slice(&compressed[..n])];
1106                let out = decode(encoding, chunks).await;
1107                assert!(out.is_err(), "{encoding:?} truncated to {n}");
1108            }
1109        }
1110    }
1111
1112    #[test]
1113    fn flate_decoder_output_fills_pages() {
1114        // some sizes end the stream exactly at the end of a page, which holds
1115        // a little less than `MAX_CHUNK_SIZE`
1116        let mut full = 0;
1117        for gzip in [false, true] {
1118            let encoding = if gzip {
1119                ContentEncoding::Gzip
1120            } else {
1121                ContentEncoding::Deflate
1122            };
1123            for len in MAX_CHUNK_SIZE - 64..=MAX_CHUNK_SIZE {
1124                let data = b"abc".repeat(len)[..len].to_vec();
1125                let mut input = Bytes::from(compress(encoding, &data));
1126                let mut decoder = FlateDecoder::new(gzip);
1127                let mut out = Vec::new();
1128                while !input.is_empty() || decoder.more {
1129                    if let Some(chunk) = decoder.feed(&mut input).unwrap() {
1130                        out.extend_from_slice(&chunk);
1131                    }
1132                    full += usize::from(decoder.done && decoder.buf.remaining_mut() == 0);
1133                }
1134                decoder.finish().unwrap();
1135                assert_eq!(out, data);
1136            }
1137        }
1138        assert!(full > 0);
1139    }
1140}