use super::{RateLimitTier, RateLimiter};
use crate::errors::AppError;
use axum::{
extract::{ConnectInfo, Request, State},
http::{HeaderMap, HeaderValue},
middleware::Next,
response::{IntoResponse, Response},
};
use oidc_auth::AuthIdentity;
use std::net::SocketAddr;
fn extract_key(request: &Request) -> String {
if let Some(identity) = request.extensions().get::<AuthIdentity>() {
return format!("user:{}", identity.subject);
}
let headers = request.headers();
if let Some(real_ip) = headers.get("x-real-ip").and_then(|v| v.to_str().ok()) {
let ip = real_ip.trim();
if !ip.is_empty() {
return format!("ip:{ip}");
}
}
if let Some(forwarded) = headers.get("x-forwarded-for").and_then(|v| v.to_str().ok())
&& let Some(ip) = forwarded.rsplit(',').next().map(|s| s.trim())
&& !ip.is_empty()
{
return format!("ip:{ip}");
}
if let Some(connect_info) = request.extensions().get::<ConnectInfo<SocketAddr>>() {
return format!("ip:{}", connect_info.0.ip());
}
tracing::warn!(
"Could not determine client IP for rate limiting - using shared 'ip:unknown' key"
);
"ip:unknown".to_string()
}
pub async fn rate_limit_middleware(
State(limiter): State<RateLimiter>,
request: Request,
next: Next,
) -> Response {
if !limiter.is_enabled() {
return next.run(request).await;
}
let tier = match request.extensions().get::<RateLimitTier>() {
Some(t) => t.clone(),
None => return next.run(request).await,
};
let key = extract_key(&request);
let limit = tier.requests_per_window;
match limiter
.check_with_config(&key, &tier.name, tier.requests_per_window, tier.window_secs)
.await
{
Ok(result) => {
if result.allowed {
let mut response = next.run(request).await;
insert_rate_limit_headers(
response.headers_mut(),
limit,
result.remaining,
result.reset_at,
);
response
} else {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.expect("system clock before epoch")
.as_secs();
let retry_after = result.reset_at.saturating_sub(now);
let mut response =
AppError::TooManyRequests("Rate limit exceeded".to_string()).into_response();
let headers = response.headers_mut();
insert_rate_limit_headers(headers, limit, 0, result.reset_at);
if let Ok(val) = HeaderValue::from_str(&retry_after.to_string()) {
headers.insert("retry-after", val);
}
response
}
}
Err(err) => {
tracing::warn!(
error = %err,
"Rate limiter Redis error - failing open"
);
next.run(request).await
}
}
}
fn insert_rate_limit_headers(headers: &mut HeaderMap, limit: u64, remaining: u64, reset_at: u64) {
if let Ok(val) = HeaderValue::from_str(&limit.to_string()) {
headers.insert("x-ratelimit-limit", val);
}
if let Ok(val) = HeaderValue::from_str(&remaining.to_string()) {
headers.insert("x-ratelimit-remaining", val);
}
if let Ok(val) = HeaderValue::from_str(&reset_at.to_string()) {
headers.insert("x-ratelimit-reset", val);
}
}
#[cfg(test)]
mod tests {
#![allow(clippy::unwrap_used)]
use super::*;
use axum::extract::ConnectInfo;
use axum::http::Request as HttpRequest;
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
fn make_request() -> Request {
HttpRequest::builder()
.uri("/test")
.body(axum::body::Body::empty())
.unwrap()
}
#[test]
fn test_extract_key_auth_identity() {
let mut req = make_request();
req.extensions_mut().insert(AuthIdentity {
subject: "user-123".to_string(),
org_id: None,
roles: vec![],
email: Some("test@example.com".to_string()),
name: Some("Test".to_string()),
session_id: None,
});
assert_eq!(extract_key(&req), "user:user-123");
}
#[test]
fn test_extract_key_x_real_ip() {
let mut req = make_request();
req.headers_mut()
.insert("x-real-ip", HeaderValue::from_static("10.0.0.1"));
assert_eq!(extract_key(&req), "ip:10.0.0.1");
}
#[test]
fn test_extract_key_x_forwarded_for_rightmost() {
let mut req = make_request();
req.headers_mut().insert(
"x-forwarded-for",
HeaderValue::from_static("spoofed, 192.168.1.1"),
);
assert_eq!(extract_key(&req), "ip:192.168.1.1");
}
#[test]
fn test_extract_key_connect_info() {
let mut req = make_request();
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(172, 16, 0, 5)), 12345);
req.extensions_mut().insert(ConnectInfo(addr));
assert_eq!(extract_key(&req), "ip:172.16.0.5");
}
#[test]
fn test_extract_key_fallback() {
let req = make_request();
assert_eq!(extract_key(&req), "ip:unknown");
}
}