Skip to main content

actix_http/encoding/
encoder.rs

1//! Stream encoders.
2
3use std::{
4    error::Error as StdError,
5    future::Future,
6    io::{self, Write as _},
7    pin::Pin,
8    task::{Context, Poll},
9};
10
11use bytes::Bytes;
12use derive_more::Display;
13#[cfg(feature = "compress-gzip")]
14use flate2::write::{GzEncoder, ZlibEncoder};
15use futures_core::ready;
16use pin_project_lite::pin_project;
17use tokio::task::{spawn_blocking, JoinHandle};
18use tracing::trace;
19#[cfg(feature = "compress-zstd")]
20use zstd::stream::write::Encoder as ZstdEncoder;
21
22use super::Writer;
23use crate::{
24    body::{self, BodySize, MessageBody},
25    header::{self, ContentEncoding, HeaderValue, CONTENT_ENCODING},
26    ResponseHead, StatusCode,
27};
28
29const MAX_CHUNK_SIZE_ENCODE_IN_PLACE: usize = 1024;
30
31pin_project! {
32    pub struct Encoder<B> {
33        #[pin]
34        body: EncoderBody<B>,
35        encoder: Option<ContentEncoder>,
36        fut: Option<JoinHandle<Result<ContentEncoder, io::Error>>>,
37        needs_flush: bool,
38        eof: bool,
39    }
40}
41
42impl<B: MessageBody> Encoder<B> {
43    fn none() -> Self {
44        Encoder {
45            body: EncoderBody::None {
46                body: body::None::new(),
47            },
48            encoder: None,
49            fut: None,
50            needs_flush: false,
51            eof: true,
52        }
53    }
54
55    fn empty() -> Self {
56        Encoder {
57            body: EncoderBody::Full { body: Bytes::new() },
58            encoder: None,
59            fut: None,
60            needs_flush: false,
61            eof: true,
62        }
63    }
64
65    pub fn response(encoding: ContentEncoding, head: &mut ResponseHead, body: B) -> Self {
66        // no need to compress empty bodies
67        match body.size() {
68            BodySize::None => return Self::none(),
69            BodySize::Sized(0) => return Self::empty(),
70            _ => {}
71        }
72
73        let should_encode = !(head.headers().contains_key(&CONTENT_ENCODING)
74            || head.status == StatusCode::SWITCHING_PROTOCOLS
75            || head.status == StatusCode::NO_CONTENT
76            || head.status == StatusCode::PARTIAL_CONTENT
77            || encoding == ContentEncoding::Identity);
78
79        let body = match body.try_into_bytes() {
80            Ok(body) => EncoderBody::Full { body },
81            Err(body) => EncoderBody::Stream { body },
82        };
83
84        if should_encode {
85            // wrap body only if encoder is feature-enabled
86            if let Some(enc) = ContentEncoder::select(encoding) {
87                update_head(encoding, head);
88
89                return Encoder {
90                    body,
91                    encoder: Some(enc),
92                    fut: None,
93                    needs_flush: false,
94                    eof: false,
95                };
96            }
97        }
98
99        Encoder {
100            body,
101            encoder: None,
102            fut: None,
103            needs_flush: false,
104            eof: false,
105        }
106    }
107}
108
109pin_project! {
110    #[project = EncoderBodyProj]
111    enum EncoderBody<B> {
112        None { body: body::None },
113        Full { body: Bytes },
114        Stream { #[pin] body: B },
115    }
116}
117
118impl<B> MessageBody for EncoderBody<B>
119where
120    B: MessageBody,
121{
122    type Error = EncoderError;
123
124    #[inline]
125    fn size(&self) -> BodySize {
126        match self {
127            EncoderBody::None { body } => body.size(),
128            EncoderBody::Full { body } => body.size(),
129            EncoderBody::Stream { body } => body.size(),
130        }
131    }
132
133    fn poll_next(
134        self: Pin<&mut Self>,
135        cx: &mut Context<'_>,
136    ) -> Poll<Option<Result<Bytes, Self::Error>>> {
137        match self.project() {
138            EncoderBodyProj::None { body } => {
139                Pin::new(body).poll_next(cx).map_err(|err| match err {})
140            }
141            EncoderBodyProj::Full { body } => {
142                Pin::new(body).poll_next(cx).map_err(|err| match err {})
143            }
144            EncoderBodyProj::Stream { body } => body
145                .poll_next(cx)
146                .map_err(|err| EncoderError::Body(err.into())),
147        }
148    }
149
150    #[inline]
151    fn try_into_bytes(self) -> Result<Bytes, Self>
152    where
153        Self: Sized,
154    {
155        match self {
156            EncoderBody::None { body } => Ok(body.try_into_bytes().unwrap()),
157            EncoderBody::Full { body } => Ok(body.try_into_bytes().unwrap()),
158            _ => Err(self),
159        }
160    }
161}
162
163impl<B> MessageBody for Encoder<B>
164where
165    B: MessageBody,
166{
167    type Error = EncoderError;
168
169    #[inline]
170    fn size(&self) -> BodySize {
171        if self.encoder.is_some() {
172            BodySize::Stream
173        } else {
174            self.body.size()
175        }
176    }
177
178    fn poll_next(
179        self: Pin<&mut Self>,
180        cx: &mut Context<'_>,
181    ) -> Poll<Option<Result<Bytes, Self::Error>>> {
182        let mut this = self.project();
183
184        loop {
185            if *this.eof {
186                return Poll::Ready(None);
187            }
188
189            if let Some(ref mut fut) = this.fut {
190                let mut encoder = ready!(Pin::new(fut).poll(cx))
191                    .map_err(|_| {
192                        EncoderError::Io(io::Error::other(
193                            "Blocking task was cancelled unexpectedly",
194                        ))
195                    })?
196                    .map_err(EncoderError::Io)?;
197
198                let chunk = encoder.take();
199                *this.encoder = Some(encoder);
200                this.fut.take();
201
202                if !chunk.is_empty() {
203                    return Poll::Ready(Some(Ok(chunk)));
204                }
205            }
206
207            let result = match this.body.as_mut().poll_next(cx) {
208                Poll::Ready(result) => result,
209
210                Poll::Pending => {
211                    if *this.needs_flush {
212                        if let Some(encoder) = this.encoder.as_mut() {
213                            // Release buffered content when the producer pauses.
214                            encoder.flush().map_err(EncoderError::Io)?;
215                            *this.needs_flush = false;
216
217                            let chunk = encoder.take();
218
219                            if !chunk.is_empty() {
220                                return Poll::Ready(Some(Ok(chunk)));
221                            }
222                        }
223                    }
224
225                    return Poll::Pending;
226                }
227            };
228
229            match result {
230                Some(Err(err)) => return Poll::Ready(Some(Err(err))),
231
232                Some(Ok(chunk)) => {
233                    if let Some(mut encoder) = this.encoder.take() {
234                        *this.needs_flush |= !chunk.is_empty();
235
236                        if chunk.len() < MAX_CHUNK_SIZE_ENCODE_IN_PLACE {
237                            encoder.write(&chunk).map_err(EncoderError::Io)?;
238                            let chunk = encoder.take();
239                            *this.encoder = Some(encoder);
240
241                            if !chunk.is_empty() {
242                                return Poll::Ready(Some(Ok(chunk)));
243                            }
244                        } else {
245                            *this.fut = Some(spawn_blocking(move || {
246                                encoder.write(&chunk)?;
247                                Ok(encoder)
248                            }));
249                        }
250                    } else {
251                        return Poll::Ready(Some(Ok(chunk)));
252                    }
253                }
254
255                None => {
256                    if let Some(encoder) = this.encoder.take() {
257                        let chunk = encoder.finish().map_err(EncoderError::Io)?;
258
259                        if chunk.is_empty() {
260                            return Poll::Ready(None);
261                        } else {
262                            *this.eof = true;
263                            return Poll::Ready(Some(Ok(chunk)));
264                        }
265                    } else {
266                        return Poll::Ready(None);
267                    }
268                }
269            }
270        }
271    }
272
273    #[inline]
274    fn try_into_bytes(mut self) -> Result<Bytes, Self>
275    where
276        Self: Sized,
277    {
278        if self.encoder.is_some() {
279            Err(self)
280        } else {
281            match self.body.try_into_bytes() {
282                Ok(body) => Ok(body),
283                Err(body) => {
284                    self.body = body;
285                    Err(self)
286                }
287            }
288        }
289    }
290}
291
292fn update_head(encoding: ContentEncoding, head: &mut ResponseHead) {
293    head.headers_mut()
294        .insert(header::CONTENT_ENCODING, encoding.to_header_value());
295    head.headers_mut()
296        .append(header::VARY, HeaderValue::from_static("accept-encoding"));
297
298    head.no_chunking(false);
299}
300
301enum ContentEncoder {
302    #[cfg(feature = "compress-gzip")]
303    Deflate(ZlibEncoder<Writer>),
304
305    #[cfg(feature = "compress-gzip")]
306    Gzip(GzEncoder<Writer>),
307
308    #[cfg(feature = "compress-brotli")]
309    Brotli(Box<brotli::CompressorWriter<Writer>>),
310
311    // Wwe need explicit 'static lifetime here because ZstdEncoder needs a lifetime argument and we
312    // use `spawn_blocking` in `Encoder::poll_next` that requires `FnOnce() -> R + Send + 'static`.
313    #[cfg(feature = "compress-zstd")]
314    Zstd(ZstdEncoder<'static, Writer>),
315}
316
317impl ContentEncoder {
318    fn select(encoding: ContentEncoding) -> Option<Self> {
319        match encoding {
320            #[cfg(feature = "compress-gzip")]
321            ContentEncoding::Deflate => Some(ContentEncoder::Deflate(ZlibEncoder::new(
322                Writer::new(),
323                flate2::Compression::fast(),
324            ))),
325
326            #[cfg(feature = "compress-gzip")]
327            ContentEncoding::Gzip => Some(ContentEncoder::Gzip(GzEncoder::new(
328                Writer::new(),
329                flate2::Compression::fast(),
330            ))),
331
332            #[cfg(feature = "compress-brotli")]
333            ContentEncoding::Brotli => Some(ContentEncoder::Brotli(new_brotli_compressor())),
334
335            #[cfg(feature = "compress-zstd")]
336            ContentEncoding::Zstd => {
337                let encoder = ZstdEncoder::new(Writer::new(), 3).ok()?;
338                Some(ContentEncoder::Zstd(encoder))
339            }
340
341            _ => None,
342        }
343    }
344
345    #[inline]
346    pub(crate) fn take(&mut self) -> Bytes {
347        match *self {
348            #[cfg(feature = "compress-brotli")]
349            ContentEncoder::Brotli(ref mut encoder) => encoder.get_mut().take(),
350
351            #[cfg(feature = "compress-gzip")]
352            ContentEncoder::Deflate(ref mut encoder) => encoder.get_mut().take(),
353
354            #[cfg(feature = "compress-gzip")]
355            ContentEncoder::Gzip(ref mut encoder) => encoder.get_mut().take(),
356
357            #[cfg(feature = "compress-zstd")]
358            ContentEncoder::Zstd(ref mut encoder) => encoder.get_mut().take(),
359        }
360    }
361
362    fn flush(&mut self) -> Result<(), io::Error> {
363        match self {
364            #[cfg(feature = "compress-brotli")]
365            ContentEncoder::Brotli(encoder) => encoder.flush(),
366
367            #[cfg(feature = "compress-gzip")]
368            ContentEncoder::Gzip(encoder) => encoder.flush(),
369
370            #[cfg(feature = "compress-gzip")]
371            ContentEncoder::Deflate(encoder) => encoder.flush(),
372
373            #[cfg(feature = "compress-zstd")]
374            ContentEncoder::Zstd(encoder) => encoder.flush(),
375        }
376    }
377
378    fn finish(self) -> Result<Bytes, io::Error> {
379        match self {
380            #[cfg(feature = "compress-brotli")]
381            ContentEncoder::Brotli(mut encoder) => match encoder.flush() {
382                Ok(()) => Ok(encoder.into_inner().buf.freeze()),
383                Err(err) => Err(err),
384            },
385
386            #[cfg(feature = "compress-gzip")]
387            ContentEncoder::Gzip(encoder) => match encoder.finish() {
388                Ok(writer) => Ok(writer.buf.freeze()),
389                Err(err) => Err(err),
390            },
391
392            #[cfg(feature = "compress-gzip")]
393            ContentEncoder::Deflate(encoder) => match encoder.finish() {
394                Ok(writer) => Ok(writer.buf.freeze()),
395                Err(err) => Err(err),
396            },
397
398            #[cfg(feature = "compress-zstd")]
399            ContentEncoder::Zstd(encoder) => match encoder.finish() {
400                Ok(writer) => Ok(writer.buf.freeze()),
401                Err(err) => Err(err),
402            },
403        }
404    }
405
406    fn write(&mut self, data: &[u8]) -> Result<(), io::Error> {
407        match *self {
408            #[cfg(feature = "compress-brotli")]
409            ContentEncoder::Brotli(ref mut encoder) => match encoder.write_all(data) {
410                Ok(_) => Ok(()),
411                Err(err) => {
412                    trace!("Error decoding br encoding: {}", err);
413                    Err(err)
414                }
415            },
416
417            #[cfg(feature = "compress-gzip")]
418            ContentEncoder::Gzip(ref mut encoder) => match encoder.write_all(data) {
419                Ok(_) => Ok(()),
420                Err(err) => {
421                    trace!("Error decoding gzip encoding: {}", err);
422                    Err(err)
423                }
424            },
425
426            #[cfg(feature = "compress-gzip")]
427            ContentEncoder::Deflate(ref mut encoder) => match encoder.write_all(data) {
428                Ok(_) => Ok(()),
429                Err(err) => {
430                    trace!("Error decoding deflate encoding: {}", err);
431                    Err(err)
432                }
433            },
434
435            #[cfg(feature = "compress-zstd")]
436            ContentEncoder::Zstd(ref mut encoder) => match encoder.write_all(data) {
437                Ok(_) => Ok(()),
438                Err(err) => {
439                    trace!("Error decoding ztsd encoding: {}", err);
440                    Err(err)
441                }
442            },
443        }
444    }
445}
446
447#[cfg(feature = "compress-brotli")]
448fn new_brotli_compressor() -> Box<brotli::CompressorWriter<Writer>> {
449    Box::new(brotli::CompressorWriter::new(
450        Writer::new(),
451        32 * 1024, // 32 KiB buffer
452        3,         // BROTLI_PARAM_QUALITY
453        22,        // BROTLI_PARAM_LGWIN
454    ))
455}
456
457#[derive(Debug, Display)]
458#[non_exhaustive]
459pub enum EncoderError {
460    /// Wrapped body stream error.
461    #[display("body")]
462    Body(Box<dyn StdError>),
463
464    /// Generic I/O error.
465    #[display("io")]
466    Io(io::Error),
467}
468
469impl StdError for EncoderError {
470    fn source(&self) -> Option<&(dyn StdError + 'static)> {
471        match self {
472            EncoderError::Body(err) => Some(&**err),
473            EncoderError::Io(err) => Some(err),
474        }
475    }
476}
477
478impl From<EncoderError> for crate::Error {
479    fn from(err: EncoderError) -> Self {
480        crate::Error::new_encoder().with_cause(err)
481    }
482}