Skip to main content

cloudillo_core/rate_limit/
middleware.rs

1// SPDX-FileCopyrightText: Szilárd Hajba
2// SPDX-License-Identifier: LGPL-3.0-or-later
3
4//! Rate Limiting Middleware
5//!
6//! Tower middleware layer for applying rate limits to Axum routes.
7
8use std::sync::Arc;
9use std::task::{Context, Poll};
10
11use axum::body::Body;
12use axum::response::IntoResponse;
13use futures::future::BoxFuture;
14use hyper::Request;
15use tower::{Layer, Service};
16
17use super::extractors::extract_client_ip;
18use super::limiter::RateLimitManager;
19use crate::app::ServerMode;
20
21/// Rate limit middleware layer
22#[derive(Clone)]
23pub struct RateLimitLayer {
24	manager: Arc<RateLimitManager>,
25	category: &'static str,
26	mode: ServerMode,
27	skip_ban: bool,
28}
29
30impl RateLimitLayer {
31	/// Create a new rate limit layer
32	pub fn new(manager: Arc<RateLimitManager>, category: &'static str, mode: ServerMode) -> Self {
33		Self { manager, category, mode, skip_ban: false }
34	}
35
36	/// Create a rate limit layer that bypasses the global ban list.
37	///
38	/// The per-category rate limit (429) still applies; only the hard ban (403)
39	/// is skipped. Used by routes that must stay reachable from a banned IP,
40	/// such as the password-recovery flow.
41	pub fn new_skip_ban(
42		manager: Arc<RateLimitManager>,
43		category: &'static str,
44		mode: ServerMode,
45	) -> Self {
46		Self { manager, category, mode, skip_ban: true }
47	}
48}
49
50impl<S> Layer<S> for RateLimitLayer {
51	type Service = RateLimitService<S>;
52
53	fn layer(&self, inner: S) -> Self::Service {
54		RateLimitService {
55			inner,
56			manager: self.manager.clone(),
57			category: self.category,
58			mode: self.mode,
59			skip_ban: self.skip_ban,
60		}
61	}
62}
63
64/// Rate limit middleware service
65#[derive(Clone)]
66pub struct RateLimitService<S> {
67	inner: S,
68	manager: Arc<RateLimitManager>,
69	category: &'static str,
70	mode: ServerMode,
71	skip_ban: bool,
72}
73
74impl<S> Service<Request<Body>> for RateLimitService<S>
75where
76	S: Service<Request<Body>, Response = axum::response::Response> + Clone + Send + 'static,
77	S::Future: Send + 'static,
78{
79	type Response = S::Response;
80	type Error = S::Error;
81	type Future = BoxFuture<'static, Result<Self::Response, Self::Error>>;
82
83	fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
84		self.inner.poll_ready(cx)
85	}
86
87	fn call(&mut self, req: Request<Body>) -> Self::Future {
88		let manager = self.manager.clone();
89		let category = self.category;
90		let mode = self.mode;
91		let skip_ban = self.skip_ban;
92		let mut inner = self.inner.clone();
93
94		Box::pin(async move {
95			// Extract client IP
96			let client_ip = extract_client_ip(&req, &mode);
97
98			if let Some(ip) = client_ip {
99				// Check rate limit
100				let result = if skip_ban {
101					manager.check_skip_ban(&ip, category)
102				} else {
103					manager.check(&ip, category)
104				};
105				if let Err(error) = result {
106					// Rate limited - return error response
107					return Ok(error.into_response());
108				}
109			}
110
111			// Not rate limited - proceed with request
112			inner.call(req).await
113		})
114	}
115}
116
117// vim: ts=4