1use 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
24pub 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 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 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
230pub 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}