systemprompt_api/services/middleware/
client_addr.rs1use std::convert::Infallible;
29use std::future::{Future, ready};
30use std::net::{IpAddr, SocketAddr};
31
32use axum::extract::{ConnectInfo, FromRequestParts};
33use axum::http::HeaderMap;
34use axum::http::request::Parts;
35use ipnet::IpNet;
36
37fn is_trusted(addr: IpAddr, trusted: &[IpNet]) -> bool {
38 trusted.iter().any(|net| net.contains(&addr))
39}
40
41#[must_use]
42pub fn resolve_client_ip(
43 headers: &HeaderMap,
44 connect_info: Option<&ConnectInfo<SocketAddr>>,
45 trusted: &[IpNet],
46) -> Option<IpAddr> {
47 let peer_ip = connect_info.map(|c| c.0.ip())?;
48
49 if !is_trusted(peer_ip, trusted) {
50 return Some(peer_ip);
51 }
52
53 if let Some(xff) = headers.get("x-forwarded-for").and_then(|v| v.to_str().ok()) {
54 let hops: Vec<&str> = xff
55 .split(',')
56 .map(str::trim)
57 .filter(|s| !s.is_empty())
58 .collect();
59 for hop in hops.iter().rev() {
60 if let Ok(addr) = hop.parse::<IpAddr>()
61 && !is_trusted(addr, trusted)
62 {
63 return Some(addr);
64 }
65 }
66 }
67
68 for header in ["x-real-ip", "fly-client-ip", "cf-connecting-ip"] {
69 if let Some(raw) = headers.get(header).and_then(|v| v.to_str().ok())
70 && let Ok(addr) = raw.trim().parse::<IpAddr>()
71 && !is_trusted(addr, trusted)
72 {
73 return Some(addr);
74 }
75 }
76
77 Some(peer_ip)
78}
79
80#[must_use]
81pub fn forwarded_headers_ignored(headers: &HeaderMap, peer_ip: IpAddr, trusted: &[IpNet]) -> bool {
82 !is_trusted(peer_ip, trusted)
83 && is_private_range(peer_ip)
84 && headers.contains_key("x-forwarded-for")
85}
86
87const fn is_private_range(ip: IpAddr) -> bool {
88 match ip {
89 IpAddr::V4(v4) => {
90 let is_cgnat = v4.octets()[0] == 100 && (v4.octets()[1] & 0xc0) == 64;
91 v4.is_loopback() || v4.is_private() || v4.is_link_local() || is_cgnat
92 },
93 IpAddr::V6(v6) => {
94 let seg0 = v6.segments()[0];
95 let is_unique_local = (seg0 & 0xfe00) == 0xfc00;
96 let is_link_local = (seg0 & 0xffc0) == 0xfe80;
97 v6.is_loopback() || is_unique_local || is_link_local
98 },
99 }
100}
101
102#[must_use]
103pub fn resolve_client_ip_from_config(
104 headers: &HeaderMap,
105 connect_info: Option<&ConnectInfo<SocketAddr>>,
106) -> Option<IpAddr> {
107 let trusted = systemprompt_models::Config::get()
108 .map(|c| c.trusted_proxies.clone())
109 .unwrap_or_default();
110 resolve_client_ip(headers, connect_info, &trusted)
111}
112
113#[must_use]
114pub fn client_ip_from_request(request: &axum::extract::Request) -> Option<IpAddr> {
115 resolve_client_ip_from_config(
116 request.headers(),
117 request.extensions().get::<ConnectInfo<SocketAddr>>(),
118 )
119}
120
121#[derive(Debug, Clone, Copy)]
122pub struct ClientIp(pub Option<IpAddr>);
123
124impl<S: Sync> FromRequestParts<S> for ClientIp {
125 type Rejection = Infallible;
126
127 fn from_request_parts(
128 parts: &mut Parts,
129 _state: &S,
130 ) -> impl Future<Output = Result<Self, Infallible>> + Send {
131 let resolved = resolve_client_ip_from_config(
132 &parts.headers,
133 parts.extensions.get::<ConnectInfo<SocketAddr>>(),
134 );
135 ready(Ok(Self(resolved)))
136 }
137}