1use std::net::{IpAddr, SocketAddr};
17
18use http::{HeaderMap, HeaderName, Request};
19
20use crate::{Error, header::extract_single_header_text};
21
22pub(crate) const CF_CONNECTING_IP: HeaderName = HeaderName::from_static("cf-connecting-ip");
23pub(crate) const CLOUDFRONT_VIEWER_ADDRESS: HeaderName =
24 HeaderName::from_static("cloudfront-viewer-address");
25pub(crate) const FLY_CLIENT_IP: HeaderName = HeaderName::from_static("fly-client-ip");
26pub(crate) const TRUE_CLIENT_IP: HeaderName = HeaderName::from_static("true-client-ip");
27pub(crate) const X_ENVOY_EXTERNAL_ADDRESS: HeaderName =
28 HeaderName::from_static("x-envoy-external-address");
29pub(crate) const X_REAL_IP: HeaderName = HeaderName::from_static("x-real-ip");
30
31pub fn extract_header_cf_connecting_ip(headers: &HeaderMap) -> Result<Option<IpAddr>, Error> {
38 extract_single_ip(headers, &CF_CONNECTING_IP)
39}
40
41pub fn extract_request_cf_connecting_ip<B>(request: &Request<B>) -> Result<Option<IpAddr>, Error> {
47 extract_header_cf_connecting_ip(request.headers())
48}
49
50pub fn extract_header_cloudfront_viewer_address(
60 headers: &HeaderMap,
61) -> Result<Option<IpAddr>, Error> {
62 extract_single_header_text(headers, &CLOUDFRONT_VIEWER_ADDRESS)?
63 .map(parse_cloudfront_viewer_address)
64 .transpose()
65}
66
67pub fn extract_request_cloudfront_viewer_address<B>(
74 request: &Request<B>,
75) -> Result<Option<IpAddr>, Error> {
76 extract_header_cloudfront_viewer_address(request.headers())
77}
78
79pub fn extract_header_fly_client_ip(headers: &HeaderMap) -> Result<Option<IpAddr>, Error> {
86 extract_single_ip(headers, &FLY_CLIENT_IP)
87}
88
89pub fn extract_request_fly_client_ip<B>(request: &Request<B>) -> Result<Option<IpAddr>, Error> {
95 extract_header_fly_client_ip(request.headers())
96}
97
98pub fn extract_header_true_client_ip(headers: &HeaderMap) -> Result<Option<IpAddr>, Error> {
106 extract_single_ip(headers, &TRUE_CLIENT_IP)
107}
108
109pub fn extract_request_true_client_ip<B>(request: &Request<B>) -> Result<Option<IpAddr>, Error> {
115 extract_header_true_client_ip(request.headers())
116}
117
118pub fn extract_header_x_envoy_external_address(
127 headers: &HeaderMap,
128) -> Result<Option<IpAddr>, Error> {
129 extract_single_ip(headers, &X_ENVOY_EXTERNAL_ADDRESS)
130}
131
132pub fn extract_request_x_envoy_external_address<B>(
138 request: &Request<B>,
139) -> Result<Option<IpAddr>, Error> {
140 extract_header_x_envoy_external_address(request.headers())
141}
142
143pub fn extract_header_x_real_ip(headers: &HeaderMap) -> Result<Option<IpAddr>, Error> {
150 extract_single_ip(headers, &X_REAL_IP)
151}
152
153pub fn extract_request_x_real_ip<B>(request: &Request<B>) -> Result<Option<IpAddr>, Error> {
159 extract_header_x_real_ip(request.headers())
160}
161
162fn extract_single_ip(headers: &HeaderMap, name: &HeaderName) -> Result<Option<IpAddr>, Error> {
163 extract_single_header_text(headers, name)?
164 .map(|value| parse_ip(value, name))
165 .transpose()
166}
167
168fn parse_ip(value: &str, name: &HeaderName) -> Result<IpAddr, Error> {
169 value
170 .trim()
171 .parse()
172 .map_err(|_| Error::invalid_header(name.clone()))
173}
174
175fn parse_cloudfront_viewer_address(value: &str) -> Result<IpAddr, Error> {
176 let value = value.trim();
177 if let Ok(address) = value.parse::<SocketAddr>() {
178 return Ok(address.ip());
179 }
180
181 let (address, port) = value
182 .rsplit_once(':')
183 .ok_or_else(|| Error::invalid_header(CLOUDFRONT_VIEWER_ADDRESS))?;
184 port.parse::<u16>()
185 .map_err(|_| Error::invalid_header(CLOUDFRONT_VIEWER_ADDRESS))?;
186 address
187 .parse()
188 .map_err(|_| Error::invalid_header(CLOUDFRONT_VIEWER_ADDRESS))
189}
190
191#[cfg(test)]
192mod tests {
193 use http::{HeaderMap, HeaderValue};
194
195 use super::*;
196
197 type Extractor = fn(&HeaderMap) -> Result<Option<IpAddr>, Error>;
198
199 fn ordinary_extractors() -> [(HeaderName, Extractor); 5] {
200 [
201 (CF_CONNECTING_IP, extract_header_cf_connecting_ip),
202 (FLY_CLIENT_IP, extract_header_fly_client_ip),
203 (TRUE_CLIENT_IP, extract_header_true_client_ip),
204 (
205 X_ENVOY_EXTERNAL_ADDRESS,
206 extract_header_x_envoy_external_address,
207 ),
208 (X_REAL_IP, extract_header_x_real_ip),
209 ]
210 }
211
212 #[test]
213 fn single_ip_fields_handle_missing_valid_and_invalid_values() {
214 for (name, extract) in ordinary_extractors() {
215 let mut headers = HeaderMap::new();
216 assert_eq!(extract(&headers).unwrap(), None);
217
218 headers.insert(&name, " 2001:db8::1 ".parse().unwrap());
219 assert_eq!(
220 extract(&headers).unwrap(),
221 Some("2001:db8::1".parse().unwrap())
222 );
223
224 headers.insert(&name, "not-an-ip".parse().unwrap());
225 let error = extract(&headers).unwrap_err();
226 assert!(matches!(error, Error::InvalidHeader { .. }));
227 assert!(!error.to_string().contains("not-an-ip"));
228 }
229 }
230
231 #[test]
232 fn every_single_ip_field_rejects_duplicates() {
233 for (name, extract) in ordinary_extractors() {
234 let mut headers = HeaderMap::new();
235 headers.append(&name, "192.0.2.1".parse().unwrap());
236 headers.append(&name, "198.51.100.2".parse().unwrap());
237 assert!(matches!(
238 extract(&headers),
239 Err(Error::DuplicateHeader { .. })
240 ));
241 }
242 }
243
244 #[test]
245 fn cloudfront_viewer_address_handles_ipv4_and_ipv6_with_ports() {
246 let mut headers = HeaderMap::new();
247 assert_eq!(
248 extract_header_cloudfront_viewer_address(&headers).unwrap(),
249 None
250 );
251
252 for (value, expected) in [
253 ("198.51.100.10:46532", "198.51.100.10"),
254 ("2001:db8::abcd:1234", "2001:db8::abcd"),
255 ("[2001:db8::17]:4711", "2001:db8::17"),
256 ] {
257 headers.insert(&CLOUDFRONT_VIEWER_ADDRESS, value.parse().unwrap());
258 assert_eq!(
259 extract_header_cloudfront_viewer_address(&headers).unwrap(),
260 Some(expected.parse().unwrap())
261 );
262 }
263
264 for value in ["198.51.100.10", "198.51.100.10:not-a-port", "not-an-ip:80"] {
265 headers.insert(&CLOUDFRONT_VIEWER_ADDRESS, value.parse().unwrap());
266 let error = extract_header_cloudfront_viewer_address(&headers).unwrap_err();
267 assert!(matches!(error, Error::InvalidHeader { .. }));
268 assert!(!error.to_string().contains(value));
269 }
270
271 headers.clear();
272 headers.append(
273 &CLOUDFRONT_VIEWER_ADDRESS,
274 "198.51.100.10:443".parse().unwrap(),
275 );
276 headers.append(
277 &CLOUDFRONT_VIEWER_ADDRESS,
278 "198.51.100.11:443".parse().unwrap(),
279 );
280 assert!(matches!(
281 extract_header_cloudfront_viewer_address(&headers),
282 Err(Error::DuplicateHeader { .. })
283 ));
284 }
285
286 #[test]
287 fn all_fields_reject_non_text_without_echoing_values() {
288 let mut cases = ordinary_extractors().to_vec();
289 cases.push((
290 CLOUDFRONT_VIEWER_ADDRESS,
291 extract_header_cloudfront_viewer_address,
292 ));
293
294 for (name, extract) in cases {
295 let mut headers = HeaderMap::new();
296 headers.insert(&name, HeaderValue::from_bytes(&[0xff]).unwrap());
297 let error = extract(&headers).unwrap_err();
298 assert!(matches!(error, Error::InvalidHeader { .. }));
299 assert!(!error.to_string().contains("255"));
300 }
301 }
302
303 #[test]
304 fn request_entry_points_delegate_to_header_extractors() {
305 let request = Request::builder()
306 .header(&CF_CONNECTING_IP, "192.0.2.1")
307 .header(&CLOUDFRONT_VIEWER_ADDRESS, "198.51.100.2:443")
308 .header(&FLY_CLIENT_IP, "203.0.113.3")
309 .header(&TRUE_CLIENT_IP, "192.0.2.4")
310 .header(&X_ENVOY_EXTERNAL_ADDRESS, "198.51.100.5")
311 .header(&X_REAL_IP, "203.0.113.6")
312 .body(())
313 .unwrap();
314
315 assert_eq!(
316 extract_request_cf_connecting_ip(&request).unwrap(),
317 Some("192.0.2.1".parse().unwrap())
318 );
319 assert_eq!(
320 extract_request_cloudfront_viewer_address(&request).unwrap(),
321 Some("198.51.100.2".parse().unwrap())
322 );
323 assert_eq!(
324 extract_request_fly_client_ip(&request).unwrap(),
325 Some("203.0.113.3".parse().unwrap())
326 );
327 assert_eq!(
328 extract_request_true_client_ip(&request).unwrap(),
329 Some("192.0.2.4".parse().unwrap())
330 );
331 assert_eq!(
332 extract_request_x_envoy_external_address(&request).unwrap(),
333 Some("198.51.100.5".parse().unwrap())
334 );
335 assert_eq!(
336 extract_request_x_real_ip(&request).unwrap(),
337 Some("203.0.113.6".parse().unwrap())
338 );
339 }
340}