1use 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
17const MIN_SIZE: u64 = 1024;
21
22const ZSTD_WINDOW_LOG: u32 = 19;
27
28pub struct Encoder<B> {
32 eof: bool,
33 body: EncoderBody<B>,
34 pending: Bytes,
36 inner: Option<ContentEncoder>,
37 fut: Option<BlockingResult<Result<ContentEncoder, io::Error>>>,
38}
39
40impl<B: MessageBody> Encoder<B> {
41 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 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 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 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 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
250const GZIP_HEADER: [u8; 10] = [0x1f, 0x8b, 8, 0, 0, 0, 0, 0, 4, 255];
253
254struct ContentEncoder {
256 codec: Codec,
257 buf: BytesMut,
258 chunks: VecDeque<Bytes>,
259 unflushed: bool,
261}
262
263#[derive(Copy, Clone, PartialEq, Eq)]
264enum Op {
265 Write,
266 Flush,
268 Finish,
270}
271
272enum Codec {
273 Deflate(Compress),
274 Gzip(Compress, Crc),
276 Zstd(CCtx<'static>),
277 Done,
279}
280
281impl ContentEncoder {
282 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 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 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 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 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 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 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 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 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 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 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 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 self.codec = Codec::Done;
423 Ok(())
424 }
425
426 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 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 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 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 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 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 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 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 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 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 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 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 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 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 assert!(window[0] < 2560 * 1024, "{}", window[0]);
937 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 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 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 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 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 assert!(poll(&mut out).unwrap() > 0, "{encoding:?}");
1100 assert_eq!(decompress_prefix(encoding, &out), small, "{encoding:?}");
1101
1102 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}