1use std::convert::Infallible;
23use std::net::{IpAddr, SocketAddr};
24
25use anyhow::{Context, bail};
26use axum::extract::{ConnectInfo, FromRequestParts, Request};
27use axum::http::HeaderMap;
28use axum::http::request::Parts;
29
30#[derive(Debug, Clone, Copy, PartialEq, Eq)]
33pub struct ClientIp(pub Option<IpAddr>);
34
35impl<S: Send + Sync> FromRequestParts<S> for ClientIp {
36 type Rejection = Infallible;
37
38 async fn from_request_parts(parts: &mut Parts, _: &S) -> Result<Self, Infallible> {
39 Ok(match parts.extensions.get::<ClientIp>() {
40 Some(ip) => *ip,
41 None => Self(peer(&parts.extensions)),
42 })
43 }
44}
45
46impl ClientIp {
47 pub(crate) fn of(req: &Request) -> Option<IpAddr> {
49 match req.extensions().get::<ClientIp>() {
50 Some(ip) => ip.0,
51 None => peer(req.extensions()),
52 }
53 }
54}
55
56#[derive(Debug, Clone, Default, PartialEq, Eq)]
60pub struct TrustedProxies {
61 any: bool,
62 nets: Vec<(IpAddr, u8)>,
63}
64
65impl TrustedProxies {
66 pub fn parse(list: &str) -> crate::Result<Self> {
69 Ok(Self::read(list)?)
70 }
71
72 pub(crate) fn read(list: &str) -> anyhow::Result<Self> {
73 let mut proxies = Self::default();
74 for item in list.split(',').map(str::trim).filter(|s| !s.is_empty()) {
75 if item == "*" {
76 proxies.any = true;
77 continue;
78 }
79 let (addr, bits) = match item.split_once('/') {
80 Some((addr, bits)) => (addr, Some(bits)),
81 None => (item, None),
82 };
83 let addr: IpAddr = addr
84 .parse()
85 .with_context(|| format!("TRUSTED_PROXIES: `{item}` is not an IP address"))?;
86 let max = if addr.is_ipv4() { 32 } else { 128 };
87 let bits = match bits {
88 Some(bits) => bits
89 .parse::<u8>()
90 .ok()
91 .filter(|b| *b <= max)
92 .with_context(|| format!("TRUSTED_PROXIES: bad prefix length in `{item}`"))?,
93 None => max,
94 };
95 proxies.nets.push((addr.to_canonical(), bits));
96 }
97 if proxies.any && !proxies.nets.is_empty() {
98 bail!("TRUSTED_PROXIES: `*` already trusts every address");
99 }
100 Ok(proxies)
101 }
102
103 pub fn contains(&self, ip: IpAddr) -> bool {
105 self.any
106 || self
107 .nets
108 .iter()
109 .any(|(net, bits)| in_net(ip.to_canonical(), *net, *bits))
110 }
111
112 fn is_empty(&self) -> bool {
113 !self.any && self.nets.is_empty()
114 }
115}
116
117fn in_net(ip: IpAddr, net: IpAddr, bits: u8) -> bool {
118 match (ip, net) {
119 (IpAddr::V4(ip), IpAddr::V4(net)) => {
120 let mask = u32::MAX.checked_shl(32 - u32::from(bits)).unwrap_or(0);
121 u32::from(ip) & mask == u32::from(net) & mask
122 }
123 (IpAddr::V6(ip), IpAddr::V6(net)) => {
124 let mask = u128::MAX.checked_shl(128 - u32::from(bits)).unwrap_or(0);
125 u128::from(ip) & mask == u128::from(net) & mask
126 }
127 _ => false,
128 }
129}
130
131fn peer(extensions: &axum::http::Extensions) -> Option<IpAddr> {
132 extensions
133 .get::<ConnectInfo<SocketAddr>>()
134 .map(|info| info.0.ip().to_canonical())
135}
136
137pub(crate) fn resolve(req: &Request, trusted: &TrustedProxies) -> ClientIp {
142 let Some(peer) = peer(req.extensions()) else {
143 return ClientIp(None);
144 };
145 if trusted.is_empty() || !trusted.contains(peer) {
146 return ClientIp(Some(peer));
147 }
148 let mut client = peer;
149 for hop in forwarded_chain(req.headers()).iter().rev() {
150 let Some(ip) = hop else {
151 break;
153 };
154 client = *ip;
155 if trusted.any || !trusted.contains(client) {
157 break;
158 }
159 }
160 ClientIp(Some(client))
161}
162
163fn forwarded_chain(headers: &HeaderMap) -> Vec<Option<IpAddr>> {
166 let xff: Vec<Option<IpAddr>> = headers
167 .get_all("x-forwarded-for")
168 .iter()
169 .filter_map(|v| v.to_str().ok())
170 .flat_map(|v| v.split(','))
171 .map(parse_node)
172 .collect();
173 if !xff.is_empty() {
174 return xff;
175 }
176 headers
177 .get_all("forwarded")
178 .iter()
179 .filter_map(|v| v.to_str().ok())
180 .flat_map(|v| v.split(','))
181 .filter_map(|element| {
182 element.split(';').find_map(|pair| {
183 let (key, value) = pair.split_once('=')?;
184 key.trim()
185 .eq_ignore_ascii_case("for")
186 .then(|| parse_node(value))
187 })
188 })
189 .collect()
190}
191
192fn parse_node(node: &str) -> Option<IpAddr> {
194 let node = node.trim().trim_matches('"');
195 if let Ok(ip) = node.parse::<IpAddr>() {
196 return Some(ip.to_canonical());
197 }
198 if let Ok(addr) = node.parse::<SocketAddr>() {
199 return Some(addr.ip().to_canonical());
200 }
201 node.strip_prefix('[')
202 .and_then(|rest| rest.strip_suffix(']'))
203 .and_then(|ip| ip.parse::<IpAddr>().ok())
204 .map(|ip| ip.to_canonical())
205}
206
207#[cfg(test)]
208mod tests {
209 use axum::body::Body;
210
211 use super::*;
212
213 fn request(peer: &str, headers: &[(&str, &str)]) -> Request {
214 let mut builder = Request::builder().uri("/");
215 for (name, value) in headers {
216 builder = builder.header(*name, *value);
217 }
218 let mut req = builder.body(Body::empty()).unwrap();
219 let addr: SocketAddr = format!("{peer}:5555").parse().unwrap();
220 req.extensions_mut().insert(ConnectInfo(addr));
221 req
222 }
223
224 fn ip(s: &str) -> Option<IpAddr> {
225 Some(s.parse().unwrap())
226 }
227
228 #[test]
229 fn parses_addresses_ranges_and_any() {
230 let proxies = TrustedProxies::parse("127.0.0.1, 10.0.0.0/8,fd00::/8").unwrap();
231 assert!(proxies.contains("10.9.8.7".parse().unwrap()));
232 assert!(proxies.contains("::ffff:127.0.0.1".parse().unwrap()));
233 assert!(proxies.contains("fd12::1".parse().unwrap()));
234 assert!(!proxies.contains("11.0.0.1".parse().unwrap()));
235 assert!(
236 TrustedProxies::parse("*")
237 .unwrap()
238 .contains("8.8.8.8".parse().unwrap())
239 );
240 assert!(TrustedProxies::parse("").unwrap().is_empty());
241 assert!(TrustedProxies::parse("10.0.0.0/33").is_err());
242 assert!(TrustedProxies::parse("proxy.local").is_err());
243 let err = TrustedProxies::parse("*, 10.0.0.1").unwrap_err();
244 assert!(
245 format!("{err:?}").contains("`*` already trusts every address"),
246 "{err:?}"
247 );
248 }
249
250 #[tokio::test]
253 async fn without_the_middleware_the_peer_is_the_client() {
254 let req = request("203.0.113.9", &[("x-forwarded-for", "1.1.1.1")]);
255 assert_eq!(ClientIp::of(&req), ip("203.0.113.9"));
256 let (mut parts, _) = req.into_parts();
257 let found = ClientIp::from_request_parts(&mut parts, &()).await.unwrap();
258 assert_eq!(found.0, ip("203.0.113.9"));
259 }
260
261 #[test]
262 fn ignores_forwarded_headers_from_untrusted_peers() {
263 let trusted = TrustedProxies::parse("10.0.0.1").unwrap();
264 let req = request("203.0.113.9", &[("x-forwarded-for", "1.1.1.1")]);
265 assert_eq!(resolve(&req, &trusted).0, ip("203.0.113.9"));
266 let req = request("10.0.0.1", &[("x-forwarded-for", "1.1.1.1")]);
267 assert_eq!(resolve(&req, &TrustedProxies::default()).0, ip("10.0.0.1"));
268 }
269
270 #[test]
271 fn takes_the_rightmost_untrusted_hop() {
272 let trusted = TrustedProxies::parse("10.0.0.0/8").unwrap();
273 let req = request(
275 "10.0.0.1",
276 &[("x-forwarded-for", "6.6.6.6, 1.2.3.4, 10.0.0.2")],
277 );
278 assert_eq!(resolve(&req, &trusted).0, ip("1.2.3.4"));
279 let req = request("10.0.0.1", &[("x-forwarded-for", "10.0.0.3")]);
280 assert_eq!(resolve(&req, &trusted).0, ip("10.0.0.3"));
281 let req = request("10.0.0.1", &[("x-forwarded-for", "unknown")]);
282 assert_eq!(resolve(&req, &trusted).0, ip("10.0.0.1"));
283 }
284
285 #[test]
286 fn reads_the_forwarded_header() {
287 let trusted = TrustedProxies::parse("*").unwrap();
288 let req = request(
289 "10.0.0.1",
290 &[(
291 "forwarded",
292 r#"for=192.0.2.60;proto=https, for="[2001:db8::1]:4711""#,
293 )],
294 );
295 assert_eq!(resolve(&req, &trusted).0, ip("2001:db8::1"));
296 let req = request("10.0.0.1", &[("x-forwarded-for", "6.6.6.6, 1.2.3.4")]);
298 assert_eq!(resolve(&req, &trusted).0, ip("1.2.3.4"));
299 }
300}