rama_http/layer/compression/stream/
layer.rs1use 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#[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 #[must_use]
63 pub fn new() -> Self {
64 Self::default()
65 }
66
67 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 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 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 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 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 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 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 let file = File::open("Cargo.toml").await.expect("file missing");
154 let stream = ReaderStream::new(file);
156 let body = Body::from_stream(stream);
158 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 let expected = tokio::fs::read("Cargo.toml").await?;
174
175 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 let decoded = decode(flate2::bufread::ZlibDecoder::new(&deflate_body[..]))?;
197 assert_eq!(decoded, expected);
198
199 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 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 async fn zeroes(_req: Request<Body>) -> Result<Response<Body>, Infallible> {
230 Ok(Response::new(Body::from(vec![0u8; 18_874_368])))
231 }
232 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)?; 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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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}