shared_framework/middleware/general.rs
1//! General-purpose middlewares for the router.
2//!
3//! Currently provides [`rate_limiter`], a per-caller request-cap hook. Attach the
4//! returned [`Middleware`](crate::controller::Middleware) to a route or pass it as a
5//! global middleware so over-limit requests fail with a 429 before the handler runs.
6//! ```ignore
7//! let mw = rate_limiter(limiter, 100);
8//! self.mount_get(&mut router, "/list", desc, handler, vec![mw]);
9//! ```
10
11use crate::logging::CorrelationContext;
12use crate::response::ErrorResult;
13use futures::FutureExt;
14
15/// Builds middleware enforcing `limit` requests per minute per caller.
16/// The caller key is the context's user id, or the `x-forwarded-for` header value
17/// falling back to `"anonymous"`. Over-limit requests fail with a 429 `ErrorResult`.
18pub fn rate_limiter(
19 limiter: std::sync::Arc<dyn crate::middleware::rate::RateLimiter>,
20 limit: u32,
21) -> crate::controller::Middleware {
22 std::sync::Arc::new(move |ctx: &mut CorrelationContext| {
23 let headers = ctx.headers();
24 let cloned = limiter.clone();
25 async move {
26 let key = ctx.user_id().unwrap_or_else(|| {
27 headers
28 .get("x-forwarded-for")
29 .and_then(|v| v.to_str().ok())
30 .unwrap_or("anonymous")
31 .to_string()
32 }) + ":"
33 + &base64::Engine::encode(&base64::engine::general_purpose::STANDARD, "path");
34 if !cloned.is_allowed(&key, limit).await {
35 return Err(ErrorResult::new("Too many requests", None, 429));
36 }
37 Ok(())
38 }
39 .boxed()
40 })
41}