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
18const MAX_CONCURRENT_DECOMPRESSIONS: usize = 8;
25
26const DECOMPRESSION_BODY_READ_TIMEOUT: Duration = Duration::from_secs(30);
28
29#[expect(missing_debug_implementations)]
32#[derive(Clone)]
33pub struct DecompressionLayer {
34 inner_layer: TowerDecompressionLayer,
35 max_body_size: usize,
37 decompression_permits: Arc<Semaphore>,
39}
40
41impl DecompressionLayer {
42 pub fn new(algos: &[impl AsRef<str>], max_body_size: usize) -> Self {
45 let mut layer = TowerDecompressionLayer::new().no_zstd().no_gzip().no_deflate().no_br();
47
48 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#[expect(missing_debug_implementations)]
86#[derive(Clone)]
87pub struct DecompressionService<S> {
88 decompression: RequestDecompression<InnerService<S>>,
89 max_body_size: usize,
90}
91
92#[derive(Clone)]
95struct InnerService<S> {
96 inner: S,
97 max_body_size: usize,
98 decompression_permits: Arc<Semaphore>,
99}
100
101#[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 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 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 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 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}