1use axum::extract::connect_info::{ConnectInfo, MockConnectInfo};
24use axum::http::request::Parts;
25use axum::http::Extensions;
26use axum::http::{HeaderMap, HeaderName};
27use std::net::{IpAddr, SocketAddr};
28
29pub static X_FORWARDED_FOR: HeaderName = HeaderName::from_static("x-forwarded-for");
31pub static X_REAL_IP: HeaderName = HeaderName::from_static("x-real-ip");
33
34#[derive(Clone, Copy, Debug)]
37pub struct TrustedProxyHops(pub u16);
38
39fn parse_forwarded_ip(raw: &str) -> Option<IpAddr> {
40 if let Ok(sa) = raw.parse::<SocketAddr>() {
42 return Some(sa.ip());
43 }
44 raw.parse::<IpAddr>().ok()
45}
46
47pub fn best_effort_client_ip(
55 headers: &HeaderMap,
56 extensions: &Extensions,
57 trusted_hops: Option<u16>,
58) -> Option<IpAddr> {
59 if let Some(MockConnectInfo(addr)) = extensions.get::<MockConnectInfo<SocketAddr>>() {
63 return Some(addr.ip());
64 }
65
66 let hops = trusted_hops.unwrap_or(0);
67
68 if hops == 0 {
71 if let Some(ConnectInfo(addr)) = extensions.get::<ConnectInfo<SocketAddr>>() {
72 return Some(addr.ip());
73 }
74 return None;
75 }
76
77 let xff_values: Vec<&str> = headers
80 .get_all(&X_FORWARDED_FOR)
81 .iter()
82 .filter_map(|v| v.to_str().ok())
83 .collect();
84 if !xff_values.is_empty() {
85 let v = xff_values.join(",");
86 let entries: Vec<&str> = v
87 .split(',')
88 .map(str::trim)
89 .filter(|s| !s.is_empty())
90 .collect();
91 if let Some(idx) = entries.len().checked_sub(hops as usize) {
95 if let Some(ip) = parse_forwarded_ip(entries[idx]) {
96 return Some(ip);
97 }
98 }
99 }
100
101 headers
102 .get(&X_REAL_IP)
103 .and_then(|v| v.to_str().ok())
104 .and_then(parse_forwarded_ip)
105}
106
107pub fn rate_limit_key_ip_or_unknown(
117 headers: &HeaderMap,
118 extensions: &Extensions,
119 trusted_hops: Option<u16>,
120) -> String {
121 best_effort_client_ip(headers, extensions, trusted_hops)
122 .map(|ip| ip.to_string())
123 .unwrap_or_else(|| "unknown".to_string())
124}
125
126#[derive(Debug, Clone, PartialEq, Eq)]
132pub struct RateLimitKey {
133 pub key: String,
134 pub fell_back: bool,
135}
136
137impl RateLimitKey {
138 pub fn resolve(
139 headers: &HeaderMap,
140 extensions: &Extensions,
141 trusted_hops: Option<u16>,
142 ) -> Self {
143 match best_effort_client_ip(headers, extensions, trusted_hops) {
144 Some(ip) => Self {
145 key: ip.to_string(),
146 fell_back: false,
147 },
148 None => Self {
149 key: "unknown".to_string(),
150 fell_back: true,
151 },
152 }
153 }
154}
155
156pub fn trusted_hops_from_parts(parts: &Parts, fallback: Option<u16>) -> Option<u16> {
159 parts
160 .extensions
161 .get::<TrustedProxyHops>()
162 .map(|h| h.0)
163 .or(fallback)
164}
165
166#[cfg(test)]
167mod tests {
168 use super::{
169 best_effort_client_ip, rate_limit_key_ip_or_unknown, trusted_hops_from_parts,
170 TrustedProxyHops, X_FORWARDED_FOR, X_REAL_IP,
171 };
172 use axum::extract::connect_info::{ConnectInfo, MockConnectInfo};
173 use axum::http::{Extensions, HeaderMap, HeaderValue};
174 use std::net::{IpAddr, SocketAddr};
175
176 fn xff() -> HeaderMap {
177 let mut headers = HeaderMap::new();
178 headers.insert(
179 &X_FORWARDED_FOR,
180 HeaderValue::from_static("203.0.113.10, 198.51.100.10"),
181 );
182 headers.insert(&X_REAL_IP, HeaderValue::from_static("198.51.100.20"));
183 headers
184 }
185
186 #[test]
187 fn configured_hops_override_connect_info_behind_proxy() {
188 let headers = xff();
191 let mut extensions = Extensions::new();
192 extensions.insert(ConnectInfo(SocketAddr::from(([127, 0, 0, 1], 4321))));
193 extensions.insert(TrustedProxyHops(2));
194
195 assert_eq!(
196 best_effort_client_ip(&headers, &extensions, Some(2)),
197 Some(IpAddr::from([203, 0, 113, 10]))
198 );
199 }
200
201 #[test]
202 fn forwarded_headers_are_ignored_without_trusted_proxies() {
203 let mut headers = HeaderMap::new();
205 headers.insert(&X_FORWARDED_FOR, HeaderValue::from_static("203.0.113.10"));
206 headers.insert(&X_REAL_IP, HeaderValue::from_static("198.51.100.20"));
207
208 assert_eq!(
209 best_effort_client_ip(&headers, &Extensions::new(), None),
210 None
211 );
212 assert_eq!(
213 best_effort_client_ip(&headers, &Extensions::new(), Some(0)),
214 None
215 );
216 }
217
218 #[test]
219 fn one_trusted_proxy_uses_rightmost_xff_entry() {
220 let headers = xff();
222 assert_eq!(
223 best_effort_client_ip(&headers, &Extensions::new(), Some(1)),
224 Some(IpAddr::from([198, 51, 100, 10]))
225 );
226 }
227
228 #[test]
229 fn two_trusted_proxies_use_second_from_right() {
230 let headers = xff();
231 assert_eq!(
232 best_effort_client_ip(&headers, &Extensions::new(), Some(2)),
233 Some(IpAddr::from([203, 0, 113, 10]))
234 );
235 }
236
237 #[test]
238 fn real_ip_fallback_when_xff_shorter_than_hop_count() {
239 let mut headers = HeaderMap::new();
240 headers.insert(&X_REAL_IP, HeaderValue::from_static("198.51.100.20"));
241 assert_eq!(
242 best_effort_client_ip(&headers, &Extensions::new(), Some(1)),
243 Some(IpAddr::from([198, 51, 100, 20]))
244 );
245 }
246
247 #[test]
248 fn spoofed_left_entries_cannot_override_trusted_resolution() {
249 let mut headers = HeaderMap::new();
251 headers.insert(
252 &X_FORWARDED_FOR,
253 HeaderValue::from_static("6.6.6.6, 203.0.113.10"),
254 );
255 assert_eq!(
256 best_effort_client_ip(&headers, &Extensions::new(), Some(1)),
257 Some(IpAddr::from([203, 0, 113, 10]))
258 );
259 }
260
261 #[test]
262 fn honest_single_hop_traffic_resolves_the_appended_entry() {
263 let mut headers = HeaderMap::new();
267 headers.insert(&X_FORWARDED_FOR, HeaderValue::from_static("198.51.100.7"));
268 assert_eq!(
269 best_effort_client_ip(&headers, &Extensions::new(), Some(1)),
270 Some(IpAddr::from([198, 51, 100, 7]))
271 );
272 }
273
274 #[test]
275 fn duplicate_xff_headers_are_joined_before_hop_indexing() {
276 let mut headers = HeaderMap::new();
279 headers.append(&X_FORWARDED_FOR, HeaderValue::from_static("6.6.6.6"));
280 headers.append(&X_FORWARDED_FOR, HeaderValue::from_static("198.51.100.10"));
281 assert_eq!(
282 best_effort_client_ip(&headers, &Extensions::new(), Some(1)),
283 Some(IpAddr::from([198, 51, 100, 10]))
284 );
285 }
286
287 #[test]
288 fn mock_connect_info_is_visible_to_best_effort() {
289 let mut extensions = Extensions::new();
290 extensions.insert(MockConnectInfo(SocketAddr::from(([127, 0, 0, 1], 4321))));
291
292 assert_eq!(
293 best_effort_client_ip(&HeaderMap::new(), &extensions, None),
294 Some(IpAddr::from([127, 0, 0, 1]))
295 );
296 }
297
298 #[test]
301 fn rate_limit_key_ip_formats_a_resolved_ip() {
302 let mut extensions = Extensions::new();
303 extensions.insert(MockConnectInfo(SocketAddr::from(([203, 0, 113, 9], 443))));
304
305 assert_eq!(
306 rate_limit_key_ip_or_unknown(&HeaderMap::new(), &extensions, None),
307 "203.0.113.9"
308 );
309 }
310
311 #[test]
312 fn rate_limit_key_ip_falls_back_to_unknown_when_unresolvable() {
313 assert_eq!(
314 rate_limit_key_ip_or_unknown(&HeaderMap::new(), &Extensions::new(), None),
315 "unknown"
316 );
317 let mut headers = HeaderMap::new();
318 headers.insert(&X_FORWARDED_FOR, HeaderValue::from_static("not-an-ip"));
319 assert_eq!(
320 rate_limit_key_ip_or_unknown(&headers, &Extensions::new(), Some(1)),
321 "unknown"
322 );
323 }
324
325 #[test]
326 fn trusted_hops_from_parts_prefers_extension_then_fallback() {
327 let mut parts = axum::http::Request::new(()).into_parts().0;
328 assert_eq!(trusted_hops_from_parts(&parts, None), None);
329 assert_eq!(trusted_hops_from_parts(&parts, Some(0)), Some(0));
330
331 parts.extensions.insert(TrustedProxyHops(3));
332 assert_eq!(trusted_hops_from_parts(&parts, None), Some(3));
333 assert_eq!(trusted_hops_from_parts(&parts, Some(0)), Some(3));
334 }
335}