1use std::net::IpAddr;
20
21use axum::http::{HeaderMap, HeaderName};
22use ipnet::IpNet;
23
24use super::{canonical, nets_contain, parse_nets};
25
26#[derive(Debug, Clone, Copy, PartialEq, Eq)]
33pub struct ClientIp(pub Option<IpAddr>);
34
35#[derive(Debug, Clone)]
37pub struct ProxyPolicy {
38 trusted: Vec<IpNet>,
39 header: HeaderName,
40}
41
42impl Default for ProxyPolicy {
43 fn default() -> Self {
45 Self {
46 trusted: Vec::new(),
47 header: HeaderName::from_static("x-forwarded-for"),
48 }
49 }
50}
51
52impl ProxyPolicy {
53 pub fn new(trusted_proxies: &[String], header: &str) -> anyhow::Result<Self> {
55 let trusted = parse_nets(trusted_proxies, "filter.trusted_proxies")?;
56 let header = HeaderName::try_from(header.to_ascii_lowercase())
57 .map_err(|error| anyhow::anyhow!("filter.forwarded_header: {error}"))?;
58 Ok(Self { trusted, header })
59 }
60
61 pub fn resolve(&self, peer: Option<IpAddr>, headers: &HeaderMap) -> Option<IpAddr> {
71 let peer = canonical(peer?);
72
73 if self.trusted.is_empty() || !nets_contain(&self.trusted, peer) {
74 return Some(peer);
75 }
76
77 let forwarded: Vec<IpAddr> = headers
78 .get_all(&self.header)
79 .iter()
80 .filter_map(|value| value.to_str().ok())
81 .flat_map(|value| value.split(','))
82 .filter_map(|entry| parse_forwarded_entry(entry.trim()))
83 .map(canonical)
84 .collect();
85
86 if let Some(client) = forwarded
88 .iter()
89 .rev()
90 .find(|ip| !nets_contain(&self.trusted, **ip))
91 {
92 return Some(*client);
93 }
94
95 Some(forwarded.first().copied().unwrap_or(peer))
98 }
99}
100
101fn parse_forwarded_entry(entry: &str) -> Option<IpAddr> {
107 if entry.is_empty() {
108 return None;
109 }
110
111 if let Some(rest) = entry.strip_prefix('[') {
113 let (host, _) = rest.split_once(']')?;
114 return host.parse().ok();
115 }
116
117 if let Ok(ip) = entry.parse::<IpAddr>() {
118 return Some(ip);
119 }
120
121 if entry.matches(':').count() == 1 {
123 let (host, _) = entry.split_once(':')?;
124 return host.parse().ok();
125 }
126
127 None
128}
129
130#[cfg(test)]
131mod tests {
132 use super::*;
133
134 fn ip(value: &str) -> IpAddr {
135 value.parse().unwrap()
136 }
137
138 fn headers(pairs: &[(&str, &str)]) -> HeaderMap {
139 let mut map = HeaderMap::new();
140 for (name, value) in pairs {
141 map.append(HeaderName::try_from(*name).unwrap(), value.parse().unwrap());
142 }
143 map
144 }
145
146 fn policy(trusted: &[&str]) -> ProxyPolicy {
147 ProxyPolicy::new(
148 &trusted
149 .iter()
150 .map(std::string::ToString::to_string)
151 .collect::<Vec<_>>(),
152 "x-forwarded-for",
153 )
154 .unwrap()
155 }
156
157 #[test]
158 fn no_peer_means_no_client() {
159 assert_eq!(policy(&[]).resolve(None, &HeaderMap::new()), None);
160 }
161
162 #[test]
163 fn peer_is_used_when_no_proxy_is_trusted() {
164 let resolved = policy(&[]).resolve(
165 Some(ip("203.0.113.9")),
166 &headers(&[("x-forwarded-for", "10.0.0.1")]),
167 );
168 assert_eq!(resolved, Some(ip("203.0.113.9")));
169 }
170
171 #[test]
172 fn forwarded_header_is_ignored_from_an_untrusted_peer() {
173 let resolved = policy(&["127.0.0.1"]).resolve(
175 Some(ip("203.0.113.9")),
176 &headers(&[("x-forwarded-for", "10.0.0.1")]),
177 );
178 assert_eq!(resolved, Some(ip("203.0.113.9")));
179 }
180
181 #[test]
182 fn forwarded_header_is_honoured_from_a_trusted_peer() {
183 let resolved = policy(&["127.0.0.1"]).resolve(
184 Some(ip("127.0.0.1")),
185 &headers(&[("x-forwarded-for", "203.0.113.9")]),
186 );
187 assert_eq!(resolved, Some(ip("203.0.113.9")));
188 }
189
190 #[test]
191 fn chained_proxies_are_walked_past() {
192 let resolved = policy(&["127.0.0.1", "10.0.0.0/8"]).resolve(
194 Some(ip("127.0.0.1")),
195 &headers(&[("x-forwarded-for", "203.0.113.9, 10.0.0.8")]),
196 );
197 assert_eq!(resolved, Some(ip("203.0.113.9")));
198 }
199
200 #[test]
201 fn repeated_header_lines_are_concatenated() {
202 let resolved = policy(&["127.0.0.1", "10.0.0.0/8"]).resolve(
203 Some(ip("127.0.0.1")),
204 &headers(&[
205 ("x-forwarded-for", "203.0.113.9"),
206 ("x-forwarded-for", "10.0.0.8"),
207 ]),
208 );
209 assert_eq!(resolved, Some(ip("203.0.113.9")));
210 }
211
212 #[test]
213 fn all_hops_trusted_falls_back_to_the_leftmost() {
214 let resolved = policy(&["127.0.0.1", "10.0.0.0/8"]).resolve(
215 Some(ip("127.0.0.1")),
216 &headers(&[("x-forwarded-for", "10.0.0.4, 10.0.0.8")]),
217 );
218 assert_eq!(resolved, Some(ip("10.0.0.4")));
219 }
220
221 #[test]
222 fn trusted_peer_without_a_usable_header_falls_back_to_the_peer() {
223 let trusted = policy(&["127.0.0.1"]);
224 assert_eq!(
225 trusted.resolve(Some(ip("127.0.0.1")), &HeaderMap::new()),
226 Some(ip("127.0.0.1"))
227 );
228 assert_eq!(
229 trusted.resolve(
230 Some(ip("127.0.0.1")),
231 &headers(&[("x-forwarded-for", "not-an-ip, also-garbage")])
232 ),
233 Some(ip("127.0.0.1"))
234 );
235 }
236
237 #[test]
238 fn ipv4_mapped_peer_is_canonicalized() {
239 let resolved = policy(&[]).resolve(Some(ip("::ffff:192.168.1.5")), &HeaderMap::new());
241 assert_eq!(resolved, Some(ip("192.168.1.5")));
242 }
243
244 #[test]
245 fn ipv4_mapped_peer_matches_a_v4_trusted_proxy() {
246 let resolved = policy(&["127.0.0.1"]).resolve(
247 Some(ip("::ffff:127.0.0.1")),
248 &headers(&[("x-forwarded-for", "203.0.113.9")]),
249 );
250 assert_eq!(resolved, Some(ip("203.0.113.9")));
251 }
252
253 #[test]
254 fn forwarded_entries_may_carry_ports() {
255 assert_eq!(parse_forwarded_entry("1.2.3.4"), Some(ip("1.2.3.4")));
256 assert_eq!(parse_forwarded_entry("1.2.3.4:8443"), Some(ip("1.2.3.4")));
257 assert_eq!(
258 parse_forwarded_entry("2001:db8::1"),
259 Some(ip("2001:db8::1"))
260 );
261 assert_eq!(
262 parse_forwarded_entry("[2001:db8::1]:443"),
263 Some(ip("2001:db8::1"))
264 );
265 assert_eq!(
266 parse_forwarded_entry("[2001:db8::1]"),
267 Some(ip("2001:db8::1"))
268 );
269 assert_eq!(parse_forwarded_entry(""), None);
270 assert_eq!(parse_forwarded_entry("unknown"), None);
271 assert_eq!(parse_forwarded_entry("[bad"), None);
272 assert_eq!(parse_forwarded_entry("1.2.3.4:notaport:x"), None);
273 }
274
275 #[test]
276 fn new_validates_its_inputs() {
277 assert!(ProxyPolicy::new(&["10.0.0.0/8".to_string()], "x-real-ip").is_ok());
278 assert!(ProxyPolicy::new(&["nope".to_string()], "x-forwarded-for").is_err());
279 assert!(ProxyPolicy::new(&[], "not a header").is_err());
280 }
281
282 #[test]
283 fn a_custom_header_name_is_used() {
284 let policy = ProxyPolicy::new(&["127.0.0.1".to_string()], "X-Real-IP").unwrap();
285 let resolved = policy.resolve(
286 Some(ip("127.0.0.1")),
287 &headers(&[("x-real-ip", "203.0.113.9")]),
288 );
289 assert_eq!(resolved, Some(ip("203.0.113.9")));
290 }
291}