Skip to main content

servers/grpc/
memory_limit.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15//! Aggregate memory admission for gRPC services.
16//!
17//! Tonic decompresses and decodes a complete request message (up to
18//! `max_recv_message_size`, 512 MiB by default) *before* the typed handler is
19//! invoked, so a handler that charges the [`ServerMemoryLimiter`] only sees
20//! the request after the memory is already allocated. For transport-compressed
21//! messages (`grpc-encoding: gzip/zstd`) a few KiB on the wire can expand to
22//! hundreds of MiB during that window, which would bypass any aggregate quota.
23//!
24//! [`MemoryLimiterExtensionLayer`] therefore reserves the configured maximum
25//! decoded message size for every compressed request *before* the inner
26//! (tonic) service runs, and holds the reservation for the whole request.
27//! Handlers can detect the reservation via [`PreDecodeMemoryReservation`] in
28//! the request extensions and skip their own post-decode charge to avoid
29//! double accounting.
30//!
31//! The reservation is intentionally conservative (it upper-bounds the decoded
32//! size before it is known); with the default unlimited limiter it is a no-op.
33
34use 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
50/// The gRPC request compression selector header.
51const GRPC_ENCODING_HEADER: &str = "grpc-encoding";
52/// The "no compression" encoding value.
53const IDENTITY_ENCODING: &str = "identity";
54
55/// Present in the request extensions when memory for the (compressed)
56/// request's decoded message was reserved before tonic decompression.
57///
58/// Handlers that charge the [`ServerMemoryLimiter`] after decoding should skip
59/// that charge when this marker is present: the reservation already covers the
60/// peak decoding memory and stays alive for the whole request.
61#[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        // Own a clone of the service so the returned future does not borrow
121        // `self` (tower's load-shedding pattern).
122        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                // A compressed message can expand up to the decoding limit
134                // inside tonic, before any handler runs. Reserve that worst
135                // case against the aggregate quota so the decoding phase is
136                // admitted as well. No-op for an unlimited limiter.
137                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        // The inner service observes the reservation; use a oneshot call and
234        // check the limiter was charged while the call is in flight via the
235        // response extensions observed by the echo service.
236        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        // The reservation is released once the inner call (request) finishes.
249        assert_eq!(0, limiter.used_bytes());
250    }
251
252    #[tokio::test]
253    async fn test_rejects_compressed_request_when_quota_exhausted() {
254        // Quota smaller than the decoding limit: the reservation cannot fit.
255        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        // gRPC errors are HTTP 200 with a grpc-status header.
266        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}