Skip to main content

rama_http/layer/compression/stream/
layer.rs

1use super::StreamCompression;
2use crate::headers::encoding::AcceptEncoding;
3use crate::layer::compression::Predicate;
4use crate::layer::compression::predicate::DefaultStreamPredicate;
5use crate::layer::util::compression::CompressionLevel;
6use rama_core::Layer;
7
8/// Compress response bodies of the underlying service.
9///
10/// This uses the `Accept-Encoding` header to pick an appropriate encoding and adds the
11/// `Content-Encoding` header to responses.
12///
13/// See the [module docs](crate::layer::compression) for more details.
14#[derive(Clone, Debug)]
15pub struct StreamCompressionLayer<P = DefaultStreamPredicate> {
16    accept: AcceptEncoding,
17    predicate: P,
18    quality: CompressionLevel,
19    enforce_not_acceptable: bool,
20}
21
22impl<P: Default> Default for StreamCompressionLayer<P> {
23    fn default() -> Self {
24        Self {
25            accept: AcceptEncoding::default(),
26            predicate: P::default(),
27            quality: CompressionLevel::default(),
28            enforce_not_acceptable: true,
29        }
30    }
31}
32
33impl<S, P> Layer<S> for StreamCompressionLayer<P>
34where
35    P: Predicate,
36{
37    type Service = StreamCompression<S, P>;
38
39    fn layer(&self, inner: S) -> Self::Service {
40        StreamCompression {
41            inner,
42            accept: self.accept,
43            predicate: self.predicate.clone(),
44            quality: self.quality,
45            enforce_not_acceptable: self.enforce_not_acceptable,
46        }
47    }
48
49    fn into_layer(self, inner: S) -> Self::Service {
50        StreamCompression {
51            inner,
52            accept: self.accept,
53            predicate: self.predicate,
54            quality: self.quality,
55            enforce_not_acceptable: self.enforce_not_acceptable,
56        }
57    }
58}
59
60impl StreamCompressionLayer {
61    /// Creates a new [`StreamCompressionLayer`].
62    #[must_use]
63    pub fn new() -> Self {
64        Self::default()
65    }
66
67    /// Replace the current compression predicate.
68    pub fn with_compress_predicate<C>(self, predicate: C) -> StreamCompressionLayer<C>
69    where
70        C: Predicate,
71    {
72        StreamCompressionLayer {
73            accept: self.accept,
74            predicate,
75            quality: self.quality,
76            enforce_not_acceptable: self.enforce_not_acceptable,
77        }
78    }
79}
80
81impl<P> StreamCompressionLayer<P> {
82    rama_utils::macros::generate_set_and_with! {
83        /// Sets whether to enable the gzip encoding.
84        pub fn gzip(mut self, enable: bool) -> Self {
85            self.accept.set_gzip(enable);
86            self
87        }
88    }
89
90    rama_utils::macros::generate_set_and_with! {
91        /// Sets whether to enable the Deflate encoding.
92        pub fn deflate(mut self, enable: bool) -> Self {
93            self.accept.set_deflate(enable);
94            self
95        }
96    }
97
98    rama_utils::macros::generate_set_and_with! {
99        /// Sets whether to enable the Brotli encoding.
100        pub fn br(mut self, enable: bool) -> Self {
101            self.accept.set_br(enable);
102            self
103        }
104    }
105
106    rama_utils::macros::generate_set_and_with! {
107        /// Sets whether to enable the Zstd encoding.
108        pub fn zstd(mut self, enable: bool) -> Self {
109            self.accept.set_zstd(enable);
110            self
111        }
112    }
113
114    rama_utils::macros::generate_set_and_with! {
115        /// Sets the compression quality.
116        pub fn quality(mut self, quality: CompressionLevel) -> Self {
117            self.quality = quality;
118            self
119        }
120    }
121
122    rama_utils::macros::generate_set_and_with! {
123        /// Sets whether to respond with `406 Not Acceptable` when the client's
124        /// `Accept-Encoding` header rejects every available representation
125        /// (e.g. `*;q=0` or a lone `identity;q=0`), as recommended by RFC 9110 §12.5.3.
126        ///
127        /// Enabled by default. Disable to opt out and instead fall back to sending an
128        /// uncompressed (identity) response regardless of the client's stated preference.
129        pub fn enforce_not_acceptable(mut self, enable: bool) -> Self {
130            self.enforce_not_acceptable = enable;
131            self
132        }
133    }
134}
135
136#[cfg(test)]
137mod tests {
138    use super::*;
139
140    use crate::layer::compression::predicate::MirrorDecompressed;
141    use crate::layer::decompression::DecompressedFrom;
142    use crate::{Request, Response, body::util::BodyExt, header::ACCEPT_ENCODING};
143    use rama_core::Service;
144    use rama_core::extensions::ExtensionsRef;
145    use rama_core::service::service_fn;
146    use rama_core::stream::io::ReaderStream;
147    use rama_http_types::Body;
148    use std::convert::Infallible;
149    use tokio::fs::File;
150
151    async fn handle(_req: Request) -> Result<Response, Infallible> {
152        // Open the file.
153        let file = File::open("Cargo.toml").await.expect("file missing");
154        // Convert the file into a `Stream`.
155        let stream = ReaderStream::new(file);
156        // Convert the `Stream` into a `Body`.
157        let body = Body::from_stream(stream);
158        // Create response.
159        Ok(Response::new(body))
160    }
161
162    #[tokio::test]
163    async fn accept_encoding_configuration_works() -> Result<(), rama_core::error::BoxError> {
164        use std::io::Read;
165
166        fn decode<R: Read>(mut r: R) -> std::io::Result<Vec<u8>> {
167            let mut buf = Vec::new();
168            r.read_to_end(&mut buf)?;
169            Ok(buf)
170        }
171
172        // Read the source file once so we can verify each response round-trips to the same bytes.
173        let expected = tokio::fs::read("Cargo.toml").await?;
174
175        // Configure a layer that only offers deflate, then confirm the response is actually
176        // deflate-encoded by decoding it and comparing to the original content.
177        let deflate_only_layer = StreamCompressionLayer::new()
178            .with_quality(CompressionLevel::Best)
179            .with_br(false)
180            .with_gzip(false);
181
182        let service = deflate_only_layer.into_layer(service_fn(handle));
183
184        let request = Request::builder()
185            .header(ACCEPT_ENCODING, "gzip, deflate, br")
186            .body(Body::empty())?;
187
188        let response = service.serve(request).await?;
189
190        assert_eq!(response.headers()["content-encoding"], "deflate");
191
192        let deflate_body = response.into_body().collect().await?.to_bytes();
193
194        // The "deflate" Content-Encoding is RFC 1950 zlib framing (2-byte header + Adler-32),
195        // not raw RFC 1951 deflate, so use ZlibDecoder rather than DeflateDecoder.
196        let decoded = decode(flate2::bufread::ZlibDecoder::new(&deflate_body[..]))?;
197        assert_eq!(decoded, expected);
198
199        // Same check for brotli.
200        let br_only_layer = StreamCompressionLayer::new()
201            .with_quality(CompressionLevel::Best)
202            .with_gzip(false)
203            .with_deflate(false);
204
205        let service = br_only_layer.into_layer(service_fn(handle));
206
207        let request = Request::builder()
208            .header(ACCEPT_ENCODING, "gzip, deflate, br")
209            .body(Body::empty())?;
210
211        let response = service.serve(request).await?;
212
213        assert_eq!(response.headers()["content-encoding"], "br");
214
215        let br_body = response.into_body().collect().await?.to_bytes();
216
217        // 4096 is the decoder's internal read-buffer size, not a content-length bound.
218        let decoded = decode(brotli::Decompressor::new(&br_body[..], 4096))?;
219        assert_eq!(decoded, expected);
220
221        Ok(())
222    }
223
224    #[tokio::test]
225    async fn zstd_is_web_safe() -> Result<(), rama_core::error::BoxError> {
226        // Test ensuring that zstd compression will not exceed an 8MiB window size; browsers do not
227        // accept responses using 16MiB+ window sizes.
228
229        async fn zeroes(_req: Request<Body>) -> Result<Response<Body>, Infallible> {
230            Ok(Response::new(Body::from(vec![0u8; 18_874_368])))
231        }
232        // zstd will (I believe) lower its window size if a larger one isn't beneficial and
233        // it knows the size of the input; use an 18MiB body to ensure it would want a
234        // >=16MiB window (though it might not be able to see the input size here).
235
236        let zstd_layer = StreamCompressionLayer::new()
237            .with_quality(CompressionLevel::Best)
238            .with_br(false)
239            .with_deflate(false)
240            .with_gzip(false);
241
242        let service = zstd_layer.into_layer(service_fn(zeroes));
243
244        let request = Request::builder()
245            .header(ACCEPT_ENCODING, "zstd")
246            .body(Body::empty())?;
247
248        let response = service.serve(request).await?;
249
250        assert_eq!(response.headers()["content-encoding"], "zstd");
251
252        let body = response.into_body();
253        let bytes = body.collect().await?.to_bytes();
254        let mut dec = zstd::Decoder::new(&*bytes)?;
255        dec.window_log_max(23)?; // Limit window size accepted by decoder to 2 ^ 23 bytes (8MiB)
256
257        std::io::copy(&mut dec, &mut std::io::sink())?;
258
259        Ok(())
260    }
261
262    #[tokio::test]
263    async fn mirror_decompressed_prefers_original_encoding()
264    -> Result<(), rama_core::error::BoxError> {
265        let service = StreamCompressionLayer::new()
266            .with_compress_predicate(MirrorDecompressed::new())
267            .into_layer(service_fn(|_: Request<Body>| async {
268                let res = Response::new(Body::from("Hello, World! Hello, World! Hello, World!"));
269                res.extensions().insert(DecompressedFrom::Brotli);
270                Ok::<_, Infallible>(res)
271            }));
272
273        let request = Request::builder()
274            .header(ACCEPT_ENCODING, "gzip, br")
275            .body(Body::empty())?;
276
277        let response = service.serve(request).await?;
278
279        assert_eq!(response.headers()["content-encoding"], "br");
280
281        Ok(())
282    }
283
284    // RFC 9110 §9.3.2: server MUST NOT send a body in response to a HEAD request.
285    #[tokio::test]
286    async fn does_not_compress_head_response() {
287        use crate::header::CONTENT_ENCODING;
288        use rama_http_types::Method;
289        let service = StreamCompressionLayer::new().into_layer(service_fn(handle));
290        let req = Request::builder()
291            .method(Method::HEAD)
292            .header(ACCEPT_ENCODING, "gzip")
293            .body(Body::empty())
294            .unwrap();
295        let res = service.serve(req).await.unwrap();
296        assert!(
297            !res.headers().contains_key(CONTENT_ENCODING),
298            "HEAD response must not carry Content-Encoding"
299        );
300    }
301
302    // RFC 9110 §9.3.6: CONNECT tunnels have no HTTP message body phase.
303    #[tokio::test]
304    async fn does_not_compress_connect_response() {
305        use crate::header::CONTENT_ENCODING;
306        use rama_http_types::Method;
307        let service = StreamCompressionLayer::new().into_layer(service_fn(handle));
308        let req = Request::builder()
309            .method(Method::CONNECT)
310            .header(ACCEPT_ENCODING, "gzip")
311            .body(Body::empty())
312            .unwrap();
313        let res = service.serve(req).await.unwrap();
314        assert!(
315            !res.headers().contains_key(CONTENT_ENCODING),
316            "CONNECT response must not carry Content-Encoding"
317        );
318    }
319
320    // RFC 9110 §15.3.5: 204 No Content responses have no body.
321    #[tokio::test]
322    async fn does_not_compress_204_response() {
323        use crate::header::CONTENT_ENCODING;
324        let service =
325            StreamCompressionLayer::new().into_layer(service_fn(async |_: Request<Body>| {
326                Ok::<_, Infallible>(Response::builder().status(204).body(Body::empty()).unwrap())
327            }));
328        let req = Request::builder()
329            .header(ACCEPT_ENCODING, "gzip")
330            .body(Body::empty())
331            .unwrap();
332        let res = service.serve(req).await.unwrap();
333        assert!(
334            !res.headers().contains_key(CONTENT_ENCODING),
335            "204 response must not carry Content-Encoding"
336        );
337    }
338
339    // RFC 9110 §15.4.5: 304 Not Modified responses have no body.
340    #[tokio::test]
341    async fn does_not_compress_304_response() {
342        use crate::header::CONTENT_ENCODING;
343        let service =
344            StreamCompressionLayer::new().into_layer(service_fn(async |_: Request<Body>| {
345                Ok::<_, Infallible>(Response::builder().status(304).body(Body::empty()).unwrap())
346            }));
347        let req = Request::builder()
348            .header(ACCEPT_ENCODING, "gzip")
349            .body(Body::empty())
350            .unwrap();
351        let res = service.serve(req).await.unwrap();
352        assert!(
353            !res.headers().contains_key(CONTENT_ENCODING),
354            "304 response must not carry Content-Encoding"
355        );
356    }
357
358    // RFC 9110 §15.2: 1xx Informational responses have no body.
359    #[tokio::test]
360    async fn does_not_compress_1xx_response() {
361        use crate::header::CONTENT_ENCODING;
362        let service =
363            StreamCompressionLayer::new().into_layer(service_fn(async |_: Request<Body>| {
364                Ok::<_, Infallible>(Response::builder().status(100).body(Body::empty()).unwrap())
365            }));
366        let req = Request::builder()
367            .header(ACCEPT_ENCODING, "gzip")
368            .body(Body::empty())
369            .unwrap();
370        let res = service.serve(req).await.unwrap();
371        assert!(
372            !res.headers().contains_key(CONTENT_ENCODING),
373            "1xx response must not carry Content-Encoding"
374        );
375    }
376
377    // RFC 9110 §15.3.6: 205 Reset Content responses have no body.
378    #[tokio::test]
379    async fn does_not_compress_205_response() {
380        use crate::header::CONTENT_ENCODING;
381        let service =
382            StreamCompressionLayer::new().into_layer(service_fn(async |_: Request<Body>| {
383                Ok::<_, Infallible>(Response::builder().status(205).body(Body::empty()).unwrap())
384            }));
385        let req = Request::builder()
386            .header(ACCEPT_ENCODING, "gzip")
387            .body(Body::empty())
388            .unwrap();
389        let res = service.serve(req).await.unwrap();
390        assert!(
391            !res.headers().contains_key(CONTENT_ENCODING),
392            "205 response must not carry Content-Encoding"
393        );
394    }
395
396    // RFC 9110 §14.2: partial-content responses carry Content-Range; compressing
397    // them would corrupt the byte-range offsets the client uses to reassemble the
398    // resource, so the service must pass them through unchanged.
399    #[tokio::test]
400    async fn does_not_compress_range_response() {
401        use crate::header::{CONTENT_ENCODING, CONTENT_RANGE};
402        let service =
403            StreamCompressionLayer::new().into_layer(service_fn(async |_: Request<Body>| {
404                Ok::<_, Infallible>(
405                    Response::builder()
406                        .status(206)
407                        .header(CONTENT_RANGE, "bytes 0-4/10")
408                        .body(Body::from("hello"))
409                        .unwrap(),
410                )
411            }));
412        let req = Request::builder()
413            .header(ACCEPT_ENCODING, "gzip")
414            .body(Body::empty())
415            .unwrap();
416        let res = service.serve(req).await.unwrap();
417        assert!(
418            !res.headers().contains_key(CONTENT_ENCODING),
419            "range response must not carry Content-Encoding"
420        );
421    }
422
423    // RFC 9110 §12.5.3: `*;q=0` rejects every representation, so the negotiation is
424    // unsatisfiable and the middleware responds 406 Not Acceptable by default.
425    #[tokio::test]
426    async fn wildcard_q_zero_returns_406() {
427        use crate::StatusCode;
428        let service = StreamCompressionLayer::new().into_layer(service_fn(handle));
429        let req = Request::builder()
430            .header(ACCEPT_ENCODING, "*;q=0")
431            .body(Body::empty())
432            .unwrap();
433        let res = service.serve(req).await.unwrap();
434        assert_eq!(res.status(), StatusCode::NOT_ACCEPTABLE);
435    }
436
437    // Disabling enforcement falls back to an uncompressed identity response instead of 406.
438    #[tokio::test]
439    async fn enforce_not_acceptable_opt_out_falls_back_to_identity() {
440        use crate::StatusCode;
441        use crate::header::CONTENT_ENCODING;
442        let service = StreamCompressionLayer::new()
443            .with_enforce_not_acceptable(false)
444            .into_layer(service_fn(handle));
445        let req = Request::builder()
446            .header(ACCEPT_ENCODING, "*;q=0")
447            .body(Body::empty())
448            .unwrap();
449        let res = service.serve(req).await.unwrap();
450        assert_eq!(res.status(), StatusCode::OK);
451        assert!(!res.headers().contains_key(CONTENT_ENCODING));
452    }
453}