1#![cfg_attr(docsrs, feature(doc_cfg))]
2#![cfg_attr(test, allow(clippy::unwrap_used))]
3
4use std::fmt::{self, Display, Formatter};
69use std::str::FromStr;
70use std::sync::LazyLock;
71
72use indexmap::IndexMap;
73use salvo_core::http::body::ResBody;
74use salvo_core::http::header::{
75 ACCEPT_ENCODING, CONTENT_ENCODING, CONTENT_LENGTH, CONTENT_TYPE, HeaderValue,
76};
77use salvo_core::http::headers::{ContentLength, HeaderMapExt};
78use salvo_core::http::{self, Method, Mime, StatusCode, append_vary_header, mime};
79use salvo_core::{Depot, FlowCtrl, Handler, Request, Response, async_trait};
80
81mod encoder;
82mod stream;
83use encoder::Encoder;
84use stream::EncodeStream;
85
86#[non_exhaustive]
88#[derive(Clone, Copy, Default, Debug, Eq, PartialEq)]
89pub enum CompressionLevel {
90 Fastest,
92 Minsize,
94 #[default]
96 Default,
97 Precise(u32),
102}
103
104#[derive(Eq, PartialEq, Clone, Copy, Debug, Hash)]
106#[non_exhaustive]
107pub enum CompressionAlgo {
108 #[cfg(feature = "brotli")]
110 #[cfg_attr(docsrs, doc(cfg(feature = "brotli")))]
111 Brotli,
112
113 #[cfg(feature = "deflate")]
115 #[cfg_attr(docsrs, doc(cfg(feature = "deflate")))]
116 Deflate,
117
118 #[cfg(feature = "gzip")]
120 #[cfg_attr(docsrs, doc(cfg(feature = "gzip")))]
121 Gzip,
122
123 #[cfg(feature = "zstd")]
125 #[cfg_attr(docsrs, doc(cfg(feature = "zstd")))]
126 Zstd,
127}
128
129impl FromStr for CompressionAlgo {
130 type Err = String;
131
132 fn from_str(s: &str) -> Result<Self, Self::Err> {
133 match s {
134 #[cfg(feature = "brotli")]
135 "br" => Ok(Self::Brotli),
136 #[cfg(feature = "brotli")]
137 "brotli" => Ok(Self::Brotli),
138
139 #[cfg(feature = "deflate")]
140 "deflate" => Ok(Self::Deflate),
141
142 #[cfg(feature = "gzip")]
143 "gzip" => Ok(Self::Gzip),
144
145 #[cfg(feature = "zstd")]
146 "zstd" => Ok(Self::Zstd),
147 _ => Err(format!("unknown compression algorithm: {s}")),
148 }
149 }
150}
151
152impl Display for CompressionAlgo {
153 #[allow(unreachable_patterns)]
154 #[allow(unused_variables)]
155 fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
156 match self {
157 #[cfg(feature = "brotli")]
158 Self::Brotli => write!(f, "br"),
159 #[cfg(feature = "deflate")]
160 Self::Deflate => write!(f, "deflate"),
161 #[cfg(feature = "gzip")]
162 Self::Gzip => write!(f, "gzip"),
163 #[cfg(feature = "zstd")]
164 Self::Zstd => write!(f, "zstd"),
165 _ => unreachable!(),
166 }
167 }
168}
169
170impl From<CompressionAlgo> for HeaderValue {
171 #[inline]
172 fn from(algo: CompressionAlgo) -> Self {
173 match algo {
174 #[cfg(feature = "brotli")]
175 CompressionAlgo::Brotli => Self::from_static("br"),
176 #[cfg(feature = "deflate")]
177 CompressionAlgo::Deflate => Self::from_static("deflate"),
178 #[cfg(feature = "gzip")]
179 CompressionAlgo::Gzip => Self::from_static("gzip"),
180 #[cfg(feature = "zstd")]
181 CompressionAlgo::Zstd => Self::from_static("zstd"),
182 }
183 }
184}
185
186#[derive(Clone, Debug)]
188#[non_exhaustive]
189pub struct Compression {
190 pub algos: IndexMap<CompressionAlgo, CompressionLevel>,
192 pub content_types: Vec<Mime>,
194 pub min_length: usize,
202 pub force_priority: bool,
205}
206
207static DEFAULT_CONTENT_TYPES: LazyLock<Vec<Mime>> = LazyLock::new(|| {
208 vec![
209 mime::TEXT_STAR,
210 mime::APPLICATION_JAVASCRIPT,
211 mime::APPLICATION_JSON,
212 mime::IMAGE_SVG,
213 "application/wasm".parse().expect("invalid mime type"),
214 "application/xml".parse().expect("invalid mime type"),
215 "application/rss+xml".parse().expect("invalid mime type"),
216 ]
217});
218
219impl Default for Compression {
220 fn default() -> Self {
221 #[allow(unused_mut)]
222 let mut algos = IndexMap::new();
223 #[cfg(feature = "zstd")]
224 algos.insert(CompressionAlgo::Zstd, CompressionLevel::Default);
225 #[cfg(feature = "gzip")]
226 algos.insert(CompressionAlgo::Gzip, CompressionLevel::Default);
227 #[cfg(feature = "deflate")]
228 algos.insert(CompressionAlgo::Deflate, CompressionLevel::Default);
229 #[cfg(feature = "brotli")]
230 algos.insert(CompressionAlgo::Brotli, CompressionLevel::Default);
231 Self {
232 algos,
233 content_types: DEFAULT_CONTENT_TYPES.clone(),
234 min_length: 1024,
235 force_priority: false,
236 }
237 }
238}
239
240impl Compression {
241 #[inline]
243 #[must_use]
244 pub fn new() -> Self {
245 Default::default()
246 }
247
248 #[inline]
250 #[must_use]
251 pub fn disable_all(mut self) -> Self {
252 self.algos.clear();
253 self
254 }
255
256 #[cfg(feature = "gzip")]
258 #[cfg_attr(docsrs, doc(cfg(feature = "gzip")))]
259 #[inline]
260 #[must_use]
261 pub fn enable_gzip(mut self, level: CompressionLevel) -> Self {
262 self.algos.insert(CompressionAlgo::Gzip, level);
263 self
264 }
265 #[cfg(feature = "gzip")]
267 #[cfg_attr(docsrs, doc(cfg(feature = "gzip")))]
268 #[inline]
269 #[must_use]
270 pub fn disable_gzip(mut self) -> Self {
271 self.algos.shift_remove(&CompressionAlgo::Gzip);
272 self
273 }
274 #[cfg(feature = "zstd")]
276 #[cfg_attr(docsrs, doc(cfg(feature = "zstd")))]
277 #[inline]
278 #[must_use]
279 pub fn enable_zstd(mut self, level: CompressionLevel) -> Self {
280 self.algos.insert(CompressionAlgo::Zstd, level);
281 self
282 }
283 #[cfg(feature = "zstd")]
285 #[cfg_attr(docsrs, doc(cfg(feature = "zstd")))]
286 #[inline]
287 #[must_use]
288 pub fn disable_zstd(mut self) -> Self {
289 self.algos.shift_remove(&CompressionAlgo::Zstd);
290 self
291 }
292 #[cfg(feature = "brotli")]
294 #[cfg_attr(docsrs, doc(cfg(feature = "brotli")))]
295 #[inline]
296 #[must_use]
297 pub fn enable_brotli(mut self, level: CompressionLevel) -> Self {
298 self.algos.insert(CompressionAlgo::Brotli, level);
299 self
300 }
301 #[cfg(feature = "brotli")]
303 #[cfg_attr(docsrs, doc(cfg(feature = "brotli")))]
304 #[inline]
305 #[must_use]
306 pub fn disable_brotli(mut self) -> Self {
307 self.algos.shift_remove(&CompressionAlgo::Brotli);
308 self
309 }
310
311 #[cfg(feature = "deflate")]
313 #[cfg_attr(docsrs, doc(cfg(feature = "deflate")))]
314 #[inline]
315 #[must_use]
316 pub fn enable_deflate(mut self, level: CompressionLevel) -> Self {
317 self.algos.insert(CompressionAlgo::Deflate, level);
318 self
319 }
320
321 #[cfg(feature = "deflate")]
323 #[cfg_attr(docsrs, doc(cfg(feature = "deflate")))]
324 #[inline]
325 #[must_use]
326 pub fn disable_deflate(mut self) -> Self {
327 self.algos.shift_remove(&CompressionAlgo::Deflate);
328 self
329 }
330
331 #[inline]
337 #[must_use]
338 pub fn min_length(mut self, size: usize) -> Self {
339 self.min_length = size;
340 self
341 }
342 #[inline]
344 #[must_use]
345 pub fn force_priority(mut self, force_priority: bool) -> Self {
346 self.force_priority = force_priority;
347 self
348 }
349
350 #[inline]
352 #[must_use]
353 pub fn content_types(mut self, content_types: &[Mime]) -> Self {
354 self.content_types = content_types.to_vec();
355 self
356 }
357
358 fn content_length_is_below_minimum(&self, res: &Response) -> bool {
359 if self.min_length == 0 {
360 return false;
361 }
362 let Some(ContentLength(length)) = res.headers().typed_get::<ContentLength>() else {
363 return false;
364 };
365 match u64::try_from(self.min_length) {
366 Ok(min_length) => length < min_length,
367 Err(_) => true,
368 }
369 }
370
371 fn negotiate(
372 &self,
373 req: &Request,
374 res: &Response,
375 ) -> Option<(CompressionAlgo, CompressionLevel)> {
376 if !self.content_types.is_empty() {
377 let content_type = res
378 .headers()
379 .get(CONTENT_TYPE)
380 .and_then(|v| v.to_str().ok())
381 .unwrap_or_default();
382 if content_type.is_empty() {
383 return None;
384 }
385 if let Ok(content_type) = content_type.parse::<Mime>() {
386 if !self.content_types.iter().any(|citem| {
387 citem.type_() == content_type.type_()
388 && (citem.subtype() == "*" || citem.subtype() == content_type.subtype())
389 }) {
390 return None;
391 }
392 } else {
393 return None;
394 }
395 }
396 let header = req
397 .headers()
398 .get(ACCEPT_ENCODING)
399 .and_then(|v| v.to_str().ok())?;
400
401 let accept_list = http::parse_accept_encoding(header);
402
403 let wildcard_q: Option<u8> = accept_list
409 .iter()
410 .find(|(name, _)| name == "*")
411 .map(|(_, q)| *q);
412
413 let parsed: smallvec::SmallVec<[(CompressionAlgo, u8); 4]> = accept_list
419 .iter()
420 .filter(|(name, _)| name != "*")
421 .filter_map(|(name, q)| name.parse::<CompressionAlgo>().ok().map(|algo| (algo, *q)))
422 .collect();
423
424 let is_accepted =
425 |algo: &CompressionAlgo| -> bool { parsed.iter().any(|(a, q)| a == algo && *q > 0) };
426 let is_rejected =
427 |algo: &CompressionAlgo| -> bool { parsed.iter().any(|(a, q)| a == algo && *q == 0) };
428
429 if self.force_priority {
430 self.algos
432 .iter()
433 .find(|(algo, _)| {
434 !is_rejected(algo) && (is_accepted(algo) || wildcard_q.is_some_and(|q| q > 0))
435 })
436 .map(|(algo, level)| (*algo, *level))
437 } else {
438 let result = parsed
440 .iter()
441 .filter(|(_, q)| *q > 0)
442 .find_map(|(algo, _)| self.algos.get(algo).map(|level| (*algo, *level)));
443
444 if result.is_some() {
445 return result;
446 }
447
448 if wildcard_q.is_some_and(|q| q > 0) {
450 self.algos
451 .iter()
452 .find(|(algo, _)| !is_rejected(algo))
453 .map(|(algo, level)| (*algo, *level))
454 } else {
455 None
456 }
457 }
458 }
459}
460
461#[async_trait]
462impl Handler for Compression {
463 async fn handle(
464 &self,
465 req: &mut Request,
466 depot: &mut Depot,
467 res: &mut Response,
468 ctrl: &mut FlowCtrl,
469 ) {
470 ctrl.call_next(req, depot, res).await;
471 if ctrl.is_ceased() || res.headers().contains_key(CONTENT_ENCODING) {
472 return;
473 }
474
475 if let Some(StatusCode::SWITCHING_PROTOCOLS | StatusCode::NO_CONTENT) = res.status_code {
476 return;
477 }
478
479 match res.take_body() {
480 ResBody::None => {
481 if req.method() != Method::HEAD {
485 return;
486 }
487 if res.headers().typed_get::<ContentLength>().is_none() {
488 return;
489 }
490 if self.content_length_is_below_minimum(res) {
491 return;
492 }
493 if let Some((algo, _level)) = self.negotiate(req, res) {
494 res.headers_mut().insert(CONTENT_ENCODING, algo.into());
495 } else {
496 return;
497 }
498 }
499 ResBody::Once(bytes) => {
500 if self.min_length > 0 && bytes.len() < self.min_length {
501 res.body(ResBody::Once(bytes));
502 return;
503 }
504 if let Some((algo, level)) = self.negotiate(req, res) {
505 res.stream(EncodeStream::new(algo, level, Some(bytes)));
506 res.headers_mut().insert(CONTENT_ENCODING, algo.into());
507 } else {
508 res.body(ResBody::Once(bytes));
509 return;
510 }
511 }
512 ResBody::Chunks(chunks) => {
513 if self.min_length > 0 {
514 let len: usize = chunks.iter().map(|c| c.len()).sum();
515 if len < self.min_length {
516 res.body(ResBody::Chunks(chunks));
517 return;
518 }
519 }
520 if let Some((algo, level)) = self.negotiate(req, res) {
521 res.stream(EncodeStream::new(algo, level, chunks));
522 res.headers_mut().insert(CONTENT_ENCODING, algo.into());
523 } else {
524 res.body(ResBody::Chunks(chunks));
525 return;
526 }
527 }
528 ResBody::Hyper(body) => {
529 if self.content_length_is_below_minimum(res) {
530 res.body(ResBody::Hyper(body));
531 return;
532 }
533 if let Some((algo, level)) = self.negotiate(req, res) {
534 res.stream(EncodeStream::new(algo, level, body));
535 res.headers_mut().insert(CONTENT_ENCODING, algo.into());
536 } else {
537 res.body(ResBody::Hyper(body));
538 return;
539 }
540 }
541 ResBody::Stream(body) => {
542 let body = body.into_inner();
543 if self.content_length_is_below_minimum(res) {
544 res.body(ResBody::stream(body));
545 return;
546 }
547 if let Some((algo, level)) = self.negotiate(req, res) {
548 res.stream(EncodeStream::new(algo, level, body));
549 res.headers_mut().insert(CONTENT_ENCODING, algo.into());
550 } else {
551 res.body(ResBody::stream(body));
552 return;
553 }
554 }
555 body => {
556 res.body(body);
557 return;
558 }
559 }
560 res.headers_mut().remove(CONTENT_LENGTH);
561 append_vary_header(res.headers_mut(), "accept-encoding");
562 }
563}
564
565#[cfg(test)]
566mod tests {
567 use salvo_core::http::header::VARY;
568 use salvo_core::prelude::*;
569 use salvo_core::test::{ResponseExt, TestClient};
570
571 use super::*;
572
573 #[handler]
574 async fn hello() -> &'static str {
575 "hello"
576 }
577
578 #[tokio::test]
579 async fn test_gzip() {
580 let comp_handler = Compression::new().min_length(1);
581 let router = Router::with_hoop(comp_handler).push(Router::with_path("hello").get(hello));
582
583 let mut res = TestClient::get("http://127.0.0.1:5801/hello")
584 .add_header(ACCEPT_ENCODING, "gzip", true)
585 .send(router)
586 .await;
587 assert_eq!(res.headers().get(CONTENT_ENCODING).unwrap(), "gzip");
588 let content = res.take_string().await.unwrap();
589 assert_eq!(content, "hello");
590 }
591
592 #[tokio::test]
593 async fn test_brotli() {
594 let comp_handler = Compression::new().min_length(1);
595 let router = Router::with_hoop(comp_handler).push(Router::with_path("hello").get(hello));
596
597 let mut res = TestClient::get("http://127.0.0.1:5801/hello")
598 .add_header(ACCEPT_ENCODING, "br", true)
599 .send(router)
600 .await;
601 assert_eq!(res.headers().get(CONTENT_ENCODING).unwrap(), "br");
602 let content = res.take_string().await.unwrap();
603 assert_eq!(content, "hello");
604 }
605
606 #[tokio::test]
607 async fn test_deflate() {
608 let comp_handler = Compression::new().min_length(1);
609 let router = Router::with_hoop(comp_handler).push(Router::with_path("hello").get(hello));
610
611 let mut res = TestClient::get("http://127.0.0.1:5801/hello")
612 .add_header(ACCEPT_ENCODING, "deflate", true)
613 .send(router)
614 .await;
615 assert_eq!(res.headers().get(CONTENT_ENCODING).unwrap(), "deflate");
616 let content = res.take_string().await.unwrap();
617 assert_eq!(content, "hello");
618 }
619
620 #[tokio::test]
621 async fn test_zstd() {
622 let comp_handler = Compression::new().min_length(1);
623 let router = Router::with_hoop(comp_handler).push(Router::with_path("hello").get(hello));
624
625 let mut res = TestClient::get("http://127.0.0.1:5801/hello")
626 .add_header(ACCEPT_ENCODING, "zstd", true)
627 .send(router)
628 .await;
629 assert_eq!(res.headers().get(CONTENT_ENCODING).unwrap(), "zstd");
630 let content = res.take_string().await.unwrap();
631 assert_eq!(content, "hello");
632 }
633
634 #[tokio::test]
635 async fn test_min_length_not_compress() {
636 let comp_handler = Compression::new().min_length(10);
637 let router = Router::with_hoop(comp_handler).push(Router::with_path("hello").get(hello));
638
639 let res = TestClient::get("http://127.0.0.1:5801/hello")
640 .add_header(ACCEPT_ENCODING, "gzip", true)
641 .send(router)
642 .await;
643 assert!(res.headers().get(CONTENT_ENCODING).is_none());
644 }
645
646 #[tokio::test]
647 async fn test_min_length_should_compress() {
648 let comp_handler = Compression::new().min_length(1);
649 let router = Router::with_hoop(comp_handler).push(Router::with_path("hello").get(hello));
650
651 let res = TestClient::get("http://127.0.0.1:5801/hello")
652 .add_header(ACCEPT_ENCODING, "gzip", true)
653 .send(router)
654 .await;
655 assert!(res.headers().get(CONTENT_ENCODING).is_some());
656 }
657
658 #[handler]
659 async fn hello_html(res: &mut Response) {
660 res.render(Text::Html("<html><body>hello</body></html>"));
661 }
662 #[tokio::test]
663 async fn test_content_types_should_compress() {
664 let comp_handler = Compression::new()
665 .min_length(1)
666 .content_types(&[mime::TEXT_HTML]);
667 let router =
668 Router::with_hoop(comp_handler).push(Router::with_path("hello").get(hello_html));
669
670 let res = TestClient::get("http://127.0.0.1:5801/hello")
671 .add_header(ACCEPT_ENCODING, "gzip", true)
672 .send(router)
673 .await;
674 assert!(res.headers().get(CONTENT_ENCODING).is_some());
675 }
676
677 #[tokio::test]
678 async fn test_content_types_not_compress() {
679 let comp_handler = Compression::new()
680 .min_length(1)
681 .content_types(&[mime::APPLICATION_JSON]);
682 let router =
683 Router::with_hoop(comp_handler).push(Router::with_path("hello").get(hello_html));
684
685 let res = TestClient::get("http://127.0.0.1:5801/hello")
686 .add_header(ACCEPT_ENCODING, "gzip", true)
687 .send(router)
688 .await;
689 assert!(res.headers().get(CONTENT_ENCODING).is_none());
690 }
691
692 #[tokio::test]
693 async fn test_q_value_preference() {
694 let comp_handler = Compression::new().min_length(1);
696 let router = Router::with_hoop(comp_handler).push(Router::with_path("hello").get(hello));
697
698 let mut res = TestClient::get("http://127.0.0.1:5801/hello")
699 .add_header(ACCEPT_ENCODING, "gzip;q=0.5, br;q=1.0", true)
700 .send(router)
701 .await;
702 assert_eq!(res.headers().get(CONTENT_ENCODING).unwrap(), "br");
703 let content = res.take_string().await.unwrap();
704 assert_eq!(content, "hello");
705 }
706
707 #[tokio::test]
708 async fn test_q_value_zero_rejects_algo() {
709 let comp_handler = Compression::new().min_length(1);
711 let router = Router::with_hoop(comp_handler).push(Router::with_path("hello").get(hello));
712
713 let mut res = TestClient::get("http://127.0.0.1:5801/hello")
714 .add_header(ACCEPT_ENCODING, "gzip;q=0, br", true)
715 .send(router)
716 .await;
717 assert_eq!(res.headers().get(CONTENT_ENCODING).unwrap(), "br");
718 let content = res.take_string().await.unwrap();
719 assert_eq!(content, "hello");
720 }
721
722 #[tokio::test]
723 async fn test_identity_only_no_compression() {
724 let comp_handler = Compression::new().min_length(1);
726 let router = Router::with_hoop(comp_handler).push(Router::with_path("hello").get(hello));
727
728 let res = TestClient::get("http://127.0.0.1:5801/hello")
729 .add_header(ACCEPT_ENCODING, "identity", true)
730 .send(router)
731 .await;
732 assert!(res.headers().get(CONTENT_ENCODING).is_none());
733 }
734
735 #[tokio::test]
736 async fn test_wildcard_uses_server_algo() {
737 let comp_handler = Compression::new()
739 .disable_all()
740 .enable_gzip(CompressionLevel::Default)
741 .min_length(1);
742 let router = Router::with_hoop(comp_handler).push(Router::with_path("hello").get(hello));
743
744 let res = TestClient::get("http://127.0.0.1:5801/hello")
745 .add_header(ACCEPT_ENCODING, "*", true)
746 .send(router)
747 .await;
748 assert_eq!(res.headers().get(CONTENT_ENCODING).unwrap(), "gzip");
749 }
750
751 #[tokio::test]
752 async fn test_wildcard_excludes_rejected_algo() {
753 let comp_handler = Compression::new()
755 .disable_all()
756 .enable_gzip(CompressionLevel::Default)
757 .enable_brotli(CompressionLevel::Default)
758 .min_length(1);
759 let router = Router::with_hoop(comp_handler).push(Router::with_path("hello").get(hello));
760
761 let res = TestClient::get("http://127.0.0.1:5801/hello")
762 .add_header(ACCEPT_ENCODING, "*, gzip;q=0", true)
763 .send(router)
764 .await;
765 assert_eq!(res.headers().get(CONTENT_ENCODING).unwrap(), "br");
766 }
767
768 #[tokio::test]
769 async fn test_multiple_wildcards_take_highest_q() {
770 let comp_handler = Compression::new()
775 .disable_all()
776 .enable_gzip(CompressionLevel::Default)
777 .min_length(1);
778 let router = Router::with_hoop(comp_handler).push(Router::with_path("hello").get(hello));
779
780 let res = TestClient::get("http://127.0.0.1:5801/hello")
781 .add_header(ACCEPT_ENCODING, "*;q=1, *;q=0", true)
782 .send(router)
783 .await;
784 assert_eq!(res.headers().get(CONTENT_ENCODING).unwrap(), "gzip");
785 }
786
787 #[tokio::test]
788 async fn test_single_content_encoding_header() {
789 let comp_handler = Compression::new().min_length(1);
791 let router = Router::with_hoop(comp_handler).push(Router::with_path("hello").get(hello));
792
793 let res = TestClient::get("http://127.0.0.1:5801/hello")
794 .add_header(ACCEPT_ENCODING, "gzip", true)
795 .send(router)
796 .await;
797 let count = res.headers().get_all(CONTENT_ENCODING).iter().count();
798 assert_eq!(count, 1, "must have exactly one Content-Encoding header");
799 }
800
801 #[tokio::test]
802 async fn test_vary_accept_encoding_is_not_duplicated() {
803 #[handler]
804 async fn hello_with_vary(res: &mut Response) {
805 res.headers_mut()
806 .insert(VARY, HeaderValue::from_static("Accept-Encoding"));
807 res.render(Text::Plain("hello"));
808 }
809
810 let comp_handler = Compression::new().min_length(1);
811 let router =
812 Router::with_hoop(comp_handler).push(Router::with_path("hello").get(hello_with_vary));
813
814 let res = TestClient::get("http://127.0.0.1:5801/hello")
815 .add_header(ACCEPT_ENCODING, "gzip", true)
816 .send(router)
817 .await;
818
819 assert_eq!(res.headers().get(CONTENT_ENCODING).unwrap(), "gzip");
820 let vary_accept_encoding_count = res
821 .headers()
822 .get_all(VARY)
823 .iter()
824 .filter_map(|value| value.to_str().ok())
825 .flat_map(|value| value.split(','))
826 .filter(|value| value.trim().eq_ignore_ascii_case("accept-encoding"))
827 .count();
828 assert_eq!(vary_accept_encoding_count, 1);
829 }
830
831 #[tokio::test]
832 async fn test_head_matches_streamed_get_compression_threshold() {
833 #[handler]
834 async fn hello_stream(res: &mut Response) {
835 res.status_code(StatusCode::OK);
836 res.headers_mut()
837 .insert(CONTENT_TYPE, HeaderValue::from_static("text/plain"));
838 res.headers_mut()
839 .insert(CONTENT_LENGTH, HeaderValue::from_static("5"));
840 res.stream(futures_util::stream::once(async {
841 Ok::<_, std::io::Error>("hello")
842 }));
843 }
844
845 #[handler]
846 async fn hello_head(res: &mut Response) {
847 res.status_code(StatusCode::OK);
848 res.headers_mut()
849 .insert(CONTENT_TYPE, HeaderValue::from_static("text/plain"));
850 res.headers_mut()
851 .insert(CONTENT_LENGTH, HeaderValue::from_static("5"));
852 }
853
854 let compressed_router = Router::with_hoop(Compression::new().min_length(1)).push(
855 Router::with_path("hello")
856 .get(hello_stream)
857 .head(hello_head),
858 );
859 let compressed_service = Service::new(compressed_router);
860
861 let compressed_get = TestClient::get("http://127.0.0.1:5801/hello")
862 .add_header(ACCEPT_ENCODING, "gzip", true)
863 .send(&compressed_service)
864 .await;
865 let mut compressed_head = TestClient::head("http://127.0.0.1:5801/hello")
866 .add_header(ACCEPT_ENCODING, "gzip", true)
867 .send(&compressed_service)
868 .await;
869
870 assert_eq!(
871 compressed_get.headers().get(CONTENT_ENCODING).unwrap(),
872 "gzip"
873 );
874 assert_eq!(
875 compressed_head.headers().get(CONTENT_ENCODING),
876 compressed_get.headers().get(CONTENT_ENCODING)
877 );
878 assert_eq!(
879 compressed_head.headers().get(VARY),
880 compressed_get.headers().get(VARY)
881 );
882 assert!(compressed_head.headers().get(CONTENT_LENGTH).is_none());
883 assert_eq!(compressed_head.take_string().await.unwrap(), "");
884
885 let uncompressed_router = Router::with_hoop(Compression::new().min_length(10)).push(
886 Router::with_path("hello")
887 .get(hello_stream)
888 .head(hello_head),
889 );
890 let uncompressed_service = Service::new(uncompressed_router);
891
892 let uncompressed_get = TestClient::get("http://127.0.0.1:5801/hello")
893 .add_header(ACCEPT_ENCODING, "gzip", true)
894 .send(&uncompressed_service)
895 .await;
896 let mut uncompressed_head = TestClient::head("http://127.0.0.1:5801/hello")
897 .add_header(ACCEPT_ENCODING, "gzip", true)
898 .send(&uncompressed_service)
899 .await;
900
901 assert!(uncompressed_get.headers().get(CONTENT_ENCODING).is_none());
902 assert!(uncompressed_head.headers().get(CONTENT_ENCODING).is_none());
903 assert_eq!(
904 uncompressed_head.headers().get(VARY),
905 uncompressed_get.headers().get(VARY)
906 );
907 assert_eq!(
908 uncompressed_head.headers().get(CONTENT_LENGTH),
909 uncompressed_get.headers().get(CONTENT_LENGTH)
910 );
911 assert_eq!(uncompressed_head.take_string().await.unwrap(), "");
912 }
913
914 #[tokio::test]
915 async fn test_head_matches_buffered_get_below_min_length() {
916 #[handler]
917 async fn hello_head(res: &mut Response) {
918 res.status_code(StatusCode::OK);
919 res.headers_mut()
920 .insert(CONTENT_TYPE, HeaderValue::from_static("text/plain"));
921 res.headers_mut()
922 .insert(CONTENT_LENGTH, HeaderValue::from_static("5"));
923 }
924
925 let router = Router::with_hoop(Compression::new().min_length(10))
926 .push(Router::with_path("hello").get(hello).head(hello_head));
927 let service = Service::new(router);
928
929 let get = TestClient::get("http://127.0.0.1:5801/hello")
930 .add_header(ACCEPT_ENCODING, "gzip", true)
931 .send(&service)
932 .await;
933 let mut head = TestClient::head("http://127.0.0.1:5801/hello")
934 .add_header(ACCEPT_ENCODING, "gzip", true)
935 .send(&service)
936 .await;
937
938 assert!(get.headers().get(CONTENT_ENCODING).is_none());
939 assert!(head.headers().get(CONTENT_ENCODING).is_none());
940 assert_eq!(head.headers().get(VARY), get.headers().get(VARY));
941 assert_eq!(head.headers().get(CONTENT_LENGTH).unwrap(), "5");
942 assert_eq!(head.take_string().await.unwrap(), "");
943 }
944
945 #[tokio::test]
946 async fn test_force_priority() {
947 let comp_handler = Compression::new()
948 .disable_all()
949 .enable_brotli(CompressionLevel::Default)
950 .enable_gzip(CompressionLevel::Default)
951 .min_length(1)
952 .force_priority(true);
953 let router = Router::with_hoop(comp_handler).push(Router::with_path("hello").get(hello));
954
955 let mut res = TestClient::get("http://127.0.0.1:5801/hello")
956 .add_header(ACCEPT_ENCODING, "gzip, br", true)
957 .send(router)
958 .await;
959 assert_eq!(res.headers().get(CONTENT_ENCODING).unwrap(), "br");
960 let content = res.take_string().await.unwrap();
961 assert_eq!(content, "hello");
962 }
963
964 #[test]
966 fn test_compression_level_default() {
967 let level: CompressionLevel = Default::default();
968 assert_eq!(level, CompressionLevel::Default);
969 }
970
971 #[test]
972 fn test_compression_level_fastest() {
973 let level = CompressionLevel::Fastest;
974 assert_eq!(level, CompressionLevel::Fastest);
975 }
976
977 #[test]
978 fn test_compression_level_minsize() {
979 let level = CompressionLevel::Minsize;
980 assert_eq!(level, CompressionLevel::Minsize);
981 }
982
983 #[test]
984 fn test_compression_level_precise() {
985 let level = CompressionLevel::Precise(5);
986 assert_eq!(level, CompressionLevel::Precise(5));
987 }
988
989 #[test]
990 fn test_compression_level_clone() {
991 let level = CompressionLevel::Fastest;
992 let cloned = level;
993 assert_eq!(level, cloned);
994 }
995
996 #[test]
997 fn test_compression_level_copy() {
998 let level = CompressionLevel::Default;
999 let copied = level;
1000 assert_eq!(level, copied);
1001 }
1002
1003 #[test]
1004 fn test_compression_level_debug() {
1005 let level = CompressionLevel::Fastest;
1006 let debug_str = format!("{level:?}");
1007 assert!(debug_str.contains("Fastest"));
1008 }
1009
1010 #[cfg(feature = "gzip")]
1012 #[test]
1013 fn test_compression_algo_gzip_from_str() {
1014 let algo: CompressionAlgo = "gzip".parse().unwrap();
1015 assert_eq!(algo, CompressionAlgo::Gzip);
1016 }
1017
1018 #[cfg(feature = "brotli")]
1019 #[test]
1020 fn test_compression_algo_brotli_from_str() {
1021 let algo: CompressionAlgo = "br".parse().unwrap();
1022 assert_eq!(algo, CompressionAlgo::Brotli);
1023
1024 let algo: CompressionAlgo = "brotli".parse().unwrap();
1025 assert_eq!(algo, CompressionAlgo::Brotli);
1026 }
1027
1028 #[cfg(feature = "deflate")]
1029 #[test]
1030 fn test_compression_algo_deflate_from_str() {
1031 let algo: CompressionAlgo = "deflate".parse().unwrap();
1032 assert_eq!(algo, CompressionAlgo::Deflate);
1033 }
1034
1035 #[cfg(feature = "zstd")]
1036 #[test]
1037 fn test_compression_algo_zstd_from_str() {
1038 let algo: CompressionAlgo = "zstd".parse().unwrap();
1039 assert_eq!(algo, CompressionAlgo::Zstd);
1040 }
1041
1042 #[test]
1043 fn test_compression_algo_unknown_from_str() {
1044 let result: Result<CompressionAlgo, _> = "unknown".parse();
1045 assert!(result.is_err());
1046 assert!(
1047 result
1048 .unwrap_err()
1049 .contains("unknown compression algorithm")
1050 );
1051 }
1052
1053 #[cfg(feature = "gzip")]
1054 #[test]
1055 fn test_compression_algo_gzip_display() {
1056 let algo = CompressionAlgo::Gzip;
1057 assert_eq!(format!("{algo}"), "gzip");
1058 }
1059
1060 #[cfg(feature = "brotli")]
1061 #[test]
1062 fn test_compression_algo_brotli_display() {
1063 let algo = CompressionAlgo::Brotli;
1064 assert_eq!(format!("{algo}"), "br");
1065 }
1066
1067 #[cfg(feature = "deflate")]
1068 #[test]
1069 fn test_compression_algo_deflate_display() {
1070 let algo = CompressionAlgo::Deflate;
1071 assert_eq!(format!("{algo}"), "deflate");
1072 }
1073
1074 #[cfg(feature = "zstd")]
1075 #[test]
1076 fn test_compression_algo_zstd_display() {
1077 let algo = CompressionAlgo::Zstd;
1078 assert_eq!(format!("{algo}"), "zstd");
1079 }
1080
1081 #[cfg(feature = "gzip")]
1082 #[test]
1083 fn test_compression_algo_into_header_value() {
1084 let algo = CompressionAlgo::Gzip;
1085 let header: HeaderValue = algo.into();
1086 assert_eq!(header, "gzip");
1087 }
1088
1089 #[test]
1090 fn test_compression_algo_debug() {
1091 #[cfg(feature = "gzip")]
1092 {
1093 let algo = CompressionAlgo::Gzip;
1094 let debug_str = format!("{algo:?}");
1095 assert!(debug_str.contains("Gzip"));
1096 }
1097 }
1098
1099 #[test]
1100 fn test_compression_algo_clone() {
1101 #[cfg(feature = "gzip")]
1102 {
1103 let algo = CompressionAlgo::Gzip;
1104 let cloned = algo;
1105 assert_eq!(algo, cloned);
1106 }
1107 }
1108
1109 #[test]
1110 fn test_compression_algo_hash() {
1111 use std::collections::HashSet;
1112 #[cfg(feature = "gzip")]
1113 {
1114 let mut set = HashSet::new();
1115 set.insert(CompressionAlgo::Gzip);
1116 assert!(set.contains(&CompressionAlgo::Gzip));
1117 }
1118 }
1119
1120 #[test]
1122 fn test_compression_new() {
1123 let comp = Compression::new();
1124 assert!(!comp.algos.is_empty());
1125 assert!(!comp.content_types.is_empty());
1126 assert_eq!(comp.min_length, 1024);
1127 assert!(!comp.force_priority);
1128 }
1129
1130 #[test]
1131 fn test_compression_default() {
1132 let comp = Compression::default();
1133 assert!(!comp.algos.is_empty());
1134 }
1135
1136 #[test]
1137 fn test_compression_disable_all() {
1138 let comp = Compression::new().disable_all();
1139 assert!(comp.algos.is_empty());
1140 }
1141
1142 #[cfg(feature = "gzip")]
1143 #[test]
1144 fn test_compression_enable_gzip() {
1145 let comp = Compression::new()
1146 .disable_all()
1147 .enable_gzip(CompressionLevel::Fastest);
1148 assert!(comp.algos.contains_key(&CompressionAlgo::Gzip));
1149 assert_eq!(
1150 comp.algos.get(&CompressionAlgo::Gzip),
1151 Some(&CompressionLevel::Fastest)
1152 );
1153 }
1154
1155 #[cfg(feature = "gzip")]
1156 #[test]
1157 fn test_compression_disable_gzip() {
1158 let comp = Compression::new().disable_gzip();
1159 assert!(!comp.algos.contains_key(&CompressionAlgo::Gzip));
1160 }
1161
1162 #[cfg(feature = "brotli")]
1163 #[test]
1164 fn test_compression_enable_brotli() {
1165 let comp = Compression::new()
1166 .disable_all()
1167 .enable_brotli(CompressionLevel::Minsize);
1168 assert!(comp.algos.contains_key(&CompressionAlgo::Brotli));
1169 }
1170
1171 #[cfg(feature = "brotli")]
1172 #[test]
1173 fn test_compression_disable_brotli() {
1174 let comp = Compression::new().disable_brotli();
1175 assert!(!comp.algos.contains_key(&CompressionAlgo::Brotli));
1176 }
1177
1178 #[cfg(feature = "zstd")]
1179 #[test]
1180 fn test_compression_enable_zstd() {
1181 let comp = Compression::new()
1182 .disable_all()
1183 .enable_zstd(CompressionLevel::Default);
1184 assert!(comp.algos.contains_key(&CompressionAlgo::Zstd));
1185 }
1186
1187 #[cfg(feature = "zstd")]
1188 #[test]
1189 fn test_compression_disable_zstd() {
1190 let comp = Compression::new().disable_zstd();
1191 assert!(!comp.algos.contains_key(&CompressionAlgo::Zstd));
1192 }
1193
1194 #[cfg(feature = "deflate")]
1195 #[test]
1196 fn test_compression_enable_deflate() {
1197 let comp = Compression::new()
1198 .disable_all()
1199 .enable_deflate(CompressionLevel::Default);
1200 assert!(comp.algos.contains_key(&CompressionAlgo::Deflate));
1201 }
1202
1203 #[cfg(feature = "deflate")]
1204 #[test]
1205 fn test_compression_disable_deflate() {
1206 let comp = Compression::new().disable_deflate();
1207 assert!(!comp.algos.contains_key(&CompressionAlgo::Deflate));
1208 }
1209
1210 #[test]
1211 fn test_compression_min_length() {
1212 let comp = Compression::new().min_length(1024);
1213 assert_eq!(comp.min_length, 1024);
1214 }
1215
1216 #[test]
1217 fn test_compression_force_priority() {
1218 let comp = Compression::new().force_priority(true);
1219 assert!(comp.force_priority);
1220 }
1221
1222 #[test]
1223 fn test_compression_content_types() {
1224 let comp = Compression::new().content_types(&[mime::TEXT_PLAIN, mime::TEXT_HTML]);
1225 assert_eq!(comp.content_types.len(), 2);
1226 assert!(comp.content_types.contains(&mime::TEXT_PLAIN));
1227 assert!(comp.content_types.contains(&mime::TEXT_HTML));
1228 }
1229
1230 #[test]
1231 fn test_compression_debug() {
1232 let comp = Compression::new();
1233 let debug_str = format!("{comp:?}");
1234 assert!(debug_str.contains("Compression"));
1235 assert!(debug_str.contains("algos"));
1236 assert!(debug_str.contains("content_types"));
1237 }
1238
1239 #[test]
1240 fn test_compression_clone() {
1241 let comp = Compression::new().min_length(100);
1242 let cloned = comp.clone();
1243 assert_eq!(comp.min_length, cloned.min_length);
1244 assert_eq!(comp.algos.len(), cloned.algos.len());
1245 }
1246
1247 #[tokio::test]
1249 async fn test_no_accept_encoding_header() {
1250 let comp_handler = Compression::new().min_length(1);
1251 let router = Router::with_hoop(comp_handler).push(Router::with_path("hello").get(hello));
1252
1253 let res = TestClient::get("http://127.0.0.1:5801/hello")
1254 .send(router)
1255 .await;
1256 assert!(res.headers().get(CONTENT_ENCODING).is_none());
1257 }
1258
1259 #[tokio::test]
1260 async fn test_unsupported_encoding() {
1261 let comp_handler = Compression::new().min_length(1);
1262 let router = Router::with_hoop(comp_handler).push(Router::with_path("hello").get(hello));
1263
1264 let res = TestClient::get("http://127.0.0.1:5801/hello")
1265 .add_header(ACCEPT_ENCODING, "unknown", true)
1266 .send(router)
1267 .await;
1268 assert!(res.headers().get(CONTENT_ENCODING).is_none());
1269 }
1270
1271 #[tokio::test]
1272 async fn test_empty_response() {
1273 #[handler]
1274 async fn empty() {}
1275
1276 let comp_handler = Compression::new();
1277 let router = Router::with_hoop(comp_handler).push(Router::with_path("empty").get(empty));
1278
1279 let res = TestClient::get("http://127.0.0.1:5801/empty")
1280 .add_header(ACCEPT_ENCODING, "gzip", true)
1281 .send(router)
1282 .await;
1283 assert!(res.headers().get(CONTENT_ENCODING).is_none());
1284 }
1285
1286 #[tokio::test]
1287 async fn test_chained_configuration() {
1288 #[cfg(all(feature = "gzip", feature = "brotli"))]
1289 {
1290 let comp_handler = Compression::new()
1291 .disable_all()
1292 .enable_gzip(CompressionLevel::Fastest)
1293 .enable_brotli(CompressionLevel::Default)
1294 .min_length(1)
1295 .force_priority(false)
1296 .content_types(&[mime::TEXT_PLAIN]);
1297
1298 assert_eq!(comp_handler.algos.len(), 2);
1299 assert_eq!(comp_handler.min_length, 1);
1300 assert!(!comp_handler.force_priority);
1301 assert_eq!(comp_handler.content_types.len(), 1);
1302 }
1303 }
1304}