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