use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use axum::extract::{ConnectInfo, MatchedPath, Request, State};
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use ipnet::IpNet;
use crate::channel::{KeyedLimiter, build_keyed_limiter};
use crate::config::RateLimitConfig;
use crate::metrics;
pub struct RateLimitState {
default_limiter: Arc<KeyedLimiter>,
admin_limiter: Option<Arc<KeyedLimiter>>,
data_limiter: Option<Arc<KeyedLimiter>>,
}
impl RateLimitState {
pub fn from_config(config: &RateLimitConfig) -> Self {
let default_limiter = Arc::new(build_keyed_limiter(
config.default_rps,
config.default_burst,
));
let admin_limiter = config
.endpoints
.admin_rps
.map(|rps| Arc::new(build_keyed_limiter(rps, rps / 2 + 1)));
let data_limiter = config
.endpoints
.data_rps
.map(|rps| Arc::new(build_keyed_limiter(rps, rps / 2 + 1)));
Self {
default_limiter,
admin_limiter,
data_limiter,
}
}
}
pub(crate) fn extract_client_ip(req: &Request, trusted_proxies: &[IpNet]) -> String {
let peer = req
.extensions()
.get::<ConnectInfo<SocketAddr>>()
.map(|ci| ci.0);
client_ip_from_parts(peer.as_ref(), req.headers(), trusted_proxies)
}
pub(crate) fn client_ip_from_parts(
peer: Option<&SocketAddr>,
headers: &axum::http::HeaderMap,
trusted_proxies: &[IpNet],
) -> String {
match peer.map(|p| p.ip().to_canonical()) {
Some(ip) if peer_is_trusted(&ip, trusted_proxies) => {
forwarded_client_ip(headers, trusted_proxies).unwrap_or_else(|| ip.to_string())
}
Some(ip) => ip.to_string(),
None => {
forwarded_client_ip(headers, trusted_proxies).unwrap_or_else(|| "unknown".to_string())
}
}
}
fn peer_is_trusted(peer: &IpAddr, trusted_proxies: &[IpNet]) -> bool {
trusted_proxies.iter().any(|net| net.contains(peer))
}
fn forwarded_client_ip(
headers: &axum::http::HeaderMap,
trusted_proxies: &[IpNet],
) -> Option<String> {
let xff_client = headers
.get("x-forwarded-for")
.and_then(|v| v.to_str().ok())
.and_then(|xff| {
let mut candidate = None;
for hop in xff
.split(',')
.map(str::trim)
.filter(|h| !h.is_empty())
.rev()
{
candidate = Some(hop);
let trusted = hop
.parse::<IpAddr>()
.is_ok_and(|ip| peer_is_trusted(&ip.to_canonical(), trusted_proxies));
if !trusted {
break;
}
}
candidate
});
if let Some(hop) = xff_client {
return Some(hop.to_string());
}
headers
.get("x-real-ip")
.and_then(|v| v.to_str().ok())
.map(str::trim)
.filter(|v| !v.is_empty())
.map(str::to_string)
}
enum RouteGroup {
Admin,
Data,
Operational,
}
impl RouteGroup {
fn label(&self) -> &'static str {
match self {
Self::Admin => "admin",
Self::Data => "data",
Self::Operational => "operational",
}
}
}
fn classify_route(path: &str) -> RouteGroup {
if path.starts_with("/api/v1/admin") {
RouteGroup::Admin
} else if path.starts_with("/api/v1/data") {
RouteGroup::Data
} else {
RouteGroup::Operational
}
}
pub async fn rate_limit_middleware(
State(state): State<crate::server::state::AppState>,
matched_path: Option<MatchedPath>,
req: Request,
next: Next,
) -> Response {
let rate_limit_state = match &state.rate_limit_state {
Some(rls) => rls,
None => return next.run(req).await,
};
let client_ip = extract_client_ip(&req, state.trusted_proxies());
let path = matched_path
.as_ref()
.map(|m: &MatchedPath| m.as_str())
.unwrap_or(req.uri().path());
let route_group = classify_route(path);
let limiter = match route_group {
RouteGroup::Admin => rate_limit_state
.admin_limiter
.as_ref()
.unwrap_or(&rate_limit_state.default_limiter),
RouteGroup::Data => rate_limit_state
.data_limiter
.as_ref()
.unwrap_or(&rate_limit_state.default_limiter),
RouteGroup::Operational => &rate_limit_state.default_limiter,
};
if limiter.check_key(&client_ip).is_err() {
metrics::record_rate_limit_rejected(route_group.label());
return rate_limited_response();
}
next.run(req).await
}
fn rate_limited_response() -> Response {
crate::errors::OrionError::RateLimited("Too many requests".to_string()).into_response()
}
#[cfg(test)]
mod tests {
use super::*;
use axum::body::Body;
use axum::http::{Request, StatusCode};
fn with_peer(mut req: Request<Body>, peer: &str) -> Request<Body> {
let addr: SocketAddr = peer.parse().expect("test");
req.extensions_mut().insert(ConnectInfo(addr));
req
}
fn nets(entries: &[&str]) -> Vec<IpNet> {
entries
.iter()
.map(|e| e.parse::<IpNet>().expect("test"))
.collect()
}
#[test]
fn test_extract_client_ip_from_xff() {
let req = Request::builder()
.header("x-forwarded-for", "192.168.1.1, 10.0.0.1")
.body(Body::empty())
.expect("test");
assert_eq!(extract_client_ip(&req, &[]), "10.0.0.1");
}
#[test]
fn test_extract_client_ip_from_xff_single() {
let req = Request::builder()
.header("x-forwarded-for", "203.0.113.5")
.body(Body::empty())
.expect("test");
assert_eq!(extract_client_ip(&req, &[]), "203.0.113.5");
}
#[test]
fn test_extract_client_ip_from_x_real_ip() {
let req = Request::builder()
.header("x-real-ip", "10.0.0.5")
.body(Body::empty())
.expect("test");
assert_eq!(extract_client_ip(&req, &[]), "10.0.0.5");
}
#[test]
fn test_extract_client_ip_xff_takes_priority_over_x_real_ip() {
let req = Request::builder()
.header("x-forwarded-for", "1.2.3.4")
.header("x-real-ip", "5.6.7.8")
.body(Body::empty())
.expect("test");
assert_eq!(extract_client_ip(&req, &[]), "1.2.3.4");
}
#[test]
fn test_extract_client_ip_no_headers() {
let req = Request::builder().body(Body::empty()).expect("test");
assert_eq!(extract_client_ip(&req, &[]), "unknown");
}
#[test]
fn test_extract_client_ip_empty_xff() {
let req = Request::builder()
.header("x-forwarded-for", "")
.body(Body::empty())
.expect("test");
assert_eq!(extract_client_ip(&req, &[]), "unknown");
}
#[test]
fn test_extract_client_ip_empty_x_real_ip() {
let req = Request::builder()
.header("x-real-ip", " ")
.body(Body::empty())
.expect("test");
assert_eq!(extract_client_ip(&req, &[]), "unknown");
}
#[test]
fn test_untrusted_peer_ignores_forwarded_headers() {
let req = Request::builder()
.header("x-forwarded-for", "1.2.3.4")
.header("x-real-ip", "5.6.7.8")
.body(Body::empty())
.expect("test");
let req = with_peer(req, "203.0.113.9:5000");
assert_eq!(extract_client_ip(&req, &[]), "203.0.113.9");
}
#[test]
fn test_trusted_peer_uses_forwarded_header() {
let req = Request::builder()
.header("x-forwarded-for", "1.2.3.4, 10.0.0.1")
.body(Body::empty())
.expect("test");
let req = with_peer(req, "10.0.0.7:5000");
assert_eq!(extract_client_ip(&req, &nets(&["10.0.0.0/8"])), "1.2.3.4");
}
#[test]
fn test_trusted_peer_ignores_client_supplied_xff_prefix() {
let req = Request::builder()
.header("x-forwarded-for", "6.6.6.6, 198.51.100.7")
.body(Body::empty())
.expect("test");
let req = with_peer(req, "10.0.0.7:5000");
assert_eq!(
extract_client_ip(&req, &nets(&["10.0.0.0/8"])),
"198.51.100.7"
);
}
#[test]
fn test_all_trusted_xff_resolves_to_leftmost() {
let req = Request::builder()
.header("x-forwarded-for", "10.0.0.2, 10.0.0.3")
.body(Body::empty())
.expect("test");
let req = with_peer(req, "10.0.0.7:5000");
assert_eq!(extract_client_ip(&req, &nets(&["10.0.0.0/8"])), "10.0.0.2");
}
#[test]
fn test_trusted_peer_without_headers_falls_back_to_peer() {
let req = Request::builder().body(Body::empty()).expect("test");
let req = with_peer(req, "10.0.0.7:5000");
assert_eq!(extract_client_ip(&req, &nets(&["10.0.0.0/8"])), "10.0.0.7");
}
#[test]
fn test_peer_outside_trusted_cidr_ignores_headers() {
let req = Request::builder()
.header("x-forwarded-for", "1.2.3.4")
.body(Body::empty())
.expect("test");
let req = with_peer(req, "192.0.2.33:5000");
assert_eq!(
extract_client_ip(&req, &nets(&["10.0.0.0/8"])),
"192.0.2.33"
);
}
#[test]
fn test_v4_mapped_v6_peer_matches_v4_cidr() {
let req = Request::builder()
.header("x-forwarded-for", "1.2.3.4")
.body(Body::empty())
.expect("test");
let req = with_peer(req, "[::ffff:10.0.0.7]:5000");
assert_eq!(extract_client_ip(&req, &nets(&["10.0.0.0/8"])), "1.2.3.4");
}
#[test]
fn test_ipv6_trusted_proxy() {
let req = Request::builder()
.header("x-real-ip", "2001:db8::1")
.body(Body::empty())
.expect("test");
let req = with_peer(req, "[fd00::1]:5000");
assert_eq!(extract_client_ip(&req, &nets(&["fd00::/8"])), "2001:db8::1");
}
#[test]
fn test_classify_route_admin() {
assert!(matches!(
classify_route("/api/v1/admin/workflows"),
RouteGroup::Admin
));
}
#[test]
fn test_classify_route_data() {
assert!(matches!(
classify_route("/api/v1/data/orders"),
RouteGroup::Data
));
}
#[test]
fn test_classify_route_operational() {
assert!(matches!(classify_route("/health"), RouteGroup::Operational));
assert!(matches!(
classify_route("/metrics"),
RouteGroup::Operational
));
}
#[test]
fn test_from_config_default() {
let config = RateLimitConfig {
enabled: true,
default_rps: 100,
default_burst: 50,
..Default::default()
};
let state = RateLimitState::from_config(&config);
assert!(
state.admin_limiter.is_some(),
"admin_rps must default to a real limit"
);
assert!(state.data_limiter.is_none());
}
#[test]
fn test_admin_rps_can_be_explicitly_unset() {
let config = RateLimitConfig {
enabled: true,
endpoints: crate::config::EndpointRateLimits {
admin_rps: None,
data_rps: None,
},
..Default::default()
};
let state = RateLimitState::from_config(&config);
assert!(state.admin_limiter.is_none());
}
#[test]
fn test_from_config_with_endpoint_limiters() {
let config = RateLimitConfig {
enabled: true,
default_rps: 100,
default_burst: 50,
endpoints: crate::config::EndpointRateLimits {
admin_rps: Some(20),
data_rps: Some(200),
},
..Default::default()
};
let state = RateLimitState::from_config(&config);
assert!(state.admin_limiter.is_some());
assert!(state.data_limiter.is_some());
}
#[test]
fn trusted_proxies_parse_with_the_platform_limiter_disabled() {
let config = RateLimitConfig {
enabled: false,
trusted_proxies: vec!["10.0.0.0/8".to_string(), "192.168.1.1".to_string()],
..Default::default()
};
assert_eq!(config.parsed_trusted_proxies().len(), 2);
}
#[test]
fn test_rate_limited_response_status() {
let response = rate_limited_response();
assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS);
assert_eq!(
response
.headers()
.get("retry-after")
.expect("test")
.to_str()
.expect("test"),
"1"
);
}
}