use std::{
net::{IpAddr, SocketAddr},
sync::Arc,
};
use axum::{
body::Body,
extract::ConnectInfo,
http::{Request, StatusCode},
middleware::Next,
response::{IntoResponse, Response},
};
use tracing::warn;
use super::{
config::RateLimitConfig,
dispatch::RateLimiter,
identity::VerifiedSubject,
key::{is_private_or_loopback, normalise_ip_key},
};
#[derive(Debug)]
pub struct RateLimitExceeded {
pub retry_after_secs: u32,
}
impl IntoResponse for RateLimitExceeded {
fn into_response(self) -> Response {
let retry = self.retry_after_secs;
let retry_str = retry.to_string();
let body = format!(
r#"{{"errors":[{{"message":"Rate limit exceeded. Please retry after {retry} second{s}."}}]}}"#,
s = if retry == 1 { "" } else { "s" }
);
(
StatusCode::TOO_MANY_REQUESTS,
[
("Content-Type", "application/json"),
("Retry-After", retry_str.as_str()),
],
body,
)
.into_response()
}
}
static PROXY_WARNING_LOGGED: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(false);
pub(super) fn extract_real_ip(
req: &Request<Body>,
trust_proxy: bool,
trusted_cidrs: &[ipnet::IpNet],
addr: &SocketAddr,
) -> String {
if trust_proxy {
if !trusted_cidrs.is_empty() {
let direct: IpAddr = addr.ip();
let from_trusted_proxy = trusted_cidrs.iter().any(|cidr| cidr.contains(&direct));
if !from_trusted_proxy {
tracing::debug!(
%direct,
"Connection not from a trusted proxy CIDR; ignoring X-Forwarded-For"
);
return direct.to_string();
}
}
if let Some(real_ip) = req
.headers()
.get("x-real-ip")
.and_then(|v| v.to_str().ok())
.map(str::trim)
.filter(|s| !s.is_empty())
{
if let Some(key) = normalise_ip_key(real_ip) {
return key;
}
}
if let Some(xff) = req.headers().get("x-forwarded-for").and_then(|v| v.to_str().ok()) {
if let Some(first) = xff.split(',').next().map(str::trim).filter(|s| !s.is_empty()) {
if let Some(key) = normalise_ip_key(first) {
return key;
}
}
}
} else if is_private_or_loopback(addr.ip())
&& !PROXY_WARNING_LOGGED.load(std::sync::atomic::Ordering::Relaxed)
&& !PROXY_WARNING_LOGGED.swap(true, std::sync::atomic::Ordering::Relaxed)
{
warn!(
peer_ip = %addr.ip(),
"Rate limiter: peer address is loopback/RFC-1918 — server appears to be \
behind a reverse proxy. All requests will share a single rate-limit bucket \
unless you set `trust_proxy_headers = true` in [security.rate_limiting]."
);
}
normalise_ip_key(&addr.ip().to_string()).unwrap_or_else(|| addr.ip().to_string())
}
#[allow(clippy::cognitive_complexity)] pub async fn rate_limit_middleware(
ConnectInfo(addr): ConnectInfo<SocketAddr>,
req: Request<Body>,
next: Next,
) -> Result<Response, RateLimitExceeded> {
let limiter = req
.extensions()
.get::<Arc<RateLimiter>>()
.cloned()
.unwrap_or_else(|| Arc::new(RateLimiter::new(RateLimitConfig::default())));
let ip = extract_real_ip(
&req,
limiter.config().trust_proxy_headers,
&limiter.config().trusted_proxy_cidrs,
&addr,
);
let path = req.uri().path().to_string();
let verified_subject = match req.extensions().get::<Arc<VerifiedSubject>>() {
Some(identity) => identity.subject(req.headers()).await,
None => None,
};
let path_result = limiter.check_path_limit(&path, &ip).await;
if !path_result.allowed {
warn!(ip = %ip, path = %path, "Per-path rate limit exceeded");
return Err(RateLimitExceeded {
retry_after_secs: path_result.retry_after_secs,
});
}
let (limit_result, limit_for_header) = if let Some(ref subject) = verified_subject {
let result = limiter.check_user_limit(subject).await;
if !result.allowed {
warn!(user_id = %subject, "Per-user rate limit exceeded");
return Err(RateLimitExceeded {
retry_after_secs: result.retry_after_secs,
});
}
(result, limiter.config().rps_per_user)
} else {
let result = limiter.check_ip_limit(&ip).await;
if !result.allowed {
warn!(ip = %ip, "IP rate limit exceeded");
return Err(RateLimitExceeded {
retry_after_secs: result.retry_after_secs,
});
}
(result, limiter.config().rps_per_ip)
};
let remaining = limit_result.remaining;
let response = next.run(req).await;
let mut response = response;
let limit = limit_for_header;
if let Ok(limit_value) = format!("{limit}").parse() {
response.headers_mut().insert("X-RateLimit-Limit", limit_value);
}
#[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)]
if let Ok(remaining_value) = format!("{}", remaining as u32).parse() {
response.headers_mut().insert("X-RateLimit-Remaining", remaining_value);
}
Ok(response)
}