use std::net::IpAddr;
use axum::http::{HeaderMap, HeaderName};
use ipnet::IpNet;
use super::{canonical, nets_contain, parse_nets};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ClientIp(pub Option<IpAddr>);
#[derive(Debug, Clone)]
pub struct ProxyPolicy {
trusted: Vec<IpNet>,
header: HeaderName,
}
impl Default for ProxyPolicy {
fn default() -> Self {
Self {
trusted: Vec::new(),
header: HeaderName::from_static("x-forwarded-for"),
}
}
}
impl ProxyPolicy {
pub fn new(trusted_proxies: &[String], header: &str) -> anyhow::Result<Self> {
let trusted = parse_nets(trusted_proxies, "filter.trusted_proxies")?;
let header = HeaderName::try_from(header.to_ascii_lowercase())
.map_err(|error| anyhow::anyhow!("filter.forwarded_header: {error}"))?;
Ok(Self { trusted, header })
}
pub fn resolve(&self, peer: Option<IpAddr>, headers: &HeaderMap) -> Option<IpAddr> {
let peer = canonical(peer?);
if self.trusted.is_empty() || !nets_contain(&self.trusted, peer) {
return Some(peer);
}
let forwarded: Vec<IpAddr> = headers
.get_all(&self.header)
.iter()
.filter_map(|value| value.to_str().ok())
.flat_map(|value| value.split(','))
.filter_map(|entry| parse_forwarded_entry(entry.trim()))
.map(canonical)
.collect();
if let Some(client) = forwarded
.iter()
.rev()
.find(|ip| !nets_contain(&self.trusted, **ip))
{
return Some(*client);
}
Some(forwarded.first().copied().unwrap_or(peer))
}
}
fn parse_forwarded_entry(entry: &str) -> Option<IpAddr> {
if entry.is_empty() {
return None;
}
if let Some(rest) = entry.strip_prefix('[') {
let (host, _) = rest.split_once(']')?;
return host.parse().ok();
}
if let Ok(ip) = entry.parse::<IpAddr>() {
return Some(ip);
}
if entry.matches(':').count() == 1 {
let (host, _) = entry.split_once(':')?;
return host.parse().ok();
}
None
}
#[cfg(test)]
mod tests {
use super::*;
fn ip(value: &str) -> IpAddr {
value.parse().unwrap()
}
fn headers(pairs: &[(&str, &str)]) -> HeaderMap {
let mut map = HeaderMap::new();
for (name, value) in pairs {
map.append(HeaderName::try_from(*name).unwrap(), value.parse().unwrap());
}
map
}
fn policy(trusted: &[&str]) -> ProxyPolicy {
ProxyPolicy::new(
&trusted
.iter()
.map(std::string::ToString::to_string)
.collect::<Vec<_>>(),
"x-forwarded-for",
)
.unwrap()
}
#[test]
fn no_peer_means_no_client() {
assert_eq!(policy(&[]).resolve(None, &HeaderMap::new()), None);
}
#[test]
fn peer_is_used_when_no_proxy_is_trusted() {
let resolved = policy(&[]).resolve(
Some(ip("203.0.113.9")),
&headers(&[("x-forwarded-for", "10.0.0.1")]),
);
assert_eq!(resolved, Some(ip("203.0.113.9")));
}
#[test]
fn forwarded_header_is_ignored_from_an_untrusted_peer() {
let resolved = policy(&["127.0.0.1"]).resolve(
Some(ip("203.0.113.9")),
&headers(&[("x-forwarded-for", "10.0.0.1")]),
);
assert_eq!(resolved, Some(ip("203.0.113.9")));
}
#[test]
fn forwarded_header_is_honoured_from_a_trusted_peer() {
let resolved = policy(&["127.0.0.1"]).resolve(
Some(ip("127.0.0.1")),
&headers(&[("x-forwarded-for", "203.0.113.9")]),
);
assert_eq!(resolved, Some(ip("203.0.113.9")));
}
#[test]
fn chained_proxies_are_walked_past() {
let resolved = policy(&["127.0.0.1", "10.0.0.0/8"]).resolve(
Some(ip("127.0.0.1")),
&headers(&[("x-forwarded-for", "203.0.113.9, 10.0.0.8")]),
);
assert_eq!(resolved, Some(ip("203.0.113.9")));
}
#[test]
fn repeated_header_lines_are_concatenated() {
let resolved = policy(&["127.0.0.1", "10.0.0.0/8"]).resolve(
Some(ip("127.0.0.1")),
&headers(&[
("x-forwarded-for", "203.0.113.9"),
("x-forwarded-for", "10.0.0.8"),
]),
);
assert_eq!(resolved, Some(ip("203.0.113.9")));
}
#[test]
fn all_hops_trusted_falls_back_to_the_leftmost() {
let resolved = policy(&["127.0.0.1", "10.0.0.0/8"]).resolve(
Some(ip("127.0.0.1")),
&headers(&[("x-forwarded-for", "10.0.0.4, 10.0.0.8")]),
);
assert_eq!(resolved, Some(ip("10.0.0.4")));
}
#[test]
fn trusted_peer_without_a_usable_header_falls_back_to_the_peer() {
let trusted = policy(&["127.0.0.1"]);
assert_eq!(
trusted.resolve(Some(ip("127.0.0.1")), &HeaderMap::new()),
Some(ip("127.0.0.1"))
);
assert_eq!(
trusted.resolve(
Some(ip("127.0.0.1")),
&headers(&[("x-forwarded-for", "not-an-ip, also-garbage")])
),
Some(ip("127.0.0.1"))
);
}
#[test]
fn ipv4_mapped_peer_is_canonicalized() {
let resolved = policy(&[]).resolve(Some(ip("::ffff:192.168.1.5")), &HeaderMap::new());
assert_eq!(resolved, Some(ip("192.168.1.5")));
}
#[test]
fn ipv4_mapped_peer_matches_a_v4_trusted_proxy() {
let resolved = policy(&["127.0.0.1"]).resolve(
Some(ip("::ffff:127.0.0.1")),
&headers(&[("x-forwarded-for", "203.0.113.9")]),
);
assert_eq!(resolved, Some(ip("203.0.113.9")));
}
#[test]
fn forwarded_entries_may_carry_ports() {
assert_eq!(parse_forwarded_entry("1.2.3.4"), Some(ip("1.2.3.4")));
assert_eq!(parse_forwarded_entry("1.2.3.4:8443"), Some(ip("1.2.3.4")));
assert_eq!(
parse_forwarded_entry("2001:db8::1"),
Some(ip("2001:db8::1"))
);
assert_eq!(
parse_forwarded_entry("[2001:db8::1]:443"),
Some(ip("2001:db8::1"))
);
assert_eq!(
parse_forwarded_entry("[2001:db8::1]"),
Some(ip("2001:db8::1"))
);
assert_eq!(parse_forwarded_entry(""), None);
assert_eq!(parse_forwarded_entry("unknown"), None);
assert_eq!(parse_forwarded_entry("[bad"), None);
assert_eq!(parse_forwarded_entry("1.2.3.4:notaport:x"), None);
}
#[test]
fn new_validates_its_inputs() {
assert!(ProxyPolicy::new(&["10.0.0.0/8".to_string()], "x-real-ip").is_ok());
assert!(ProxyPolicy::new(&["nope".to_string()], "x-forwarded-for").is_err());
assert!(ProxyPolicy::new(&[], "not a header").is_err());
}
#[test]
fn a_custom_header_name_is_used() {
let policy = ProxyPolicy::new(&["127.0.0.1".to_string()], "X-Real-IP").unwrap();
let resolved = policy.resolve(
Some(ip("127.0.0.1")),
&headers(&[("x-real-ip", "203.0.113.9")]),
);
assert_eq!(resolved, Some(ip("203.0.113.9")));
}
}