1use std::net::SocketAddr;
4use url::Url;
5
6pub fn is_private_ip(ip: &str) -> bool {
20 if let Ok(addr) = ip.parse::<std::net::IpAddr>() {
21 match addr {
22 std::net::IpAddr::V4(ipv4) => {
23 ipv4.is_private()
24 || ipv4.is_loopback()
25 || ipv4.is_link_local()
26 || ipv4.is_unspecified()
27 || ipv4.is_multicast()
28 }
29 std::net::IpAddr::V6(ipv6) => {
30 ipv6.is_loopback()
31 || ipv6.is_unspecified()
32 || ipv6.is_unicast_link_local()
33 || ipv6.is_unique_local()
34 || ipv6.is_multicast()
35 }
36 }
37 } else {
38 false
39 }
40}
41
42pub fn extract_client_ip(headers: &http::HeaderMap, connect_info: &SocketAddr) -> String {
63 fn is_valid_ip(s: &str) -> bool {
64 s.trim().parse::<std::net::IpAddr>().is_ok()
65 }
66
67 if let Some(ip) = headers
69 .get("X-Real-IP")
70 .and_then(|h| h.to_str().ok())
71 .filter(|s| is_valid_ip(s))
72 {
73 return ip.trim().to_string();
74 }
75
76 if let Some(ip) = headers
78 .get("X-Forwarded-For")
79 .and_then(|h| h.to_str().ok())
80 .and_then(|h| h.split(',').next_back())
81 .filter(|s| is_valid_ip(s))
82 {
83 return ip.trim().to_string();
84 }
85
86 let connect_ip = connect_info.ip().to_string();
88 if !is_private_ip(&connect_ip) {
89 return connect_ip;
90 }
91
92 "0.0.0.0".to_string()
93}
94
95pub struct WebExt;
99
100impl WebExt {
101 pub fn domain(url_str: &str) -> Option<String> {
103 Url::parse(url_str)
104 .ok()
105 .and_then(|u| u.host_str().map(|s| s.to_string()))
106 }
107
108 pub fn path(url_str: &str) -> Option<String> {
110 Url::parse(url_str).ok().map(|u| u.path().to_string())
111 }
112
113 pub fn real_ip(headers: &[(String, String)], remote_addr: &str) -> String {
115 for (key, val) in headers {
116 if key.to_lowercase() == "x-forwarded-for" {
117 return val.split(',').next().unwrap_or("").trim().to_string();
118 }
119 if key.to_lowercase() == "x-real-ip" {
120 return val.clone();
121 }
122 }
123 remote_addr
124 .split(':')
125 .next()
126 .unwrap_or(remote_addr)
127 .to_string()
128 }
129 pub fn is_private_ip(ip: &str) -> bool {
131 is_private_ip(ip)
132 }
133 pub fn build_query(params: &[(&str, &str)]) -> String {
135 if params.is_empty() {
136 return String::new();
137 }
138 let parts: Vec<String> = params
139 .iter()
140 .map(|(k, v)| format!("{}={}", urlencoding(k), urlencoding(v)))
141 .collect();
142 format!("?{}", parts.join("&"))
143 }
144}
145
146fn urlencoding(s: &str) -> String {
147 let mut result = String::with_capacity(s.len() * 3);
148 for byte in s.bytes() {
149 match byte {
150 b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'_' | b'.' | b'~' => {
151 result.push(byte as char);
152 }
153 _ => {
154 result.push_str(&format!("%{:02X}", byte));
155 }
156 }
157 }
158 result
159}
160
161#[cfg(test)]
162mod tests {
163 use super::*;
164 use http::{HeaderMap, HeaderName, HeaderValue};
165 use std::net::SocketAddr;
166
167 fn header_maps(entries: &[(&str, &str)]) -> HeaderMap {
169 let mut map = HeaderMap::new();
170 for (k, v) in entries {
171 let name: HeaderName = k.parse().unwrap();
172 let value: HeaderValue = v.parse().unwrap();
173 map.insert(name, value);
174 }
175 map
176 }
177
178 fn connect(ip: &str) -> SocketAddr {
180 format!("{}:54321", ip).parse().unwrap()
181 }
182
183 #[test]
184 fn x_real_ip_takes_priority_over_xff_first_segment() {
185 let headers = header_maps(&[
187 ("x-forwarded-for", "1.2.3.4, 203.0.113.9"),
188 ("x-real-ip", "203.0.113.9"),
189 ]);
190 assert_eq!(
191 extract_client_ip(&headers, &connect("10.0.0.5")),
192 "203.0.113.9"
193 );
194 }
195
196 #[test]
197 fn xff_uses_last_segment_not_attacker_first_segment() {
198 let headers = header_maps(&[("x-forwarded-for", "1.2.3.4, 198.51.100.7")]);
200 assert_eq!(
201 extract_client_ip(&headers, &connect("10.0.0.5")),
202 "198.51.100.7"
203 );
204 }
205
206 #[test]
207 fn single_hop_xff_returns_itself() {
208 let headers = header_maps(&[("x-forwarded-for", "203.0.113.42")]);
210 assert_eq!(
211 extract_client_ip(&headers, &connect("10.0.0.5")),
212 "203.0.113.42"
213 );
214 }
215
216 #[test]
217 fn invalid_ip_in_headers_falls_back() {
218 let headers = header_maps(&[("x-forwarded-for", "not-an-ip, 198.51.100.7")]);
220 assert_eq!(
221 extract_client_ip(&headers, &connect("10.0.0.5")),
222 "198.51.100.7"
223 );
224 let bad = header_maps(&[("x-forwarded-for", "evil")]);
226 assert_eq!(extract_client_ip(&bad, &connect("10.0.0.5")), "0.0.0.0");
227 }
228
229 #[test]
230 fn private_connect_ip_does_not_leak() {
231 let headers = header_maps(&[]);
233 assert_eq!(extract_client_ip(&headers, &connect("10.1.1.1")), "0.0.0.0");
234 }
235}