servers/grpc/
memory_limit.rs1use std::convert::Infallible;
35use std::sync::Arc;
36use std::task::{Context, Poll};
37
38use axum::response::IntoResponse;
39use common_memory_manager::MemoryGuard;
40use futures::future::BoxFuture;
41use http::Request;
42use tonic::Status;
43use tonic::codegen::Service;
44use tonic::server::NamedService;
45use tower::Layer;
46
47use crate::request_memory_limiter::ServerMemoryLimiter;
48use crate::request_memory_metrics::RequestMemoryMetrics;
49
50const GRPC_ENCODING_HEADER: &str = "grpc-encoding";
52const IDENTITY_ENCODING: &str = "identity";
54
55#[derive(Clone)]
62pub(crate) struct PreDecodeMemoryReservation {
63 _guard: Arc<MemoryGuard<RequestMemoryMetrics>>,
64}
65
66#[derive(Clone)]
67pub struct MemoryLimiterExtensionLayer {
68 limiter: ServerMemoryLimiter,
69 max_decoding_message_size: usize,
70}
71
72impl MemoryLimiterExtensionLayer {
73 pub fn new(limiter: ServerMemoryLimiter, max_decoding_message_size: usize) -> Self {
74 Self {
75 limiter,
76 max_decoding_message_size,
77 }
78 }
79}
80
81impl<S> Layer<S> for MemoryLimiterExtensionLayer {
82 type Service = MemoryLimiterExtensionService<S>;
83
84 fn layer(&self, service: S) -> Self::Service {
85 MemoryLimiterExtensionService {
86 inner: service,
87 limiter: self.limiter.clone(),
88 max_decoding_message_size: self.max_decoding_message_size,
89 }
90 }
91}
92
93#[derive(Clone)]
94pub struct MemoryLimiterExtensionService<S> {
95 inner: S,
96 limiter: ServerMemoryLimiter,
97 max_decoding_message_size: usize,
98}
99
100impl<S: NamedService> NamedService for MemoryLimiterExtensionService<S> {
101 const NAME: &'static str = S::NAME;
102}
103
104impl<S, ReqBody> Service<Request<ReqBody>> for MemoryLimiterExtensionService<S>
105where
106 S: Service<Request<ReqBody>, Error = Infallible> + Clone + Send + 'static,
107 S::Response: axum::response::IntoResponse,
108 S::Future: Send + 'static,
109 ReqBody: Send + 'static,
110{
111 type Response = axum::response::Response;
112 type Error = Infallible;
113 type Future = BoxFuture<'static, Result<Self::Response, Self::Error>>;
114
115 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
116 self.inner.poll_ready(cx)
117 }
118
119 fn call(&mut self, mut req: Request<ReqBody>) -> Self::Future {
120 let mut this = self.clone();
123 Box::pin(async move {
124 req.extensions_mut().insert(this.limiter.clone());
125
126 let compressed = req
127 .headers()
128 .get(GRPC_ENCODING_HEADER)
129 .and_then(|value| value.to_str().ok())
130 .is_some_and(|value| !value.eq_ignore_ascii_case(IDENTITY_ENCODING));
131
132 if compressed {
133 let reservation = this.max_decoding_message_size as u64;
138 match this.limiter.acquire(reservation).await {
139 Ok(guard) => {
140 req.extensions_mut().insert(PreDecodeMemoryReservation {
141 _guard: Arc::new(guard),
142 });
143 }
144 Err(e) => {
145 return Ok(Status::resource_exhausted(format!(
146 "request memory limit exceeded: {e}"
147 ))
148 .into_http::<axum::body::Body>()
149 .into_response());
150 }
151 }
152 }
153
154 match this.inner.call(req).await {
155 Ok(response) => Ok(response.into_response()),
156 Err(e) => match e {},
157 }
158 })
159 }
160}
161
162#[cfg(test)]
163mod tests {
164 use std::convert::Infallible;
165
166 use axum::body::Body;
167 use axum::response::IntoResponse;
168 use common_memory_manager::OnExhaustedPolicy;
169 use futures_util::future::BoxFuture;
170 use http::{HeaderValue, StatusCode};
171 use tower::ServiceExt;
172
173 use super::*;
174 use crate::grpc::memory_limit::MemoryLimiterExtensionLayer;
175
176 #[derive(Clone)]
177 struct EchoService;
178
179 impl Service<Request<Body>> for EchoService {
180 type Response = axum::response::Response;
181 type Error = Infallible;
182 type Future = BoxFuture<'static, Result<Self::Response, Self::Error>>;
183
184 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
185 Poll::Ready(Ok(()))
186 }
187
188 fn call(&mut self, req: Request<Body>) -> Self::Future {
189 let saw_reservation = req
190 .extensions()
191 .get::<PreDecodeMemoryReservation>()
192 .is_some();
193 let saw_limiter = req.extensions().get::<ServerMemoryLimiter>().is_some();
194 Box::pin(async move {
195 Ok((
196 StatusCode::OK,
197 format!("reservation={saw_reservation} limiter={saw_limiter}"),
198 )
199 .into_response())
200 })
201 }
202 }
203
204 impl NamedService for EchoService {
205 const NAME: &'static str = "test.Echo";
206 }
207
208 #[tokio::test]
209 async fn test_inserts_limiter_for_uncompressed_requests() {
210 let limiter = ServerMemoryLimiter::new(1024 * 1024, OnExhaustedPolicy::Fail);
211 let mut svc =
212 MemoryLimiterExtensionLayer::new(limiter.clone(), 512 * 1024 * 1024).layer(EchoService);
213
214 let req = Request::builder().body(Body::empty()).unwrap();
215 let res = svc.ready().await.unwrap().call(req).await.unwrap();
216 assert_eq!(res.status(), StatusCode::OK);
217 let body = axum::body::to_bytes(res.into_body(), 1024).await.unwrap();
218 assert_eq!(
219 &body[..],
220 &b"reservation=false limiter=true"[..],
221 "uncompressed requests must not be pre-reserved"
222 );
223 assert_eq!(0, limiter.used_bytes());
224 }
225
226 #[tokio::test]
227 async fn test_reserves_max_message_size_for_compressed_requests() {
228 let max_size = 512 * 1024 * 1024;
229 let limiter = ServerMemoryLimiter::new(1024 * 1024 * 1024, OnExhaustedPolicy::Fail);
230 let mut svc =
231 MemoryLimiterExtensionLayer::new(limiter.clone(), max_size).layer(EchoService);
232
233 let req = Request::builder()
237 .header(GRPC_ENCODING_HEADER, HeaderValue::from_static("gzip"))
238 .body(Body::empty())
239 .unwrap();
240 let res = svc.ready().await.unwrap().call(req).await.unwrap();
241 assert_eq!(res.status(), StatusCode::OK);
242 let body = axum::body::to_bytes(res.into_body(), 1024).await.unwrap();
243 assert_eq!(
244 &body[..],
245 &b"reservation=true limiter=true"[..],
246 "compressed requests must carry a pre-decode reservation"
247 );
248 assert_eq!(0, limiter.used_bytes());
250 }
251
252 #[tokio::test]
253 async fn test_rejects_compressed_request_when_quota_exhausted() {
254 let limiter = ServerMemoryLimiter::new(1024, OnExhaustedPolicy::Fail);
256 let mut svc =
257 MemoryLimiterExtensionLayer::new(limiter.clone(), 512 * 1024 * 1024).layer(EchoService);
258
259 let req = Request::builder()
260 .header(GRPC_ENCODING_HEADER, HeaderValue::from_static("zstd"))
261 .body(Body::empty())
262 .unwrap();
263 let res = svc.ready().await.unwrap().call(req).await.unwrap();
264 assert_eq!(res.status(), StatusCode::OK);
265 assert_eq!(
267 res.headers()
268 .get("grpc-status")
269 .and_then(|v| v.to_str().ok()),
270 Some("8"),
271 "expected RESOURCE_EXHAUSTED grpc status"
272 );
273 }
274
275 #[tokio::test]
276 async fn test_identity_encoding_is_not_reserved() {
277 let limiter = ServerMemoryLimiter::new(1024 * 1024, OnExhaustedPolicy::Fail);
278 let mut svc =
279 MemoryLimiterExtensionLayer::new(limiter.clone(), 512 * 1024 * 1024).layer(EchoService);
280
281 let req = Request::builder()
282 .header(
283 GRPC_ENCODING_HEADER,
284 HeaderValue::from_static(IDENTITY_ENCODING),
285 )
286 .body(Body::empty())
287 .unwrap();
288 let res = svc.ready().await.unwrap().call(req).await.unwrap();
289 let body = axum::body::to_bytes(res.into_body(), 1024).await.unwrap();
290 assert_eq!(&body[..], &b"reservation=false limiter=true"[..]);
291 }
292}