use std::net::{IpAddr, SocketAddr};
use axum::extract::{ConnectInfo, FromRequestParts, Request};
use axum::http::request::Parts;
use axum::middleware::Next;
use axum::response::Response;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum TrustedEntry {
Exact(std::net::IpAddr),
Net(std::net::IpAddr, u8),
}
#[derive(Debug, thiserror::Error, PartialEq, Eq)]
pub enum CidrParseError {
#[error("missing '/' in CIDR: {0}")]
MissingPrefix(String),
#[error("invalid address in CIDR: {0}")]
Address(String),
#[error("invalid prefix length in CIDR: {0}")]
Prefix(String),
#[error("CIDR names a host, not a network (host bits set below the prefix): {0}")]
HostBitsSet(String),
}
fn parse_cidr(spec: &str) -> Result<(std::net::IpAddr, u8), CidrParseError> {
let (addr, prefix) = spec
.split_once('/')
.ok_or_else(|| CidrParseError::MissingPrefix(spec.to_string()))?;
let addr: std::net::IpAddr = addr
.trim()
.parse()
.map_err(|_| CidrParseError::Address(spec.to_string()))?;
let prefix: u8 = prefix
.trim()
.parse()
.map_err(|_| CidrParseError::Prefix(spec.to_string()))?;
let max = if addr.is_ipv4() { 32 } else { 128 };
if prefix > max {
return Err(CidrParseError::Prefix(spec.to_string()));
}
if !host_bits_clear(addr, prefix) {
return Err(CidrParseError::HostBitsSet(spec.to_string()));
}
Ok((addr, prefix))
}
fn host_bits_clear(addr: std::net::IpAddr, prefix: u8) -> bool {
match addr {
std::net::IpAddr::V4(a) => prefix == 32 || u32::from(a) & (u32::MAX >> prefix) == 0,
std::net::IpAddr::V6(a) => prefix == 128 || u128::from(a) & (u128::MAX >> prefix) == 0,
}
}
fn in_network(addr: std::net::IpAddr, net: std::net::IpAddr, prefix: u8) -> bool {
match (addr, net) {
(std::net::IpAddr::V4(a), std::net::IpAddr::V4(n)) => {
prefix == 0 || {
let shift = 32 - u32::from(prefix);
u32::from(a) >> shift == u32::from(n) >> shift
}
}
(std::net::IpAddr::V6(a), std::net::IpAddr::V6(n)) => {
prefix == 0 || {
let shift = 128 - u32::from(prefix);
u128::from(a) >> shift == u128::from(n) >> shift
}
}
_ => false,
}
}
#[derive(Debug, Clone, Default)]
pub struct TrustedProxies {
entries: std::sync::Arc<[TrustedEntry]>,
peer: PeerTrust,
real_ip: bool,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
enum PeerTrust {
#[default]
Enumerated,
AnyPeer,
}
impl TrustedProxies {
pub fn new(addrs: impl IntoIterator<Item = std::net::IpAddr>) -> Self {
Self {
entries: addrs.into_iter().map(TrustedEntry::Exact).collect(),
peer: PeerTrust::Enumerated,
real_ip: false,
}
}
pub fn from_cidrs<I, S>(specs: I) -> Result<Self, CidrParseError>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
let entries = specs
.into_iter()
.map(|spec| parse_cidr(spec.as_ref()).map(|(net, p)| TrustedEntry::Net(net, p)))
.collect::<Result<Vec<_>, _>>()?;
Ok(Self {
entries: entries.into(),
peer: PeerTrust::Enumerated,
real_ip: false,
})
}
pub fn with_cidrs<I, S>(self, specs: I) -> Result<Self, CidrParseError>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
let mut entries: Vec<TrustedEntry> = self.entries.iter().copied().collect();
for spec in specs {
let (net, prefix) = parse_cidr(spec.as_ref())?;
entries.push(TrustedEntry::Net(net, prefix));
}
Ok(Self {
entries: entries.into(),
peer: self.peer,
real_ip: self.real_ip,
})
}
pub fn loopback_only() -> Self {
Self::new([
std::net::IpAddr::V4(std::net::Ipv4Addr::LOCALHOST),
std::net::IpAddr::V6(std::net::Ipv6Addr::LOCALHOST),
])
}
pub fn private_transport() -> Self {
Self {
entries: Vec::new().into(),
peer: PeerTrust::AnyPeer,
real_ip: false,
}
}
pub fn proxy_overwrites_real_ip(mut self) -> Self {
self.real_ip = true;
self
}
pub fn trusts_peer(&self, peer: std::net::IpAddr) -> bool {
match self.peer {
PeerTrust::AnyPeer => true,
PeerTrust::Enumerated => self.trusts_hop(peer),
}
}
pub fn trusts_hop(&self, addr: std::net::IpAddr) -> bool {
self.entries.iter().any(|entry| match *entry {
TrustedEntry::Exact(address) => address == addr,
TrustedEntry::Net(net, prefix) => in_network(addr, net, prefix),
})
}
pub fn client_ip(
&self,
headers: &axum::http::HeaderMap,
peer: Option<std::net::IpAddr>,
) -> ClientIp {
let found = |addr, source| ClientIp {
addr: Some(addr),
source,
};
match peer {
Some(peer) if !self.trusts_peer(peer) => found(peer, Source::Peer),
Some(peer) => match client_from_chain(headers, self) {
ChainOutcome::Client(addr) => found(addr, Source::Forwarded),
ChainOutcome::NoAnswer => found(peer, Source::Peer),
},
None if self.peer != PeerTrust::AnyPeer => ClientIp::default(),
None => match client_from_chain(headers, self) {
ChainOutcome::Client(addr) => found(addr, Source::Forwarded),
ChainOutcome::NoAnswer => ClientIp::default(),
},
}
}
}
enum ChainOutcome {
Client(std::net::IpAddr),
NoAnswer,
}
fn client_from_chain(headers: &axum::http::HeaderMap, trusted: &TrustedProxies) -> ChainOutcome {
let mut lines = headers.get_all("X-Forwarded-For").iter().peekable();
if lines.peek().is_some() {
let mut entries: Vec<&str> = Vec::new();
for value in lines {
let Ok(text) = value.to_str() else {
return ChainOutcome::NoAnswer;
};
entries.extend(text.split(','));
}
for entry in entries.iter().rev() {
match entry.trim().parse::<std::net::IpAddr>() {
Ok(addr) if trusted.trusts_hop(addr) => continue,
Ok(addr) => return ChainOutcome::Client(addr),
Err(_) => return ChainOutcome::NoAnswer,
}
}
return ChainOutcome::NoAnswer;
}
if !trusted.real_ip {
return ChainOutcome::NoAnswer;
}
let mut real_ip = headers.get_all("X-Real-IP").iter();
match (real_ip.next(), real_ip.next()) {
(Some(value), None) => value
.to_str()
.ok()
.and_then(|value| value.trim().parse().ok())
.map(ChainOutcome::Client)
.unwrap_or(ChainOutcome::NoAnswer),
_ => ChainOutcome::NoAnswer,
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct ClientIp {
addr: Option<IpAddr>,
source: Source,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
#[cfg_attr(
feature = "rkyv",
derive(rkyv::Archive, rkyv::Serialize, rkyv::Deserialize)
)]
#[serde(rename_all = "snake_case")]
pub enum Source {
#[default]
Unknown,
Peer,
Forwarded,
Supplied,
}
impl Source {
pub fn as_str(self) -> &'static str {
match self {
Self::Unknown => "unknown",
Self::Peer => "peer",
Self::Forwarded => "forwarded",
Self::Supplied => "supplied",
}
}
}
impl std::fmt::Display for Source {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
impl ClientIp {
pub fn get(self) -> Option<IpAddr> {
self.addr
}
pub fn source(self) -> Source {
self.source
}
pub fn resolved(addr: Option<IpAddr>) -> Self {
Self {
addr,
source: Source::Supplied,
}
}
#[cfg(any(test, feature = "testing"))]
pub fn for_test(addr: Option<IpAddr>) -> Self {
Self {
addr,
source: Source::Supplied,
}
}
}
async fn run(trusted: TrustedProxies, mut request: Request, next: Next) -> Response {
let peer = request
.extensions()
.get::<ConnectInfo<SocketAddr>>()
.map(|info| info.0.ip());
let resolved = trusted.client_ip(request.headers(), peer);
request.extensions_mut().insert(resolved);
next.run(request).await
}
pub fn layer(router: axum::Router, trusted: TrustedProxies) -> axum::Router {
router.layer(axum::middleware::from_fn(move |request, next| {
let trusted = trusted.clone();
async move { run(trusted, request, next).await }
}))
}
impl<S> FromRequestParts<S> for ClientIp
where
S: Send + Sync,
{
type Rejection = std::convert::Infallible;
async fn from_request_parts(parts: &mut Parts, _: &S) -> Result<Self, Self::Rejection> {
Ok(parts
.extensions
.get::<ClientIp>()
.copied()
.unwrap_or_default())
}
}
#[cfg(test)]
mod tests;