1use 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 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 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 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 #[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, 3, 22, ))
455}
456
457#[derive(Debug, Display)]
458#[non_exhaustive]
459pub enum EncoderError {
460 #[display("body")]
462 Body(Box<dyn StdError>),
463
464 #[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}