use std::net::{IpAddr, SocketAddr};
use axum::extract::ConnectInfo;
use axum::http::{HeaderMap, HeaderName, Request};
use ipnetwork::IpNetwork;
use tower_governor::errors::GovernorError;
use tower_governor::key_extractor::KeyExtractor;
const X_FORWARDED_FOR: HeaderName = HeaderName::from_static("x-forwarded-for");
const MAX_FORWARDED_ENTRIES: usize = 64;
#[derive(Debug, Clone)]
pub struct TrustedProxyKeyExtractor {
trusted_cidrs: Vec<IpNetwork>,
}
impl TrustedProxyKeyExtractor {
pub fn new(trusted_cidrs: Vec<IpNetwork>) -> Self {
Self { trusted_cidrs }
}
fn is_trusted(&self, ip: IpAddr) -> bool {
self.trusted_cidrs.iter().any(|cidr| cidr.contains(ip))
}
fn client_from_chain(&self, headers: &HeaderMap) -> Option<IpAddr> {
let mut entries: Vec<&str> = Vec::new();
for value in headers.get_all(X_FORWARDED_FOR) {
let value = value.to_str().ok()?;
for entry in value.split(',').map(str::trim).filter(|e| !e.is_empty()) {
if entries.len() == MAX_FORWARDED_ENTRIES {
return None;
}
entries.push(entry);
}
}
for entry in entries.iter().rev() {
let ip = parse_forwarded_ip(entry)?;
if !self.is_trusted(ip) {
return Some(ip);
}
}
None
}
}
impl KeyExtractor for TrustedProxyKeyExtractor {
type Key = IpAddr;
fn extract<T>(&self, req: &Request<T>) -> Result<IpAddr, GovernorError> {
let peer = peer_ip(req).ok_or(GovernorError::UnableToExtractKey)?;
if !self.is_trusted(peer) {
return Ok(peer);
}
Ok(self.client_from_chain(req.headers()).unwrap_or(peer))
}
}
fn peer_ip<T>(req: &Request<T>) -> Option<IpAddr> {
req.extensions()
.get::<ConnectInfo<SocketAddr>>()
.map(|ci| canonical(ci.0.ip()))
}
fn parse_forwarded_ip(entry: &str) -> Option<IpAddr> {
if let Ok(ip) = entry.parse::<IpAddr>() {
return Some(canonical(ip));
}
if let Ok(addr) = entry.parse::<SocketAddr>() {
return Some(canonical(addr.ip()));
}
let bracketed = entry.strip_prefix('[')?.strip_suffix(']')?;
bracketed.parse::<IpAddr>().ok().map(canonical)
}
fn canonical(ip: IpAddr) -> IpAddr {
ip.to_canonical()
}
pub async fn insert_default_connect_info_if_missing(
mut request: axum::extract::Request,
next: axum::middleware::Next,
) -> axum::response::Response {
use std::net::Ipv4Addr;
if request
.extensions()
.get::<ConnectInfo<SocketAddr>>()
.is_none()
{
let synthetic = ConnectInfo(SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)), 0));
request.extensions_mut().insert(synthetic);
}
next.run(request).await
}
#[cfg(test)]
mod tests {
use super::*;
fn req_with(peer: Option<&str>, xff: Option<&str>) -> Request<()> {
req_with_all(peer, xff.into_iter().collect())
}
fn req_with_all(peer: Option<&str>, xffs: Vec<&str>) -> Request<()> {
let mut builder = Request::builder().uri("/");
for xff in xffs {
builder = builder.header("x-forwarded-for", xff);
}
let mut req = builder.body(()).unwrap();
if let Some(peer) = peer {
req.extensions_mut()
.insert(ConnectInfo(peer.parse::<SocketAddr>().unwrap()));
}
req
}
fn cidrs(nets: &[&str]) -> Vec<IpNetwork> {
nets.iter().map(|n| n.parse().unwrap()).collect()
}
fn ip(s: &str) -> IpAddr {
s.parse().unwrap()
}
fn loopback() -> TrustedProxyKeyExtractor {
TrustedProxyKeyExtractor::new(cidrs(&["127.0.0.1/32"]))
}
#[test]
fn untrusted_peer_keys_on_peer_ignoring_forged_xff() {
let req = req_with(Some("203.0.113.9:1234"), Some("9.9.9.9"));
assert_eq!(loopback().extract(&req).unwrap(), ip("203.0.113.9"));
}
#[test]
fn trusted_peer_keys_on_the_client_the_proxy_appended() {
let req = req_with(Some("127.0.0.1:1234"), Some("9.9.9.9, 203.0.113.9"));
assert_eq!(loopback().extract(&req).unwrap(), ip("203.0.113.9"));
}
#[test]
fn distinct_clients_through_trusted_peer_key_differently() {
let a = req_with(Some("127.0.0.1:1"), Some("198.51.100.1"));
let b = req_with(Some("127.0.0.1:2"), Some("198.51.100.2"));
let extractor = loopback();
assert_ne!(
extractor.extract(&a).unwrap(),
extractor.extract(&b).unwrap()
);
}
#[test]
fn empty_cidr_list_never_trusts_xff() {
let extractor = TrustedProxyKeyExtractor::new(Vec::new());
let req = req_with(Some("127.0.0.1:1234"), Some("9.9.9.9"));
assert_eq!(extractor.extract(&req).unwrap(), ip("127.0.0.1"));
}
#[test]
fn trusted_peer_without_xff_falls_back_to_peer() {
let req = req_with(Some("127.0.0.1:1234"), None);
assert_eq!(loopback().extract(&req).unwrap(), ip("127.0.0.1"));
}
#[test]
fn missing_connect_info_is_unable_to_extract() {
let req = req_with(None, Some("9.9.9.9"));
assert!(matches!(
loopback().extract(&req),
Err(GovernorError::UnableToExtractKey)
));
}
#[test]
fn chained_trusted_proxies_are_skipped_to_reach_the_client() {
let extractor = TrustedProxyKeyExtractor::new(cidrs(&["127.0.0.1/32", "10.0.0.0/8"]));
let req = req_with(
Some("127.0.0.1:1234"),
Some("198.51.100.7, 10.0.0.5, 10.0.0.6"),
);
assert_eq!(extractor.extract(&req).unwrap(), ip("198.51.100.7"));
}
#[test]
fn forged_trusted_looking_prefix_cannot_reach_past_the_real_hop() {
let extractor = TrustedProxyKeyExtractor::new(cidrs(&["127.0.0.1/32", "10.0.0.0/8"]));
let req = req_with(
Some("127.0.0.1:1234"),
Some("10.0.0.1, 10.0.0.2, 203.0.113.9"),
);
assert_eq!(extractor.extract(&req).unwrap(), ip("203.0.113.9"));
}
#[test]
fn multiple_header_lines_concatenate_in_order() {
let extractor = TrustedProxyKeyExtractor::new(cidrs(&["127.0.0.1/32", "10.0.0.0/8"]));
let req = req_with_all(Some("127.0.0.1:1"), vec!["198.51.100.7", "10.0.0.5"]);
assert_eq!(extractor.extract(&req).unwrap(), ip("198.51.100.7"));
}
#[test]
fn malformed_entry_falls_back_to_peer() {
let req = req_with(Some("127.0.0.1:1234"), Some("not-an-ip"));
assert_eq!(loopback().extract(&req).unwrap(), ip("127.0.0.1"));
}
#[test]
fn empty_entries_do_not_end_the_walk() {
for header in ["203.0.113.9,", "9.9.9.9,, 203.0.113.9", " , 203.0.113.9"] {
let req = req_with(Some("127.0.0.1:1234"), Some(header));
assert_eq!(
loopback().extract(&req).unwrap(),
ip("203.0.113.9"),
"{header}"
);
}
}
#[test]
fn chain_trusted_end_to_end_falls_back_to_peer() {
let extractor = TrustedProxyKeyExtractor::new(cidrs(&["127.0.0.1/32", "10.0.0.0/8"]));
let req = req_with(Some("127.0.0.1:1"), Some("10.0.0.5, 10.0.0.6"));
assert_eq!(extractor.extract(&req).unwrap(), ip("127.0.0.1"));
}
#[test]
fn absurdly_long_chain_falls_back_to_peer() {
let chain = (0..MAX_FORWARDED_ENTRIES + 1)
.map(|_| "203.0.113.9")
.collect::<Vec<_>>()
.join(", ");
let req = req_with(Some("127.0.0.1:1"), Some(&chain));
assert_eq!(loopback().extract(&req).unwrap(), ip("127.0.0.1"));
}
#[test]
fn port_and_bracket_forms_are_accepted() {
let extractor = TrustedProxyKeyExtractor::new(cidrs(&["127.0.0.1/32"]));
for (entry, expected) in [
("203.0.113.9:443", "203.0.113.9"),
("[2001:db8::1]:443", "2001:db8::1"),
("[2001:db8::1]", "2001:db8::1"),
("2001:db8::1", "2001:db8::1"),
] {
let req = req_with(Some("127.0.0.1:1"), Some(entry));
assert_eq!(extractor.extract(&req).unwrap(), ip(expected), "{entry}");
}
}
#[test]
fn ipv4_mapped_peer_matches_an_ipv4_cidr() {
let req = req_with(Some("[::ffff:127.0.0.1]:1234"), Some("203.0.113.9"));
assert_eq!(loopback().extract(&req).unwrap(), ip("203.0.113.9"));
}
#[tokio::test]
async fn synthetic_connect_info_inserted_only_when_missing() {
use axum::body::Body;
use axum::routing::get;
use tower::ServiceExt;
async fn handler(ConnectInfo(addr): ConnectInfo<SocketAddr>) -> String {
addr.ip().to_string()
}
let app = axum::Router::new()
.route("/", get(handler))
.layer(axum::middleware::from_fn(
insert_default_connect_info_if_missing,
));
let resp = app
.oneshot(Request::builder().uri("/").body(Body::empty()).unwrap())
.await
.unwrap();
let body = axum::body::to_bytes(resp.into_body(), usize::MAX)
.await
.unwrap();
assert_eq!(&body[..], b"127.0.0.1");
}
}