Skip to main content

rama_http/layer/compression/
mod.rs

1//! Middleware that compresses response bodies.
2//!
3//! If you require streaming compression (e.g. for SSE etc),
4//! you most likely want to use the compression middleware from [`stream`] instead.
5//!
6//! # Example
7//!
8//! Example showing how to respond with the compressed contents of a file.
9//!
10//! ```rust
11//! use rama_core::bytes::Bytes;
12//! use rama_core::futures::stream::StreamExt;
13//! use rama_core::error::BoxError;
14//! use rama_http::body::Frame;
15//! use rama_http::body::util::{BodyExt, StreamBody};
16//! use rama_http::body::util::combinators::BoxBody as InnerBoxBody;
17//! use rama_http::layer::compression::CompressionLayer;
18//! use rama_http::{Body, Request, Response, header::ACCEPT_ENCODING};
19//! use rama_core::service::service_fn;
20//! use rama_core::{Service, Layer};
21//! use std::convert::Infallible;
22//! use tokio::fs::{self, File};
23//! use rama_core::stream::io::ReaderStream;
24//!
25//! type BoxBody = InnerBoxBody<Bytes, std::io::Error>;
26//!
27//! # #[tokio::main]
28//! # async fn main() -> Result<(), BoxError> {
29//! async fn handle(req: Request) -> Result<Response<BoxBody>, Infallible> {
30//!     // Open the file.
31//!     let file = File::open("Cargo.toml").await.expect("file missing");
32//!     // Convert the file into a `Stream` of `Bytes`.
33//!     let stream = ReaderStream::new(file);
34//!     // Convert the stream into a stream of data `Frame`s.
35//!     let stream = stream.map(|res| match res {
36//!         Ok(v) => Ok(Frame::data(v)),
37//!         Err(e) => Err(e),
38//!     });
39//!     // Convert the `Stream` into a `Body`.
40//!     let body = StreamBody::new(stream);
41//!     // Erase the type because it's very hard to name in the function signature.
42//!     let body = BodyExt::boxed(body);
43//!     // Create response.
44//!     Ok(Response::new(body))
45//! }
46//!
47//! let mut service = (
48//!     // Compress responses based on the `Accept-Encoding` header.
49//!     CompressionLayer::new(),
50//! ).into_layer(service_fn(handle));
51//!
52//! // Call the service.
53//! let request = Request::builder()
54//!     .header(ACCEPT_ENCODING, "gzip")
55//!     .body(Body::default())?;
56//!
57//! let response = service
58//!     .serve(request)
59//!     .await?;
60//!
61//! assert_eq!(response.headers()["content-encoding"], "gzip");
62//!
63//! // Read the body
64//! let bytes = response
65//!     .into_body()
66//!     .collect()
67//!     .await?
68//!     .to_bytes();
69//!
70//! // The compressed body should be smaller 🤞
71//! let uncompressed_len = fs::read_to_string("Cargo.toml").await?.len();
72//! assert!(bytes.len() < uncompressed_len);
73//! #
74//! # Ok(())
75//! # }
76//! ```
77//!
78
79pub mod predicate;
80pub mod stream;
81
82pub(crate) mod body;
83mod layer;
84mod pin_project_cfg;
85mod service;
86
87#[doc(inline)]
88pub use self::{
89    body::CompressionBody,
90    layer::CompressionLayer,
91    predicate::{DefaultPredicate, MirrorDecompressed, Predicate, PreferredEncoding},
92    service::Compression,
93};
94#[doc(inline)]
95pub use crate::layer::util::compression::CompressionLevel;
96
97#[cfg(test)]
98mod tests {
99    use super::*;
100
101    use crate::layer::compression::predicate::{MirrorDecompressed, PreferredEncoding, SizeAbove};
102    use crate::layer::decompression::DecompressedFrom;
103
104    use crate::header::{
105        ACCEPT_ENCODING, ACCEPT_RANGES, CONTENT_ENCODING, CONTENT_RANGE, CONTENT_TYPE, RANGE,
106    };
107    use crate::{HeaderMap, HeaderValue, Request, Response, StreamingBody, body::util::BodyExt};
108    use async_compression::tokio::write::{BrotliDecoder, BrotliEncoder};
109    use flate2::read::GzDecoder;
110    use rama_core::Service;
111    use rama_core::bytes::Bytes;
112    use rama_core::error::BoxError;
113    use rama_core::extensions::ExtensionsRef;
114    use rama_core::service::service_fn;
115    use rama_core::stream::io::StreamReader;
116    use rama_http_types::Body;
117    use std::convert::Infallible;
118    use std::io::Read;
119    use std::sync::{Arc, RwLock};
120    use tokio::io::{AsyncReadExt, AsyncWriteExt};
121
122    // Compression filter allows every other request to be compressed
123    #[derive(Clone)]
124    struct Always;
125
126    impl Predicate for Always {
127        fn should_compress<B>(&self, _: &mut rama_http_types::Response<B>) -> bool
128        where
129            B: StreamingBody,
130        {
131            true
132        }
133    }
134
135    #[tokio::test]
136    async fn gzip_works() {
137        let svc = service_fn(handle);
138        let svc = Compression::new(svc).with_compress_predicate(Always);
139
140        // call the service
141        let req = Request::builder()
142            .header("accept-encoding", "gzip")
143            .body(Body::empty())
144            .unwrap();
145        let res = svc.serve(req).await.unwrap();
146
147        // read the compressed body
148        let collected = res.into_body().collect().await.unwrap();
149        let compressed_data = collected.to_bytes();
150
151        // decompress the body
152        // doing this with flate2 as that is much easier than async-compression and blocking during
153        // tests is fine
154        let mut decoder = GzDecoder::new(&compressed_data[..]);
155        let mut decompressed = String::new();
156        decoder.read_to_string(&mut decompressed).unwrap();
157
158        assert_eq!(decompressed, "Hello, World!");
159    }
160
161    #[tokio::test]
162    async fn x_gzip_works() {
163        let svc = service_fn(handle);
164        let svc = Compression::new(svc).with_compress_predicate(Always);
165
166        // call the service
167        let req = Request::builder()
168            .header("accept-encoding", "x-gzip")
169            .body(Body::empty())
170            .unwrap();
171        let res = svc.serve(req).await.unwrap();
172
173        // we treat x-gzip as equivalent to gzip and don't have to return x-gzip
174        // taking extra caution by checking all headers with this name
175        assert_eq!(
176            res.headers()
177                .get_all("content-encoding")
178                .iter()
179                .collect::<Vec<&HeaderValue>>(),
180            vec!(HeaderValue::from_static("gzip"))
181        );
182
183        // read the compressed body
184        let collected = res.into_body().collect().await.unwrap();
185        let compressed_data = collected.to_bytes();
186
187        // decompress the body
188        // doing this with flate2 as that is much easier than async-compression and blocking during
189        // tests is fine
190        let mut decoder = GzDecoder::new(&compressed_data[..]);
191        let mut decompressed = String::new();
192        decoder.read_to_string(&mut decompressed).unwrap();
193
194        assert_eq!(decompressed, "Hello, World!");
195    }
196
197    #[tokio::test]
198    async fn zstd_works() {
199        let svc = service_fn(handle);
200        let svc = Compression::new(svc).with_compress_predicate(Always);
201
202        // call the service
203        let req = Request::builder()
204            .header("accept-encoding", "zstd")
205            .body(Body::empty())
206            .unwrap();
207        let res = svc.serve(req).await.unwrap();
208
209        // read the compressed body
210        let body = res.into_body();
211        let compressed_data = body.collect().await.unwrap().to_bytes();
212
213        // decompress the body
214        let decompressed = zstd::stream::decode_all(std::io::Cursor::new(compressed_data)).unwrap();
215        let decompressed = String::from_utf8(decompressed).unwrap();
216
217        assert_eq!(decompressed, "Hello, World!");
218    }
219
220    #[tokio::test]
221    async fn predicate_only_compresses_previously_decompressed_responses() {
222        let svc = service_fn(async |_| {
223            let res = Response::new(Body::from("Hello, World!"));
224            res.extensions().insert(DecompressedFrom::Gzip);
225            Ok::<_, Infallible>(res)
226        });
227        let svc = Compression::new(svc).with_compress_predicate(MirrorDecompressed::new());
228
229        let req = Request::builder()
230            .header("accept-encoding", "gzip")
231            .body(Body::empty())
232            .unwrap();
233        let res = svc.serve(req).await.unwrap();
234
235        assert_eq!(res.headers()[CONTENT_ENCODING], "gzip");
236
237        let collected = res.into_body().collect().await.unwrap();
238        let compressed_data = collected.to_bytes();
239
240        let mut decoder = GzDecoder::new(&compressed_data[..]);
241        let mut decompressed = String::new();
242        decoder.read_to_string(&mut decompressed).unwrap();
243
244        assert_eq!(decompressed, "Hello, World!");
245    }
246
247    #[tokio::test]
248    async fn predicate_skips_responses_that_were_not_decompressed() {
249        let svc =
250            service_fn(async |_| Ok::<_, Infallible>(Response::new(Body::from("Hello, World!"))));
251        let svc = Compression::new(svc).with_compress_predicate(MirrorDecompressed::new());
252
253        let req = Request::builder()
254            .header("accept-encoding", "gzip")
255            .body(Body::empty())
256            .unwrap();
257        let res = svc.serve(req).await.unwrap();
258
259        assert!(!res.headers().contains_key(CONTENT_ENCODING));
260
261        let collected = res.into_body().collect().await.unwrap();
262        assert_eq!(collected.to_bytes().as_ref(), b"Hello, World!");
263    }
264
265    #[tokio::test]
266    async fn mirror_decompressed_sets_preferred_encoding() {
267        let mut res = Response::new(Body::from("Hello, World!"));
268        res.extensions().insert(DecompressedFrom::Brotli);
269
270        let predicate = MirrorDecompressed::new();
271        assert!(predicate.should_compress(&mut res));
272        assert_eq!(
273            res.extensions().get_ref::<PreferredEncoding>(),
274            Some(&PreferredEncoding::Brotli)
275        );
276    }
277
278    #[tokio::test]
279    async fn respect_content_encoding_overrides_predicate_preference() {
280        let svc = service_fn(async |_| {
281            let mut res = Response::new(Body::from("Hello, World! Hello, World! Hello, World!"));
282            res.headers_mut()
283                .insert(CONTENT_ENCODING, HeaderValue::from_static("gzip"));
284            res.extensions().insert(PreferredEncoding::Brotli);
285            Ok::<_, Infallible>(res)
286        });
287        let svc = Compression::new(svc)
288            .with_respect_content_encoding_if_possible()
289            .with_compress_predicate(Always);
290
291        let req = Request::builder()
292            .header("accept-encoding", "gzip, br")
293            .body(Body::empty())
294            .unwrap();
295        let res = svc.serve(req).await.unwrap();
296
297        assert_eq!(res.headers()[CONTENT_ENCODING], "gzip");
298    }
299
300    #[tokio::test]
301    async fn no_recompress() {
302        const DATA: &str = "Hello, World! I'm already compressed with br!";
303
304        let svc = service_fn(async |_| {
305            let buf = {
306                let mut buf = Vec::new();
307
308                let mut enc = BrotliEncoder::new(&mut buf);
309                enc.write_all(DATA.as_bytes()).await?;
310                enc.flush().await?;
311                buf
312            };
313
314            let resp = Response::builder()
315                .header("content-encoding", "br")
316                .body(Body::from(buf))
317                .unwrap();
318            Ok::<_, std::io::Error>(resp)
319        });
320        let svc = Compression::new(svc);
321
322        // call the service
323        //
324        // note: the accept-encoding doesn't match the content-encoding above, so that
325        // we're able to see if the compression layer triggered or not
326        let req = Request::builder()
327            .header("accept-encoding", "gzip")
328            .body(Body::empty())
329            .unwrap();
330        let res = svc.serve(req).await.unwrap();
331
332        // check we didn't recompress
333        assert_eq!(
334            res.headers()
335                .get("content-encoding")
336                .and_then(|h| h.to_str().ok())
337                .unwrap_or_default(),
338            "br",
339        );
340
341        // read the compressed body
342        let body = res.into_body();
343        let data = body.collect().await.unwrap().to_bytes();
344
345        // decompress the body
346        let data = {
347            let mut output_buf = Vec::new();
348            let mut decoder = BrotliDecoder::new(&mut output_buf);
349            decoder
350                .write_all(&data)
351                .await
352                .expect("couldn't brotli-decode");
353            decoder.flush().await.expect("couldn't flush");
354            output_buf
355        };
356
357        assert_eq!(data, DATA.as_bytes());
358    }
359
360    async fn handle(_req: Request) -> Result<Response, Infallible> {
361        let body = Body::from("Hello, World!");
362        Ok(Response::builder().body(body).unwrap())
363    }
364
365    #[tokio::test]
366    async fn will_not_compress_if_filtered_out() {
367        use predicate::Predicate;
368
369        const DATA: &str = "Hello world uncompressed";
370
371        let svc_fn = service_fn(async |_| {
372            let resp = Response::builder()
373                // .header("content-encoding", "br")
374                .body(Body::from(DATA.as_bytes()))
375                .unwrap();
376            Ok::<_, std::io::Error>(resp)
377        });
378
379        // Compression filter allows every other request to be compressed
380        #[derive(Default, Clone)]
381        struct EveryOtherResponse(Arc<RwLock<u64>>);
382
383        impl Predicate for EveryOtherResponse {
384            fn should_compress<B>(&self, _: &mut rama_http_types::Response<B>) -> bool
385            where
386                B: StreamingBody,
387            {
388                let mut guard = self.0.write().unwrap();
389                let should_compress = !(*guard).is_multiple_of(2);
390                *guard += 1;
391                should_compress
392            }
393        }
394
395        let svc = Compression::new(svc_fn).with_compress_predicate(EveryOtherResponse::default());
396        let req = Request::builder()
397            .header("accept-encoding", "br")
398            .body(Body::empty())
399            .unwrap();
400        let res = svc.serve(req).await.unwrap();
401
402        // read the uncompressed body
403        let body = res.into_body();
404        let data = body.collect().await.unwrap().to_bytes();
405        let still_uncompressed = String::from_utf8(data.to_vec()).unwrap();
406        assert_eq!(DATA, &still_uncompressed);
407
408        // Compression filter will compress the next body
409        let req = Request::builder()
410            .header("accept-encoding", "br")
411            .body(Body::empty())
412            .unwrap();
413        let res = svc.serve(req).await.unwrap();
414
415        // read the compressed body
416        let body = res.into_body();
417        let data = body.collect().await.unwrap().to_bytes();
418        String::from_utf8(data.to_vec()).unwrap_err();
419    }
420
421    #[tokio::test]
422    async fn doesnt_compress_images() {
423        async fn handle(_req: Request<Body>) -> Result<Response<Body>, Infallible> {
424            let mut res = Response::new(Body::from(
425                "a".repeat((SizeAbove::DEFAULT_MIN_SIZE * 2) as usize),
426            ));
427            res.headers_mut()
428                .insert(CONTENT_TYPE, "image/png".parse().unwrap());
429            Ok(res)
430        }
431
432        let svc = Compression::new(service_fn(handle));
433
434        let res = svc
435            .serve(
436                Request::builder()
437                    .header(ACCEPT_ENCODING, "gzip")
438                    .body(Body::empty())
439                    .unwrap(),
440            )
441            .await
442            .unwrap();
443        assert!(res.headers().get(CONTENT_ENCODING).is_none());
444    }
445
446    #[tokio::test]
447    async fn does_compress_svg() {
448        async fn handle(_req: Request<Body>) -> Result<Response<Body>, Infallible> {
449            let mut res = Response::new(Body::from(
450                "a".repeat((SizeAbove::DEFAULT_MIN_SIZE * 2) as usize),
451            ));
452            res.headers_mut()
453                .insert(CONTENT_TYPE, "image/svg+xml".parse().unwrap());
454            Ok(res)
455        }
456
457        let svc = Compression::new(service_fn(handle));
458
459        let res = svc
460            .serve(
461                Request::builder()
462                    .header(ACCEPT_ENCODING, "gzip")
463                    .body(Body::empty())
464                    .unwrap(),
465            )
466            .await
467            .unwrap();
468        assert_eq!(res.headers()[CONTENT_ENCODING], "gzip");
469    }
470
471    #[tokio::test]
472    async fn does_compress_grpc_web() {
473        async fn handle(_req: Request<Body>) -> Result<Response<Body>, Infallible> {
474            let mut res = Response::new(Body::from(
475                "a".repeat((SizeAbove::DEFAULT_MIN_SIZE * 2) as usize),
476            ));
477            res.headers_mut()
478                .insert(CONTENT_TYPE, "application/grpc-web+proto".parse().unwrap());
479            Ok(res)
480        }
481
482        let svc = Compression::new(service_fn(handle));
483
484        let res = svc
485            .serve(
486                Request::builder()
487                    .header(ACCEPT_ENCODING, "gzip")
488                    .body(Body::empty())
489                    .unwrap(),
490            )
491            .await
492            .unwrap();
493        assert_eq!(res.headers()[CONTENT_ENCODING], "gzip");
494    }
495
496    #[tokio::test]
497    async fn compress_with_quality() {
498        const DATA: &str = "Check compression quality level! Check compression quality level! Check compression quality level!";
499        let level = CompressionLevel::Best;
500
501        let svc = service_fn(async |_| {
502            let resp = Response::builder()
503                .body(Body::from(DATA.as_bytes()))
504                .unwrap();
505            Ok::<_, std::io::Error>(resp)
506        });
507
508        let svc = Compression::new(svc).with_quality(level);
509
510        // call the service
511        let req = Request::builder()
512            .header("accept-encoding", "br")
513            .body(Body::empty())
514            .unwrap();
515        let res = svc.serve(req).await.unwrap();
516
517        // read the compressed body
518        let body = res.into_body();
519        let compressed_data = body.collect().await.unwrap().to_bytes();
520
521        // build the compressed body with the same quality level
522        let compressed_with_level = {
523            use async_compression::tokio::bufread::BrotliEncoder;
524
525            let stream = Box::pin(rama_core::futures::stream::once(async {
526                Ok::<_, std::io::Error>(DATA.as_bytes())
527            }));
528            let reader = StreamReader::new(stream);
529            let mut enc = BrotliEncoder::with_quality(reader, level.into_async_compression());
530
531            let mut buf = Vec::new();
532            enc.read_to_end(&mut buf).await.unwrap();
533            buf
534        };
535
536        assert_eq!(
537            compressed_data,
538            compressed_with_level.as_slice(),
539            "Compression level is not respected"
540        );
541    }
542
543    #[tokio::test]
544    async fn should_not_compress_ranges() {
545        let svc = service_fn(async |_| {
546            let mut res = Response::new(Body::from("Hello"));
547            let headers = res.headers_mut();
548            headers.insert(ACCEPT_RANGES, "bytes".parse().unwrap());
549            headers.insert(CONTENT_RANGE, "bytes 0-4/*".parse().unwrap());
550            Ok::<_, std::io::Error>(res)
551        });
552        let svc = Compression::new(svc).with_compress_predicate(Always);
553
554        // call the service
555        let req = Request::builder()
556            .header(ACCEPT_ENCODING, "gzip")
557            .header(RANGE, "bytes=0-4")
558            .body(Body::empty())
559            .unwrap();
560        let res = svc.serve(req).await.unwrap();
561        let headers = res.headers().clone();
562
563        // read the uncompressed body
564        let collected = res.into_body().collect().await.unwrap().to_bytes();
565
566        assert_eq!(headers[ACCEPT_RANGES], "bytes");
567        assert!(!headers.contains_key(CONTENT_ENCODING));
568        assert_eq!(collected, "Hello");
569    }
570
571    #[tokio::test]
572    async fn should_strip_accept_ranges_header_when_compressing() {
573        let svc = service_fn(async |_| {
574            let mut res = Response::new(Body::from("Hello, World!"));
575            res.headers_mut()
576                .insert(ACCEPT_RANGES, "bytes".parse().unwrap());
577            Ok::<_, std::io::Error>(res)
578        });
579        let svc = Compression::new(svc).with_compress_predicate(Always);
580
581        // call the service
582        let req = Request::builder()
583            .header(ACCEPT_ENCODING, "gzip")
584            .body(Body::empty())
585            .unwrap();
586        let res = svc.serve(req).await.unwrap();
587        let headers = res.headers().clone();
588
589        // read the compressed body
590        let collected = res.into_body().collect().await.unwrap();
591        let compressed_data = collected.to_bytes();
592
593        // decompress the body
594        // doing this with flate2 as that is much easier than async-compression and blocking during
595        // tests is fine
596        let mut decoder = GzDecoder::new(&compressed_data[..]);
597        let mut decompressed = String::new();
598        decoder.read_to_string(&mut decompressed).unwrap();
599
600        assert!(!headers.contains_key(ACCEPT_RANGES));
601        assert_eq!(headers[CONTENT_ENCODING], "gzip");
602        assert_eq!(decompressed, "Hello, World!");
603    }
604
605    #[tokio::test]
606    async fn trailers_with_empty_body() {
607        let svc = service_fn(|_req: Request<Body>| async {
608            let mut trailers = HeaderMap::new();
609            trailers.insert("grpc-status", "0".parse().unwrap());
610            trailers.insert("grpc-message", "OK".parse().unwrap());
611            let body = Body::empty().with_trailer_headers(trailers);
612            Ok::<_, Infallible>(Response::builder().body(body).unwrap())
613        });
614        let svc = Compression::new(svc).with_compress_predicate(Always);
615
616        let req = Request::builder()
617            .header("accept-encoding", "gzip")
618            .body(Body::empty())
619            .unwrap();
620        let res = svc.serve(req).await.unwrap();
621
622        let collected = res.into_body().collect().await.unwrap();
623        let trailers = collected.trailers().cloned().unwrap();
624        assert_eq!(trailers["grpc-status"], "0");
625        assert_eq!(trailers["grpc-message"], "OK");
626    }
627
628    #[tokio::test]
629    async fn trailers_with_streamed_body() {
630        // Simulate a gRPC-like streamed response: multiple data frames followed by trailers
631        let svc = service_fn(|_req: Request<Body>| async {
632            let stream = rama_core::stream::iter(vec![
633                Ok::<_, BoxError>(Bytes::from("chunk1")),
634                Ok(Bytes::from("chunk2")),
635                Ok(Bytes::from("chunk3")),
636            ]);
637            let mut trailers = HeaderMap::new();
638            trailers.insert("grpc-status", "0".parse().unwrap());
639            let body = Body::from_stream(stream).with_trailer_headers(trailers);
640            Ok::<_, Infallible>(Response::builder().body(body).unwrap())
641        });
642        let svc = Compression::new(svc).with_compress_predicate(Always);
643
644        let req = Request::builder()
645            .header("accept-encoding", "gzip")
646            .body(Body::empty())
647            .unwrap();
648        let res = svc.serve(req).await.unwrap();
649
650        let collected = res.into_body().collect().await.unwrap();
651        let trailers = collected.trailers().cloned().unwrap();
652        let compressed_data = collected.to_bytes();
653
654        let mut decoder = GzDecoder::new(&compressed_data[..]);
655        let mut decompressed = String::new();
656        decoder.read_to_string(&mut decompressed).unwrap();
657
658        assert_eq!(decompressed, "chunk1chunk2chunk3");
659        assert_eq!(trailers["grpc-status"], "0");
660    }
661
662    #[tokio::test]
663    async fn trailers_with_grpc_web_content_type() {
664        let svc = service_fn(|_req: Request<Body>| async {
665            let mut trailers = HeaderMap::new();
666            trailers.insert("grpc-status", "0".parse().unwrap());
667            let body = Body::from("a".repeat((SizeAbove::DEFAULT_MIN_SIZE * 2) as usize))
668                .with_trailer_headers(trailers);
669            let mut res = Response::new(body);
670            res.headers_mut()
671                .insert(CONTENT_TYPE, "application/grpc-web+proto".parse().unwrap());
672            Ok::<_, Infallible>(res)
673        });
674        let svc = Compression::new(svc).with_compress_predicate(Always);
675
676        let req = Request::builder()
677            .header("accept-encoding", "gzip")
678            .body(Body::empty())
679            .unwrap();
680        let res = svc.serve(req).await.unwrap();
681
682        let collected = res.into_body().collect().await.unwrap();
683        let trailers = collected.trailers().cloned().unwrap();
684        assert_eq!(trailers["grpc-status"], "0");
685    }
686
687    #[tokio::test]
688    async fn size_hint_identity() {
689        const MSG: &str = "Hello, world!";
690        let svc = service_fn(async |_| Ok::<_, std::io::Error>(Response::new(Body::from(MSG))));
691        let svc = Compression::new(svc);
692
693        let req = Request::new(Body::empty());
694        let res = svc.serve(req).await.unwrap();
695        let body = res.into_body();
696        assert_eq!(body.size_hint().exact().unwrap(), MSG.len() as u64);
697    }
698
699    // RFC 9110 §9.3.2: server MUST NOT send a body in response to a HEAD request.
700    // Compressing an absent body would produce a spurious Content-Encoding header.
701    #[tokio::test]
702    async fn does_not_compress_head_response() {
703        use rama_http_types::Method;
704        let svc = Compression::new(service_fn(handle)).with_compress_predicate(Always);
705        let req = Request::builder()
706            .method(Method::HEAD)
707            .header(ACCEPT_ENCODING, "gzip")
708            .body(Body::empty())
709            .unwrap();
710        let res = svc.serve(req).await.unwrap();
711        assert!(
712            !res.headers().contains_key(CONTENT_ENCODING),
713            "HEAD response must not carry Content-Encoding"
714        );
715    }
716
717    // RFC 9110 §9.3.6: CONNECT tunnels have no HTTP message body phase.
718    #[tokio::test]
719    async fn does_not_compress_connect_response() {
720        use rama_http_types::Method;
721        let svc = Compression::new(service_fn(handle)).with_compress_predicate(Always);
722        let req = Request::builder()
723            .method(Method::CONNECT)
724            .header(ACCEPT_ENCODING, "gzip")
725            .body(Body::empty())
726            .unwrap();
727        let res = svc.serve(req).await.unwrap();
728        assert!(
729            !res.headers().contains_key(CONTENT_ENCODING),
730            "CONNECT response must not carry Content-Encoding"
731        );
732    }
733
734    // RFC 9110 §15.3.5: 204 No Content responses have no body.
735    #[tokio::test]
736    async fn does_not_compress_204_response() {
737        let svc = Compression::new(service_fn(async |_| {
738            Ok::<_, Infallible>(Response::builder().status(204).body(Body::empty()).unwrap())
739        }))
740        .with_compress_predicate(Always);
741        let req = Request::builder()
742            .header(ACCEPT_ENCODING, "gzip")
743            .body(Body::empty())
744            .unwrap();
745        let res = svc.serve(req).await.unwrap();
746        assert!(
747            !res.headers().contains_key(CONTENT_ENCODING),
748            "204 response must not carry Content-Encoding"
749        );
750    }
751
752    // RFC 9110 §15.4.5: 304 Not Modified responses have no body.
753    #[tokio::test]
754    async fn does_not_compress_304_response() {
755        let svc = Compression::new(service_fn(async |_| {
756            Ok::<_, Infallible>(Response::builder().status(304).body(Body::empty()).unwrap())
757        }))
758        .with_compress_predicate(Always);
759        let req = Request::builder()
760            .header(ACCEPT_ENCODING, "gzip")
761            .body(Body::empty())
762            .unwrap();
763        let res = svc.serve(req).await.unwrap();
764        assert!(
765            !res.headers().contains_key(CONTENT_ENCODING),
766            "304 response must not carry Content-Encoding"
767        );
768    }
769
770    // RFC 9110 §15.2: 1xx Informational responses have no body.
771    #[tokio::test]
772    async fn does_not_compress_1xx_response() {
773        let svc = Compression::new(service_fn(async |_| {
774            Ok::<_, Infallible>(Response::builder().status(100).body(Body::empty()).unwrap())
775        }))
776        .with_compress_predicate(Always);
777        let req = Request::builder()
778            .header(ACCEPT_ENCODING, "gzip")
779            .body(Body::empty())
780            .unwrap();
781        let res = svc.serve(req).await.unwrap();
782        assert!(
783            !res.headers().contains_key(CONTENT_ENCODING),
784            "1xx response must not carry Content-Encoding"
785        );
786    }
787
788    // RFC 9110 §15.3.6: 205 Reset Content responses have no body.
789    #[tokio::test]
790    async fn does_not_compress_205_response() {
791        let svc = Compression::new(service_fn(async |_| {
792            Ok::<_, Infallible>(Response::builder().status(205).body(Body::empty()).unwrap())
793        }))
794        .with_compress_predicate(Always);
795        let req = Request::builder()
796            .header(ACCEPT_ENCODING, "gzip")
797            .body(Body::empty())
798            .unwrap();
799        let res = svc.serve(req).await.unwrap();
800        assert!(
801            !res.headers().contains_key(CONTENT_ENCODING),
802            "205 response must not carry Content-Encoding"
803        );
804    }
805
806    #[tokio::test]
807    async fn wildcard_q_zero_returns_406() {
808        let svc = Compression::new(service_fn(handle)).with_compress_predicate(Always);
809        let req = Request::builder()
810            .header(ACCEPT_ENCODING, "*;q=0")
811            .body(Body::empty())
812            .unwrap();
813        let res = svc.serve(req).await.unwrap();
814
815        assert_eq!(res.status(), crate::StatusCode::NOT_ACCEPTABLE);
816        assert!(
817            res.headers()
818                .get_all(crate::header::VARY)
819                .iter()
820                .any(|v| v.to_str().unwrap().contains("accept-encoding"))
821        );
822    }
823
824    #[tokio::test]
825    async fn wildcard_q_zero_with_gzip_picks_gzip() {
826        let svc = Compression::new(service_fn(handle)).with_compress_predicate(Always);
827        let req = Request::builder()
828            .header(ACCEPT_ENCODING, "*;q=0,gzip")
829            .body(Body::empty())
830            .unwrap();
831        let res = svc.serve(req).await.unwrap();
832
833        assert_eq!(res.status(), crate::StatusCode::OK);
834        assert_eq!(res.headers().get(CONTENT_ENCODING).unwrap(), "gzip");
835    }
836
837    #[tokio::test]
838    async fn wildcard_alone_compresses() {
839        let svc = Compression::new(service_fn(handle)).with_compress_predicate(Always);
840        let req = Request::builder()
841            .header(ACCEPT_ENCODING, "*")
842            .body(Body::empty())
843            .unwrap();
844        let res = svc.serve(req).await.unwrap();
845
846        assert_eq!(res.status(), crate::StatusCode::OK);
847        // `*` makes every encoding acceptable, so the best supported one is picked.
848        assert!(res.headers().contains_key(CONTENT_ENCODING));
849    }
850
851    #[tokio::test]
852    async fn identity_q_zero_alone_returns_406() {
853        let svc = Compression::new(service_fn(handle)).with_compress_predicate(Always);
854        let req = Request::builder()
855            .header(ACCEPT_ENCODING, "identity;q=0")
856            .body(Body::empty())
857            .unwrap();
858        let res = svc.serve(req).await.unwrap();
859
860        assert_eq!(res.status(), crate::StatusCode::NOT_ACCEPTABLE);
861    }
862
863    #[tokio::test]
864    async fn identity_q_zero_with_gzip_picks_gzip() {
865        let svc = Compression::new(service_fn(handle)).with_compress_predicate(Always);
866        let req = Request::builder()
867            .header(ACCEPT_ENCODING, "identity;q=0,gzip")
868            .body(Body::empty())
869            .unwrap();
870        let res = svc.serve(req).await.unwrap();
871
872        assert_eq!(res.status(), crate::StatusCode::OK);
873        assert_eq!(res.headers().get(CONTENT_ENCODING).unwrap(), "gzip");
874    }
875
876    #[tokio::test]
877    async fn enforce_not_acceptable_opt_out_falls_back_to_identity() {
878        // With 406 enforcement disabled, an otherwise-unsatisfiable Accept-Encoding falls back
879        // to an uncompressed identity response instead of 406.
880        let svc = Compression::new(service_fn(handle))
881            .with_compress_predicate(Always)
882            .with_enforce_not_acceptable(false);
883        let req = Request::builder()
884            .header(ACCEPT_ENCODING, "*;q=0")
885            .body(Body::empty())
886            .unwrap();
887        let res = svc.serve(req).await.unwrap();
888
889        assert_eq!(res.status(), crate::StatusCode::OK);
890        assert!(!res.headers().contains_key(CONTENT_ENCODING));
891    }
892}