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