Skip to main content

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}