Skip to main content

reth_rpc_layer/
decompression_layer.rs

1use http::StatusCode;
2use http_body_util::{BodyExt, LengthLimitError, Limited};
3use jsonrpsee_http_client::{HttpBody, HttpRequest, HttpResponse};
4use std::{
5    future::{poll_fn, Future},
6    pin::Pin,
7    sync::Arc,
8    task::{Context, Poll},
9};
10use tokio::sync::Semaphore;
11use tower::{Layer, Service};
12use tower_http::decompression::{
13    RequestDecompression, RequestDecompressionLayer as TowerDecompressionLayer,
14};
15use tracing::debug;
16
17/// Maximum number of request bodies that may be decompressed concurrently.
18///
19/// Each in-flight decompression can materialize up to the configured maximum body size, so this
20/// bounds the worst-case memory usage at `MAX_CONCURRENT_DECOMPRESSIONS * max_body_size` (120 MiB
21/// with the default 15 MiB request size limit). Excess requests wait for a permit instead of
22/// being rejected.
23const MAX_CONCURRENT_DECOMPRESSIONS: usize = 8;
24
25/// This layer is a wrapper around [`tower_http::decompression::RequestDecompressionLayer`] that
26/// integrates with jsonrpsee's HTTP types.
27#[expect(missing_debug_implementations)]
28#[derive(Clone)]
29pub struct DecompressionLayer {
30    inner_layer: TowerDecompressionLayer,
31    /// Maximum size in bytes for both compressed and decompressed bodies.
32    max_body_size: usize,
33    /// Bounds concurrent decompression work across all services created from this layer.
34    decompression_permits: Arc<Semaphore>,
35}
36
37impl DecompressionLayer {
38    /// Creates a new decompression layer from a list of algorithm names.
39    /// Supported: zstd, gzip, deflate, br
40    pub fn new(algos: &[impl AsRef<str>], max_body_size: usize) -> Self {
41        // Start with all algorithms explicitly disabled
42        let mut layer = TowerDecompressionLayer::new().no_zstd().no_gzip().no_deflate().no_br();
43
44        // Only enable the algorithms that were explicitly passed.
45        for algo in algos {
46            match algo.as_ref() {
47                "zstd" => layer = layer.zstd(true),
48                "gzip" => layer = layer.gzip(true),
49                "deflate" => layer = layer.deflate(true),
50                "br" | "brotli" => layer = layer.br(true),
51                _ => {}
52            }
53        }
54
55        Self {
56            inner_layer: layer,
57            max_body_size,
58            decompression_permits: Arc::new(Semaphore::new(MAX_CONCURRENT_DECOMPRESSIONS)),
59        }
60    }
61}
62
63impl<S> Layer<S> for DecompressionLayer {
64    type Service = DecompressionService<S>;
65
66    fn layer(&self, inner: S) -> Self::Service {
67        DecompressionService {
68            decompression: self.inner_layer.layer(InnerService {
69                inner,
70                max_body_size: self.max_body_size,
71                decompression_permits: self.decompression_permits.clone(),
72            }),
73            max_body_size: self.max_body_size,
74        }
75    }
76}
77
78/// Service that performs request decompression with body size limiting.
79///
80/// Created by [`DecompressionLayer`].
81#[expect(missing_debug_implementations)]
82#[derive(Clone)]
83pub struct DecompressionService<S> {
84    decompression: RequestDecompression<InnerService<S>>,
85    max_body_size: usize,
86}
87
88/// Inner service wrapper to handle type conversion between jsonrpsee and `tower_http`
89/// with body size limiting.
90#[derive(Clone)]
91struct InnerService<S> {
92    inner: S,
93    max_body_size: usize,
94    decompression_permits: Arc<Semaphore>,
95}
96
97/// Marks a request whose body Tower will decompress before it reaches [`InnerService`].
98#[derive(Clone, Copy, Debug)]
99struct CompressedBody {
100    content_length: Option<u64>,
101}
102
103impl<S> Service<http::Request<tower_http::decompression::DecompressionBody<HttpBody>>>
104    for InnerService<S>
105where
106    S: Service<HttpRequest, Response = HttpResponse> + Clone + Send + 'static,
107    S::Future: Send + 'static,
108{
109    type Response = http::Response<HttpBody>;
110    type Error = S::Error;
111    type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
112
113    fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
114        Poll::Ready(Ok(()))
115    }
116
117    fn call(
118        &mut self,
119        req: http::Request<tower_http::decompression::DecompressionBody<HttpBody>>,
120    ) -> Self::Future {
121        let mut inner = self.inner.clone();
122        let max_body_size = self.max_body_size;
123        let decompression_permits = self.decompression_permits.clone();
124        Box::pin(async move {
125            let (mut parts, body) = req.into_parts();
126
127            let Some(compressed) = parts.extensions.remove::<CompressedBody>() else {
128                poll_fn(|cx| inner.poll_ready(cx)).await?;
129                return inner.call(HttpRequest::from_parts(parts, HttpBody::new(body))).await;
130            };
131
132            if compressed.content_length.is_some_and(|length| length > max_body_size as u64) {
133                return Ok(err_response(StatusCode::PAYLOAD_TOO_LARGE, "Payload Too Large"));
134            }
135
136            // HTTP/2 allows many concurrent streams per connection, so bound how many bodies are
137            // decompressed and materialized at once to cap memory and CPU usage.
138            let permit = decompression_permits
139                .acquire_owned()
140                .await
141                .expect("decompression semaphore is never closed");
142
143            let body = match Limited::new(body, max_body_size).collect().await {
144                Ok(body) => body,
145                Err(err) if err.is::<LengthLimitError>() => {
146                    return Ok(err_response(StatusCode::PAYLOAD_TOO_LARGE, "Payload Too Large"));
147                }
148                Err(err) => {
149                    debug!(target: "rpc::decompression", %err, "Failed to decompress request body");
150                    return Ok(err_response(StatusCode::BAD_REQUEST, "Invalid compressed body"));
151                }
152            };
153
154            // Decompression is done; release the permit before dispatching the request.
155            drop(permit);
156
157            poll_fn(|cx| inner.poll_ready(cx)).await?;
158            inner.call(HttpRequest::from_parts(parts, HttpBody::new(body))).await
159        })
160    }
161}
162
163impl<S> Service<HttpRequest> for DecompressionService<S>
164where
165    S: Service<HttpRequest, Response = HttpResponse> + Clone + Send + 'static,
166    S::Future: Send + 'static,
167{
168    type Response = HttpResponse;
169    type Error = S::Error;
170    type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
171
172    fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
173        self.decompression.poll_ready(cx)
174    }
175
176    fn call(&mut self, mut req: HttpRequest) -> Self::Future {
177        // RFC 9110 ยง8.4.1: content-coding tokens are case-insensitive, but `tower_http` matches
178        // lowercase bytes exactly, so normalize the header before decompression.
179        if let Some(encoding) = req.headers().get(http::header::CONTENT_ENCODING) &&
180            encoding.as_bytes().iter().any(u8::is_ascii_uppercase) &&
181            let Ok(normalized) =
182                http::HeaderValue::from_bytes(&encoding.as_bytes().to_ascii_lowercase())
183        {
184            req.headers_mut().insert(http::header::CONTENT_ENCODING, normalized);
185        }
186
187        if req
188            .headers()
189            .get(http::header::CONTENT_ENCODING)
190            .is_some_and(|encoding| encoding.as_bytes() != b"identity")
191        {
192            let content_length = req
193                .headers()
194                .get(http::header::CONTENT_LENGTH)
195                .and_then(|value| value.to_str().ok())
196                .and_then(|value| value.parse().ok());
197            req.extensions_mut().insert(CompressedBody { content_length });
198
199            let (parts, body) = req.into_parts();
200            req = HttpRequest::from_parts(
201                parts,
202                HttpBody::new(Limited::new(body, self.max_body_size)),
203            );
204        }
205
206        let fut = self.decompression.call(req);
207
208        Box::pin(async move { Ok(fut.await?.map(HttpBody::new)) })
209    }
210}
211
212#[inline]
213fn err_response(status: StatusCode, msg: &'static str) -> HttpResponse {
214    http::Response::builder()
215        .status(status)
216        .header(http::header::CONTENT_TYPE, "text/plain")
217        .body(HttpBody::from(msg))
218        .expect("static error response is valid")
219}
220
221#[cfg(test)]
222mod tests {
223    use super::*;
224    use http::header::CONTENT_ENCODING;
225    use http_body_util::BodyExt;
226    use jsonrpsee_http_client::{HttpRequest, HttpResponse};
227    use std::{convert::Infallible, future::ready, io::Write};
228
229    const TEST_DATA: &str = r#"{"method":"test","params":["test data"],"id":1}"#;
230    const DEFAULT_MAX_SIZE: usize = 15 * 1024 * 1024;
231
232    type Compressor = fn(&[u8]) -> Vec<u8>;
233
234    #[derive(Clone)]
235    struct MockEchoService;
236
237    impl Service<HttpRequest> for MockEchoService {
238        type Response = HttpResponse;
239        type Error = Infallible;
240        type Future = std::future::Ready<Result<Self::Response, Self::Error>>;
241
242        fn poll_ready(
243            &mut self,
244            _: &mut std::task::Context<'_>,
245        ) -> std::task::Poll<Result<(), Self::Error>> {
246            std::task::Poll::Ready(Ok(()))
247        }
248
249        fn call(&mut self, req: HttpRequest) -> Self::Future {
250            let (_parts, body) = req.into_parts();
251            ready(Ok(HttpResponse::builder().status(200).body(body).unwrap()))
252        }
253    }
254
255    fn setup_service(
256        algorithms: &[&str],
257        max_size: usize,
258    ) -> DecompressionService<MockEchoService> {
259        DecompressionLayer::new(algorithms, max_size).layer(MockEchoService)
260    }
261
262    async fn get_response_body(response: HttpResponse) -> Vec<u8> {
263        response.into_body().collect().await.unwrap().to_bytes().to_vec()
264    }
265
266    fn compress_gzip(data: &[u8]) -> Vec<u8> {
267        let mut encoder = flate2::write::GzEncoder::new(Vec::new(), flate2::Compression::default());
268        encoder.write_all(data).unwrap();
269        encoder.finish().unwrap()
270    }
271
272    fn compress_deflate(data: &[u8]) -> Vec<u8> {
273        let mut encoder =
274            flate2::write::ZlibEncoder::new(Vec::new(), flate2::Compression::default());
275        encoder.write_all(data).unwrap();
276        encoder.finish().unwrap()
277    }
278
279    fn compress_brotli(data: &[u8]) -> Vec<u8> {
280        let mut compressed = Vec::new();
281        {
282            let mut encoder = brotli::CompressorWriter::new(&mut compressed, 4096, 11, 22);
283            encoder.write_all(data).unwrap();
284            encoder.flush().unwrap();
285        }
286        compressed
287    }
288
289    fn compress_zstd(data: &[u8]) -> Vec<u8> {
290        let mut compressed = Vec::new();
291        {
292            let mut encoder = zstd::Encoder::new(&mut compressed, 3).unwrap();
293            encoder.write_all(data).unwrap();
294            encoder.finish().unwrap();
295        }
296        compressed
297    }
298
299    fn build_compressed_request(encoding: &str, body: Vec<u8>) -> HttpRequest {
300        HttpRequest::builder()
301            .header(CONTENT_ENCODING, encoding)
302            .body(HttpBody::from(body))
303            .unwrap()
304    }
305
306    #[tokio::test]
307    async fn configured_algorithms_are_decompressed_without_content_length() {
308        let cases: [(&str, Compressor); 4] = [
309            ("zstd", compress_zstd),
310            ("gzip", compress_gzip),
311            ("deflate", compress_deflate),
312            ("br", compress_brotli),
313        ];
314
315        for (algorithm, compress) in cases {
316            let mut service = setup_service(&[algorithm], DEFAULT_MAX_SIZE);
317            let request = build_compressed_request(algorithm, compress(TEST_DATA.as_bytes()));
318            let response = service.call(request).await.unwrap();
319
320            assert_eq!(response.status(), StatusCode::OK, "algorithm: {algorithm}");
321            assert_eq!(get_response_body(response).await, TEST_DATA.as_bytes());
322        }
323    }
324
325    #[tokio::test]
326    async fn oversized_decompressed_body_is_rejected() {
327        const MAX_SIZE: usize = 1024;
328        let mut service = setup_service(&["gzip"], MAX_SIZE);
329        let body = vec![0; MAX_SIZE + 1];
330        let response =
331            service.call(build_compressed_request("gzip", compress_gzip(&body))).await.unwrap();
332
333        assert_eq!(response.status(), StatusCode::PAYLOAD_TOO_LARGE);
334    }
335
336    #[tokio::test]
337    async fn oversized_compressed_body_without_content_length_is_rejected() {
338        let compressed = compress_gzip(TEST_DATA.as_bytes());
339        let max_size = compressed.len() - 1;
340        assert!(TEST_DATA.len() <= max_size);
341
342        let mut service = setup_service(&["gzip"], max_size);
343        let response = service.call(build_compressed_request("gzip", compressed)).await.unwrap();
344
345        assert_eq!(response.status(), StatusCode::PAYLOAD_TOO_LARGE);
346    }
347
348    #[tokio::test]
349    async fn oversized_compressed_content_length_is_rejected() {
350        let compressed = compress_gzip(TEST_DATA.as_bytes());
351        let max_size = compressed.len();
352        let mut request = build_compressed_request("gzip", compressed);
353        request.headers_mut().insert(http::header::CONTENT_LENGTH, (max_size + 1).into());
354
355        let mut service = setup_service(&["gzip"], max_size);
356        let response = service.call(request).await.unwrap();
357
358        assert_eq!(response.status(), StatusCode::PAYLOAD_TOO_LARGE);
359    }
360
361    #[tokio::test]
362    async fn malformed_compressed_body_is_rejected() {
363        let mut service = setup_service(&["gzip"], DEFAULT_MAX_SIZE);
364        let response = service
365            .call(build_compressed_request("gzip", b"not a gzip stream".to_vec()))
366            .await
367            .unwrap();
368
369        assert_eq!(response.status(), StatusCode::BAD_REQUEST);
370    }
371
372    #[tokio::test]
373    async fn disabled_algorithm_is_rejected() {
374        let mut service = setup_service(&["gzip"], DEFAULT_MAX_SIZE);
375        let mut request = build_compressed_request("zstd", compress_zstd(TEST_DATA.as_bytes()));
376        request.headers_mut().insert(http::header::CONTENT_LENGTH, (DEFAULT_MAX_SIZE + 1).into());
377
378        let response = service.call(request).await.unwrap();
379        assert_eq!(response.status(), StatusCode::UNSUPPORTED_MEDIA_TYPE);
380    }
381
382    #[tokio::test]
383    async fn mixed_case_content_codings_are_accepted() {
384        let cases: [(&str, Compressor); 3] =
385            [("GZip", compress_gzip), ("ZSTD", compress_zstd), ("Br", compress_brotli)];
386
387        for (encoding, compress) in cases {
388            let mut service =
389                setup_service(&[encoding.to_ascii_lowercase().as_str()], DEFAULT_MAX_SIZE);
390            let request = build_compressed_request(encoding, compress(TEST_DATA.as_bytes()));
391            let response = service.call(request).await.unwrap();
392
393            assert_eq!(response.status(), StatusCode::OK, "encoding: {encoding}");
394            assert_eq!(get_response_body(response).await, TEST_DATA.as_bytes());
395        }
396    }
397
398    #[tokio::test]
399    async fn concurrent_compressed_requests_all_complete() {
400        let service = setup_service(&["gzip"], DEFAULT_MAX_SIZE);
401        let mut tasks = tokio::task::JoinSet::new();
402
403        // Spawn more requests than there are decompression permits to ensure permits are
404        // released and waiting requests complete.
405        for _ in 0..4 * MAX_CONCURRENT_DECOMPRESSIONS {
406            let mut service = service.clone();
407            tasks.spawn(async move {
408                let request = build_compressed_request("gzip", compress_gzip(TEST_DATA.as_bytes()));
409                service.call(request).await.unwrap()
410            });
411        }
412
413        while let Some(response) = tasks.join_next().await {
414            assert_eq!(response.unwrap().status(), StatusCode::OK);
415        }
416    }
417
418    #[tokio::test]
419    async fn identity_and_unencoded_bodies_are_not_eagerly_collected() {
420        for encoding in [None, Some("identity"), Some("Identity")] {
421            let mut service = setup_service(&["gzip"], DEFAULT_MAX_SIZE);
422            let body = Limited::new(HttpBody::from(TEST_DATA), 0);
423            let mut request = HttpRequest::builder();
424            if let Some(encoding) = encoding {
425                request = request.header(CONTENT_ENCODING, encoding);
426            }
427
428            let response = service.call(request.body(HttpBody::new(body)).unwrap()).await.unwrap();
429            assert_eq!(response.status(), StatusCode::OK);
430        }
431    }
432}