use axum::extract::connect_info::{ConnectInfo, MockConnectInfo};
use axum::http::request::Parts;
use axum::http::Extensions;
use axum::http::{HeaderMap, HeaderName};
use std::net::{IpAddr, SocketAddr};
pub static X_FORWARDED_FOR: HeaderName = HeaderName::from_static("x-forwarded-for");
pub static X_REAL_IP: HeaderName = HeaderName::from_static("x-real-ip");
#[derive(Clone, Copy, Debug)]
pub struct TrustedProxyHops(pub u16);
fn parse_forwarded_ip(raw: &str) -> Option<IpAddr> {
if let Ok(sa) = raw.parse::<SocketAddr>() {
return Some(sa.ip());
}
raw.parse::<IpAddr>().ok()
}
pub fn best_effort_client_ip(
headers: &HeaderMap,
extensions: &Extensions,
trusted_hops: Option<u16>,
) -> Option<IpAddr> {
if let Some(MockConnectInfo(addr)) = extensions.get::<MockConnectInfo<SocketAddr>>() {
return Some(addr.ip());
}
let hops = trusted_hops.unwrap_or(0);
if hops == 0 {
if let Some(ConnectInfo(addr)) = extensions.get::<ConnectInfo<SocketAddr>>() {
return Some(addr.ip());
}
return None;
}
let xff_values: Vec<&str> = headers
.get_all(&X_FORWARDED_FOR)
.iter()
.filter_map(|v| v.to_str().ok())
.collect();
if !xff_values.is_empty() {
let v = xff_values.join(",");
let entries: Vec<&str> = v
.split(',')
.map(str::trim)
.filter(|s| !s.is_empty())
.collect();
if let Some(idx) = entries.len().checked_sub(hops as usize) {
if let Some(ip) = parse_forwarded_ip(entries[idx]) {
return Some(ip);
}
}
}
headers
.get(&X_REAL_IP)
.and_then(|v| v.to_str().ok())
.and_then(parse_forwarded_ip)
}
pub fn rate_limit_key_ip_or_unknown(
headers: &HeaderMap,
extensions: &Extensions,
trusted_hops: Option<u16>,
) -> String {
best_effort_client_ip(headers, extensions, trusted_hops)
.map(|ip| ip.to_string())
.unwrap_or_else(|| "unknown".to_string())
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RateLimitKey {
pub key: String,
pub fell_back: bool,
}
impl RateLimitKey {
pub fn resolve(
headers: &HeaderMap,
extensions: &Extensions,
trusted_hops: Option<u16>,
) -> Self {
match best_effort_client_ip(headers, extensions, trusted_hops) {
Some(ip) => Self {
key: ip.to_string(),
fell_back: false,
},
None => Self {
key: "unknown".to_string(),
fell_back: true,
},
}
}
}
pub fn trusted_hops_from_parts(parts: &Parts, fallback: Option<u16>) -> Option<u16> {
parts
.extensions
.get::<TrustedProxyHops>()
.map(|h| h.0)
.or(fallback)
}
#[cfg(test)]
mod tests {
use super::{
best_effort_client_ip, rate_limit_key_ip_or_unknown, trusted_hops_from_parts,
TrustedProxyHops, X_FORWARDED_FOR, X_REAL_IP,
};
use axum::extract::connect_info::{ConnectInfo, MockConnectInfo};
use axum::http::{Extensions, HeaderMap, HeaderValue};
use std::net::{IpAddr, SocketAddr};
fn xff() -> HeaderMap {
let mut headers = HeaderMap::new();
headers.insert(
&X_FORWARDED_FOR,
HeaderValue::from_static("203.0.113.10, 198.51.100.10"),
);
headers.insert(&X_REAL_IP, HeaderValue::from_static("198.51.100.20"));
headers
}
#[test]
fn configured_hops_override_connect_info_behind_proxy() {
let headers = xff();
let mut extensions = Extensions::new();
extensions.insert(ConnectInfo(SocketAddr::from(([127, 0, 0, 1], 4321))));
extensions.insert(TrustedProxyHops(2));
assert_eq!(
best_effort_client_ip(&headers, &extensions, Some(2)),
Some(IpAddr::from([203, 0, 113, 10]))
);
}
#[test]
fn forwarded_headers_are_ignored_without_trusted_proxies() {
let mut headers = HeaderMap::new();
headers.insert(&X_FORWARDED_FOR, HeaderValue::from_static("203.0.113.10"));
headers.insert(&X_REAL_IP, HeaderValue::from_static("198.51.100.20"));
assert_eq!(
best_effort_client_ip(&headers, &Extensions::new(), None),
None
);
assert_eq!(
best_effort_client_ip(&headers, &Extensions::new(), Some(0)),
None
);
}
#[test]
fn one_trusted_proxy_uses_rightmost_xff_entry() {
let headers = xff();
assert_eq!(
best_effort_client_ip(&headers, &Extensions::new(), Some(1)),
Some(IpAddr::from([198, 51, 100, 10]))
);
}
#[test]
fn two_trusted_proxies_use_second_from_right() {
let headers = xff();
assert_eq!(
best_effort_client_ip(&headers, &Extensions::new(), Some(2)),
Some(IpAddr::from([203, 0, 113, 10]))
);
}
#[test]
fn real_ip_fallback_when_xff_shorter_than_hop_count() {
let mut headers = HeaderMap::new();
headers.insert(&X_REAL_IP, HeaderValue::from_static("198.51.100.20"));
assert_eq!(
best_effort_client_ip(&headers, &Extensions::new(), Some(1)),
Some(IpAddr::from([198, 51, 100, 20]))
);
}
#[test]
fn spoofed_left_entries_cannot_override_trusted_resolution() {
let mut headers = HeaderMap::new();
headers.insert(
&X_FORWARDED_FOR,
HeaderValue::from_static("6.6.6.6, 203.0.113.10"),
);
assert_eq!(
best_effort_client_ip(&headers, &Extensions::new(), Some(1)),
Some(IpAddr::from([203, 0, 113, 10]))
);
}
#[test]
fn honest_single_hop_traffic_resolves_the_appended_entry() {
let mut headers = HeaderMap::new();
headers.insert(&X_FORWARDED_FOR, HeaderValue::from_static("198.51.100.7"));
assert_eq!(
best_effort_client_ip(&headers, &Extensions::new(), Some(1)),
Some(IpAddr::from([198, 51, 100, 7]))
);
}
#[test]
fn duplicate_xff_headers_are_joined_before_hop_indexing() {
let mut headers = HeaderMap::new();
headers.append(&X_FORWARDED_FOR, HeaderValue::from_static("6.6.6.6"));
headers.append(&X_FORWARDED_FOR, HeaderValue::from_static("198.51.100.10"));
assert_eq!(
best_effort_client_ip(&headers, &Extensions::new(), Some(1)),
Some(IpAddr::from([198, 51, 100, 10]))
);
}
#[test]
fn mock_connect_info_is_visible_to_best_effort() {
let mut extensions = Extensions::new();
extensions.insert(MockConnectInfo(SocketAddr::from(([127, 0, 0, 1], 4321))));
assert_eq!(
best_effort_client_ip(&HeaderMap::new(), &extensions, None),
Some(IpAddr::from([127, 0, 0, 1]))
);
}
#[test]
fn rate_limit_key_ip_formats_a_resolved_ip() {
let mut extensions = Extensions::new();
extensions.insert(MockConnectInfo(SocketAddr::from(([203, 0, 113, 9], 443))));
assert_eq!(
rate_limit_key_ip_or_unknown(&HeaderMap::new(), &extensions, None),
"203.0.113.9"
);
}
#[test]
fn rate_limit_key_ip_falls_back_to_unknown_when_unresolvable() {
assert_eq!(
rate_limit_key_ip_or_unknown(&HeaderMap::new(), &Extensions::new(), None),
"unknown"
);
let mut headers = HeaderMap::new();
headers.insert(&X_FORWARDED_FOR, HeaderValue::from_static("not-an-ip"));
assert_eq!(
rate_limit_key_ip_or_unknown(&headers, &Extensions::new(), Some(1)),
"unknown"
);
}
#[test]
fn trusted_hops_from_parts_prefers_extension_then_fallback() {
let mut parts = axum::http::Request::new(()).into_parts().0;
assert_eq!(trusted_hops_from_parts(&parts, None), None);
assert_eq!(trusted_hops_from_parts(&parts, Some(0)), Some(0));
parts.extensions.insert(TrustedProxyHops(3));
assert_eq!(trusted_hops_from_parts(&parts, None), Some(3));
assert_eq!(trusted_hops_from_parts(&parts, Some(0)), Some(3));
}
}