1use std::net::{IpAddr, SocketAddr};
13
14use axum::extract::{ConnectInfo, FromRequestParts};
15use axum::http::request::Parts;
16
17use crate::config::ForwardedHeader;
18use crate::state::AppState;
19
20#[derive(Debug, Clone, Copy)]
23pub struct ClientIp(pub IpAddr);
24
25impl std::fmt::Display for ClientIp {
26 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
27 write!(f, "{}", self.0)
28 }
29}
30
31impl FromRequestParts<AppState> for ClientIp {
32 type Rejection = std::convert::Infallible;
33
34 async fn from_request_parts(
35 parts: &mut Parts,
36 state: &AppState,
37 ) -> Result<Self, Self::Rejection> {
38 let peer = ConnectInfo::<SocketAddr>::from_request_parts(parts, state)
44 .await
45 .map_or(
46 IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED),
47 |ConnectInfo(addr)| addr.ip(),
48 );
49
50 let cfg = &state.config;
51 if matches!(cfg.forwarded_header, ForwardedHeader::None)
52 || !is_trusted_proxy(&cfg.trusted_proxies, peer)
53 {
54 return Ok(Self(peer));
55 }
56
57 let resolved = match cfg.forwarded_header {
58 ForwardedHeader::XForwardedFor => extract_x_forwarded_for(parts, &cfg.trusted_proxies),
59 ForwardedHeader::Forwarded => extract_forwarded(parts, &cfg.trusted_proxies),
60 ForwardedHeader::None => None,
61 };
62
63 Ok(Self(resolved.unwrap_or(peer)))
64 }
65}
66
67fn is_trusted_proxy(networks: &[ipnet::IpNet], ip: IpAddr) -> bool {
69 networks.iter().any(|net| net.contains(&ip))
70}
71
72fn extract_x_forwarded_for(parts: &Parts, trusted: &[ipnet::IpNet]) -> Option<IpAddr> {
77 let header = parts.headers.get_all("x-forwarded-for");
78 let combined: Vec<String> = header
81 .iter()
82 .filter_map(|v| v.to_str().ok())
83 .flat_map(|s| s.split(','))
84 .map(|s| s.trim().to_owned())
85 .filter(|s| !s.is_empty())
86 .collect();
87 let mut iter = combined.into_iter().rev();
88 while let Some(entry) = iter.next() {
89 match entry.parse::<IpAddr>() {
90 Ok(ip) => {
91 if is_trusted_proxy(trusted, ip) {
92 continue;
93 }
94 return Some(ip);
95 }
96 Err(err) => {
97 let safe: String = entry.chars().flat_map(char::escape_debug).collect();
104 tracing::debug!(
105 "ignoring un-parseable X-Forwarded-For entry {safe:?}: {err}; remaining chain ignored"
106 );
107 drop(iter);
108 return None;
109 }
110 }
111 }
112 None
113}
114
115fn extract_forwarded(parts: &Parts, trusted: &[ipnet::IpNet]) -> Option<IpAddr> {
118 let entries: Vec<String> = parts
119 .headers
120 .get_all("forwarded")
121 .iter()
122 .filter_map(|v| v.to_str().ok())
123 .flat_map(|s| s.split(','))
124 .map(|s| s.trim().to_owned())
125 .filter(|s| !s.is_empty())
126 .collect();
127 let mut iter = entries.into_iter().rev();
128 while let Some(entry) = iter.next() {
129 let Some(for_value) = forwarded_for_param(&entry) else {
130 continue;
131 };
132 match parse_forwarded_for_node(&for_value) {
133 Some(ip) => {
134 if is_trusted_proxy(trusted, ip) {
135 continue;
136 }
137 return Some(ip);
138 }
139 None => {
140 let safe: String = for_value.chars().flat_map(char::escape_debug).collect();
143 tracing::debug!(
144 "ignoring un-parseable Forwarded for= node {safe:?}; remaining chain ignored"
145 );
146 drop(iter);
147 return None;
148 }
149 }
150 }
151 None
152}
153
154fn forwarded_for_param(entry: &str) -> Option<String> {
158 for raw in entry.split(';') {
159 let raw = raw.trim();
160 let Some((name, value)) = raw.split_once('=') else {
161 continue;
162 };
163 if name.trim().eq_ignore_ascii_case("for") {
164 let v = value.trim();
165 let stripped = v.strip_prefix('"').and_then(|s| s.strip_suffix('"'));
166 return Some(stripped.unwrap_or(v).to_owned());
167 }
168 }
169 None
170}
171
172fn parse_forwarded_for_node(node: &str) -> Option<IpAddr> {
175 let trimmed = node.trim();
176 if trimmed.is_empty() || trimmed.eq_ignore_ascii_case("unknown") {
177 return None;
178 }
179 if let Some(rest) = trimmed.strip_prefix('[') {
181 if let Some((inside, _after)) = rest.split_once(']') {
182 return inside.parse().ok();
183 }
184 return None;
185 }
186 if let Ok(ip) = trimmed.parse::<IpAddr>() {
189 return Some(ip);
190 }
191 if let Some((host, _port)) = trimmed.rsplit_once(':')
193 && let Ok(ip) = host.parse::<IpAddr>()
194 {
195 return Some(ip);
196 }
197 None
198}