Skip to main content

ntex/web/middleware/
compress.rs

1//! `Middleware` for compressing response body.
2use crate::http::encoding::Encoder;
3use crate::http::header::{ACCEPT_ENCODING, ContentEncoding};
4use crate::service::{Ctx, Middleware, Service};
5use crate::web::{BodyEncoding, State, WebRequest, WebResponse};
6
7#[derive(Debug, Clone)]
8/// `Middleware` for compressing response body.
9///
10/// Bodies with a known size below 1KiB are sent uncompressed.
11/// Use `BodyEncoding` trait for overriding response compression.
12/// To disable compression set encoding to `ContentEncoding::Identity` value.
13///
14/// ```rust
15/// use ntex::web::{self, middleware, App, HttpResponse};
16///
17/// fn main() {
18///     let app = App::default()
19///         .middleware(middleware::Compress::default())
20///         .service(
21///             web::resource("/test")
22///                 .route(web::get().to(async || { HttpResponse::Ok() }))
23///                 .route(web::head().to(async || { HttpResponse::MethodNotAllowed() }))
24///         );
25/// }
26/// ```
27pub struct Compress {
28    enc: ContentEncoding,
29}
30
31impl Compress {
32    /// Create new `Compress` middleware with the specified encoding.
33    ///
34    /// Use `Compress::default()` to select the encoding automatically
35    /// ([`ContentEncoding::Auto`]).
36    pub fn new(encoding: ContentEncoding) -> Self {
37        Compress { enc: encoding }
38    }
39}
40
41impl Default for Compress {
42    fn default() -> Self {
43        Compress::new(ContentEncoding::Auto)
44    }
45}
46
47impl<S, St> Middleware<S, St> for Compress {
48    type Service = CompressMiddleware<S>;
49
50    fn create(&self, _: &St, service: S) -> Self::Service {
51        CompressMiddleware {
52            service,
53            encoding: self.enc,
54        }
55    }
56}
57
58#[derive(Debug)]
59pub struct CompressMiddleware<S> {
60    service: S,
61    encoding: ContentEncoding,
62}
63
64impl<S, St, In> Service<St, WebRequest<In>> for CompressMiddleware<S>
65where
66    S: Service<St, WebRequest<In>, Res = WebResponse>,
67    St: State,
68{
69    type Res = WebResponse;
70    type Error = S::Error;
71
72    crate::forward_ready!(St, service);
73    crate::forward_shutdown!(St, service);
74
75    async fn call(
76        &self,
77        req: WebRequest<In>,
78        ctx: Ctx<'_, Self, St>,
79    ) -> Result<WebResponse, S::Error> {
80        // negotiate content-encoding
81        let values = req
82            .headers()
83            .get_all(&ACCEPT_ENCODING)
84            .filter_map(|val| val.to_str().ok());
85        let encoding = AcceptEncoding::parse(values, self.encoding);
86
87        let resp = ctx.call(&self.service, req).await?;
88
89        let enc = if let Some(enc) = resp.response().get_encoding() {
90            enc
91        } else {
92            encoding
93        };
94
95        Ok(resp.map_body(move |head, body| Encoder::response(enc, head, body)))
96    }
97}
98
99/// Encodings that can be picked for a `*` entry.
100const WILDCARD: [ContentEncoding; 3] = [
101    ContentEncoding::Zstd,
102    ContentEncoding::Gzip,
103    ContentEncoding::Deflate,
104];
105
106#[derive(Debug, PartialEq)]
107struct AcceptEncoding {
108    /// `None` for the `*` entry
109    encoding: Option<ContentEncoding>,
110    /// Client weight in thousandths, `0..=1000`
111    q: u16,
112}
113
114impl AcceptEncoding {
115    fn new(tag: &str) -> Option<AcceptEncoding> {
116        let mut parts = tag.split(';');
117        let name = parts.next()?.trim();
118        if name.is_empty() {
119            return None;
120        }
121        let encoding = if name == "*" {
122            None
123        } else {
124            Some(ContentEncoding::from(name))
125        };
126
127        // a malformed weight makes the entry unacceptable
128        let q = parts
129            .filter_map(|param| param.split_once('='))
130            .find(|(name, _)| name.trim().eq_ignore_ascii_case("q"))
131            .map_or(1000, |(_, val)| parse_q(val.trim()).unwrap_or(0));
132        Some(AcceptEncoding { encoding, q })
133    }
134
135    /// Pick a response encoding from the `Accept-Encoding` header values.
136    ///
137    /// Entries are ranked by the client's `q` value, ties go to the
138    /// server's preference. `q=0` means "not acceptable" and `*` covers
139    /// every encoding the client did not list.
140    fn parse<'a, I>(values: I, encoding: ContentEncoding) -> ContentEncoding
141    where
142        I: Iterator<Item = &'a str>,
143    {
144        let mut listed: Vec<(ContentEncoding, u16)> = Vec::new();
145        let mut wildcard = None;
146        for item in values.flat_map(|val| val.split(',')) {
147            match AcceptEncoding::new(item) {
148                Some(AcceptEncoding {
149                    encoding: Some(enc),
150                    q,
151                }) => listed.push((enc, q)),
152                Some(AcceptEncoding { encoding: None, q }) => {
153                    wildcard.get_or_insert(q);
154                }
155                None => {}
156            }
157        }
158        let weight = |enc: ContentEncoding| {
159            listed
160                .iter()
161                .find(|(e, _)| *e == enc)
162                .map(|(_, q)| *q)
163                .or(wildcard)
164                .unwrap_or(0)
165        };
166
167        if encoding != ContentEncoding::Auto {
168            return if weight(encoding) > 0 {
169                encoding
170            } else {
171                ContentEncoding::Identity
172            };
173        }
174
175        let mut best: Option<(ContentEncoding, u16)> = None;
176        let candidates = listed.iter().map(|(enc, _)| *enc).chain(WILDCARD);
177        for enc in candidates.filter(|enc| Encoder::can_encode(*enc)) {
178            let q = weight(enc);
179            let better = match best {
180                None => q > 0,
181                Some((b, bq)) => q > bq || (q == bq && enc.quality() > b.quality()),
182            };
183            if better {
184                best = Some((enc, q));
185            }
186        }
187        best.map_or(ContentEncoding::Identity, |(enc, _)| enc)
188    }
189}
190
191/// Parse an RFC 9110 `qvalue` into thousandths.
192fn parse_q(val: &str) -> Option<u16> {
193    let (int, frac) = match val.split_once('.') {
194        Some((int, frac)) => (int, frac),
195        None => (val, ""),
196    };
197    if frac.len() > 3 || !frac.bytes().all(|b| b.is_ascii_digit()) {
198        return None;
199    }
200    let int = match int {
201        "0" => 0,
202        "1" => 1000,
203        _ => return None,
204    };
205    let frac = frac
206        .bytes()
207        .zip([100, 10, 1])
208        .map(|(b, scale)| u16::from(b - b'0') * scale)
209        .sum::<u16>();
210    if int + frac > 1000 { None } else { Some(int + frac) }
211}
212
213#[cfg(test)]
214mod tests {
215    use super::*;
216
217    fn parse(raw: &str, encoding: ContentEncoding) -> ContentEncoding {
218        AcceptEncoding::parse([raw].into_iter(), encoding)
219    }
220
221    #[test]
222    fn test_parse_q() {
223        for (val, q) in [
224            ("0", 0),
225            ("0.", 0),
226            ("0.5", 500),
227            ("0.05", 50),
228            ("0.125", 125),
229            ("0.999", 999),
230            ("1", 1000),
231            ("1.", 1000),
232            ("1.000", 1000),
233        ] {
234            assert_eq!(parse_q(val), Some(q), "{val}");
235        }
236        for val in [
237            "", ".5", "0.1234", "1.001", "1.5", "2", "-0", "0.a", "0.00a", "01", "abc",
238        ] {
239            assert_eq!(parse_q(val), None, "{val}");
240        }
241    }
242
243    #[test]
244    fn test_accept_encoding_entry() {
245        let entry = |tag| AcceptEncoding::new(tag).unwrap();
246        assert_eq!(
247            entry("gzip"),
248            AcceptEncoding {
249                encoding: Some(ContentEncoding::Gzip),
250                q: 1000
251            }
252        );
253        assert_eq!(entry(" gzip ; q=0.5 ").q, 500);
254        assert_eq!(entry("gzip;Q=0.5").q, 500);
255        assert_eq!(entry("gzip;level=1;q=0.25").q, 250);
256        assert_eq!(entry("gzip;q=abc").q, 0);
257        assert_eq!(entry("gzip;0.8").q, 1000);
258        assert_eq!(entry("*;q=0.1").encoding, None);
259        assert!(AcceptEncoding::new(" ").is_none());
260    }
261
262    #[test]
263    fn test_auto_skips_unsupported_encodings() {
264        let auto = ContentEncoding::Auto;
265        assert_eq!(
266            parse("gzip, deflate, br, zstd", auto),
267            ContentEncoding::Zstd
268        );
269        assert_eq!(parse("gzip, deflate, br", auto), ContentEncoding::Gzip);
270        assert_eq!(parse("br, deflate", auto), ContentEncoding::Deflate);
271        assert_eq!(parse("br", auto), ContentEncoding::Identity);
272        assert_eq!(parse("gzip, br", ContentEncoding::Br), ContentEncoding::Br);
273        assert_eq!(parse(" br ,  gzip ; q=1.0 ", auto), ContentEncoding::Gzip);
274        assert_eq!(parse("", auto), ContentEncoding::Identity);
275    }
276
277    #[test]
278    fn test_auto_honours_q() {
279        let auto = ContentEncoding::Auto;
280        assert_eq!(parse("gzip;q=0.5, zstd;q=0.1", auto), ContentEncoding::Gzip);
281        assert_eq!(
282            parse("deflate;q=0.9, gzip;q=0.8", auto),
283            ContentEncoding::Deflate
284        );
285        assert_eq!(parse("zstd;q=0, gzip", auto), ContentEncoding::Gzip);
286        assert_eq!(parse("gzip;q=0", auto), ContentEncoding::Identity);
287        assert_eq!(parse("gzip;q=abc", auto), ContentEncoding::Identity);
288        assert_eq!(parse("gzip;q=2, deflate", auto), ContentEncoding::Deflate);
289        assert_eq!(parse("gzip, zstd;q=0", auto), ContentEncoding::Gzip);
290        // the first entry for an encoding wins
291        assert_eq!(parse("gzip;q=0, gzip", auto), ContentEncoding::Identity);
292        // equal weights go to the server's preference
293        assert_eq!(
294            parse("deflate;q=0.5, gzip;q=0.5, zstd;q=0.5", auto),
295            ContentEncoding::Zstd
296        );
297        assert_eq!(parse("deflate, gzip", auto), ContentEncoding::Gzip);
298        assert_eq!(parse("x-gzip", auto), ContentEncoding::Gzip);
299    }
300
301    #[test]
302    fn test_auto_wildcard() {
303        let auto = ContentEncoding::Auto;
304        assert_eq!(parse("*", auto), ContentEncoding::Zstd);
305        assert_eq!(parse("*;q=0", auto), ContentEncoding::Identity);
306        assert_eq!(parse("*;q=0.5, zstd;q=0", auto), ContentEncoding::Gzip);
307        assert_eq!(
308            parse("zstd;q=0, gzip;q=0, *", auto),
309            ContentEncoding::Deflate
310        );
311        assert_eq!(parse("*;q=0.5, deflate", auto), ContentEncoding::Deflate);
312        assert_eq!(parse("gzip;q=0.1, *;q=0.5", auto), ContentEncoding::Zstd);
313        assert_eq!(parse("*;q=0, gzip;q=0.1", auto), ContentEncoding::Gzip);
314        assert_eq!(parse("*;q=0.5, *", auto), ContentEncoding::Zstd);
315        assert_eq!(parse("*;q=0, *, zstd;q=0", auto), ContentEncoding::Identity);
316    }
317
318    #[test]
319    fn test_explicit_encoding() {
320        let gzip = ContentEncoding::Gzip;
321        assert_eq!(parse("gzip", gzip), gzip);
322        assert_eq!(parse("zstd", gzip), ContentEncoding::Identity);
323        assert_eq!(parse("gzip;q=0", gzip), ContentEncoding::Identity);
324        assert_eq!(parse("gzip;q=0.001", gzip), gzip);
325        assert_eq!(parse("*", gzip), gzip);
326        assert_eq!(parse("*;q=0", gzip), ContentEncoding::Identity);
327        assert_eq!(parse("*, gzip;q=0", gzip), ContentEncoding::Identity);
328        assert_eq!(parse("*;q=0, gzip", gzip), gzip);
329    }
330
331    #[test]
332    fn test_repeated_headers() {
333        let values = ["zstd;q=0", "gzip;q=0.5, deflate;q=0.1"];
334        assert_eq!(
335            AcceptEncoding::parse(values.into_iter(), ContentEncoding::Auto),
336            ContentEncoding::Gzip
337        );
338    }
339
340    #[crate::rt_test]
341    async fn test_compress_accept_encoding() {
342        use crate::http::header::{CONTENT_ENCODING, HeaderValue};
343        use crate::web::test::{TestRequest, call_service, init_service};
344        use crate::web::{self, App, HttpResponse};
345
346        let srv = init_service(App::new().middleware(Compress::default()).route(
347            "/",
348            web::get().to(async || HttpResponse::Ok().body("a".repeat(1024))),
349        ))
350        .await;
351
352        let req = TestRequest::default()
353            .header(ACCEPT_ENCODING, "gzip")
354            .to_request();
355        let resp = call_service(&srv, req).await;
356        assert_eq!(resp.headers().get(CONTENT_ENCODING).unwrap(), "gzip");
357
358        let req = TestRequest::default()
359            .header(
360                ACCEPT_ENCODING,
361                HeaderValue::from_bytes(b"gzip\xff").unwrap(),
362            )
363            .to_request();
364        let resp = call_service(&srv, req).await;
365        assert!(resp.headers().get(CONTENT_ENCODING).is_none());
366
367        let req = TestRequest::default()
368            .header(ACCEPT_ENCODING, "zstd;q=0")
369            .header(ACCEPT_ENCODING, "gzip;q=0.5")
370            .to_request();
371        let resp = call_service(&srv, req).await;
372        assert_eq!(resp.headers().get(CONTENT_ENCODING).unwrap(), "gzip");
373
374        let req = TestRequest::default()
375            .header(
376                ACCEPT_ENCODING,
377                HeaderValue::from_bytes(b"zstd\xff").unwrap(),
378            )
379            .header(ACCEPT_ENCODING, "deflate")
380            .to_request();
381        let resp = call_service(&srv, req).await;
382        assert_eq!(resp.headers().get(CONTENT_ENCODING).unwrap(), "deflate");
383
384        let req = TestRequest::default()
385            .header(ACCEPT_ENCODING, "gzip;q=0")
386            .to_request();
387        let resp = call_service(&srv, req).await;
388        assert!(resp.headers().get(CONTENT_ENCODING).is_none());
389    }
390}