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
17const MAX_CONCURRENT_DECOMPRESSIONS: usize = 8;
24
25#[expect(missing_debug_implementations)]
28#[derive(Clone)]
29pub struct DecompressionLayer {
30 inner_layer: TowerDecompressionLayer,
31 max_body_size: usize,
33 decompression_permits: Arc<Semaphore>,
35}
36
37impl DecompressionLayer {
38 pub fn new(algos: &[impl AsRef<str>], max_body_size: usize) -> Self {
41 let mut layer = TowerDecompressionLayer::new().no_zstd().no_gzip().no_deflate().no_br();
43
44 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#[expect(missing_debug_implementations)]
82#[derive(Clone)]
83pub struct DecompressionService<S> {
84 decompression: RequestDecompression<InnerService<S>>,
85 max_body_size: usize,
86}
87
88#[derive(Clone)]
91struct InnerService<S> {
92 inner: S,
93 max_body_size: usize,
94 decompression_permits: Arc<Semaphore>,
95}
96
97#[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 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 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 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 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}