Skip to main content

salvo_compression/
lib.rs

1#![cfg_attr(docsrs, feature(doc_cfg))]
2#![cfg_attr(test, allow(clippy::unwrap_used))]
3
4//! Compression middleware for the Salvo web framework.
5//!
6//! This middleware automatically compresses HTTP responses using various algorithms,
7//! reducing bandwidth usage and improving load times for clients.
8//!
9//! # Supported Algorithms
10//!
11//! | Algorithm | Feature | Content-Encoding |
12//! |-----------|---------|------------------|
13//! | Gzip | `gzip` | `gzip` |
14//! | Brotli | `brotli` | `br` |
15//! | Deflate | `deflate` | `deflate` |
16//! | Zstd | `zstd` | `zstd` |
17//!
18//! # Example
19//!
20//! ```no_run
21//! use salvo_compression::{Compression, CompressionLevel};
22//! use salvo_core::prelude::*;
23//!
24//! #[handler]
25//! async fn hello() -> &'static str {
26//!     "hello"
27//! }
28//!
29//! let compression = Compression::new()
30//!     .enable_gzip(CompressionLevel::Default)
31//!     .min_length(1024); // Only compress responses > 1KB
32//!
33//! let _router = Router::new().hoop(compression).get(hello);
34//! ```
35//!
36//! # Algorithm Negotiation
37//!
38//! The middleware negotiates the compression algorithm based on the client's
39//! `Accept-Encoding` header. By default, it respects the client's preference order.
40//! Use `force_priority(true)` to use the server's configured priority instead.
41//!
42//! # Compression Levels
43//!
44//! - [`CompressionLevel::Fastest`]: Fastest compression, larger output
45//! - [`CompressionLevel::Default`]: Balanced compression (recommended)
46//! - [`CompressionLevel::Minsize`]: Best compression, slower
47//! - `CompressionLevel::Precise(u32)`: Fine-grained control
48//!
49//! # Default Content Types
50//!
51//! By default, the middleware compresses:
52//! - `text/*` (HTML, CSS, plain text, etc.)
53//! - `application/javascript`
54//! - `application/json`
55//! - `application/xml`, `application/rss+xml`
56//! - `application/wasm`
57//! - `image/svg+xml`
58//!
59//! Use `.content_types()` to customize which MIME types are compressed.
60//!
61//! # Minimum Length
62//!
63//! Small responses may not benefit from compression. Use `.min_length(bytes)`
64//! to skip compression for responses smaller than the specified size.
65//!
66//! Read more: <https://salvo.rs>
67
68use 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/// Level of compression data should be compressed with.
87#[non_exhaustive]
88#[derive(Clone, Copy, Default, Debug, Eq, PartialEq)]
89pub enum CompressionLevel {
90    /// Fastest quality of compression, usually produces a bigger size.
91    Fastest,
92    /// Best quality of compression, usually produces the smallest size.
93    Minsize,
94    /// Default quality of compression defined by the selected compression algorithm.
95    #[default]
96    Default,
97    /// Precise quality based on the underlying compression algorithms'
98    /// qualities. The interpretation of this depends on the algorithm chosen
99    /// and the specific implementation backing it.
100    /// Qualities are implicitly clamped to the algorithm's maximum.
101    Precise(u32),
102}
103
104/// CompressionAlgo
105#[derive(Eq, PartialEq, Clone, Copy, Debug, Hash)]
106#[non_exhaustive]
107pub enum CompressionAlgo {
108    /// Compress use Brotli algo.
109    #[cfg(feature = "brotli")]
110    #[cfg_attr(docsrs, doc(cfg(feature = "brotli")))]
111    Brotli,
112
113    /// Compress use Deflate algo.
114    #[cfg(feature = "deflate")]
115    #[cfg_attr(docsrs, doc(cfg(feature = "deflate")))]
116    Deflate,
117
118    /// Compress use Gzip algo.
119    #[cfg(feature = "gzip")]
120    #[cfg_attr(docsrs, doc(cfg(feature = "gzip")))]
121    Gzip,
122
123    /// Compress use Zstd algo.
124    #[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/// Compression
187#[derive(Clone, Debug)]
188#[non_exhaustive]
189pub struct Compression {
190    /// Compression algorithms to use.
191    pub algos: IndexMap<CompressionAlgo, CompressionLevel>,
192    /// Content types to compress.
193    pub content_types: Vec<Mime>,
194    /// Minimum body size to compress; bodies smaller than this value are not compressed.
195    ///
196    /// This threshold only applies to bodies whose length is known up front
197    /// (in-memory `Once`/`Chunks` bodies). Streaming bodies (`Hyper`/`Stream`)
198    /// have no known length and are always compressed regardless of
199    /// `min_length`, so a tiny streamed body can still end up larger after
200    /// adding `Content-Encoding` and framing overhead.
201    pub min_length: usize,
202    /// Ignore the client's algorithm order in `Accept-Encoding` and always use the server's
203    /// configured priority.
204    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    /// Create a new `Compression`.
242    #[inline]
243    #[must_use]
244    pub fn new() -> Self {
245        Default::default()
246    }
247
248    /// Remove all compression algorithms.
249    #[inline]
250    #[must_use]
251    pub fn disable_all(mut self) -> Self {
252        self.algos.clear();
253        self
254    }
255
256    /// Sets `Compression` with algos.
257    #[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    /// Disable gzip compression.
266    #[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    /// Enable zstd compression.
275    #[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    /// Disable zstd compression.
284    #[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    /// Enable brotli compression.
293    #[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    /// Disable brotli compression.
302    #[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    /// Enable deflate compression.
312    #[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    /// Disable deflate compression.
322    #[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    /// Sets minimum compression size, if body is less than this value, no compression.
332    /// Default is 1kb.
333    ///
334    /// Streaming bodies with an explicit `Content-Length` also respect this threshold. Streaming
335    /// bodies with an unknown length are always compressed.
336    #[inline]
337    #[must_use]
338    pub fn min_length(mut self, size: usize) -> Self {
339        self.min_length = size;
340        self
341    }
342    /// Sets `Compression` with force_priority.
343    #[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    /// Sets `Compression` with content types list.
351    #[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        // `parse_accept_encoding` sorts entries by descending q-value, so the
404        // first wildcard token is also the one with the highest q. Mirror the
405        // original `.find(...)` semantics rather than overwriting `wildcard_q`
406        // on every iteration — see https://github.com/salvo-rs/salvo/pull/1489
407        // for the multi-wildcard regression this guards against.
408        let wildcard_q: Option<u8> = accept_list
409            .iter()
410            .find(|(name, _)| name == "*")
411            .map(|(_, q)| *q);
412
413        // Parse each non-wildcard `Accept-Encoding` entry once, dropping tokens
414        // that do not name a known compression algorithm. The lookups below
415        // (`is_accepted` / `is_rejected` / client-preference) then never have
416        // to call `.parse::<CompressionAlgo>()` again. Entries keep the order
417        // produced by `parse_accept_encoding` (descending q-value).
418        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            // Server preference: pick the highest-priority server algo the client accepts.
431            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            // Client preference: pick the highest q-value algo the server supports.
439            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            // Wildcard `*`: use the server's top algo that is not explicitly rejected.
449            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                // A HEAD-aware handler can intentionally omit the body before this middleware
482                // runs. Preserve the representation headers that the equivalent GET response
483                // would receive when its uncompressed length is known.
484                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        // Client prefers br (q=1.0) over gzip (q=0.5)
695        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        // gzip is explicitly rejected (q=0), only br is acceptable
710        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        // identity means no encoding; server must not compress
725        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        // `*` means accept any encoding; server picks its preferred algo
738        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        // `*` but gzip;q=0 — server must not use gzip, falls back to next algo
754        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        // When the client lists multiple wildcards, the effective q-value
771        // is the highest one — matching the pre-refactor behaviour. With
772        // `*;q=1, *;q=0` (q-sorted by `parse_accept_encoding`, so the q=1
773        // wildcard comes first), wildcard compression must still be allowed.
774        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        // Ensure only one Content-Encoding header is set (no duplicates via append)
790        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    // Tests for CompressionLevel
965    #[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    // Tests for CompressionAlgo
1011    #[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    // Tests for Compression struct
1121    #[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    // Tests for no compression scenarios
1248    #[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}