Skip to main content

ntex/http/encoding/
encoder.rs

1//! Stream encoder
2use std::{
3    collections::VecDeque, fmt, future::Future, io, pin::Pin, rc::Rc, task::Context, task::Poll,
4};
5
6use flate2::{Compress, Compression, Crc, FlushCompress, Status};
7use zstd::zstd_safe::{CCtx, CParameter, InBuffer, OutBuffer, zstd_sys::ZSTD_EndDirective};
8
9use crate::http::body::{Body, BodySize, MessageBody, ResponseBody};
10use crate::http::header::{CONTENT_ENCODING, ContentEncoding, HeaderValue};
11use crate::http::{ResponseHead, StatusCode};
12use crate::rt::BlockingResult;
13use crate::util::{BufMut, BytePageSize, Bytes, BytesMut, dyn_rc_err};
14
15use super::{Spare, offload, zstd_error};
16
17/// Bodies of a known size below this are sent without compression.
18///
19/// Small bodies barely shrink, while every encoder allocates its state.
20const MIN_SIZE: u64 = 1024;
21
22/// The largest `zstd` window of the encoder, 512KiB.
23///
24/// Without it a body of unknown size gets a 2MiB window and its encoder
25/// allocates about 3.6MiB, for barely better compression.
26const ZSTD_WINDOW_LOG: u32 = 19;
27
28/// Response body encoder.
29///
30/// Compresses a response body with the selected content encoding.
31pub struct Encoder<B> {
32    eof: bool,
33    body: EncoderBody<B>,
34    /// The part of a large chunk that is not encoded yet
35    pending: Bytes,
36    inner: Option<ContentEncoder>,
37    fut: Option<BlockingResult<Result<ContentEncoder, io::Error>>>,
38}
39
40impl<B: MessageBody> Encoder<B> {
41    /// Wrap a response body in an encoder for `encoding`.
42    ///
43    /// On success the `Content-Encoding` header is set and chunked
44    /// transfer-encoding is enabled. The body is returned unchanged if the
45    /// encoding is not supported, is `Identity` or `Auto`, if the response
46    /// already has a `Content-Encoding` header, if the status is
47    /// `101 Switching Protocols` or `204 No Content`, or if the body is empty
48    /// or its size is known to be below 1KiB.
49    pub fn response(
50        encoding: ContentEncoding,
51        head: &mut ResponseHead,
52        body: ResponseBody<B>,
53    ) -> ResponseBody<B> {
54        let size = match body.size() {
55            BodySize::None | BodySize::Empty => return body,
56            BodySize::Sized(size) if size < MIN_SIZE => return body,
57            BodySize::Sized(size) => Some(size),
58            BodySize::Stream => None,
59        };
60        if head.headers().contains_key(&CONTENT_ENCODING)
61            || head.status == StatusCode::SWITCHING_PROTOCOLS
62            || head.status == StatusCode::NO_CONTENT
63        {
64            return body;
65        }
66        let Some(encoder) = ContentEncoder::new(encoding, size) else {
67            return body;
68        };
69
70        let body = match body {
71            ResponseBody::Other(b) => match b {
72                Body::None | Body::Empty => unreachable!(),
73                Body::Bytes(buf) => EncoderBody::Bytes(buf),
74                Body::Message(stream) => EncoderBody::BoxedStream(stream),
75            },
76            ResponseBody::Body(stream) => EncoderBody::Stream(stream),
77        };
78        head.headers_mut().insert(
79            CONTENT_ENCODING,
80            HeaderValue::from_static(encoding.as_str()),
81        );
82        head.no_chunking(false);
83        ResponseBody::Other(Body::from_message(Encoder {
84            body,
85            eof: false,
86            pending: Bytes::new(),
87            fut: None,
88            inner: Some(encoder),
89        }))
90    }
91}
92
93impl Encoder<()> {
94    /// Returns true if the encoder can produce `encoding`.
95    pub(crate) fn can_encode(encoding: ContentEncoding) -> bool {
96        matches!(
97            encoding,
98            ContentEncoding::Deflate | ContentEncoding::Gzip | ContentEncoding::Zstd
99        )
100    }
101}
102
103impl<B: fmt::Debug> fmt::Debug for Encoder<B> {
104    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
105        f.debug_struct("Encoder")
106            .field("eof", &self.eof)
107            .field("body", &self.body)
108            .field("pending", &self.pending.len())
109            .field("encoder", &self.inner)
110            .field("fut", &self.fut.as_ref().map(|_| "JoinHandle(_)"))
111            .finish()
112    }
113}
114
115enum EncoderBody<B> {
116    Bytes(Bytes),
117    Stream(B),
118    BoxedStream(Box<dyn MessageBody>),
119}
120
121impl<B> fmt::Debug for EncoderBody<B> {
122    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
123        match self {
124            EncoderBody::Bytes(b) => write!(f, "EncoderBody::Bytes({b:?})"),
125            EncoderBody::Stream(_) => write!(f, "EncoderBody::Stream(_)"),
126            EncoderBody::BoxedStream(_) => write!(f, "EncoderBody::BoxedStream(_)"),
127        }
128    }
129}
130
131impl<B: MessageBody> MessageBody for Encoder<B> {
132    fn size(&self) -> BodySize {
133        BodySize::Stream
134    }
135
136    fn poll_next_chunk(
137        &mut self,
138        cx: &mut Context<'_>,
139    ) -> Poll<Option<Result<Bytes, Rc<dyn std::error::Error>>>> {
140        let result = self.poll_encoded(cx);
141        if let Poll::Ready(Some(Err(_))) = result {
142            // the encoder state is lost, the stream must not continue with raw data
143            self.eof = true;
144            self.inner = None;
145            self.fut = None;
146            self.pending = Bytes::new();
147        }
148        result
149    }
150}
151
152impl<B: MessageBody> Encoder<B> {
153    fn poll_encoded(
154        &mut self,
155        cx: &mut Context<'_>,
156    ) -> Poll<Option<Result<Bytes, Rc<dyn std::error::Error>>>> {
157        loop {
158            if let Some(chunk) = self.inner.as_mut().and_then(ContentEncoder::take) {
159                return Poll::Ready(Some(Ok(chunk)));
160            }
161
162            if self.eof {
163                self.inner = None;
164                return Poll::Ready(None);
165            }
166
167            if let Some(ref mut fut) = self.fut {
168                let encoder = match Pin::new(fut).poll(cx) {
169                    Poll::Ready(Ok(Ok(item))) => item,
170                    Poll::Ready(Ok(Err(e))) => return Poll::Ready(Some(Err(Rc::new(e)))),
171                    Poll::Ready(Err(_)) => {
172                        return Poll::Ready(Some(Err(Rc::new(io::Error::new(
173                            io::ErrorKind::Interrupted,
174                            "Canceled",
175                        )))));
176                    }
177                    Poll::Pending => return Poll::Pending,
178                };
179                self.inner = Some(encoder);
180                self.fut = None;
181                continue;
182            }
183
184            if !self.pending.is_empty() {
185                let chunk = std::mem::take(&mut self.pending);
186                self.encode(chunk)?;
187                continue;
188            }
189
190            let result = match self.body {
191                EncoderBody::Bytes(ref mut b) => {
192                    if b.is_empty() {
193                        Poll::Ready(None)
194                    } else {
195                        Poll::Ready(Some(Ok(std::mem::take(b))))
196                    }
197                }
198                EncoderBody::Stream(ref mut b) => b.poll_next_chunk(cx),
199                EncoderBody::BoxedStream(ref mut b) => b.poll_next_chunk(cx),
200            };
201            match result {
202                Poll::Ready(Some(Ok(chunk))) => self.encode(chunk)?,
203                Poll::Ready(None) => {
204                    self.eof = true;
205                    if let Some(encoder) = self.inner.as_mut() {
206                        encoder.finish().map_err(dyn_rc_err)?;
207                    }
208                }
209                // the body has nothing more for now, so the client gets what
210                // was written so far instead of waiting for the next chunk
211                Poll::Pending => {
212                    if let Some(encoder) = self.inner.as_mut() {
213                        encoder.flush().map_err(dyn_rc_err)?;
214                        if let Some(chunk) = encoder.take() {
215                            return Poll::Ready(Some(Ok(chunk)));
216                        }
217                    }
218                    return Poll::Pending;
219                }
220                Poll::Ready(Some(Err(e))) => return Poll::Ready(Some(Err(e))),
221            }
222        }
223    }
224}
225
226impl<B> Encoder<B> {
227    /// Encodes a small chunk in place.
228    ///
229    /// A large chunk is encoded on the blocking thread pool, one part at a
230    /// time, so the output of each part is sent before the next is encoded.
231    fn encode(&mut self, mut chunk: Bytes) -> Result<(), Rc<dyn std::error::Error>> {
232        let Some(mut encoder) = self.inner.take() else {
233            return Ok(());
234        };
235        if chunk.len() < encoder.limit() {
236            encoder.write(&chunk).map_err(dyn_rc_err)?;
237            self.inner = Some(encoder);
238        } else {
239            let part = chunk.split_to(chunk.len().min(encoder.task_size()));
240            self.pending = chunk;
241            self.fut = Some(offload(move || {
242                encoder.write(&part)?;
243                Ok(encoder)
244            }));
245        }
246        Ok(())
247    }
248}
249
250/// The `gzip` header written by the encoder, without a name or modification
251/// time. The extra flags byte marks the fastest compression level.
252const GZIP_HEADER: [u8; 10] = [0x1f, 0x8b, 8, 0, 0, 0, 0, 0, 4, 255];
253
254/// Compresses into pages of a buffer, full pages are queued for sending.
255struct ContentEncoder {
256    codec: Codec,
257    buf: BytesMut,
258    chunks: VecDeque<Bytes>,
259    /// Input was written since the last flush
260    unflushed: bool,
261}
262
263#[derive(Copy, Clone, PartialEq, Eq)]
264enum Op {
265    Write,
266    /// Write the output of all input so far
267    Flush,
268    /// Write the end of the stream
269    Finish,
270}
271
272enum Codec {
273    Deflate(Compress),
274    /// The checksum and size of the input for the trailer
275    Gzip(Compress, Crc),
276    Zstd(CCtx<'static>),
277    /// The stream is finished, the encoder state is released
278    Done,
279}
280
281impl ContentEncoder {
282    /// Chunks of this size and larger are encoded on the blocking thread pool.
283    ///
284    /// `gzip` and `deflate` are several times slower than `zstd`, so they are
285    /// offloaded much earlier.
286    const fn limit(&self) -> usize {
287        match self.codec {
288            Codec::Deflate(_) | Codec::Gzip(..) => 16 * 1024,
289            Codec::Zstd(_) | Codec::Done => 512 * 1024,
290        }
291    }
292
293    /// The most input a single blocking task encodes.
294    const fn task_size(&self) -> usize {
295        match self.codec {
296            Codec::Deflate(_) | Codec::Gzip(..) => 256 * 1024,
297            Codec::Zstd(_) | Codec::Done => 1024 * 1024,
298        }
299    }
300
301    /// Creates an encoder, `size` is the length of the body if it is known.
302    fn new(encoding: ContentEncoding, size: Option<u64>) -> Option<Self> {
303        let mut buf = BytesMut::with_page_size(BytePageSize::Size32);
304        let codec = match encoding {
305            ContentEncoding::Deflate => Codec::Deflate(Compress::new(Compression::fast(), true)),
306            ContentEncoding::Gzip => {
307                buf.extend_from_slice(&GZIP_HEADER);
308                Codec::Gzip(Compress::new(Compression::fast(), false), Crc::new())
309            }
310            // with a known size zstd picks a smaller window and stores the size in the frame
311            ContentEncoding::Zstd => {
312                let mut ctx = CCtx::try_create()?;
313                ctx.set_parameter(CParameter::CompressionLevel(0)).ok()?;
314                ctx.set_parameter(CParameter::WindowLog(ZSTD_WINDOW_LOG))
315                    .ok()?;
316                ctx.set_pledged_src_size(size).ok()?;
317                Codec::Zstd(ctx)
318            }
319            _ => return None,
320        };
321        Some(ContentEncoder {
322            codec,
323            buf,
324            chunks: VecDeque::new(),
325            unflushed: false,
326        })
327    }
328
329    /// Returns the encoded output, a full page first.
330    fn take(&mut self) -> Option<Bytes> {
331        self.chunks
332            .pop_front()
333            .or_else(|| (!self.buf.is_empty()).then(|| self.buf.take()))
334    }
335
336    /// Queues a full page and starts a new one.
337    ///
338    /// A page is never grown, so its output is not copied.
339    fn reserve(&mut self) {
340        if self.buf.remaining_mut() == 0 {
341            if !self.buf.is_empty() {
342                self.chunks.push_back(self.buf.take());
343            }
344            self.buf.reserve_more();
345        }
346    }
347
348    /// Compresses `data` into the buffer, returns `true` once a flush or the
349    /// end of the stream is written completely.
350    fn compress(&mut self, data: &mut &[u8], op: Op) -> io::Result<bool> {
351        self.reserve();
352        match &mut self.codec {
353            Codec::Deflate(inner) | Codec::Gzip(inner, _) => {
354                let flush = match op {
355                    Op::Write => FlushCompress::None,
356                    Op::Flush => FlushCompress::Sync,
357                    Op::Finish => FlushCompress::Finish,
358                };
359                let (total_in, total_out) = (inner.total_in(), inner.total_out());
360                // SAFETY: the encoder only writes to the slice, and `advance_mut`
361                // covers just the bytes it has written
362                let status = unsafe {
363                    let spare = self.buf.chunk_mut().as_uninit_slice_mut();
364                    inner
365                        .compress_uninit(data, spare, flush)
366                        .map_err(io::Error::other)?
367                };
368                let read = usize::try_from(inner.total_in() - total_in).unwrap();
369                let written = usize::try_from(inner.total_out() - total_out).unwrap();
370                unsafe { self.buf.advance_mut(written) };
371                *data = &data[read..];
372
373                if status == Status::StreamEnd {
374                    Ok(true)
375                } else if read == 0 && written == 0 {
376                    Err(io::ErrorKind::WriteZero.into())
377                } else {
378                    // a flush is complete once it leaves space in the page
379                    Ok(op == Op::Flush && self.buf.remaining_mut() > 0)
380                }
381            }
382            Codec::Zstd(ctx) => {
383                let end = match op {
384                    Op::Write => ZSTD_EndDirective::ZSTD_e_continue,
385                    Op::Flush => ZSTD_EndDirective::ZSTD_e_flush,
386                    Op::Finish => ZSTD_EndDirective::ZSTD_e_end,
387                };
388                let mut src = InBuffer::around(data);
389                let mut spare = Spare::new(&mut self.buf);
390                let mut dst = OutBuffer::around(&mut spare);
391                // the size of the output the context still holds
392                let left = ctx
393                    .compress_stream2(&mut dst, &mut src, end)
394                    .map_err(zstd_error)?;
395                let read = src.pos();
396                *data = &data[read..];
397                Ok(op != Op::Write && left == 0)
398            }
399            Codec::Done => Err(io::Error::other("the stream is finished")),
400        }
401    }
402
403    /// Appends `data` across pages.
404    fn put(&mut self, mut data: &[u8]) {
405        while !data.is_empty() {
406            self.reserve();
407            let size = data.len().min(self.buf.remaining_mut());
408            self.buf.extend_from_slice(&data[..size]);
409            data = &data[size..];
410        }
411    }
412
413    /// Writes the end of the stream.
414    fn finish(&mut self) -> io::Result<()> {
415        while !self.compress(&mut &[][..], Op::Finish)? {}
416        if let Codec::Gzip(_, crc) = &self.codec {
417            let (sum, amount) = (crc.sum(), crc.amount());
418            self.put(&sum.to_le_bytes());
419            self.put(&amount.to_le_bytes());
420        }
421        // only the output is left, which may take a while to send
422        self.codec = Codec::Done;
423        Ok(())
424    }
425
426    /// Writes the output of all input so far, the stream continues after it.
427    ///
428    /// `gzip` and `deflate` add about 5 bytes for each flush.
429    fn flush(&mut self) -> io::Result<()> {
430        if self.unflushed {
431            self.unflushed = false;
432            while !self.compress(&mut &[][..], Op::Flush)? {}
433        }
434        Ok(())
435    }
436
437    fn write(&mut self, mut data: &[u8]) -> io::Result<()> {
438        self.unflushed |= !data.is_empty();
439        if let Codec::Gzip(_, crc) = &mut self.codec {
440            crc.update(data);
441        }
442        while !data.is_empty() {
443            self.compress(&mut data, Op::Write)
444                .inspect_err(|err| log::trace!("Failed to encode to {self:?}: {err}"))?;
445        }
446        Ok(())
447    }
448}
449
450impl fmt::Debug for ContentEncoder {
451    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
452        match self.codec {
453            Codec::Deflate(_) => write!(f, "ContentEncoder::Deflate"),
454            Codec::Gzip(..) => write!(f, "ContentEncoder::Gzip"),
455            Codec::Zstd(_) => write!(f, "ContentEncoder::Zstd"),
456            Codec::Done => write!(f, "ContentEncoder::Done"),
457        }
458    }
459}
460
461#[cfg(test)]
462mod tests {
463    use std::future::poll_fn;
464
465    use super::*;
466    use crate::rt::spawn_blocking;
467
468    #[crate::rt_test]
469    async fn encoder_is_fused_after_error() {
470        let mut enc = Encoder::<Body> {
471            eof: false,
472            pending: Bytes::new(),
473            body: EncoderBody::Bytes(Bytes::from_static(b"raw data")),
474            inner: None,
475            fut: Some(spawn_blocking(|| Err(io::Error::other("encode failed")))),
476        };
477        assert_eq!(enc.size(), BodySize::Stream);
478
479        let res = poll_fn(|cx| enc.poll_next_chunk(cx)).await;
480        assert!(matches!(res, Some(Err(_))));
481        assert_eq!(enc.size(), BodySize::Stream);
482        assert!(poll_fn(|cx| enc.poll_next_chunk(cx)).await.is_none());
483    }
484
485    struct EndOnce(bool);
486
487    impl MessageBody for EndOnce {
488        fn size(&self) -> BodySize {
489            BodySize::Stream
490        }
491
492        fn poll_next_chunk(
493            &mut self,
494            _: &mut Context<'_>,
495        ) -> Poll<Option<Result<Bytes, Rc<dyn std::error::Error>>>> {
496            assert!(!self.0, "body polled after end of stream");
497            self.0 = true;
498            Poll::Ready(None)
499        }
500    }
501
502    #[crate::rt_test]
503    async fn encoder_is_fused_after_end_of_stream() {
504        let mut enc = Encoder::<EndOnce> {
505            eof: false,
506            pending: Bytes::new(),
507            body: EncoderBody::Stream(EndOnce(false)),
508            inner: None,
509            fut: None,
510        };
511        assert!(poll_fn(|cx| enc.poll_next_chunk(cx)).await.is_none());
512        assert!(poll_fn(|cx| enc.poll_next_chunk(cx)).await.is_none());
513    }
514
515    async fn collect(body: &mut ResponseBody<Body>) -> Vec<u8> {
516        let mut buf = Vec::new();
517        while let Some(chunk) = poll_fn(|cx| body.poll_next_chunk(cx)).await {
518            buf.extend_from_slice(&chunk.unwrap());
519        }
520        buf
521    }
522
523    fn gunzip(data: &[u8]) -> Vec<u8> {
524        use std::io::Read;
525
526        let mut buf = Vec::new();
527        flate2::read::GzDecoder::new(data)
528            .read_to_end(&mut buf)
529            .unwrap();
530        buf
531    }
532
533    #[crate::rt_test]
534    async fn encoder_response_bodies() {
535        let mut head = ResponseHead::new(StatusCode::OK, crate::http::Version::HTTP_11);
536        let body = Encoder::<Body>::response(ContentEncoding::Gzip, &mut head, Body::None.into());
537        assert!(matches!(body, ResponseBody::Other(Body::None)));
538        let body = Encoder::<Body>::response(ContentEncoding::Gzip, &mut head, Body::Empty.into());
539        assert!(matches!(body, ResponseBody::Other(Body::Empty)));
540
541        // identity encoding is not applied
542        let body = Encoder::<Body>::response(
543            ContentEncoding::Identity,
544            &mut head,
545            Body::from("data").into(),
546        );
547        assert!(matches!(body, ResponseBody::Other(Body::Bytes(_))));
548        assert!(!head.headers().contains_key(CONTENT_ENCODING));
549
550        // nor are `Auto`, responses without a body and encoded responses
551        let data = "data".repeat(256);
552        let body = Encoder::<Body>::response(
553            ContentEncoding::Auto,
554            &mut head,
555            Body::from(data.clone()).into(),
556        );
557        assert!(matches!(body, ResponseBody::Other(Body::Bytes(_))));
558        assert!(!head.headers().contains_key(CONTENT_ENCODING));
559        for status in [StatusCode::SWITCHING_PROTOCOLS, StatusCode::NO_CONTENT] {
560            let mut head = ResponseHead::new(status, crate::http::Version::HTTP_11);
561            let body = Encoder::<Body>::response(
562                ContentEncoding::Gzip,
563                &mut head,
564                Body::from(data.clone()).into(),
565            );
566            assert!(matches!(body, ResponseBody::Other(Body::Bytes(_))));
567            assert!(!head.headers().contains_key(CONTENT_ENCODING));
568        }
569        let mut head = ResponseHead::new(StatusCode::OK, crate::http::Version::HTTP_11);
570        head.headers_mut()
571            .insert(CONTENT_ENCODING, HeaderValue::from_static("br"));
572        let body =
573            Encoder::<Body>::response(ContentEncoding::Gzip, &mut head, Body::from(data).into());
574        assert!(matches!(body, ResponseBody::Other(Body::Bytes(_))));
575        assert_eq!(head.headers().get(CONTENT_ENCODING).unwrap(), "br");
576
577        // boxed message stream
578        let data = "data".repeat(256);
579        let mut head = ResponseHead::new(StatusCode::OK, crate::http::Version::HTTP_11);
580        let inner = Body::from_message(Body::from(data.clone()));
581        let mut body = Encoder::<Body>::response(ContentEncoding::Gzip, &mut head, inner.into());
582        assert_eq!(head.headers().get(CONTENT_ENCODING).unwrap(), "gzip");
583        assert_eq!(gunzip(&collect(&mut body).await), data.as_bytes());
584
585        // typed body stream
586        let typed = "typed".repeat(205);
587        let mut head = ResponseHead::new(StatusCode::OK, crate::http::Version::HTTP_11);
588        let body = Encoder::<Body>::response(
589            ContentEncoding::Gzip,
590            &mut head,
591            ResponseBody::Body(Body::from(typed.clone())),
592        );
593        let ResponseBody::Other(Body::Message(mut msg)) = body else {
594            panic!("expected encoded body")
595        };
596        let mut buf = Vec::new();
597        while let Some(chunk) = poll_fn(|cx| msg.poll_next_chunk(cx)).await {
598            buf.extend_from_slice(&chunk.unwrap());
599        }
600        assert_eq!(gunzip(&buf), typed.as_bytes());
601    }
602
603    struct Chunks(Vec<Bytes>, BodySize);
604
605    impl MessageBody for Chunks {
606        fn size(&self) -> BodySize {
607            self.1
608        }
609
610        fn poll_next_chunk(
611            &mut self,
612            _: &mut Context<'_>,
613        ) -> Poll<Option<Result<Bytes, Rc<dyn std::error::Error>>>> {
614            Poll::Ready(self.0.pop().map(Ok))
615        }
616    }
617
618    #[crate::rt_test]
619    async fn encoder_skips_small_bodies() {
620        for encoding in [
621            ContentEncoding::Gzip,
622            ContentEncoding::Deflate,
623            ContentEncoding::Zstd,
624        ] {
625            let small = Bytes::from(vec![b'x'; 1023]);
626            let large = Bytes::from(vec![b'x'; 1024]);
627
628            let mut head = ResponseHead::new(StatusCode::OK, crate::http::Version::HTTP_11);
629            let body =
630                Encoder::<Body>::response(encoding, &mut head, Body::from(small.clone()).into());
631            assert!(matches!(body, ResponseBody::Other(Body::Bytes(ref b)) if *b == small));
632            assert!(!head.headers().contains_key(CONTENT_ENCODING));
633
634            let mut head = ResponseHead::new(StatusCode::OK, crate::http::Version::HTTP_11);
635            let stream = Chunks(vec![small.clone()], BodySize::Sized(1023));
636            let body = Encoder::response(encoding, &mut head, ResponseBody::Body(stream));
637            assert!(matches!(body, ResponseBody::Body(_)));
638            assert!(!head.headers().contains_key(CONTENT_ENCODING));
639
640            let mut head = ResponseHead::new(StatusCode::OK, crate::http::Version::HTTP_11);
641            let mut body =
642                Encoder::<Body>::response(encoding, &mut head, Body::from(large.clone()).into());
643            assert_eq!(
644                head.headers().get(CONTENT_ENCODING).unwrap(),
645                encoding.as_str()
646            );
647            assert_eq!(decompress(encoding, &collect(&mut body).await), large);
648
649            // a stream of unknown size is encoded
650            let mut head = ResponseHead::new(StatusCode::OK, crate::http::Version::HTTP_11);
651            let stream = Chunks(vec![Bytes::from_static(b"x")], BodySize::Stream);
652            let mut body =
653                Encoder::<Body>::response(encoding, &mut head, Body::from_message(stream).into());
654            assert_eq!(
655                head.headers().get(CONTENT_ENCODING).unwrap(),
656                encoding.as_str()
657            );
658            assert_eq!(decompress(encoding, &collect(&mut body).await), b"x");
659        }
660    }
661
662    fn zstd_content_size(data: &[u8]) -> Option<u64> {
663        zstd::zstd_safe::get_frame_content_size(data).unwrap()
664    }
665
666    #[crate::rt_test]
667    async fn encoder_zstd_known_size() {
668        let data = Bytes::from(vec![b'x'; 4096]);
669
670        // the frame stores the size of a sized body
671        let mut head = ResponseHead::new(StatusCode::OK, crate::http::Version::HTTP_11);
672        let mut body = Encoder::<Body>::response(
673            ContentEncoding::Zstd,
674            &mut head,
675            Body::from(data.clone()).into(),
676        );
677        let frame = collect(&mut body).await;
678        assert_eq!(zstd_content_size(&frame), Some(4096));
679        assert_eq!(decompress(ContentEncoding::Zstd, &frame), data);
680
681        let mut head = ResponseHead::new(StatusCode::OK, crate::http::Version::HTTP_11);
682        let chunks = vec![data.slice(2048..), data.slice(..2048)];
683        let stream = Body::from_message(Chunks(chunks, BodySize::Sized(4096)));
684        let mut body = Encoder::<Body>::response(ContentEncoding::Zstd, &mut head, stream.into());
685        let frame = collect(&mut body).await;
686        assert_eq!(zstd_content_size(&frame), Some(4096));
687        assert_eq!(decompress(ContentEncoding::Zstd, &frame), data);
688
689        // the size of a stream is unknown
690        let mut head = ResponseHead::new(StatusCode::OK, crate::http::Version::HTTP_11);
691        let stream = Body::from_message(Chunks(vec![data.clone()], BodySize::Stream));
692        let mut body = Encoder::<Body>::response(ContentEncoding::Zstd, &mut head, stream.into());
693        let frame = collect(&mut body).await;
694        assert_eq!(zstd_content_size(&frame), None);
695        assert_eq!(decompress(ContentEncoding::Zstd, &frame), data);
696
697        // a body larger than its size fails
698        let mut head = ResponseHead::new(StatusCode::OK, crate::http::Version::HTTP_11);
699        let stream = Body::from_message(Chunks(vec![data.clone()], BodySize::Sized(2048)));
700        let mut body = Encoder::<Body>::response(ContentEncoding::Zstd, &mut head, stream.into());
701        assert!(matches!(
702            poll_fn(|cx| body.poll_next_chunk(cx)).await,
703            Some(Err(_))
704        ));
705    }
706
707    fn decompress(encoding: ContentEncoding, data: &[u8]) -> Vec<u8> {
708        use std::io::Read;
709
710        let mut buf = Vec::new();
711        match encoding {
712            ContentEncoding::Gzip => return gunzip(data),
713            ContentEncoding::Deflate => {
714                flate2::read::ZlibDecoder::new(data)
715                    .read_to_end(&mut buf)
716                    .unwrap();
717            }
718            ContentEncoding::Zstd => buf = zstd::decode_all(data).unwrap(),
719            _ => unreachable!(),
720        }
721        buf
722    }
723
724    #[crate::rt_test]
725    async fn encoder_offloads_large_chunks() {
726        for (encoding, limit) in [
727            (ContentEncoding::Gzip, 16 * 1024),
728            (ContentEncoding::Deflate, 16 * 1024),
729            (ContentEncoding::Zstd, 512 * 1024),
730        ] {
731            for (len, offloaded) in [(limit - 1, 0), (limit, 1)] {
732                let data: Vec<u8> = (0..len).map(|i: usize| (i % 251) as u8).collect();
733                let before = super::super::offloaded();
734                let mut head = ResponseHead::new(StatusCode::OK, crate::http::Version::HTTP_11);
735                let mut body =
736                    Encoder::<Body>::response(encoding, &mut head, Body::from(data.clone()).into());
737                assert_eq!(
738                    head.headers().get(CONTENT_ENCODING).unwrap(),
739                    encoding.as_str()
740                );
741                assert_eq!(decompress(encoding, &collect(&mut body).await), data);
742                assert_eq!(
743                    super::super::offloaded() - before,
744                    offloaded,
745                    "{encoding:?} {len}"
746                );
747            }
748        }
749    }
750
751    /// Incompressible data, so the encoded size is close to `len`.
752    fn random(len: usize) -> Vec<u8> {
753        let mut x = 0x2545_f491_4f6c_dd1d_u64;
754        (0..len)
755            .map(|_| {
756                x ^= x << 13;
757                x ^= x >> 7;
758                x ^= x << 17;
759                x as u8
760            })
761            .collect()
762    }
763
764    #[crate::rt_test]
765    async fn encoder_splits_large_chunks() {
766        for (encoding, task, limit) in [
767            (ContentEncoding::Gzip, 256 * 1024, 16 * 1024),
768            (ContentEncoding::Deflate, 256 * 1024, 16 * 1024),
769            (ContentEncoding::Zstd, 1024 * 1024, 512 * 1024),
770        ] {
771            // two parts on the pool, the rest is below the limit and encoded in place
772            let data = random(2 * task + limit / 2);
773            let before = super::super::offloaded();
774            let mut head = ResponseHead::new(StatusCode::OK, crate::http::Version::HTTP_11);
775            let mut body =
776                Encoder::<Body>::response(encoding, &mut head, Body::from(data.clone()).into());
777
778            let mut encoded = Vec::new();
779            let mut chunks = 0;
780            while let Some(chunk) = poll_fn(|cx| body.poll_next_chunk(cx)).await {
781                let chunk = chunk.unwrap();
782                assert!(chunk.len() <= 32 * 1024, "{encoding:?} {}", chunk.len());
783                encoded.extend_from_slice(&chunk);
784                chunks += 1;
785            }
786            assert!(
787                chunks > encoded.len() / (32 * 1024),
788                "{encoding:?} {chunks}"
789            );
790            assert_eq!(decompress(encoding, &encoded), data);
791            assert_eq!(super::super::offloaded() - before, 2, "{encoding:?}");
792        }
793    }
794
795    #[crate::rt_test]
796    async fn encoder_drops_pending_on_error() {
797        let mut enc = Encoder::<Body> {
798            eof: false,
799            pending: Bytes::from_static(b"rest"),
800            body: EncoderBody::Bytes(Bytes::new()),
801            inner: None,
802            fut: Some(spawn_blocking(|| Err(io::Error::other("encode failed")))),
803        };
804        assert!(matches!(
805            poll_fn(|cx| enc.poll_next_chunk(cx)).await,
806            Some(Err(_))
807        ));
808        assert!(enc.pending.is_empty());
809        assert!(poll_fn(|cx| enc.poll_next_chunk(cx)).await.is_none());
810    }
811
812    #[crate::rt_test]
813    async fn encoder_blocking_task_canceled() {
814        let mut enc = Encoder::<Body> {
815            eof: false,
816            pending: Bytes::new(),
817            body: EncoderBody::Bytes(Bytes::new()),
818            inner: None,
819            fut: Some(spawn_blocking(|| panic!("encoder panic"))),
820        };
821        let res = poll_fn(|cx| enc.poll_next_chunk(cx)).await;
822        assert!(matches!(res, Some(Err(ref e)) if e.to_string() == "Canceled"));
823        assert!(poll_fn(|cx| enc.poll_next_chunk(cx)).await.is_none());
824    }
825
826    #[crate::rt_test]
827    async fn encoder_without_content_encoder() {
828        let mut enc = Encoder::<Body> {
829            eof: false,
830            pending: Bytes::new(),
831            body: EncoderBody::Bytes(Bytes::from_static(b"data")),
832            inner: None,
833            fut: None,
834        };
835        assert!(poll_fn(|cx| enc.poll_next_chunk(cx)).await.is_none());
836    }
837
838    #[test]
839    fn encoder_debug() {
840        assert!(ContentEncoder::new(ContentEncoding::Identity, None).is_none());
841        for (encoding, name) in [
842            (ContentEncoding::Gzip, "ContentEncoder::Gzip"),
843            (ContentEncoding::Deflate, "ContentEncoder::Deflate"),
844            (ContentEncoding::Zstd, "ContentEncoder::Zstd"),
845        ] {
846            let enc = Encoder::<Body> {
847                eof: false,
848                pending: Bytes::new(),
849                body: EncoderBody::Bytes(Bytes::from_static(b"data")),
850                inner: ContentEncoder::new(encoding, None),
851                fut: None,
852            };
853            let s = format!("{enc:?}");
854            assert!(s.contains(name) && s.contains("EncoderBody::Bytes"), "{s}");
855        }
856
857        let s = format!("{:?}", EncoderBody::Stream(Body::Empty));
858        assert_eq!(s, "EncoderBody::Stream(_)");
859        let s = format!(
860            "{:?}",
861            EncoderBody::<Body>::BoxedStream(Box::new(Body::Empty))
862        );
863        assert_eq!(s, "EncoderBody::BoxedStream(_)");
864    }
865
866    fn encode_all(encoding: ContentEncoding, data: &[u8]) -> Vec<Bytes> {
867        let mut enc = ContentEncoder::new(encoding, None).unwrap();
868        enc.write(data).unwrap();
869        enc.finish().unwrap();
870        std::iter::from_fn(|| enc.take()).collect()
871    }
872
873    #[test]
874    fn encoder_writes_pages() {
875        // the output of a large write is queued page by page
876        for encoding in [
877            ContentEncoding::Gzip,
878            ContentEncoding::Deflate,
879            ContentEncoding::Zstd,
880        ] {
881            let data = random(200 * 1024);
882            let chunks = encode_all(encoding, &data);
883            assert!(chunks.len() > 6, "{encoding:?} {}", chunks.len());
884            assert!(chunks.iter().all(|c| c.len() <= 32 * 1024 && !c.is_empty()));
885            assert_eq!(decompress(encoding, &chunks.concat()), data);
886        }
887
888        // gzip has the same header as flate2
889        let chunks = encode_all(ContentEncoding::Gzip, b"data");
890        let mut e = flate2::write::GzEncoder::new(Vec::new(), flate2::Compression::fast());
891        std::io::Write::write_all(&mut e, b"data").unwrap();
892        assert_eq!(chunks[0][..10], e.finish().unwrap()[..10]);
893    }
894
895    #[test]
896    fn encoder_finish_across_pages() {
897        // some sizes end the compressed data a few bytes before the end of a
898        // page, the end of the stream continues on the next one
899        let data = random(33 * 1024);
900        for encoding in [
901            ContentEncoding::Gzip,
902            ContentEncoding::Deflate,
903            ContentEncoding::Zstd,
904        ] {
905            let mut split = 0;
906            for len in 32 * 1024 - 64..32 * 1024 {
907                let chunks = encode_all(encoding, &data[..len]);
908                assert!(chunks.iter().all(|c| c.len() <= 32 * 1024));
909                if chunks.len() > 1 && chunks.last().unwrap().len() < 8 {
910                    split += 1;
911                }
912                assert_eq!(decompress(encoding, &chunks.concat()), &data[..len]);
913            }
914            assert!(split > 0, "{encoding:?}");
915        }
916    }
917
918    #[test]
919    fn encoder_zstd_window() {
920        let data = random(4096);
921        let mut window = Vec::new();
922        for size in [None, Some(4096)] {
923            let mut enc = ContentEncoder::new(ContentEncoding::Zstd, size).unwrap();
924            enc.write(&data).unwrap();
925            let Codec::Zstd(ctx) = &enc.codec else {
926                unreachable!()
927            };
928            window.push(ctx.sizeof());
929            enc.finish().unwrap();
930            let frame = std::iter::from_fn(|| enc.take())
931                .collect::<Vec<_>>()
932                .concat();
933            assert_eq!(decompress(ContentEncoding::Zstd, &frame), data);
934        }
935        // 3.6MiB with the default 2MiB window of a stream
936        assert!(window[0] < 2560 * 1024, "{}", window[0]);
937        // a known size keeps the window small
938        assert!(window[1] < 256 * 1024, "{}", window[1]);
939    }
940
941    #[test]
942    fn encoder_releases_state_on_finish() {
943        for encoding in [
944            ContentEncoding::Gzip,
945            ContentEncoding::Deflate,
946            ContentEncoding::Zstd,
947        ] {
948            let data = random(64 * 1024);
949            let mut enc = ContentEncoder::new(encoding, None).unwrap();
950            enc.write(&data).unwrap();
951            enc.finish().unwrap();
952            assert!(matches!(enc.codec, Codec::Done), "{encoding:?}");
953            assert_eq!(format!("{enc:?}"), "ContentEncoder::Done");
954            assert!(enc.write(b"more").is_err());
955            assert!(enc.finish().is_err());
956
957            let chunks: Vec<_> = std::iter::from_fn(|| enc.take()).collect();
958            assert!(chunks.len() > 1);
959            assert_eq!(decompress(encoding, &chunks.concat()), data);
960        }
961    }
962
963    #[test]
964    fn encoder_put_across_pages() {
965        let mut enc = ContentEncoder::new(ContentEncoding::Deflate, None).unwrap();
966        enc.reserve();
967        let room = enc.buf.remaining_mut();
968        enc.put(&vec![1; room - 3]);
969        enc.put(&[2; 8]);
970        let page = enc.take().unwrap();
971        assert_eq!(page.len(), room);
972        assert_eq!(page[room - 3..], [2; 3]);
973        assert_eq!(enc.take().unwrap(), [2; 5][..]);
974        assert!(enc.take().is_none());
975    }
976
977    #[test]
978    fn encoder_write_after_full_page() {
979        // a write can leave a full page, `zstd` does with these sizes, the
980        // next write starts a new page and does not queue an empty chunk
981        let mut full = 0;
982        for encoding in [
983            ContentEncoding::Gzip,
984            ContentEncoding::Deflate,
985            ContentEncoding::Zstd,
986        ] {
987            let data = random(256 * 1024);
988            let mut enc = ContentEncoder::new(encoding, None).unwrap();
989            let mut out = Vec::new();
990            for part in data.chunks(16 * 1024) {
991                enc.write(part).unwrap();
992                full += usize::from(enc.buf.remaining_mut() == 0);
993                while let Some(chunk) = enc.take() {
994                    assert!(!chunk.is_empty(), "{encoding:?}");
995                    out.extend_from_slice(&chunk);
996                }
997            }
998            enc.finish().unwrap();
999            while let Some(chunk) = enc.take() {
1000                assert!(!chunk.is_empty(), "{encoding:?}");
1001                out.extend_from_slice(&chunk);
1002            }
1003            assert_eq!(decompress(encoding, &out), data);
1004        }
1005        assert!(full > 0);
1006    }
1007
1008    /// A stream body, it returns `Pending` at `None` until the test removes it.
1009    struct Parts(Rc<std::cell::RefCell<VecDeque<Option<Bytes>>>>);
1010
1011    impl MessageBody for Parts {
1012        fn size(&self) -> BodySize {
1013            BodySize::Stream
1014        }
1015
1016        fn poll_next_chunk(
1017            &mut self,
1018            _: &mut Context<'_>,
1019        ) -> Poll<Option<Result<Bytes, Rc<dyn std::error::Error>>>> {
1020            let mut parts = self.0.borrow_mut();
1021            match parts.front() {
1022                Some(None) => Poll::Pending,
1023                Some(Some(_)) => Poll::Ready(parts.pop_front().unwrap().map(Ok)),
1024                None => Poll::Ready(None),
1025            }
1026        }
1027    }
1028
1029    /// Decodes the output so far, the stream does not have to be complete.
1030    fn decompress_prefix(encoding: ContentEncoding, data: &[u8]) -> Vec<u8> {
1031        use std::io::Write;
1032
1033        match encoding {
1034            ContentEncoding::Gzip => {
1035                let mut d = flate2::write::GzDecoder::new(Vec::new());
1036                d.write_all(data).unwrap();
1037                d.flush().unwrap();
1038                d.get_ref().clone()
1039            }
1040            ContentEncoding::Deflate => {
1041                let mut d = flate2::write::ZlibDecoder::new(Vec::new());
1042                d.write_all(data).unwrap();
1043                d.flush().unwrap();
1044                d.get_ref().clone()
1045            }
1046            ContentEncoding::Zstd => {
1047                let mut d = zstd::stream::write::Decoder::new(Vec::new()).unwrap();
1048                d.write_all(data).unwrap();
1049                d.flush().unwrap();
1050                d.get_ref().clone()
1051            }
1052            _ => unreachable!(),
1053        }
1054    }
1055
1056    #[test]
1057    fn encoder_flushes_when_body_is_pending() {
1058        let mut cx = Context::from_waker(std::task::Waker::noop());
1059        let small = Bytes::from_static(b"event: one\n\n");
1060        // inline writes, the flushed output is larger than a page
1061        let large = Bytes::from(random(48 * 1024));
1062
1063        for encoding in [
1064            ContentEncoding::Gzip,
1065            ContentEncoding::Deflate,
1066            ContentEncoding::Zstd,
1067        ] {
1068            let parts = Rc::new(std::cell::RefCell::new(VecDeque::new()));
1069            parts
1070                .borrow_mut()
1071                .extend([Some(small.clone()), None, Some(Bytes::new()), None]);
1072            parts
1073                .borrow_mut()
1074                .extend(large.chunks(12 * 1024).map(|c| Some(large.slice_ref(c))));
1075            parts.borrow_mut().push_back(None);
1076            let mut enc = Encoder {
1077                eof: false,
1078                pending: Bytes::new(),
1079                body: EncoderBody::Stream(Parts(parts.clone())),
1080                inner: ContentEncoder::new(encoding, None),
1081                fut: None,
1082            };
1083            let mut out = Vec::new();
1084            let mut poll = |out: &mut Vec<u8>| {
1085                let mut chunks = 0;
1086                loop {
1087                    match enc.poll_next_chunk(&mut cx) {
1088                        Poll::Ready(Some(chunk)) => {
1089                            chunks += 1;
1090                            out.extend_from_slice(&chunk.unwrap());
1091                        }
1092                        Poll::Ready(None) => return None,
1093                        Poll::Pending => return Some(chunks),
1094                    }
1095                }
1096            };
1097
1098            // everything written before the body waits can be decoded
1099            assert!(poll(&mut out).unwrap() > 0, "{encoding:?}");
1100            assert_eq!(decompress_prefix(encoding, &out), small, "{encoding:?}");
1101
1102            // no new input, nothing to flush
1103            assert_eq!(poll(&mut out), Some(0), "{encoding:?}");
1104            parts.borrow_mut().pop_front();
1105            assert_eq!(poll(&mut out), Some(0), "{encoding:?}");
1106
1107            parts.borrow_mut().pop_front();
1108            assert!(poll(&mut out).unwrap() > 1, "{encoding:?}");
1109            let data = [&small[..], &large[..]].concat();
1110            assert_eq!(decompress_prefix(encoding, &out), data, "{encoding:?}");
1111
1112            parts.borrow_mut().pop_front();
1113            assert_eq!(poll(&mut out), None);
1114            assert_eq!(decompress(encoding, &out), data, "{encoding:?}");
1115        }
1116    }
1117}