Skip to main content

tower_http/decompression/
mod.rs

1//! Middleware that decompresses request and response bodies.
2//!
3//! # Examples
4//!
5//! #### Request
6//!
7//! ```rust
8//! use bytes::Bytes;
9//! use flate2::{write::GzEncoder, Compression};
10//! use http::{header, HeaderValue, Request, Response};
11//! use http_body_util::{Full, BodyExt};
12//! use std::{error::Error, io::Write};
13//! use tower::{Service, ServiceBuilder, service_fn, ServiceExt};
14//! use tower_http::{BoxError, decompression::{DecompressionBody, RequestDecompressionLayer}};
15//!
16//! # #[tokio::main]
17//! # async fn main() -> Result<(), BoxError> {
18//! // A request encoded with gzip coming from some HTTP client.
19//! let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
20//! encoder.write_all(b"Hello?")?;
21//! let request = Request::builder()
22//!     .header(header::CONTENT_ENCODING, "gzip")
23//!     .body(Full::from(encoder.finish()?))?;
24//!
25//! // Our HTTP server
26//! let mut server = ServiceBuilder::new()
27//!     // Automatically decompress request bodies.
28//!     .layer(RequestDecompressionLayer::new())
29//!     .service(service_fn(handler));
30//!
31//! // Send the request, with the gzip encoded body, to our server.
32//! let _response = server.ready().await?.call(request).await?;
33//!
34//! // Handler receives request whose body is decoded when read
35//! async fn handler(
36//!     mut req: Request<DecompressionBody<Full<Bytes>>>,
37//! ) -> Result<Response<Full<Bytes>>, BoxError>{
38//!     let data = req.into_body().collect().await?.to_bytes();
39//!     assert_eq!(&data[..], b"Hello?");
40//!     Ok(Response::new(Full::from("Hello, World!")))
41//! }
42//! # Ok(())
43//! # }
44//! ```
45//!
46//! #### Response
47//!
48//! ```rust
49//! use bytes::Bytes;
50//! use http::{Request, Response};
51//! use http_body_util::{Full, BodyExt};
52//! use std::convert::Infallible;
53//! use tower::{Service, ServiceExt, ServiceBuilder, service_fn};
54//! use tower_http::{compression::Compression, decompression::DecompressionLayer, BoxError};
55//! #
56//! # #[tokio::main]
57//! # async fn main() -> Result<(), tower_http::BoxError> {
58//! # async fn handle(req: Request<Full<Bytes>>) -> Result<Response<Full<Bytes>>, Infallible> {
59//! #     let body = Full::from("Hello, World!");
60//! #     Ok(Response::new(body))
61//! # }
62//!
63//! // Some opaque service that applies compression.
64//! let service = Compression::new(service_fn(handle));
65//!
66//! // Our HTTP client.
67//! let mut client = ServiceBuilder::new()
68//!     // Automatically decompress response bodies.
69//!     .layer(DecompressionLayer::new())
70//!     .service(service);
71//!
72//! // Call the service.
73//! //
74//! // `DecompressionLayer` takes care of setting `Accept-Encoding`.
75//! let request = Request::new(Full::<Bytes>::default());
76//!
77//! let response = client
78//!     .ready()
79//!     .await?
80//!     .call(request)
81//!     .await?;
82//!
83//! // Read the body
84//! let body = response.into_body();
85//! let bytes = body.collect().await?.to_bytes().to_vec();
86//! let body = String::from_utf8(bytes).map_err(Into::<BoxError>::into)?;
87//!
88//! assert_eq!(body, "Hello, World!");
89//! #
90//! # Ok(())
91//! # }
92//! ```
93
94mod request;
95
96mod body;
97mod future;
98mod layer;
99mod service;
100
101pub use self::{
102    body::DecompressionBody, future::ResponseFuture, layer::DecompressionLayer,
103    service::Decompression,
104};
105
106pub use self::request::future::RequestDecompressionFuture;
107pub use self::request::layer::RequestDecompressionLayer;
108pub use self::request::service::RequestDecompression;
109
110#[cfg(test)]
111mod tests {
112    use std::convert::Infallible;
113    use std::io::Write;
114    use std::time::Duration;
115
116    use super::*;
117    use crate::test_helpers::Body;
118    use crate::{compression::Compression, test_helpers::WithTrailers};
119    use bytes::Bytes;
120    use futures_util::StreamExt;
121    use http::Response;
122    use http::{HeaderMap, HeaderName, Request};
123    use http_body_util::BodyExt;
124    use tower::{service_fn, Service, ServiceExt};
125
126    #[tokio::test]
127    async fn works() {
128        let mut client = Decompression::new(Compression::new(service_fn(handle)));
129
130        let req = Request::builder()
131            .header("accept-encoding", "gzip")
132            .body(Body::empty())
133            .unwrap();
134        let res = client.ready().await.unwrap().call(req).await.unwrap();
135
136        // read the body, it will be decompressed automatically
137        let body = res.into_body();
138        let collected = body.collect().await.unwrap();
139        let trailers = collected.trailers().cloned().unwrap();
140        let decompressed_data = String::from_utf8(collected.to_bytes().to_vec()).unwrap();
141
142        assert_eq!(decompressed_data, "Hello, World!");
143
144        // maintains trailers
145        assert_eq!(trailers["foo"], "bar");
146    }
147
148    async fn handle(_req: Request<Body>) -> Result<Response<WithTrailers<Body>>, Infallible> {
149        let mut trailers = HeaderMap::new();
150        trailers.insert(HeaderName::from_static("foo"), "bar".parse().unwrap());
151        let body = Body::from("Hello, World!").with_trailers(trailers);
152        Ok(Response::builder().body(body).unwrap())
153    }
154
155    #[tokio::test]
156    async fn decompress_multi_zstd() {
157        let mut client = Decompression::new(service_fn(handle_multi_zstd));
158
159        let req = Request::builder()
160            .header("accept-encoding", "zstd")
161            .body(Body::empty())
162            .unwrap();
163        let res = client.ready().await.unwrap().call(req).await.unwrap();
164
165        // read the body, it will be decompressed automatically
166        let body = res.into_body();
167        let decompressed_data =
168            String::from_utf8(body.collect().await.unwrap().to_bytes().to_vec()).unwrap();
169
170        assert_eq!(decompressed_data, "Hello, World!");
171    }
172
173    async fn handle_multi_zstd(_req: Request<Body>) -> Result<Response<Body>, Infallible> {
174        let mut buf = Vec::new();
175        let mut enc1 = zstd::Encoder::new(&mut buf, Default::default()).unwrap();
176        enc1.write_all(b"Hello, ").unwrap();
177        enc1.finish().unwrap();
178
179        let mut enc2 = zstd::Encoder::new(&mut buf, Default::default()).unwrap();
180        enc2.write_all(b"World!").unwrap();
181        enc2.finish().unwrap();
182
183        let mut res = Response::new(Body::from(buf));
184        res.headers_mut()
185            .insert("content-encoding", "zstd".parse().unwrap());
186        Ok(res)
187    }
188
189    #[allow(dead_code)]
190    async fn is_compatible_with_hyper() {
191        let client =
192            hyper_util::client::legacy::Client::builder(hyper_util::rt::TokioExecutor::new())
193                .build_http();
194        let mut client = Decompression::new(client);
195
196        let req = Request::new(Body::empty());
197
198        let _: Response<DecompressionBody<_>> =
199            client.ready().await.unwrap().call(req).await.unwrap();
200    }
201
202    #[tokio::test]
203    async fn decompress_empty() {
204        let mut client = Decompression::new(Compression::new(service_fn(handle_empty)));
205
206        let req = Request::builder()
207            .header("accept-encoding", "gzip")
208            .body(Body::empty())
209            .unwrap();
210        let res = client.ready().await.unwrap().call(req).await.unwrap();
211
212        let body = res.into_body();
213        let decompressed_data =
214            String::from_utf8(body.collect().await.unwrap().to_bytes().to_vec()).unwrap();
215
216        assert_eq!(decompressed_data, "");
217    }
218
219    async fn handle_empty(_req: Request<Body>) -> Result<Response<Body>, Infallible> {
220        let mut res = Response::new(Body::empty());
221        res.headers_mut()
222            .insert("content-encoding", "gzip".parse().unwrap());
223        Ok(res)
224    }
225
226    #[tokio::test]
227    async fn decompress_empty_with_trailers() {
228        let mut client =
229            Decompression::new(Compression::new(service_fn(handle_empty_with_trailers)));
230
231        let req = Request::builder()
232            .header("accept-encoding", "gzip")
233            .body(Body::empty())
234            .unwrap();
235        let res = client.ready().await.unwrap().call(req).await.unwrap();
236
237        let body = res.into_body();
238        let collected = body.collect().await.unwrap();
239        let trailers = collected.trailers().cloned().unwrap();
240        let decompressed_data = String::from_utf8(collected.to_bytes().to_vec()).unwrap();
241
242        assert_eq!(decompressed_data, "");
243        assert_eq!(trailers["foo"], "bar");
244    }
245
246    async fn handle_empty_with_trailers(
247        _req: Request<Body>,
248    ) -> Result<Response<WithTrailers<Body>>, Infallible> {
249        let mut trailers = HeaderMap::new();
250        trailers.insert(HeaderName::from_static("foo"), "bar".parse().unwrap());
251        let body = Body::empty().with_trailers(trailers);
252        Ok(Response::builder()
253            .header("content-encoding", "gzip")
254            .body(body)
255            .unwrap())
256    }
257
258    #[cfg(feature = "decompression-br")]
259    #[tokio::test]
260    async fn brotli_rejects_extra_data_without_waiting_for_end_of_body() {
261        let mut compressed = Vec::new();
262        {
263            let mut encoder = brotli::CompressorWriter::new(&mut compressed, 4096, 5, 20);
264            encoder.write_all(b"Hello, World!").unwrap();
265        }
266
267        let svc = service_fn(move |_req: Request<Body>| {
268            let compressed = compressed.clone();
269            async move {
270                let stream = futures_util::stream::iter([
271                    Ok::<_, Infallible>(Bytes::from(compressed)),
272                    Ok(Bytes::from_static(b"extra")),
273                ])
274                .chain(futures_util::stream::pending());
275
276                Ok::<_, Infallible>(
277                    Response::builder()
278                        .header("content-encoding", "br")
279                        .body(Body::from_stream(stream))
280                        .unwrap(),
281                )
282            }
283        });
284        let mut client = Decompression::new(svc);
285
286        let res = client
287            .ready()
288            .await
289            .unwrap()
290            .call(Request::new(Body::empty()))
291            .await
292            .unwrap();
293
294        let result = tokio::time::timeout(Duration::from_secs(1), res.into_body().collect())
295            .await
296            .expect("extra data should produce an error without waiting for the body to end");
297        let error = result.unwrap_err();
298
299        assert_eq!(
300            error.to_string(),
301            "there are extra bytes after body has been decompressed"
302        );
303    }
304
305    #[cfg(feature = "decompression-br")]
306    #[tokio::test]
307    async fn brotli_keeps_trailers_that_follow_an_empty_data_frame() {
308        let mut compressed = Vec::new();
309        {
310            let mut encoder = brotli::CompressorWriter::new(&mut compressed, 4096, 5, 20);
311            encoder.write_all(b"Hello, World!").unwrap();
312        }
313
314        let svc = service_fn(move |_req: Request<Body>| {
315            let compressed = compressed.clone();
316            async move {
317                // An empty data frame between the payload and the trailers is
318                // legal, and must not be read as the end of the body.
319                let stream = futures_util::stream::iter([
320                    Ok::<_, Infallible>(Bytes::from(compressed)),
321                    Ok(Bytes::new()),
322                ]);
323                let mut trailers = HeaderMap::new();
324                trailers.insert(HeaderName::from_static("foo"), "bar".parse().unwrap());
325
326                Ok::<_, Infallible>(
327                    Response::builder()
328                        .header("content-encoding", "br")
329                        .body(Body::from_stream(stream).with_trailers(trailers))
330                        .unwrap(),
331                )
332            }
333        });
334        let mut client = Decompression::new(svc);
335
336        let res = client
337            .ready()
338            .await
339            .unwrap()
340            .call(Request::new(Body::empty()))
341            .await
342            .unwrap();
343
344        let collected = res.into_body().collect().await.unwrap();
345        let trailers = collected
346            .trailers()
347            .cloned()
348            .expect("trailers following an empty data frame should still arrive");
349
350        assert_eq!(trailers["foo"], "bar");
351        assert_eq!(
352            String::from_utf8(collected.to_bytes().to_vec()).unwrap(),
353            "Hello, World!"
354        );
355    }
356}