use std::{
net::{IpAddr, SocketAddr},
str::FromStr,
};
use http::{HeaderMap, HeaderName};
use crate::Error;
macro_rules! client_ip_headers {
(
$(
$(#[$docs:meta])*
($variant:ident, $name:literal, $extractor:path);
)+
) => {
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ClientIpHeader {
$(
$(#[$docs])*
$variant,
)+
}
impl FromStr for ClientIpHeader {
type Err = Error;
fn from_str(source: &str) -> Result<Self, Self::Err> {
let name = HeaderName::from_bytes(source.as_bytes())
.map_err(|_| Error::unsupported_header_name(source))?;
match name.as_str() {
$(
$name => Ok(Self::$variant),
)+
_ => Err(Error::unsupported_header_name(source)),
}
}
}
impl ClientIpHeader {
fn extract(self, headers: &HeaderMap) -> Result<Option<IpAddr>, Error> {
match self {
$(
Self::$variant => ($extractor)(headers),
)+
}
}
}
};
}
client_ip_headers! {
(
CfConnectingIp,
"cf-connecting-ip",
crate::client_ip_headers::extract_header_cf_connecting_ip
);
(XRealIp, "x-real-ip", crate::client_ip_headers::extract_header_x_real_ip);
(Forwarded, "forwarded", crate::forwarded::extract_rightmost_forwarded);
(
XForwardedFor,
"x-forwarded-for",
crate::x_forwarded::extract_rightmost_x_forwarded_for
);
(
CloudFrontViewerAddress,
"cloudfront-viewer-address",
crate::client_ip_headers::extract_header_cloudfront_viewer_address
);
(
FlyClientIp,
"fly-client-ip",
crate::client_ip_headers::extract_header_fly_client_ip
);
(
TrueClientIp,
"true-client-ip",
crate::client_ip_headers::extract_header_true_client_ip
);
(
XEnvoyExternalAddress,
"x-envoy-external-address",
crate::client_ip_headers::extract_header_x_envoy_external_address
);
}
impl TryFrom<&str> for ClientIpHeader {
type Error = Error;
fn try_from(source: &str) -> Result<Self, Self::Error> {
source.parse()
}
}
pub const CLIENT_IP_HEADERS: &[ClientIpHeader] = &[
ClientIpHeader::Forwarded,
ClientIpHeader::XForwardedFor,
ClientIpHeader::XRealIp,
ClientIpHeader::CfConnectingIp,
];
pub const fn extract_peer_ip(peer: SocketAddr) -> IpAddr {
extract_peer_address(peer).ip()
}
pub const fn extract_peer_address(peer: SocketAddr) -> SocketAddr {
peer
}
#[cfg(feature = "axum")]
pub fn extract_axum_peer_address<B>(request: &http::Request<B>) -> Option<SocketAddr> {
request
.extensions()
.get::<axum::extract::ConnectInfo<SocketAddr>>()
.map(|info| info.0)
}
#[cfg(feature = "axum")]
pub fn extract_axum_peer_ip<B>(request: &http::Request<B>) -> Option<IpAddr> {
extract_axum_peer_address(request).map(|peer| peer.ip())
}
pub fn extract_client_ip(headers: &HeaderMap) -> Result<Option<IpAddr>, Error> {
extract_client_ip_with_headers(headers, CLIENT_IP_HEADERS)
}
pub fn extract_client_ip_with_headers(
headers: &HeaderMap,
sources: &[ClientIpHeader],
) -> Result<Option<IpAddr>, Error> {
for source in sources {
if let Some(ip) = source.extract(headers)? {
return Ok(Some(ip));
}
}
Ok(None)
}
#[cfg(test)]
mod tests {
use http::HeaderMap;
use crate::{forwarded::FORWARDED, x_forwarded::X_FORWARDED_FOR};
use super::*;
#[test]
fn peer_functions_preserve_the_transport_address() {
let peer: SocketAddr = "203.0.113.8:443".parse().unwrap();
assert_eq!(extract_peer_address(peer), peer);
assert_eq!(
extract_peer_ip(peer),
"203.0.113.8".parse::<IpAddr>().unwrap()
);
}
#[cfg(feature = "axum")]
#[test]
fn axum_peer_address_reads_connect_info_extension() {
let peer: SocketAddr = "203.0.113.8:443".parse().unwrap();
let mut request = http::Request::new(());
request
.extensions_mut()
.insert(axum::extract::ConnectInfo(peer));
assert_eq!(extract_axum_peer_address(&request), Some(peer));
}
#[cfg(feature = "axum")]
#[test]
fn axum_peer_address_returns_none_without_connect_info() {
let request = http::Request::new(());
assert_eq!(extract_axum_peer_address(&request), None);
}
#[cfg(feature = "axum")]
#[test]
fn axum_peer_ip_reads_connect_info_extension() {
let peer: SocketAddr = "203.0.113.8:443".parse().unwrap();
let mut request = http::Request::new(());
request
.extensions_mut()
.insert(axum::extract::ConnectInfo(peer));
assert_eq!(extract_axum_peer_ip(&request), Some(peer.ip()));
}
#[cfg(feature = "axum")]
#[test]
fn axum_peer_ip_returns_none_without_connect_info() {
let request = http::Request::new(());
assert_eq!(extract_axum_peer_ip(&request), None);
}
#[test]
fn default_order_is_stable_and_first_present_header_wins() {
assert_eq!(
CLIENT_IP_HEADERS,
&[
ClientIpHeader::Forwarded,
ClientIpHeader::XForwardedFor,
ClientIpHeader::XRealIp,
ClientIpHeader::CfConnectingIp,
]
);
let mut headers = HeaderMap::new();
headers.insert("cf-connecting-ip", "192.0.2.1".parse().unwrap());
headers.insert("x-real-ip", "192.0.2.2".parse().unwrap());
headers.insert(&FORWARDED, "for=192.0.2.3".parse().unwrap());
headers.insert(&X_FORWARDED_FOR, "192.0.2.4".parse().unwrap());
assert_eq!(
extract_client_ip(&headers).unwrap(),
Some("192.0.2.3".parse().unwrap())
);
}
#[test]
fn default_order_falls_through_only_when_a_header_is_absent() {
let mut headers = HeaderMap::new();
headers.insert(&FORWARDED, "for=192.0.2.3".parse().unwrap());
headers.insert(&X_FORWARDED_FOR, "192.0.2.4".parse().unwrap());
assert_eq!(
extract_client_ip(&headers).unwrap(),
Some("192.0.2.3".parse().unwrap())
);
headers.insert("x-real-ip", "not-an-ip".parse().unwrap());
assert_eq!(
extract_client_ip(&headers).unwrap(),
Some("192.0.2.3".parse().unwrap())
);
headers.insert(&FORWARDED, "for=unknown".parse().unwrap());
assert!(matches!(
extract_client_ip(&headers),
Err(Error::InvalidHeader { .. })
));
}
#[test]
fn chain_headers_return_the_rightmost_address() {
let mut headers = HeaderMap::new();
headers.insert(
&FORWARDED,
"for=192.0.2.1, for=198.51.100.2".parse().unwrap(),
);
assert_eq!(
extract_client_ip_with_headers(&headers, &[ClientIpHeader::Forwarded]).unwrap(),
Some("198.51.100.2".parse().unwrap())
);
headers.remove(&FORWARDED);
headers.insert(&X_FORWARDED_FOR, "192.0.2.1, 198.51.100.3".parse().unwrap());
assert_eq!(
extract_client_ip_with_headers(&headers, &[ClientIpHeader::XForwardedFor]).unwrap(),
Some("198.51.100.3".parse().unwrap())
);
}
#[test]
fn custom_order_changes_precedence() {
let mut headers = HeaderMap::new();
headers.insert("cf-connecting-ip", "192.0.2.1".parse().unwrap());
headers.insert("x-real-ip", "192.0.2.2".parse().unwrap());
let custom_headers = [ClientIpHeader::XRealIp, ClientIpHeader::CfConnectingIp];
assert_eq!(
extract_client_ip_with_headers(&headers, &custom_headers).unwrap(),
Some("192.0.2.2".parse().unwrap())
);
assert_eq!(extract_client_ip_with_headers(&headers, &[]).unwrap(), None);
}
#[test]
fn custom_order_supports_every_documented_single_value_header() {
for (source, header) in [
(ClientIpHeader::CfConnectingIp, "cf-connecting-ip"),
(ClientIpHeader::XRealIp, "x-real-ip"),
(ClientIpHeader::FlyClientIp, "fly-client-ip"),
(ClientIpHeader::TrueClientIp, "true-client-ip"),
(
ClientIpHeader::XEnvoyExternalAddress,
"x-envoy-external-address",
),
] {
let mut headers = HeaderMap::new();
headers.insert(header, "192.0.2.10".parse().unwrap());
assert_eq!(
extract_client_ip_with_headers(&headers, &[source]).unwrap(),
Some("192.0.2.10".parse().unwrap()),
"source {header}",
);
}
let mut headers = HeaderMap::new();
headers.insert(
"cloudfront-viewer-address",
"192.0.2.10:443".parse().unwrap(),
);
assert_eq!(
extract_client_ip_with_headers(&headers, &[ClientIpHeader::CloudFrontViewerAddress],)
.unwrap(),
Some("192.0.2.10".parse().unwrap())
);
}
#[test]
fn parses_supported_header_names() {
for (name, expected) in [
("cf-connecting-ip", ClientIpHeader::CfConnectingIp),
("X-Real-IP", ClientIpHeader::XRealIp),
("forwarded", ClientIpHeader::Forwarded),
("x-forwarded-for", ClientIpHeader::XForwardedFor),
(
"cloudfront-viewer-address",
ClientIpHeader::CloudFrontViewerAddress,
),
("fly-client-ip", ClientIpHeader::FlyClientIp),
("true-client-ip", ClientIpHeader::TrueClientIp),
(
"x-envoy-external-address",
ClientIpHeader::XEnvoyExternalAddress,
),
] {
assert_eq!(ClientIpHeader::try_from(name).unwrap(), expected);
assert_eq!(name.parse::<ClientIpHeader>().unwrap(), expected);
}
for header in ["forwarded-for", "not a header"] {
let error = header.parse::<ClientIpHeader>().unwrap_err();
assert!(matches!(error, Error::UnsupportedHeaderName { .. }));
assert!(error.to_string().contains(header));
}
}
#[test]
fn malformed_selected_source_does_not_fall_through() {
let mut headers = HeaderMap::new();
headers.insert(&FORWARDED, "for=unknown".parse().unwrap());
headers.insert(&X_FORWARDED_FOR, "192.0.2.4".parse().unwrap());
assert!(matches!(
extract_client_ip_with_headers(
&headers,
&[ClientIpHeader::Forwarded, ClientIpHeader::XForwardedFor,],
),
Err(Error::InvalidHeader { .. })
));
}
}