Skip to main content

renox_core/
client_ip.rs

1//! The visitor's IP address, also behind a reverse proxy.
2//!
3//! Behind Caddy, nginx or a load balancer every connection comes from the
4//! proxy, which passes the visitor's address in `X-Forwarded-For` or
5//! `Forwarded`. Those headers are believed only from the addresses in
6//! `TRUSTED_PROXIES`, since anyone can send them:
7//!
8//! ```text
9//! TRUSTED_PROXIES=127.0.0.1,10.0.0.0/8   # Caddy on the same host, a private network
10//! TRUSTED_PROXIES=*                      # whoever connects, e.g. a platform's load balancer
11//! ```
12//!
13//! Rate limits, the login lock and [`ClientIp`] all use the same address.
14//!
15//! ```
16//! # use renox::prelude::*;
17//! async fn whoami(ClientIp(ip): ClientIp) -> String {
18//!     ip.map_or("unknown".into(), |ip| ip.to_string())
19//! }
20//! ```
21
22use 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/// The client's IP address, when the server knows it: the connection's
31/// address, or the one a trusted proxy forwarded.
32#[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    /// The address of a request that went through the app's middleware.
48    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/// Proxies whose `X-Forwarded-For` and `Forwarded` headers are believed,
57/// from `TRUSTED_PROXIES`: addresses and CIDR ranges separated by commas, or
58/// `*` to believe whoever connects, for the last hop only.
59#[derive(Debug, Clone, Default, PartialEq, Eq)]
60pub struct TrustedProxies {
61    any: bool,
62    nets: Vec<(IpAddr, u8)>,
63}
64
65impl TrustedProxies {
66    /// Parses a `TRUSTED_PROXIES` list; fails on a bad address or prefix, or `*`
67    /// mixed with addresses.
68    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    /// Whether `ip` is a trusted proxy.
104    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
137/// The client's address: the connection's, unless it is a trusted proxy.
138/// Then the forwarded chain is walked from the right (the hop nearest to
139/// us), skipping trusted proxies, since only the entries they appended can
140/// be believed.
141pub(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            // Garbage the previous hop didn't vouch for; stop at that hop.
152            break;
153        };
154        client = *ip;
155        // `*` vouches for the connecting proxy only: one hop.
156        if trusted.any || !trusted.contains(client) {
157            break;
158        }
159    }
160    ClientIp(Some(client))
161}
162
163/// Addresses in `X-Forwarded-For`, or else in `Forwarded`, left to right;
164/// `None` where an entry is not an address (`unknown`, an obfuscated name).
165fn 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
192/// `1.2.3.4`, `1.2.3.4:80`, `"[2001:db8::1]:4711"` or `2001:db8::1`.
193fn 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    /// Outside the app's middleware (no resolved address yet), the
251    /// connection's address is the client's.
252    #[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        // The visitor made up 6.6.6.6; the proxy appended the real 1.2.3.4.
274        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        // `*` reads one hop, so a made-up first entry doesn't count.
297        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}