use std::convert::Infallible;
use std::future::{Future, ready};
use std::net::{IpAddr, SocketAddr};
use axum::extract::{ConnectInfo, FromRequestParts};
use axum::http::HeaderMap;
use axum::http::request::Parts;
use ipnet::IpNet;
fn is_trusted(addr: IpAddr, trusted: &[IpNet]) -> bool {
trusted.iter().any(|net| net.contains(&addr))
}
#[must_use]
pub fn resolve_client_ip(
headers: &HeaderMap,
connect_info: Option<&ConnectInfo<SocketAddr>>,
trusted: &[IpNet],
) -> Option<IpAddr> {
let peer_ip = connect_info.map(|c| c.0.ip())?;
if !is_trusted(peer_ip, trusted) {
return Some(peer_ip);
}
if let Some(xff) = headers.get("x-forwarded-for").and_then(|v| v.to_str().ok()) {
let hops: Vec<&str> = xff
.split(',')
.map(str::trim)
.filter(|s| !s.is_empty())
.collect();
for hop in hops.iter().rev() {
if let Ok(addr) = hop.parse::<IpAddr>()
&& !is_trusted(addr, trusted)
{
return Some(addr);
}
}
}
for header in ["x-real-ip", "cf-connecting-ip"] {
if let Some(raw) = headers.get(header).and_then(|v| v.to_str().ok())
&& let Ok(addr) = raw.trim().parse::<IpAddr>()
&& !is_trusted(addr, trusted)
{
return Some(addr);
}
}
Some(peer_ip)
}
#[must_use]
pub fn resolve_client_ip_from_config(
headers: &HeaderMap,
connect_info: Option<&ConnectInfo<SocketAddr>>,
) -> Option<IpAddr> {
let trusted = systemprompt_models::Config::get()
.map(|c| c.trusted_proxies.clone())
.unwrap_or_default();
resolve_client_ip(headers, connect_info, &trusted)
}
#[must_use]
pub fn client_ip_from_request(request: &axum::extract::Request) -> Option<IpAddr> {
resolve_client_ip_from_config(
request.headers(),
request.extensions().get::<ConnectInfo<SocketAddr>>(),
)
}
#[derive(Debug, Clone, Copy)]
pub struct ClientIp(pub Option<IpAddr>);
impl<S: Sync> FromRequestParts<S> for ClientIp {
type Rejection = Infallible;
fn from_request_parts(
parts: &mut Parts,
_state: &S,
) -> impl Future<Output = Result<Self, Infallible>> + Send {
let resolved = resolve_client_ip_from_config(
&parts.headers,
parts.extensions.get::<ConnectInfo<SocketAddr>>(),
);
ready(Ok(Self(resolved)))
}
}