1use std::net::IpAddr;
26
27use axum::http::{HeaderMap, HeaderName};
28use ipnet::IpNet;
29
30#[derive(Debug, Clone, Copy, PartialEq, Eq)]
37pub struct ClientIp(pub Option<IpAddr>);
38
39#[derive(Debug, Clone)]
41pub struct ProxyPolicy {
42 trusted: Vec<IpNet>,
43 header: HeaderName,
44}
45
46impl Default for ProxyPolicy {
47 fn default() -> Self {
49 Self {
50 trusted: Vec::new(),
51 header: HeaderName::from_static("x-forwarded-for"),
52 }
53 }
54}
55
56impl ProxyPolicy {
57 pub fn new(trusted_proxies: &[String], header: &str) -> anyhow::Result<Self> {
59 let trusted = parse_nets(trusted_proxies, "filter.trusted_proxies")?;
60 let header = HeaderName::try_from(header.to_ascii_lowercase())
61 .map_err(|error| anyhow::anyhow!("filter.forwarded_header: {error}"))?;
62 Ok(Self { trusted, header })
63 }
64
65 pub fn resolve(&self, peer: Option<IpAddr>, headers: &HeaderMap) -> Option<IpAddr> {
75 let peer = canonical(peer?);
76
77 if self.trusted.is_empty() || !nets_contain(&self.trusted, peer) {
78 return Some(peer);
79 }
80
81 let forwarded: Vec<IpAddr> = headers
82 .get_all(&self.header)
83 .iter()
84 .filter_map(|value| value.to_str().ok())
85 .flat_map(|value| value.split(','))
86 .filter_map(|entry| parse_forwarded_entry(entry.trim()))
87 .map(canonical)
88 .collect();
89
90 if let Some(client) = forwarded
92 .iter()
93 .rev()
94 .find(|ip| !nets_contain(&self.trusted, **ip))
95 {
96 return Some(*client);
97 }
98
99 Some(forwarded.first().copied().unwrap_or(peer))
102 }
103}
104
105fn parse_forwarded_entry(entry: &str) -> Option<IpAddr> {
111 if entry.is_empty() {
112 return None;
113 }
114
115 if let Some(rest) = entry.strip_prefix('[') {
117 let (host, _) = rest.split_once(']')?;
118 return host.parse().ok();
119 }
120
121 if let Ok(ip) = entry.parse::<IpAddr>() {
122 return Some(ip);
123 }
124
125 if entry.matches(':').count() == 1 {
127 let (host, _) = entry.split_once(':')?;
128 return host.parse().ok();
129 }
130
131 None
132}
133
134pub fn parse_net(entry: &str) -> anyhow::Result<IpNet> {
140 if let Ok(net) = entry.parse::<IpNet>() {
141 return Ok(net);
142 }
143 match entry.parse::<IpAddr>() {
144 Ok(addr) => Ok(IpNet::from(addr)),
145 Err(_) => anyhow::bail!("invalid network or address: {entry}"),
146 }
147}
148
149pub fn parse_nets(entries: &[String], setting: &str) -> anyhow::Result<Vec<IpNet>> {
151 entries
152 .iter()
153 .map(|entry| parse_net(entry).map_err(|error| anyhow::anyhow!("{setting}: {error}")))
154 .collect()
155}
156
157pub fn canonical(ip: IpAddr) -> IpAddr {
164 ip.to_canonical()
165}
166
167pub(crate) fn nets_contain(nets: &[IpNet], ip: IpAddr) -> bool {
169 let ip = canonical(ip);
170 nets.iter().any(|net| net.contains(&ip))
171}
172
173#[derive(Debug, Clone)]
175pub struct RequestId(pub String);
176
177#[cfg(test)]
178mod tests {
179 use super::*;
180
181 fn ip(value: &str) -> IpAddr {
182 value.parse().unwrap()
183 }
184
185 fn headers(pairs: &[(&str, &str)]) -> HeaderMap {
186 let mut map = HeaderMap::new();
187 for (name, value) in pairs {
188 map.append(HeaderName::try_from(*name).unwrap(), value.parse().unwrap());
189 }
190 map
191 }
192
193 fn policy(trusted: &[&str]) -> ProxyPolicy {
194 ProxyPolicy::new(
195 &trusted
196 .iter()
197 .map(std::string::ToString::to_string)
198 .collect::<Vec<_>>(),
199 "x-forwarded-for",
200 )
201 .unwrap()
202 }
203
204 #[test]
205 fn no_peer_means_no_client() {
206 assert_eq!(policy(&[]).resolve(None, &HeaderMap::new()), None);
207 }
208
209 #[test]
210 fn peer_is_used_when_no_proxy_is_trusted() {
211 let resolved = policy(&[]).resolve(
212 Some(ip("203.0.113.9")),
213 &headers(&[("x-forwarded-for", "10.0.0.1")]),
214 );
215 assert_eq!(resolved, Some(ip("203.0.113.9")));
216 }
217
218 #[test]
219 fn forwarded_header_is_ignored_from_an_untrusted_peer() {
220 let resolved = policy(&["127.0.0.1"]).resolve(
222 Some(ip("203.0.113.9")),
223 &headers(&[("x-forwarded-for", "10.0.0.1")]),
224 );
225 assert_eq!(resolved, Some(ip("203.0.113.9")));
226 }
227
228 #[test]
229 fn forwarded_header_is_honoured_from_a_trusted_peer() {
230 let resolved = policy(&["127.0.0.1"]).resolve(
231 Some(ip("127.0.0.1")),
232 &headers(&[("x-forwarded-for", "203.0.113.9")]),
233 );
234 assert_eq!(resolved, Some(ip("203.0.113.9")));
235 }
236
237 #[test]
238 fn chained_proxies_are_walked_past() {
239 let resolved = policy(&["127.0.0.1", "10.0.0.0/8"]).resolve(
241 Some(ip("127.0.0.1")),
242 &headers(&[("x-forwarded-for", "203.0.113.9, 10.0.0.8")]),
243 );
244 assert_eq!(resolved, Some(ip("203.0.113.9")));
245 }
246
247 #[test]
248 fn repeated_header_lines_are_concatenated() {
249 let resolved = policy(&["127.0.0.1", "10.0.0.0/8"]).resolve(
250 Some(ip("127.0.0.1")),
251 &headers(&[
252 ("x-forwarded-for", "203.0.113.9"),
253 ("x-forwarded-for", "10.0.0.8"),
254 ]),
255 );
256 assert_eq!(resolved, Some(ip("203.0.113.9")));
257 }
258
259 #[test]
260 fn all_hops_trusted_falls_back_to_the_leftmost() {
261 let resolved = policy(&["127.0.0.1", "10.0.0.0/8"]).resolve(
262 Some(ip("127.0.0.1")),
263 &headers(&[("x-forwarded-for", "10.0.0.4, 10.0.0.8")]),
264 );
265 assert_eq!(resolved, Some(ip("10.0.0.4")));
266 }
267
268 #[test]
269 fn trusted_peer_without_a_usable_header_falls_back_to_the_peer() {
270 let trusted = policy(&["127.0.0.1"]);
271 assert_eq!(
272 trusted.resolve(Some(ip("127.0.0.1")), &HeaderMap::new()),
273 Some(ip("127.0.0.1"))
274 );
275 assert_eq!(
276 trusted.resolve(
277 Some(ip("127.0.0.1")),
278 &headers(&[("x-forwarded-for", "not-an-ip, also-garbage")])
279 ),
280 Some(ip("127.0.0.1"))
281 );
282 }
283
284 #[test]
285 fn ipv4_mapped_peer_is_canonicalized() {
286 let resolved = policy(&[]).resolve(Some(ip("::ffff:192.168.1.5")), &HeaderMap::new());
288 assert_eq!(resolved, Some(ip("192.168.1.5")));
289 }
290
291 #[test]
292 fn ipv4_mapped_peer_matches_a_v4_trusted_proxy() {
293 let resolved = policy(&["127.0.0.1"]).resolve(
294 Some(ip("::ffff:127.0.0.1")),
295 &headers(&[("x-forwarded-for", "203.0.113.9")]),
296 );
297 assert_eq!(resolved, Some(ip("203.0.113.9")));
298 }
299
300 #[test]
301 fn forwarded_entries_may_carry_ports() {
302 assert_eq!(parse_forwarded_entry("1.2.3.4"), Some(ip("1.2.3.4")));
303 assert_eq!(parse_forwarded_entry("1.2.3.4:8443"), Some(ip("1.2.3.4")));
304 assert_eq!(
305 parse_forwarded_entry("2001:db8::1"),
306 Some(ip("2001:db8::1"))
307 );
308 assert_eq!(
309 parse_forwarded_entry("[2001:db8::1]:443"),
310 Some(ip("2001:db8::1"))
311 );
312 assert_eq!(
313 parse_forwarded_entry("[2001:db8::1]"),
314 Some(ip("2001:db8::1"))
315 );
316 assert_eq!(parse_forwarded_entry(""), None);
317 assert_eq!(parse_forwarded_entry("unknown"), None);
318 assert_eq!(parse_forwarded_entry("[bad"), None);
319 assert_eq!(parse_forwarded_entry("1.2.3.4:notaport:x"), None);
320 }
321
322 #[test]
323 fn new_validates_its_inputs() {
324 assert!(ProxyPolicy::new(&["10.0.0.0/8".to_string()], "x-real-ip").is_ok());
325 assert!(ProxyPolicy::new(&["nope".to_string()], "x-forwarded-for").is_err());
326 assert!(ProxyPolicy::new(&[], "not a header").is_err());
327 }
328
329 #[test]
330 fn a_custom_header_name_is_used() {
331 let policy = ProxyPolicy::new(&["127.0.0.1".to_string()], "X-Real-IP").unwrap();
332 let resolved = policy.resolve(
333 Some(ip("127.0.0.1")),
334 &headers(&[("x-real-ip", "203.0.113.9")]),
335 );
336 assert_eq!(resolved, Some(ip("203.0.113.9")));
337 }
338
339 #[test]
340 fn parse_net_accepts_cidr_and_bare_addresses() {
341 assert!(parse_net("192.168.1.0/24").is_ok());
342 assert!(parse_net("fd00::/8").is_ok());
343
344 let host = parse_net("203.0.113.7").unwrap();
347 assert_eq!(host.prefix_len(), 32);
348 assert!(host.contains(&"203.0.113.7".parse::<IpAddr>().unwrap()));
349 assert!(!host.contains(&"203.0.113.8".parse::<IpAddr>().unwrap()));
350
351 let host6 = parse_net("2001:db8::1").unwrap();
352 assert_eq!(host6.prefix_len(), 128);
353
354 assert!(parse_net("not-a-network").is_err());
355 assert!(parse_net("192.168.1.0/99").is_err());
356 }
357
358 #[test]
359 fn nets_contain_canonicalizes_ipv4_mapped_addresses() {
360 let nets = parse_nets(&["192.168.1.0/24".to_string()], "test").unwrap();
361 assert!(nets_contain(&nets, "192.168.1.5".parse().unwrap()));
362 assert!(nets_contain(&nets, "::ffff:192.168.1.5".parse().unwrap()));
363 assert!(!nets_contain(&nets, "10.0.0.1".parse().unwrap()));
364 }
365
366 #[test]
367 fn parse_nets_names_the_offending_setting() {
368 let error = parse_nets(&["garbage".to_string()], "filter.check.net.allow")
369 .unwrap_err()
370 .to_string();
371 assert!(error.contains("filter.check.net.allow"), "{error}");
372 assert!(error.contains("garbage"), "{error}");
373 }
374
375 #[test]
376 fn canonical_unmaps_ipv4_in_ipv6() {
377 assert_eq!(
378 canonical("::ffff:10.0.0.1".parse().unwrap()),
379 "10.0.0.1".parse::<IpAddr>().unwrap()
380 );
381 }
382}