http_extract/
x_forwarded.rs1use std::net::{IpAddr, SocketAddr};
12
13use http::{HeaderMap, HeaderName, Request};
14
15use crate::Error;
16
17pub const X_FORWARDED_FOR: HeaderName = HeaderName::from_static("x-forwarded-for");
19
20pub const X_FORWARDED_PROTO: HeaderName = HeaderName::from_static("x-forwarded-proto");
22
23pub fn extract_header_x_forwarded_for(headers: &HeaderMap) -> Result<Option<Vec<IpAddr>>, Error> {
34 extract_comma_values(headers, &X_FORWARDED_FOR, |value| {
35 parse_ip(value, X_FORWARDED_FOR)
36 })
37}
38
39pub fn extract_request_x_forwarded_for<B>(
45 request: &Request<B>,
46) -> Result<Option<Vec<IpAddr>>, Error> {
47 extract_header_x_forwarded_for(request.headers())
48}
49
50pub fn extract_header_x_forwarded_proto(headers: &HeaderMap) -> Result<Option<Vec<String>>, Error> {
58 extract_comma_values(headers, &X_FORWARDED_PROTO, |value| {
59 if is_scheme(value) {
60 Ok(value.to_ascii_lowercase())
61 } else {
62 Err(Error::invalid_header(X_FORWARDED_PROTO))
63 }
64 })
65}
66
67pub fn extract_request_x_forwarded_proto<B>(
74 request: &Request<B>,
75) -> Result<Option<Vec<String>>, Error> {
76 extract_header_x_forwarded_proto(request.headers())
77}
78
79pub fn extract_rightmost_x_forwarded_for(headers: &HeaderMap) -> Result<Option<IpAddr>, Error> {
85 Ok(extract_header_x_forwarded_for(headers)?.and_then(|ips| ips.last().copied()))
86}
87
88fn parse_ip(value: &str, name: HeaderName) -> Result<IpAddr, Error> {
90 if let Ok(address) = value.parse() {
91 return Ok(address);
92 }
93 if let Ok(address) = value.parse::<SocketAddr>() {
94 return Ok(address.ip());
95 }
96 Err(Error::invalid_header(name))
97}
98
99fn extract_comma_values<T>(
101 headers: &HeaderMap,
102 name: &HeaderName,
103 mut parse: impl FnMut(&str) -> Result<T, Error>,
104) -> Result<Option<Vec<T>>, Error> {
105 let mut output = Vec::new();
106 let mut present = false;
107 for value in headers.get_all(name) {
108 present = true;
109 let value = value
110 .to_str()
111 .map_err(|_| Error::invalid_header(name.clone()))?;
112 for item in value.split(',') {
113 let item = item.trim();
114 if item.is_empty() {
115 return Err(Error::invalid_header(name.clone()));
116 }
117 output.push(parse(item)?);
118 }
119 }
120 Ok(present.then_some(output))
121}
122
123fn is_scheme(value: &str) -> bool {
124 let mut bytes = value.bytes();
125 matches!(bytes.next(), Some(byte) if byte.is_ascii_alphabetic())
126 && bytes.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'+' | b'-' | b'.'))
127}
128
129#[cfg(test)]
130mod tests {
131 use std::net::{IpAddr, Ipv4Addr};
132
133 use http::{HeaderMap, HeaderValue};
134
135 use super::*;
136
137 #[test]
138 fn extracts_x_forwarded_for_across_field_lines() {
139 let mut headers = HeaderMap::new();
140 headers.append(&X_FORWARDED_FOR, "192.0.2.1, 198.51.100.2".parse().unwrap());
141 headers.append(&X_FORWARDED_FOR, "203.0.113.3".parse().unwrap());
142
143 assert_eq!(
144 extract_header_x_forwarded_for(&headers).unwrap().unwrap(),
145 vec![
146 IpAddr::V4(Ipv4Addr::new(192, 0, 2, 1)),
147 IpAddr::V4(Ipv4Addr::new(198, 51, 100, 2)),
148 IpAddr::V4(Ipv4Addr::new(203, 0, 113, 3)),
149 ]
150 );
151 }
152
153 #[test]
154 fn rejects_invalid_x_forwarded_for_values() {
155 let mut headers = HeaderMap::new();
156 for value in [
157 "[2001:db8::1]",
158 "[2001:db8::1]junk",
159 "[2001:db8::1]:65536",
160 "[2001:db8::1]:99999",
161 ] {
162 headers.insert(&X_FORWARDED_FOR, value.parse().unwrap());
163 assert!(
164 matches!(
165 extract_header_x_forwarded_for(&headers),
166 Err(Error::InvalidHeader { .. })
167 ),
168 "unexpectedly accepted {value:?}",
169 );
170 }
171
172 headers.insert(&X_FORWARDED_FOR, HeaderValue::from_bytes(&[0xff]).unwrap());
173 assert!(matches!(
174 extract_header_x_forwarded_for(&headers),
175 Err(Error::InvalidHeader { .. })
176 ));
177 }
178
179 #[test]
180 fn accepts_bracketed_ipv6_with_valid_port() {
181 let mut headers = HeaderMap::new();
182 headers.insert(&X_FORWARDED_FOR, "[2001:db8::1]:65535".parse().unwrap());
183 assert_eq!(
184 extract_header_x_forwarded_for(&headers).unwrap(),
185 Some(vec!["2001:db8::1".parse().unwrap()]),
186 );
187 }
188
189 #[test]
190 fn extracts_and_normalizes_x_forwarded_proto() {
191 let mut headers = HeaderMap::new();
192 assert_eq!(extract_header_x_forwarded_proto(&headers).unwrap(), None);
193
194 headers.append(&X_FORWARDED_PROTO, "HTTPS, Web+TLS".parse().unwrap());
195 assert_eq!(
196 extract_header_x_forwarded_proto(&headers).unwrap().unwrap(),
197 vec!["https".to_owned(), "web+tls".to_owned()]
198 );
199
200 headers.insert(&X_FORWARDED_PROTO, "http_2".parse().unwrap());
201 assert!(matches!(
202 extract_header_x_forwarded_proto(&headers),
203 Err(Error::InvalidHeader { .. })
204 ));
205 }
206
207 #[test]
208 fn request_entry_points_delegate_to_headers() {
209 let request = Request::builder()
210 .header(&X_FORWARDED_FOR, "192.0.2.1")
211 .header(&X_FORWARDED_PROTO, "HTTPS")
212 .body(())
213 .unwrap();
214 assert_eq!(
215 extract_request_x_forwarded_for(&request).unwrap(),
216 Some(vec!["192.0.2.1".parse().unwrap()])
217 );
218 assert_eq!(
219 extract_request_x_forwarded_proto(&request).unwrap(),
220 Some(vec!["https".to_owned()])
221 );
222 }
223}