use axum::{body::Body, extract::State, http::Request, middleware::Next, response::Response};
use serde::{Deserialize, Serialize};
use std::net::{IpAddr, Ipv4Addr};
#[derive(Clone, Debug, Copy)]
pub struct ClientIp(pub IpAddr);
fn ip_in_cidr(ip: &IpAddr, network: &IpAddr, prefix_len: u8) -> bool {
match (ip, network) {
(IpAddr::V4(ip), IpAddr::V4(net)) => {
if prefix_len == 0 {
return true;
}
if prefix_len > 32 {
return false;
}
let shift = 32 - prefix_len;
let mask = if shift == 32 { 0u32 } else { u32::MAX << shift };
(u32::from_be_bytes(ip.octets()) & mask) == (u32::from_be_bytes(net.octets()) & mask)
}
(IpAddr::V6(ip), IpAddr::V6(net)) => {
if prefix_len == 0 {
return true;
}
if prefix_len > 128 {
return false;
}
let shift = 128 - prefix_len;
let mask = if shift == 128 {
0u128
} else {
u128::MAX << shift
};
(u128::from_be_bytes(ip.octets()) & mask) == (u128::from_be_bytes(net.octets()) & mask)
}
_ => false,
}
}
fn canonicalize_ip(ip: IpAddr) -> IpAddr {
match ip {
IpAddr::V6(v6) => match v6.to_ipv4_mapped() {
Some(v4) => IpAddr::V4(v4),
None => IpAddr::V6(v6),
},
v4 => v4,
}
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct TrustedProxies {
exact: Vec<IpAddr>,
cidrs: Vec<(IpAddr, u8)>,
}
impl Default for TrustedProxies {
fn default() -> Self {
Self {
exact: vec![],
cidrs: vec![
("127.0.0.0".parse().unwrap(), 8), ("10.0.0.0".parse().unwrap(), 8), ("172.16.0.0".parse().unwrap(), 12), ("192.168.0.0".parse().unwrap(), 16), ("::1".parse().unwrap(), 128), ("fc00::".parse().unwrap(), 7), ],
}
}
}
impl TrustedProxies {
pub fn new(exact: Vec<IpAddr>, cidrs: Vec<(IpAddr, u8)>) -> Self {
Self { exact, cidrs }
}
pub fn none() -> Self {
Self {
exact: vec![],
cidrs: vec![],
}
}
pub fn default_for_edge(acme_edge: bool) -> Self {
if acme_edge {
Self::none()
} else {
Self::default()
}
}
pub fn is_trusted(&self, ip: &IpAddr) -> bool {
if self.exact.contains(ip) {
return true;
}
self.cidrs
.iter()
.any(|(net, prefix)| ip_in_cidr(ip, net, *prefix))
}
pub fn extract_client_ip(&self, req: &Request<Body>, conn_ip: Option<IpAddr>) -> IpAddr {
let Some(conn_ip_value) = conn_ip else {
return IpAddr::V4(Ipv4Addr::LOCALHOST);
};
let conn_ip_value = canonicalize_ip(conn_ip_value);
if !self.is_trusted(&conn_ip_value) {
return conn_ip_value;
}
let xff: Vec<IpAddr> = req
.headers()
.get("x-forwarded-for")
.and_then(|v| v.to_str().ok())
.map(|s| {
s.split(',')
.filter_map(|ip| ip.trim().parse::<IpAddr>().ok())
.map(canonicalize_ip)
.collect()
})
.unwrap_or_default();
if xff.is_empty() {
return conn_ip_value;
}
for ip in xff.iter().rev() {
if !self.is_trusted(ip) {
return *ip;
}
}
xff.into_iter().next().unwrap_or(conn_ip_value)
}
}
use crate::utils::aliases::AEngine;
pub async fn trusted_proxies_middleware(
State(engine): State<AEngine>,
mut req: Request<Body>,
next: Next,
) -> Response {
use axum::extract::ConnectInfo;
use std::net::SocketAddr;
let conn_ip = req
.extensions()
.get::<ConnectInfo<SocketAddr>>()
.map(|ci| ci.0.ip());
let client_ip = engine.trusted_proxies.extract_client_ip(&req, conn_ip);
req.extensions_mut().insert(ClientIp(client_ip));
next.run(req).await
}
#[cfg(test)]
mod tests {
use super::*;
fn ip(s: &str) -> IpAddr {
s.parse().unwrap()
}
fn req_with_xff(xff: Option<&str>) -> Request<Body> {
let mut b = Request::builder();
if let Some(v) = xff {
b = b.header("x-forwarded-for", v);
}
b.body(Body::empty()).unwrap()
}
#[test]
fn default_for_edge_trusts_private_ranges_behind_proxy() {
let tp = TrustedProxies::default_for_edge(false);
assert!(tp.is_trusted(&ip("10.0.0.1")), "10/8 trusted behind proxy");
assert!(tp.is_trusted(&ip("127.0.0.1")), "loopback trusted");
}
#[test]
fn default_for_edge_trusts_nothing_when_acme_edge() {
let tp = TrustedProxies::default_for_edge(true);
assert!(
!tp.is_trusted(&ip("10.0.0.1")),
"private-range peer must NOT be trusted in edge mode"
);
assert!(
!tp.is_trusted(&ip("127.0.0.1")),
"loopback not trusted in edge mode"
);
}
#[test]
fn edge_mode_ignores_spoofed_xff() {
let tp = TrustedProxies::default_for_edge(true);
let req = req_with_xff(Some("1.2.3.4"));
assert_eq!(
tp.extract_client_ip(&req, Some(ip("10.0.0.9"))),
ip("10.0.0.9")
);
}
#[test]
fn canonicalize_maps_ipv4_mapped_v6() {
assert_eq!(canonicalize_ip(ip("::ffff:10.0.0.5")), ip("10.0.0.5"));
assert_eq!(canonicalize_ip(ip("::ffff:8.8.8.8")), ip("8.8.8.8"));
}
#[test]
fn canonicalize_leaves_plain_ipv4_untouched() {
assert_eq!(canonicalize_ip(ip("1.2.3.4")), ip("1.2.3.4"));
}
#[test]
fn canonicalize_leaves_genuine_ipv6_untouched() {
assert_eq!(canonicalize_ip(ip("::1")), ip("::1"));
assert_eq!(canonicalize_ip(ip("2001:db8::1")), ip("2001:db8::1"));
}
#[test]
fn no_conn_info_never_trusts_xff() {
let tp = TrustedProxies::default();
let req = req_with_xff(Some("9.9.9.9"));
assert_eq!(tp.extract_client_ip(&req, None), ip("127.0.0.1"));
}
#[test]
fn direct_untrusted_peer_ignores_xff() {
let tp = TrustedProxies::default();
let req = req_with_xff(Some("9.9.9.9")); assert_eq!(
tp.extract_client_ip(&req, Some(ip("8.8.8.8"))),
ip("8.8.8.8")
);
}
#[test]
fn trusted_proxy_uses_xff_client() {
let tp = TrustedProxies::default();
let req = req_with_xff(Some("9.9.9.9"));
assert_eq!(
tp.extract_client_ip(&req, Some(ip("10.0.0.1"))),
ip("9.9.9.9")
);
}
#[test]
fn trusted_proxy_walks_right_to_left() {
let tp = TrustedProxies::default();
let req = req_with_xff(Some("9.9.9.9, 10.0.0.2"));
assert_eq!(
tp.extract_client_ip(&req, Some(ip("10.0.0.1"))),
ip("9.9.9.9")
);
}
#[test]
fn trusted_proxy_no_xff_returns_peer() {
let tp = TrustedProxies::default();
let req = req_with_xff(None);
assert_eq!(
tp.extract_client_ip(&req, Some(ip("10.0.0.1"))),
ip("10.0.0.1")
);
}
#[test]
fn all_xff_trusted_returns_leftmost() {
let tp = TrustedProxies::default();
let req = req_with_xff(Some("10.0.0.9, 10.0.0.2"));
assert_eq!(
tp.extract_client_ip(&req, Some(ip("10.0.0.1"))),
ip("10.0.0.9")
);
}
#[test]
fn ipv4_mapped_proxy_is_recognized_as_trusted() {
let tp = TrustedProxies::default();
let req = req_with_xff(Some("9.9.9.9"));
assert_eq!(
tp.extract_client_ip(&req, Some(ip("::ffff:10.0.0.1"))),
ip("9.9.9.9")
);
}
#[test]
fn ipv4_mapped_xff_entry_is_canonicalized() {
let tp = TrustedProxies::default();
let req = req_with_xff(Some("::ffff:9.9.9.9"));
assert_eq!(
tp.extract_client_ip(&req, Some(ip("10.0.0.1"))),
ip("9.9.9.9")
);
}
}