use std::fmt;
use std::net::{IpAddr, SocketAddr};
use std::str::FromStr;
use axum::http::{HeaderMap, HeaderName};
pub const X_FORWARDED_FOR: HeaderName = HeaderName::from_static("x-forwarded-for");
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub struct ProxyNet {
base: IpAddr,
prefix: u8,
}
impl ProxyNet {
pub fn new(base: IpAddr, prefix: u8) -> Result<Self, ProxyNetError> {
let base = base.to_canonical();
let width = Self::width(base);
if prefix > width {
return Err(ProxyNetError::Prefix { prefix, width });
}
Ok(Self {
base: Self::mask(base, prefix),
prefix,
})
}
#[must_use]
pub fn host(addr: IpAddr) -> Self {
let base = addr.to_canonical();
Self {
prefix: Self::width(base),
base,
}
}
#[must_use]
pub fn contains(&self, addr: IpAddr) -> bool {
let addr = addr.to_canonical();
if self.base.is_ipv4() != addr.is_ipv4() {
return false;
}
Self::mask(addr, self.prefix) == self.base
}
fn width(addr: IpAddr) -> u8 {
if addr.is_ipv4() { 32 } else { 128 }
}
fn mask(addr: IpAddr, prefix: u8) -> IpAddr {
match addr {
IpAddr::V4(v4) => {
let bits = u32::from(v4);
let kept = if prefix == 0 {
0
} else {
bits & (u32::MAX << (32 - u32::from(prefix)))
};
IpAddr::V4(kept.into())
}
IpAddr::V6(v6) => {
let bits = u128::from(v6);
let kept = if prefix == 0 {
0
} else {
bits & (u128::MAX << (128 - u32::from(prefix)))
};
IpAddr::V6(kept.into())
}
}
}
}
impl fmt::Display for ProxyNet {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(formatter, "{}/{}", self.base, self.prefix)
}
}
impl From<IpAddr> for ProxyNet {
fn from(addr: IpAddr) -> Self {
Self::host(addr)
}
}
impl FromStr for ProxyNet {
type Err = ProxyNetError;
fn from_str(text: &str) -> Result<Self, Self::Err> {
let text = text.trim();
let (addr, prefix) = match text.split_once('/') {
Some((addr, prefix)) => (addr.trim(), Some(prefix.trim())),
None => (text, None),
};
let addr: IpAddr = addr
.parse()
.map_err(|_| ProxyNetError::Address(addr.to_string()))?;
match prefix {
None => Ok(Self::host(addr)),
Some(prefix) => {
let bits: u8 = prefix
.parse()
.map_err(|_| ProxyNetError::Address(text.to_string()))?;
Self::new(addr, bits)
}
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum ProxyNetError {
Address(String),
Prefix {
prefix: u8,
width: u8,
},
}
impl fmt::Display for ProxyNetError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Address(text) => {
write!(formatter, "`{text}` is not an IP address or CIDR block")
}
Self::Prefix { prefix, width } => write!(
formatter,
"a /{prefix} prefix does not fit an address of {width} bits"
),
}
}
}
impl std::error::Error for ProxyNetError {}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
#[non_exhaustive]
pub struct TrustedProxies {
nets: Vec<ProxyNet>,
}
impl TrustedProxies {
#[must_use]
pub fn none() -> Self {
Self::default()
}
#[must_use]
pub fn trust(mut self, net: impl Into<ProxyNet>) -> Self {
self.nets.push(net.into());
self
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.nets.is_empty()
}
#[must_use]
pub fn nets(&self) -> &[ProxyNet] {
&self.nets
}
#[must_use]
pub fn contains(&self, addr: IpAddr) -> bool {
self.nets.iter().any(|net| net.contains(addr))
}
}
impl FromIterator<ProxyNet> for TrustedProxies {
fn from_iter<I: IntoIterator<Item = ProxyNet>>(iter: I) -> Self {
Self {
nets: iter.into_iter().collect(),
}
}
}
impl FromStr for TrustedProxies {
type Err = ProxyNetError;
fn from_str(text: &str) -> Result<Self, Self::Err> {
text.split([',', ' ', '\t', '\n'])
.map(str::trim)
.filter(|entry| !entry.is_empty())
.map(ProxyNet::from_str)
.collect::<Result<Vec<_>, _>>()
.map(|nets| Self { nets })
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub struct ClientIp(IpAddr);
impl ClientIp {
#[must_use]
pub fn resolve(peer: IpAddr, headers: &HeaderMap, trusted: &TrustedProxies) -> Self {
let peer = peer.to_canonical();
if !trusted.contains(peer) {
return Self(peer);
}
let chain: Vec<&str> = headers
.get_all(X_FORWARDED_FOR)
.iter()
.filter_map(|value| value.to_str().ok())
.flat_map(|value| value.split(','))
.collect();
let mut nearest = peer;
for entry in chain.into_iter().rev() {
let Some(addr) = parse_forwarded(entry) else {
break;
};
if trusted.contains(addr) {
nearest = addr;
continue;
}
return Self(addr);
}
Self(nearest)
}
#[must_use]
pub fn addr(&self) -> IpAddr {
self.0
}
}
impl fmt::Display for ClientIp {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
self.0.fmt(formatter)
}
}
impl From<ClientIp> for IpAddr {
fn from(client: ClientIp) -> Self {
client.0
}
}
fn parse_forwarded(entry: &str) -> Option<IpAddr> {
let entry = entry.trim();
if entry.is_empty() {
return None;
}
if let Ok(addr) = entry.parse::<IpAddr>() {
return Some(addr.to_canonical());
}
if let Ok(socket) = entry.parse::<SocketAddr>() {
return Some(socket.ip().to_canonical());
}
entry
.strip_prefix('[')
.and_then(|rest| rest.strip_suffix(']'))
.and_then(|inner| inner.parse::<IpAddr>().ok())
.map(|addr| addr.to_canonical())
}
#[cfg(test)]
mod tests {
use super::*;
fn ip(text: &str) -> IpAddr {
text.parse().expect("a literal address")
}
fn headers(chain: &[&str]) -> HeaderMap {
let mut map = HeaderMap::new();
for value in chain {
map.append(
X_FORWARDED_FOR,
value.parse().expect("a literal header value"),
);
}
map
}
#[test]
fn a_block_masks_its_host_bits() {
let net: ProxyNet = "10.4.1.9/8".parse().expect("a literal block");
assert_eq!(net.to_string(), "10.0.0.0/8");
assert!(net.contains(ip("10.255.255.255")));
assert!(!net.contains(ip("11.0.0.1")));
}
#[test]
fn a_bare_address_is_a_single_host() {
let net: ProxyNet = "192.0.2.5".parse().expect("a literal address");
assert_eq!(net.to_string(), "192.0.2.5/32");
assert!(net.contains(ip("192.0.2.5")));
assert!(!net.contains(ip("192.0.2.6")));
}
#[test]
fn the_two_address_families_do_not_match_each_other() {
let v4: ProxyNet = "0.0.0.0/0".parse().expect("a literal block");
assert!(v4.contains(ip("203.0.113.1")));
assert!(!v4.contains(ip("2001:db8::1")));
let v6: ProxyNet = "2001:db8::/32".parse().expect("a literal block");
assert!(v6.contains(ip("2001:db8::dead:beef")));
assert!(!v6.contains(ip("2001:db9::1")));
}
#[test]
fn a_mapped_v4_address_is_compared_as_v4() {
let net: ProxyNet = "10.0.0.0/8".parse().expect("a literal block");
assert!(net.contains(ip("::ffff:10.1.2.3")));
}
#[test]
fn an_over_long_prefix_is_refused() {
assert_eq!(
"10.0.0.0/33".parse::<ProxyNet>(),
Err(ProxyNetError::Prefix {
prefix: 33,
width: 32
})
);
assert!(matches!(
"not-an-address/8".parse::<ProxyNet>(),
Err(ProxyNetError::Address(_))
));
}
#[test]
fn a_list_parses_from_one_string() {
let trusted: TrustedProxies = "10.0.0.0/8, 127.0.0.1"
.parse()
.expect("two literal entries");
assert_eq!(trusted.nets().len(), 2);
assert!(trusted.contains(ip("10.9.9.9")));
assert!(trusted.contains(ip("127.0.0.1")));
assert!(!trusted.contains(ip("203.0.113.1")));
assert!(
"".parse::<TrustedProxies>()
.expect("empty is a list")
.is_empty()
);
}
#[test]
fn an_untrusted_peer_cannot_forge_a_forwarded_header() {
let resolved = ClientIp::resolve(
ip("203.0.113.9"),
&headers(&["1.2.3.4"]),
&TrustedProxies::none(),
);
assert_eq!(resolved.addr(), ip("203.0.113.9"));
}
#[test]
fn a_trusted_proxy_hands_over_the_client() {
let trusted = TrustedProxies::none().trust(ip("10.0.0.1"));
let resolved = ClientIp::resolve(ip("10.0.0.1"), &headers(&["203.0.113.9"]), &trusted);
assert_eq!(resolved.addr(), ip("203.0.113.9"));
}
#[test]
fn the_walk_stops_at_the_first_untrusted_hop() {
let trusted: TrustedProxies = "10.0.0.0/8".parse().expect("a literal block");
let resolved = ClientIp::resolve(
ip("10.0.0.1"),
&headers(&["9.9.9.9, 203.0.113.9, 10.0.0.2"]),
&trusted,
);
assert_eq!(resolved.addr(), ip("203.0.113.9"));
}
#[test]
fn several_header_lines_are_one_chain() {
let trusted: TrustedProxies = "10.0.0.0/8".parse().expect("a literal block");
let resolved = ClientIp::resolve(
ip("10.0.0.1"),
&headers(&["203.0.113.9", "10.0.0.2"]),
&trusted,
);
assert_eq!(resolved.addr(), ip("203.0.113.9"));
}
#[test]
fn an_all_trusted_chain_resolves_to_its_leftmost_entry() {
let trusted: TrustedProxies = "10.0.0.0/8".parse().expect("a literal block");
let resolved =
ClientIp::resolve(ip("10.0.0.1"), &headers(&["10.0.0.7, 10.0.0.2"]), &trusted);
assert_eq!(resolved.addr(), ip("10.0.0.7"));
}
#[test]
fn an_unreadable_hop_ends_the_walk() {
let trusted: TrustedProxies = "10.0.0.0/8".parse().expect("a literal block");
let resolved = ClientIp::resolve(
ip("10.0.0.1"),
&headers(&["203.0.113.9, unknown, 10.0.0.2"]),
&trusted,
);
assert_eq!(
resolved.addr(),
ip("10.0.0.2"),
"an entry we cannot read must not let the one behind it through"
);
}
#[test]
fn a_hop_with_a_port_is_still_an_address() {
let trusted: TrustedProxies = "10.0.0.0/8".parse().expect("a literal block");
for entry in ["203.0.113.9:4711", "[2001:db8::1]:443", "[2001:db8::1]"] {
let resolved = ClientIp::resolve(ip("10.0.0.1"), &headers(&[entry]), &trusted);
assert_ne!(
resolved.addr(),
ip("10.0.0.1"),
"`{entry}` should have parsed"
);
}
}
#[test]
fn a_trusted_peer_with_no_header_stays_the_peer() {
let trusted: TrustedProxies = "10.0.0.0/8".parse().expect("a literal block");
let resolved = ClientIp::resolve(ip("10.0.0.1"), &HeaderMap::new(), &trusted);
assert_eq!(resolved.addr(), ip("10.0.0.1"));
}
}