Skip to main content

lean_ctx/core/ocla/
wire_middleware.rs

1//! Tower middleware for the OCLA HTTP wire contract.
2
3use std::{
4    convert::Infallible,
5    future::Future,
6    num::NonZeroUsize,
7    pin::Pin,
8    sync::{Arc, Mutex},
9    task::{Context, Poll},
10    time::{Duration, Instant},
11};
12
13use axum::{
14    body::{Body, to_bytes},
15    http::{HeaderMap, Method, Request, StatusCode},
16    response::Response,
17};
18use lru::LruCache;
19use tower::{
20    Layer, Service, ServiceBuilder,
21    layer::util::{Identity, Stack},
22};
23
24/// Maximum request body accepted by the OCLA wire API.
25pub const MAX_REQUEST_BYTES: usize = 256 * 1024;
26
27const CACHE_CAPACITY: NonZeroUsize = NonZeroUsize::new(1024).unwrap();
28const CACHE_TTL: Duration = Duration::from_mins(5);
29
30#[derive(Clone)]
31struct CachedResponse {
32    status: StatusCode,
33    headers: HeaderMap,
34    body: Vec<u8>,
35    expires_at: Instant,
36}
37
38impl CachedResponse {
39    fn into_response(self) -> Response {
40        let mut response = Response::new(Body::from(self.body));
41        *response.status_mut() = self.status;
42        *response.headers_mut() = self.headers;
43        response
44    }
45}
46
47#[derive(Clone)]
48pub struct IdempotencyLayer {
49    cache: Arc<Mutex<LruCache<String, CachedResponse>>>,
50}
51
52impl Default for IdempotencyLayer {
53    fn default() -> Self {
54        Self::new()
55    }
56}
57impl IdempotencyLayer {
58    /// Creates an idempotency cache with a 1024-entry, five-minute policy.
59    pub fn new() -> Self {
60        Self {
61            cache: Arc::new(Mutex::new(LruCache::new(CACHE_CAPACITY))),
62        }
63    }
64
65    fn cached_response(&self, key: &str) -> Option<Response> {
66        let mut cache = self
67            .cache
68            .lock()
69            .unwrap_or_else(std::sync::PoisonError::into_inner);
70        let cached = cache.get(key).cloned();
71        match cached {
72            Some(response) if response.expires_at > Instant::now() => {
73                Some(response.into_response())
74            }
75            Some(_) => {
76                cache.pop(key);
77                None
78            }
79            None => None,
80        }
81    }
82
83    fn store_response(&self, key: String, response: &Response, body: &[u8]) {
84        let cached = CachedResponse {
85            status: response.status(),
86            headers: response.headers().clone(),
87            body: body.to_vec(),
88            expires_at: Instant::now() + CACHE_TTL,
89        };
90        self.cache
91            .lock()
92            .unwrap_or_else(std::sync::PoisonError::into_inner)
93            .put(key, cached);
94    }
95}
96
97impl<S> Layer<S> for IdempotencyLayer {
98    type Service = IdempotencyService<S>;
99
100    fn layer(&self, inner: S) -> Self::Service {
101        IdempotencyService {
102            inner,
103            layer: self.clone(),
104        }
105    }
106}
107
108#[derive(Clone)]
109pub struct IdempotencyService<S> {
110    inner: S,
111    layer: IdempotencyLayer,
112}
113
114impl<S> Service<Request<Body>> for IdempotencyService<S>
115where
116    S: Service<Request<Body>, Response = Response, Error = Infallible> + Clone + Send + 'static,
117    S::Future: Send + 'static,
118{
119    type Response = Response;
120    type Error = Infallible;
121    type Future = Pin<Box<dyn Future<Output = Result<Response, Infallible>> + Send>>;
122
123    fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
124        self.inner.poll_ready(cx)
125    }
126
127    fn call(&mut self, request: Request<Body>) -> Self::Future {
128        let key = request
129            .method()
130            .eq(&Method::POST)
131            .then(|| request.headers().get("Idempotency-Key"))
132            .flatten()
133            .and_then(|value| value.to_str().ok())
134            .filter(|value| !value.is_empty())
135            .map(str::to_owned);
136
137        if let Some(key) = key.as_deref()
138            && let Some(response) = self.layer.cached_response(key)
139        {
140            return Box::pin(async move { Ok(response) });
141        }
142
143        let layer = self.layer.clone();
144        let mut inner = self.inner.clone();
145        Box::pin(async move {
146            let response = inner.call(request).await?;
147            let Some(key) = key else {
148                return Ok(response);
149            };
150
151            let (parts, body) = response.into_parts();
152            let Ok(body) = to_bytes(body, usize::MAX).await else {
153                return Ok(Response::from_parts(parts, Body::empty()));
154            };
155            let response = Response::from_parts(parts, Body::from(body.clone()));
156            layer.store_response(key, &response, &body);
157            Ok(response)
158        })
159    }
160}
161
162#[derive(Clone, Copy, Debug, Default)]
163pub struct RequestSizeLimit;
164
165impl RequestSizeLimit {
166    /// Creates a 256 KiB request-size limit layer.
167    pub const fn new() -> Self {
168        Self
169    }
170}
171
172impl<S> Layer<S> for RequestSizeLimit {
173    type Service = RequestSizeLimitService<S>;
174
175    fn layer(&self, inner: S) -> Self::Service {
176        RequestSizeLimitService { inner }
177    }
178}
179
180#[derive(Clone)]
181pub struct RequestSizeLimitService<S> {
182    inner: S,
183}
184
185impl<S> Service<Request<Body>> for RequestSizeLimitService<S>
186where
187    S: Service<Request<Body>, Response = Response, Error = Infallible> + Clone + Send + 'static,
188    S::Future: Send + 'static,
189{
190    type Response = Response;
191    type Error = Infallible;
192    type Future = Pin<Box<dyn Future<Output = Result<Response, Infallible>> + Send>>;
193
194    fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
195        self.inner.poll_ready(cx)
196    }
197
198    fn call(&mut self, request: Request<Body>) -> Self::Future {
199        let too_large = request
200            .headers()
201            .get("Content-Length")
202            .and_then(|value| value.to_str().ok())
203            .and_then(|value| value.parse::<u64>().ok())
204            .is_some_and(|length| length > MAX_REQUEST_BYTES as u64);
205        if too_large {
206            return Box::pin(async { Ok(payload_too_large()) });
207        }
208
209        let (parts, body) = request.into_parts();
210        let mut inner = self.inner.clone();
211        Box::pin(async move {
212            let body = match to_bytes(body, MAX_REQUEST_BYTES + 1).await {
213                Ok(body) if body.len() <= MAX_REQUEST_BYTES => body,
214                _ => return Ok(payload_too_large()),
215            };
216            inner
217                .call(Request::from_parts(parts, Body::from(body)))
218                .await
219        })
220    }
221}
222
223fn payload_too_large() -> Response {
224    Response::builder()
225        .status(StatusCode::PAYLOAD_TOO_LARGE)
226        .body(Body::empty())
227        .expect("valid payload-too-large response")
228}
229
230/// Builds the OCLA middleware stack: idempotency outside request-size checks.
231pub fn ocla_middleware()
232-> ServiceBuilder<Stack<IdempotencyLayer, Stack<RequestSizeLimit, Identity>>> {
233    ServiceBuilder::new()
234        .layer(RequestSizeLimit::new())
235        .layer(IdempotencyLayer::new())
236}
237
238#[cfg(test)]
239mod tests {
240    use super::*;
241    use axum::body::to_bytes;
242    use tower::{ServiceExt, service_fn};
243
244    fn response(body: &'static str) -> Response {
245        Response::new(Body::from(body))
246    }
247
248    #[tokio::test]
249    async fn idempotency_cache_hit_and_miss() {
250        let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
251        let calls_for_service = calls.clone();
252        let service = service_fn(move |_request: Request<Body>| {
253            let calls = calls_for_service.clone();
254            async move {
255                calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
256                Ok::<_, Infallible>(response("cached"))
257            }
258        });
259        let mut service = ocla_middleware().service(service);
260        let request = |key| {
261            Request::builder()
262                .method(Method::POST)
263                .header("Idempotency-Key", key)
264                .body(Body::empty())
265                .expect("request")
266        };
267
268        for key in ["same-key", "same-key", "new-key"] {
269            let response = service
270                .ready()
271                .await
272                .expect("ready")
273                .call(request(key))
274                .await
275                .expect("response");
276            assert_eq!(
277                to_bytes(response.into_body(), 1024).await.unwrap(),
278                "cached"
279            );
280        }
281        assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 2);
282    }
283
284    #[tokio::test]
285    async fn request_size_limit_rejects_oversized_body() {
286        let service = service_fn(|_request: Request<Body>| async {
287            Ok::<_, Infallible>(response("processed"))
288        });
289        let mut service = RequestSizeLimit::new().layer(service);
290        let request = Request::builder()
291            .method(Method::POST)
292            .body(Body::from(vec![0_u8; MAX_REQUEST_BYTES + 1]))
293            .expect("request");
294
295        let response = service
296            .ready()
297            .await
298            .expect("ready")
299            .call(request)
300            .await
301            .unwrap();
302        assert_eq!(response.status(), StatusCode::PAYLOAD_TOO_LARGE);
303    }
304}