use crate::utils::trad::t;
use axum::{
body::Body,
extract::State,
http::{Method, Request, StatusCode, header},
middleware::Next,
response::{IntoResponse, Response},
};
use std::{
collections::HashMap,
sync::{Arc, Mutex},
time::{Duration, Instant},
};
use tokio::time::interval;
type Store = Arc<Mutex<HashMap<String, (u32, Instant)>>>;
#[derive(Clone)]
pub struct RateLimiter {
store: Store,
pub max_requests: u32,
pub window: Duration,
pub methods: Option<Vec<Method>>,
}
impl RateLimiter {
pub fn new() -> Self {
Self {
store: Arc::new(Mutex::new(HashMap::new())),
max_requests: 60,
window: Duration::from_secs(60),
methods: None,
}
}
#[must_use]
pub fn only_methods(mut self, methods: Vec<Method>) -> Self {
self.methods = Some(methods);
self
}
#[must_use]
pub fn max_requests(mut self, max: u32) -> Self {
self.max_requests = max;
self
}
#[must_use]
pub fn retry_after(mut self, secs: u64) -> Self {
self.window = Duration::from_secs(secs);
self
}
pub fn spawn_cleanup(&self, period: tokio::time::Duration) {
let store = self.store.clone();
let window = self.window;
tokio::spawn(async move {
let mut ticker = interval(period);
loop {
ticker.tick().await;
let mut guard = match store.lock() {
Ok(g) => g,
Err(p) => p.into_inner(),
};
let now = Instant::now();
guard.retain(|_, (_, start)| now.duration_since(*start) < window);
}
});
}
#[must_use]
pub fn retry_after_secs(&self, key: &str) -> u64 {
let store = match self.store.lock() {
Ok(s) => s,
Err(p) => p.into_inner(),
};
match store.get(key) {
Some((_, start)) => {
let interval = Instant::now().duration_since(*start);
self.window.saturating_sub(interval).as_secs()
}
None => 0,
}
}
#[must_use]
pub fn is_allowed(&self, key: &str) -> bool {
let mut store = match self.store.lock() {
Ok(s) => s,
Err(p) => p.into_inner(),
};
let now = Instant::now();
let entry = store.entry(key.to_string()).or_insert((0, now));
if now.duration_since(entry.1) >= self.window {
*entry = (1, now);
true
} else if entry.0 < self.max_requests {
entry.0 = entry.0.saturating_add(1);
true
} else {
false
}
}
}
impl Default for RateLimiter {
fn default() -> Self {
Self::new()
}
}
impl std::fmt::Debug for RateLimiter {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RateLimiter")
.field("max_requests", &self.max_requests)
.field("window", &self.window)
.field("methods", &self.methods)
.finish()
}
}
fn extract_ip(req: &Request<Body>) -> String {
req.extensions()
.get::<axum::extract::ConnectInfo<std::net::SocketAddr>>()
.map(|ci| ci.0.ip().to_string())
.unwrap_or_else(|| "unknown".to_string())
}
pub async fn rate_limit_middleware(
State(limiter): State<Arc<RateLimiter>>,
req: Request<Body>,
next: Next,
) -> Response {
if let Some(ref methods) = limiter.methods
&& !methods.contains(req.method())
{
return next.run(req).await;
}
let ip = req
.extensions()
.get::<crate::middleware::security::trusted_proxies::ClientIp>()
.map(|c| c.0.to_string())
.unwrap_or_else(|| extract_ip(&req));
if limiter.is_allowed(&ip) {
next.run(req).await
} else {
let retry_after_secs = limiter.retry_after_secs(&ip);
if let Some(level) = crate::utils::runique_log::get_log()
.middleware
.as_ref()
.and_then(|m| m.rate_limit)
{
crate::runique_log!(level, %ip, retry_after = retry_after_secs, "rate limited");
}
(
StatusCode::TOO_MANY_REQUESTS,
[(header::RETRY_AFTER, retry_after_secs.to_string())],
t("html.429_text").into_owned(),
)
.into_response()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_is_allowed_basic() {
let limiter = RateLimiter::new().max_requests(3).retry_after(60);
assert!(limiter.is_allowed("ip1"));
assert!(limiter.is_allowed("ip1"));
assert!(limiter.is_allowed("ip1"));
assert!(!limiter.is_allowed("ip1"));
assert!(limiter.is_allowed("ip2"));
}
#[test]
fn test_only_methods_skips_unlisted() {
let limiter = RateLimiter::new()
.max_requests(2)
.retry_after(60)
.only_methods(vec![Method::POST]);
for _ in 0..10 {
let method = Method::GET;
let methods = limiter.methods.as_ref().unwrap();
assert!(!methods.contains(&method), "GET should not be in the list");
}
assert!(limiter.is_allowed("ip1"));
assert!(limiter.is_allowed("ip1"));
assert!(!limiter.is_allowed("ip1"));
}
#[test]
fn test_only_methods_counts_listed() {
let limiter = RateLimiter::new()
.max_requests(1)
.retry_after(60)
.only_methods(vec![Method::POST, Method::PUT]);
let methods = limiter.methods.as_ref().unwrap();
assert!(methods.contains(&Method::POST));
assert!(methods.contains(&Method::PUT));
assert!(!methods.contains(&Method::GET));
assert!(!methods.contains(&Method::DELETE));
}
#[test]
fn test_no_methods_filter_counts_all() {
let limiter = RateLimiter::new().max_requests(1).retry_after(60);
assert!(limiter.methods.is_none());
assert!(limiter.is_allowed("ip1"));
assert!(!limiter.is_allowed("ip1"));
}
#[test]
fn test_retry_after_secs_unknown_key() {
let limiter = RateLimiter::new().max_requests(5).retry_after(60);
assert_eq!(limiter.retry_after_secs("unknown"), 0);
}
#[test]
fn test_retry_after_secs_known_key() {
let limiter = RateLimiter::new().max_requests(1).retry_after(60);
let _ = limiter.is_allowed("ip1");
let _ = limiter.is_allowed("ip1"); let remaining = limiter.retry_after_secs("ip1");
assert!(remaining > 0 && remaining <= 60);
}
}