1use std::{
18 fmt,
19 net::{IpAddr, SocketAddr},
20 sync::Arc,
21};
22
23use axum::extract::{ConnectInfo, FromRequestParts};
24use http::{Extensions, HeaderMap, StatusCode, request::Parts};
25use ipnet::IpNet;
26use tuwunel_core::config::IpSource;
27
28#[derive(Clone, Copy, Debug)]
30pub(crate) struct ClientIp(pub(crate) IpAddr);
31
32#[derive(Clone, Copy, Debug)]
35pub struct ConfiguredIpSource(pub IpSource);
36
37#[derive(Clone, Debug)]
41pub struct TrustedPeerSubnets(pub Arc<[IpNet]>);
42
43impl<S> FromRequestParts<S> for ClientIp
44where
45 S: Sync,
46{
47 type Rejection = (StatusCode, &'static str);
48
49 async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
50 const ERROR: StatusCode = StatusCode::INTERNAL_SERVER_ERROR;
51
52 if let Some(&ConfiguredIpSource(source)) = parts.extensions.get::<ConfiguredIpSource>()
53 && !peer_is_trusted(&parts.extensions)
54 {
55 return secure_extract(source, &parts.headers, &parts.extensions)
56 .map(Self)
57 .ok_or((ERROR, "Can't extract client IP from configured ip_source"));
58 }
59
60 insecure_fallback(&parts.headers, &parts.extensions)
61 .map(Self)
62 .ok_or((ERROR, "Can't extract `ClientIp`, provide `axum::extract::ConnectInfo`"))
63 }
64}
65
66impl fmt::Display for ClientIp {
67 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { fmt::Display::fmt(&self.0, f) }
68}
69
70fn peer_is_trusted(extensions: &Extensions) -> bool {
71 let Some(ConnectInfo(addr)) = extensions.get::<ConnectInfo<SocketAddr>>() else {
72 return false;
73 };
74
75 let peer = addr.ip().to_canonical();
76
77 peer.is_loopback()
78 || extensions
79 .get::<TrustedPeerSubnets>()
80 .is_some_and(|TrustedPeerSubnets(nets)| nets.iter().any(|net| net.contains(&peer)))
81}
82
83fn secure_extract(
84 source: IpSource,
85 headers: &HeaderMap,
86 extensions: &Extensions,
87) -> Option<IpAddr> {
88 match source {
89 | IpSource::ConnectInfo => extensions
90 .get::<ConnectInfo<SocketAddr>>()
91 .map(|ConnectInfo(addr)| addr.ip()),
92 | IpSource::RightmostXForwardedFor => rightmost_x_forwarded_for(headers),
93 | IpSource::RightmostForwarded => rightmost_forwarded(headers),
94 | IpSource::XRealIp => single_ip_header(headers, "x-real-ip"),
95 | IpSource::CfConnectingIp => single_ip_header(headers, "cf-connecting-ip"),
96 | IpSource::TrueClientIp => single_ip_header(headers, "true-client-ip"),
97 | IpSource::FlyClientIp => single_ip_header(headers, "fly-client-ip"),
98 | IpSource::CloudFrontViewerAddress => cloudfront_viewer_address(headers),
99 }
100}
101
102fn rightmost_x_forwarded_for(headers: &HeaderMap) -> Option<IpAddr> {
103 headers
104 .get_all("x-forwarded-for")
105 .iter()
106 .filter_map(|v| v.to_str().ok())
107 .flat_map(|s| s.split(','))
108 .filter_map(|s| s.trim().parse::<IpAddr>().ok())
109 .next_back()
110}
111
112fn rightmost_forwarded(headers: &HeaderMap) -> Option<IpAddr> {
113 headers
114 .get_all("forwarded")
115 .iter()
116 .filter_map(|v| v.to_str().ok())
117 .flat_map(|s| s.split(','))
118 .filter_map(parse_forwarded_for)
119 .next_back()
120}
121
122fn insecure_fallback(headers: &HeaderMap, extensions: &Extensions) -> Option<IpAddr> {
124 leftmost_x_forwarded_for(headers)
125 .or_else(|| leftmost_forwarded(headers))
126 .or_else(|| single_ip_header(headers, "x-real-ip"))
127 .or_else(|| single_ip_header(headers, "fly-client-ip"))
128 .or_else(|| single_ip_header(headers, "true-client-ip"))
129 .or_else(|| single_ip_header(headers, "cf-connecting-ip"))
130 .or_else(|| cloudfront_viewer_address(headers))
131 .or_else(|| {
132 extensions
133 .get::<ConnectInfo<SocketAddr>>()
134 .map(|ConnectInfo(addr)| addr.ip())
135 })
136}
137
138fn leftmost_x_forwarded_for(headers: &HeaderMap) -> Option<IpAddr> {
139 headers
140 .get_all("x-forwarded-for")
141 .iter()
142 .filter_map(|v| v.to_str().ok())
143 .flat_map(|s| s.split(','))
144 .find_map(|s| s.trim().parse::<IpAddr>().ok())
145}
146
147fn leftmost_forwarded(headers: &HeaderMap) -> Option<IpAddr> {
150 headers
151 .get_all("forwarded")
152 .iter()
153 .filter_map(|v| v.to_str().ok())
154 .flat_map(|s| s.split(','))
155 .find_map(parse_forwarded_for)
156}
157
158fn parse_forwarded_for(stanza: &str) -> Option<IpAddr> {
159 let for_value = stanza
160 .split(';')
161 .find_map(|part| {
162 let (k, v) = part.split_once('=')?;
163 k.trim()
164 .eq_ignore_ascii_case("for")
165 .then_some(v.trim())
166 })?
167 .trim_matches('"');
168
169 let bracketed = for_value
170 .strip_prefix('[')
171 .and_then(|s| s.split_once(']'))
172 .map(|(ip, _rest)| ip);
173
174 let candidate = bracketed
175 .or_else(|| for_value.rsplit_once(':').map(|(ip, _port)| ip))
176 .unwrap_or(for_value);
177
178 candidate.trim().parse::<IpAddr>().ok()
179}
180
181fn single_ip_header(headers: &HeaderMap, name: &'static str) -> Option<IpAddr> {
182 headers
183 .get(name)
184 .and_then(|v| v.to_str().ok())
185 .and_then(|s| s.trim().parse::<IpAddr>().ok())
186}
187
188fn cloudfront_viewer_address(headers: &HeaderMap) -> Option<IpAddr> {
189 headers
190 .get("cloudfront-viewer-address")
191 .and_then(|v| v.to_str().ok())
192 .and_then(|s| s.rsplit_once(':').map(|(ip, _port)| ip))
193 .and_then(|s| s.trim().parse::<IpAddr>().ok())
194}
195
196#[cfg(test)]
197mod tests {
198 use std::{iter, net::SocketAddr, sync::Arc};
199
200 use axum::{
201 extract::{ConnectInfo, FromRequestParts},
202 http::{Request, StatusCode, request::Parts},
203 };
204 use ipnet::IpNet;
205 use tuwunel_core::config::IpSource;
206
207 use super::{ClientIp, ConfiguredIpSource, TrustedPeerSubnets};
208
209 fn trusted(nets: &[&str]) -> TrustedPeerSubnets {
210 let nets: Arc<[IpNet]> = nets
211 .iter()
212 .map(|s| s.parse().expect("test CIDR"))
213 .collect();
214
215 TrustedPeerSubnets(nets)
216 }
217
218 fn parts(headers: impl IntoIterator<Item = (&'static str, &'static str)>) -> Parts {
219 let mut request = Request::builder().uri("/");
220 for (name, value) in headers {
221 request = request.header(name, value);
222 }
223 let (parts, ()) = request.body(()).unwrap().into_parts();
224 parts
225 }
226
227 async fn extract_client_ip(
228 parts: &mut Parts,
229 ) -> Result<ClientIp, (StatusCode, &'static str)> {
230 ClientIp::from_request_parts(parts, &()).await
231 }
232
233 #[tokio::test]
234 async fn x_forwarded_for_uses_leftmost_ip() {
235 let mut parts = parts([("X-Forwarded-For", "1.1.1.1, 2.2.2.2")]);
236 let ClientIp(ip) = extract_client_ip(&mut parts).await.unwrap();
237 assert_eq!(ip.to_string(), "1.1.1.1");
238 }
239
240 #[tokio::test]
241 async fn x_forwarded_for_takes_priority_over_x_real_ip() {
242 let mut parts =
243 parts([("X-Forwarded-For", "1.1.1.1, 2.2.2.2"), ("X-Real-Ip", "3.3.3.3")]);
244 let ClientIp(ip) = extract_client_ip(&mut parts).await.unwrap();
245 assert_eq!(ip.to_string(), "1.1.1.1");
246 }
247
248 #[tokio::test]
249 async fn x_forwarded_for_accepts_ipv6() {
250 let mut parts = parts([("X-Forwarded-For", "2001:db8::1, 2001:db8::2")]);
251 let ClientIp(ip) = extract_client_ip(&mut parts).await.unwrap();
252 assert_eq!(ip.to_string(), "2001:db8::1");
253 }
254
255 #[tokio::test]
256 async fn x_real_ip_works() {
257 let mut parts = parts([("X-Real-Ip", "1.2.3.4")]);
258 let ClientIp(ip) = extract_client_ip(&mut parts).await.unwrap();
259 assert_eq!(ip.to_string(), "1.2.3.4");
260 }
261
262 #[tokio::test]
263 async fn malformed_headers_fall_through_to_next_valid_source() {
264 let mut parts = parts([
265 ("X-Forwarded-For", "foo"),
266 ("X-Real-Ip", "foo"),
267 ("Forwarded", "foo"),
268 ("Forwarded", "for=1.1.1.1;proto=https;by=2.2.2.2"),
269 ]);
270 let ClientIp(ip) = extract_client_ip(&mut parts).await.unwrap();
271 assert_eq!(ip.to_string(), "1.1.1.1");
272 }
273
274 #[tokio::test]
275 async fn no_headers_or_connect_info_rejects() {
276 let mut parts = parts(iter::empty());
277 let err = extract_client_ip(&mut parts).await.unwrap_err();
278 assert_eq!(err.0, StatusCode::INTERNAL_SERVER_ERROR);
279 assert!(err.1.contains("ConnectInfo"), "{err:?}");
280 }
281
282 #[tokio::test]
283 async fn configured_source_uses_secure_extraction() {
284 let mut parts = parts([("X-Forwarded-For", "1.1.1.1, 2.2.2.2")]);
285 parts
286 .extensions
287 .insert(ConfiguredIpSource(IpSource::RightmostXForwardedFor));
288 let ClientIp(ip) = extract_client_ip(&mut parts).await.unwrap();
289 assert_eq!(ip.to_string(), "2.2.2.2");
290 }
291
292 #[tokio::test]
293 async fn configured_source_without_matching_header_rejects() {
294 let mut parts = parts(iter::empty());
295 parts
296 .extensions
297 .insert(ConfiguredIpSource(IpSource::RightmostXForwardedFor));
298 let err = extract_client_ip(&mut parts).await.unwrap_err();
299 assert_eq!(err.0, StatusCode::INTERNAL_SERVER_ERROR);
300 assert_eq!(err.1, "Can't extract client IP from configured ip_source");
301 }
302
303 #[tokio::test]
304 async fn connect_info_fallback_uses_real_socket_addr_without_config() {
305 let socket_addr = SocketAddr::from(([203, 0, 113, 9], 4567));
306 let mut parts = parts(iter::empty());
307 parts.extensions.insert(ConnectInfo(socket_addr));
308
309 let ClientIp(ip) = extract_client_ip(&mut parts).await.unwrap();
310 assert_eq!(ip, socket_addr.ip());
311 }
312
313 #[tokio::test]
314 async fn loopback_peer_bypasses_configured_source_for_locally_connected_bridges() {
315 let socket_addr = SocketAddr::from(([127, 0, 0, 1], 38000));
316 let mut parts = parts(iter::empty());
317 parts.extensions.insert(ConnectInfo(socket_addr));
318 parts
319 .extensions
320 .insert(ConfiguredIpSource(IpSource::RightmostXForwardedFor));
321
322 let ClientIp(ip) = extract_client_ip(&mut parts).await.unwrap();
323 assert_eq!(ip, socket_addr.ip());
324 }
325
326 #[tokio::test]
327 async fn loopback_peer_with_proxy_header_still_uses_insecure_fallback() {
328 let socket_addr = SocketAddr::from(([127, 0, 0, 1], 38000));
333 let mut parts = parts([("X-Forwarded-For", "9.9.9.9")]);
334 parts.extensions.insert(ConnectInfo(socket_addr));
335 parts
336 .extensions
337 .insert(ConfiguredIpSource(IpSource::RightmostXForwardedFor));
338
339 let ClientIp(ip) = extract_client_ip(&mut parts).await.unwrap();
340 assert_eq!(ip.to_string(), "9.9.9.9");
341 }
342
343 #[tokio::test]
344 async fn ipv6_loopback_peer_also_bypasses_configured_source() {
345 let socket_addr = SocketAddr::from(([0_u16, 0, 0, 0, 0, 0, 0, 1], 38000));
346 let mut parts = parts(iter::empty());
347 parts.extensions.insert(ConnectInfo(socket_addr));
348 parts
349 .extensions
350 .insert(ConfiguredIpSource(IpSource::RightmostXForwardedFor));
351
352 let ClientIp(ip) = extract_client_ip(&mut parts).await.unwrap();
353 assert_eq!(ip, socket_addr.ip());
354 }
355
356 #[tokio::test]
357 async fn non_loopback_peer_with_configured_source_still_rejects() {
358 let socket_addr = SocketAddr::from(([203, 0, 113, 9], 38000));
359 let mut parts = parts(iter::empty());
360 parts.extensions.insert(ConnectInfo(socket_addr));
361 parts
362 .extensions
363 .insert(ConfiguredIpSource(IpSource::RightmostXForwardedFor));
364
365 let err = extract_client_ip(&mut parts).await.unwrap_err();
366 assert_eq!(err.0, StatusCode::INTERNAL_SERVER_ERROR);
367 assert_eq!(err.1, "Can't extract client IP from configured ip_source");
368 }
369
370 #[tokio::test]
371 async fn trusted_subnet_peer_bypasses_configured_source() {
372 let socket_addr = SocketAddr::from(([172, 18, 0, 5], 38000));
373 let mut parts = parts(iter::empty());
374 parts.extensions.insert(ConnectInfo(socket_addr));
375 parts
376 .extensions
377 .insert(ConfiguredIpSource(IpSource::RightmostXForwardedFor));
378 parts
379 .extensions
380 .insert(trusted(&["172.18.0.0/16"]));
381
382 let ClientIp(ip) = extract_client_ip(&mut parts).await.unwrap();
383 assert_eq!(ip, socket_addr.ip());
384 }
385
386 #[tokio::test]
387 async fn trusted_subnet_peer_with_proxy_header_uses_insecure_fallback() {
388 let socket_addr = SocketAddr::from(([172, 18, 0, 5], 38000));
389 let mut parts = parts([("X-Forwarded-For", "9.9.9.9")]);
390 parts.extensions.insert(ConnectInfo(socket_addr));
391 parts
392 .extensions
393 .insert(ConfiguredIpSource(IpSource::RightmostXForwardedFor));
394 parts
395 .extensions
396 .insert(trusted(&["172.18.0.0/16"]));
397
398 let ClientIp(ip) = extract_client_ip(&mut parts).await.unwrap();
399 assert_eq!(ip.to_string(), "9.9.9.9");
400 }
401
402 #[tokio::test]
403 async fn non_trusted_peer_with_subnets_configured_still_rejects() {
404 let socket_addr = SocketAddr::from(([203, 0, 113, 9], 38000));
405 let mut parts = parts(iter::empty());
406 parts.extensions.insert(ConnectInfo(socket_addr));
407 parts
408 .extensions
409 .insert(ConfiguredIpSource(IpSource::RightmostXForwardedFor));
410 parts
411 .extensions
412 .insert(trusted(&["172.18.0.0/16"]));
413
414 let err = extract_client_ip(&mut parts).await.unwrap_err();
415 assert_eq!(err.0, StatusCode::INTERNAL_SERVER_ERROR);
416 assert_eq!(err.1, "Can't extract client IP from configured ip_source");
417 }
418
419 #[tokio::test]
420 async fn ipv6_trusted_subnet_peer_bypasses_configured_source() {
421 let socket_addr = SocketAddr::from(([0xFD00_u16, 0, 0, 0, 0, 0, 0, 1], 38000));
422 let mut parts = parts(iter::empty());
423 parts.extensions.insert(ConnectInfo(socket_addr));
424 parts
425 .extensions
426 .insert(ConfiguredIpSource(IpSource::RightmostXForwardedFor));
427 parts.extensions.insert(trusted(&["fd00::/8"]));
428
429 let ClientIp(ip) = extract_client_ip(&mut parts).await.unwrap();
430 assert_eq!(ip, socket_addr.ip());
431 }
432
433 #[tokio::test]
434 async fn trusted_single_host_cidr_matches_only_that_address() {
435 let configured = ConfiguredIpSource(IpSource::RightmostXForwardedFor);
436
437 let mut listed = parts(iter::empty());
438 listed
439 .extensions
440 .insert(ConnectInfo(SocketAddr::from(([10, 0, 0, 5], 38000))));
441 listed.extensions.insert(configured);
442 listed
443 .extensions
444 .insert(trusted(&["10.0.0.5/32"]));
445
446 let ClientIp(ip) = extract_client_ip(&mut listed).await.unwrap();
447 assert_eq!(ip.to_string(), "10.0.0.5");
448
449 let mut neighbour = parts(iter::empty());
450 neighbour
451 .extensions
452 .insert(ConnectInfo(SocketAddr::from(([10, 0, 0, 6], 38000))));
453 neighbour.extensions.insert(configured);
454 neighbour
455 .extensions
456 .insert(trusted(&["10.0.0.5/32"]));
457
458 let err = extract_client_ip(&mut neighbour)
459 .await
460 .unwrap_err();
461 assert_eq!(err.0, StatusCode::INTERNAL_SERVER_ERROR);
462 }
463
464 #[tokio::test]
465 async fn loopback_still_bypasses_when_trusted_subnets_extension_absent() {
466 let socket_addr = SocketAddr::from(([127, 0, 0, 1], 38000));
467 let mut parts = parts(iter::empty());
468 parts.extensions.insert(ConnectInfo(socket_addr));
469 parts
470 .extensions
471 .insert(ConfiguredIpSource(IpSource::RightmostXForwardedFor));
472
473 let ClientIp(ip) = extract_client_ip(&mut parts).await.unwrap();
474 assert_eq!(ip, socket_addr.ip());
475 }
476}