ip_geolocation.rs (11309B)
1 use std::{ 2 net::{IpAddr, Ipv4Addr, Ipv6Addr}, 3 sync::OnceLock, 4 }; 5 6 use serde::{Serialize, Serializer}; 7 8 const DATABASE_MAGIC: &[u8; 8] = b"IUNAGEO2"; 9 // Compile-time input: the bytes become part of every executable that uses this crate. 10 const DATABASE: &[u8] = include_bytes!("ip_geolocation/embedded/ip-country.bin"); 11 12 static BUNDLED_GEOLOCATION: OnceLock<IpGeolocation> = OnceLock::new(); 13 14 #[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] 15 pub struct CountryCode([u8; 2]); 16 17 impl CountryCode { 18 fn new(bytes: [u8; 2]) -> Option<Self> { 19 bytes 20 .iter() 21 .all(u8::is_ascii_uppercase) 22 .then_some(Self(bytes)) 23 } 24 25 pub fn as_str(&self) -> &str { 26 std::str::from_utf8(&self.0).expect("country codes contain ASCII only") 27 } 28 } 29 30 impl Serialize for CountryCode { 31 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error> 32 where 33 S: Serializer, 34 { 35 serializer.serialize_str(self.as_str()) 36 } 37 } 38 39 #[derive(Clone, Copy, Default)] 40 struct PrefixGroup { 41 offset: u32, 42 count: u32, 43 } 44 45 struct EmbeddedDatabase { 46 bytes: &'static [u8], 47 ipv4: [PrefixGroup; 33], 48 ipv6: [PrefixGroup; 129], 49 } 50 51 impl EmbeddedDatabase { 52 fn parse(bytes: &'static [u8]) -> Result<Self, &'static str> { 53 const HEADER_SIZE: usize = 8 + (33 + 129) * 4; 54 if bytes.len() < HEADER_SIZE || &bytes[..8] != DATABASE_MAGIC { 55 return Err("invalid database header"); 56 } 57 58 let mut counts_cursor = 8; 59 let mut read_count = || { 60 let count = 61 u32::from_le_bytes(bytes[counts_cursor..counts_cursor + 4].try_into().unwrap()); 62 counts_cursor += 4; 63 count 64 }; 65 let ipv4_counts = std::array::from_fn::<_, 33, _>(|_| read_count()); 66 let ipv6_counts = std::array::from_fn::<_, 129, _>(|_| read_count()); 67 let mut data_cursor = HEADER_SIZE; 68 let mut ipv4 = [PrefixGroup::default(); 33]; 69 let mut ipv6 = [PrefixGroup::default(); 129]; 70 71 for (group, count) in ipv4.iter_mut().zip(ipv4_counts) { 72 *group = prefix_group(data_cursor, count, 6, bytes.len())?; 73 data_cursor += count as usize * 6; 74 } 75 for (group, count) in ipv6.iter_mut().zip(ipv6_counts) { 76 *group = prefix_group(data_cursor, count, 18, bytes.len())?; 77 data_cursor += count as usize * 18; 78 } 79 if data_cursor != bytes.len() { 80 return Err("trailing database data"); 81 } 82 Ok(Self { bytes, ipv4, ipv6 }) 83 } 84 85 fn lookup_ipv4(&self, address: Ipv4Addr) -> Option<CountryCode> { 86 let address = u32::from(address); 87 for prefix in (0..=32).rev() { 88 let network = address & prefix_mask_u32(prefix); 89 let group = self.ipv4[prefix as usize]; 90 if let Some(country) = binary_search_group(self.bytes, group, 4, u128::from(network)) { 91 return Some(country); 92 } 93 } 94 None 95 } 96 97 fn lookup_ipv6(&self, address: Ipv6Addr) -> Option<CountryCode> { 98 let address = u128::from(address); 99 for prefix in (0..=128).rev() { 100 let network = address & prefix_mask_u128(prefix); 101 let group = self.ipv6[prefix as usize]; 102 if let Some(country) = binary_search_group(self.bytes, group, 16, network) { 103 return Some(country); 104 } 105 } 106 None 107 } 108 } 109 110 fn prefix_group( 111 offset: usize, 112 count: u32, 113 record_size: usize, 114 database_len: usize, 115 ) -> Result<PrefixGroup, &'static str> { 116 let byte_len = (count as usize) 117 .checked_mul(record_size) 118 .ok_or("database size overflow")?; 119 offset 120 .checked_add(byte_len) 121 .filter(|end| *end <= database_len) 122 .ok_or("truncated database")?; 123 Ok(PrefixGroup { 124 offset: offset.try_into().map_err(|_| "database offset overflow")?, 125 count, 126 }) 127 } 128 129 fn binary_search_group( 130 database: &[u8], 131 group: PrefixGroup, 132 address_len: usize, 133 target: u128, 134 ) -> Option<CountryCode> { 135 let record_size = address_len + 2; 136 let mut left = 0usize; 137 let mut right = group.count as usize; 138 while left < right { 139 let middle = left + (right - left) / 2; 140 let offset = group.offset as usize + middle * record_size; 141 let record = &database[offset..offset + record_size]; 142 let value = match address_len { 143 4 => u128::from(u32::from_be_bytes(record[..4].try_into().unwrap())), 144 16 => u128::from_be_bytes(record[..16].try_into().unwrap()), 145 _ => unreachable!(), 146 }; 147 match value.cmp(&target) { 148 std::cmp::Ordering::Less => left = middle + 1, 149 std::cmp::Ordering::Greater => right = middle, 150 std::cmp::Ordering::Equal => { 151 return CountryCode::new([record[address_len], record[address_len + 1]]); 152 } 153 } 154 } 155 None 156 } 157 158 fn prefix_mask_u32(prefix: u32) -> u32 { 159 if prefix == 0 { 160 0 161 } else { 162 u32::MAX << (32 - prefix) 163 } 164 } 165 166 fn prefix_mask_u128(prefix: u32) -> u128 { 167 if prefix == 0 { 168 0 169 } else { 170 u128::MAX << (128 - prefix) 171 } 172 } 173 174 pub struct IpGeolocation { 175 database: EmbeddedDatabase, 176 } 177 178 impl IpGeolocation { 179 pub fn bundled() -> &'static Self { 180 BUNDLED_GEOLOCATION.get_or_init(|| { 181 Self::from_database(DATABASE).expect("bundled IP geolocation database must be valid") 182 }) 183 } 184 185 pub fn country_for_ip(&self, address: IpAddr) -> Option<CountryCode> { 186 match address { 187 IpAddr::V4(address) if ipv4_is_geolocatable(address) => { 188 self.database.lookup_ipv4(address) 189 } 190 IpAddr::V6(address) => { 191 if let Some(mapped) = address.to_ipv4_mapped() { 192 return self.country_for_ip(IpAddr::V4(mapped)); 193 } 194 ipv6_is_geolocatable(address).then_some(())?; 195 self.database.lookup_ipv6(address) 196 } 197 IpAddr::V4(_) => None, 198 } 199 } 200 201 fn from_database(database: &'static [u8]) -> Result<Self, &'static str> { 202 Ok(Self { 203 database: EmbeddedDatabase::parse(database)?, 204 }) 205 } 206 207 #[cfg(test)] 208 pub(crate) fn from_entries(entries: &[(IpAddr, u8, &str)]) -> Self { 209 let mut ipv4 = std::array::from_fn::<_, 33, _>(|_| Vec::new()); 210 let mut ipv6 = std::array::from_fn::<_, 129, _>(|_| Vec::new()); 211 for (address, prefix, country) in entries { 212 let bytes: [u8; 2] = country.as_bytes().try_into().unwrap(); 213 match address { 214 IpAddr::V4(address) => { 215 ipv4[*prefix as usize].push((address.octets().to_vec(), bytes)); 216 } 217 IpAddr::V6(address) => { 218 ipv6[*prefix as usize].push((address.octets().to_vec(), bytes)); 219 } 220 } 221 } 222 let mut database = DATABASE_MAGIC.to_vec(); 223 for count in ipv4.iter().chain(ipv6.iter()).map(Vec::len) { 224 database.extend_from_slice(&(count as u32).to_le_bytes()); 225 } 226 for group in ipv4.iter_mut().chain(ipv6.iter_mut()) { 227 group.sort_unstable(); 228 for (address, country) in group { 229 database.extend_from_slice(address); 230 database.extend_from_slice(country); 231 } 232 } 233 Self::from_database(Box::leak(database.into_boxed_slice())).unwrap() 234 } 235 } 236 237 fn ipv4_is_geolocatable(address: Ipv4Addr) -> bool { 238 let [a, b, c, _] = address.octets(); 239 !(a == 0 240 || a == 10 241 || a == 127 242 || (a == 100 && (64..=127).contains(&b)) 243 || (a == 169 && b == 254) 244 || (a == 172 && (16..=31).contains(&b)) 245 || (a == 192 && b == 0 && c == 0) 246 || (a == 192 && b == 0 && c == 2) 247 || (a == 192 && b == 88 && c == 99) 248 || (a == 192 && b == 168) 249 || (a == 198 && (b == 18 || b == 19)) 250 || (a == 198 && b == 51 && c == 100) 251 || (a == 203 && b == 0 && c == 113) 252 || a >= 224) 253 } 254 255 fn ipv6_is_geolocatable(address: Ipv6Addr) -> bool { 256 let octets = address.octets(); 257 !(address.is_unspecified() 258 || address.is_loopback() 259 || octets[0] == 0xff 260 || octets[0] & 0xfe == 0xfc 261 || (octets[0] == 0xfe && octets[1] & 0xc0 == 0x80) 262 || (octets[0] == 0xfe && octets[1] & 0xc0 == 0xc0) 263 || (octets[0] == 0x20 && octets[1] == 0x01 && octets[2] == 0x0d && octets[3] == 0xb8)) 264 } 265 266 #[cfg(test)] 267 mod tests { 268 use super::*; 269 270 fn fixture() -> IpGeolocation { 271 IpGeolocation::from_entries(&[ 272 ("0.0.0.0".parse().unwrap(), 8, "US"), 273 ("8.0.0.0".parse().unwrap(), 8, "US"), 274 ("8.8.0.0".parse().unwrap(), 16, "NL"), 275 ("10.0.0.0".parse().unwrap(), 8, "US"), 276 ("127.0.0.0".parse().unwrap(), 8, "US"), 277 ("169.254.0.0".parse().unwrap(), 16, "US"), 278 ("172.16.0.0".parse().unwrap(), 12, "US"), 279 ("192.168.0.0".parse().unwrap(), 16, "US"), 280 ("224.0.0.0".parse().unwrap(), 4, "US"), 281 ("::".parse().unwrap(), 0, "US"), 282 ("2001:4860::".parse().unwrap(), 32, "US"), 283 ("fc00::".parse().unwrap(), 7, "US"), 284 ]) 285 } 286 287 #[test] 288 fn finds_ipv4_country() { 289 assert_eq!( 290 fixture() 291 .country_for_ip("8.1.2.3".parse().unwrap()) 292 .unwrap() 293 .as_str(), 294 "US" 295 ); 296 } 297 298 #[test] 299 fn finds_ipv6_country() { 300 assert_eq!( 301 fixture() 302 .country_for_ip("2001:4860:4860::8888".parse().unwrap()) 303 .unwrap() 304 .as_str(), 305 "US" 306 ); 307 } 308 309 #[test] 310 fn uses_longest_prefix_match() { 311 assert_eq!( 312 fixture() 313 .country_for_ip("8.8.8.8".parse().unwrap()) 314 .unwrap() 315 .as_str(), 316 "NL" 317 ); 318 } 319 320 #[test] 321 fn unknown_address_has_no_country() { 322 assert_eq!(fixture().country_for_ip("11.0.0.1".parse().unwrap()), None); 323 } 324 325 #[test] 326 fn non_geographic_addresses_have_no_country() { 327 let geolocation = fixture(); 328 for address in [ 329 "10.0.0.1", 330 "172.16.0.1", 331 "192.168.1.10", 332 "127.0.0.1", 333 "169.254.1.1", 334 "0.0.0.0", 335 "224.0.0.1", 336 "::", 337 "::1", 338 "fe80::1", 339 "fc00::1", 340 "ff02::1", 341 ] { 342 assert_eq!( 343 geolocation.country_for_ip(address.parse().unwrap()), 344 None, 345 "{address}" 346 ); 347 } 348 } 349 350 #[test] 351 fn bundled_database_loads_both_address_families() { 352 let geolocation = IpGeolocation::bundled(); 353 assert_eq!( 354 geolocation 355 .country_for_ip("8.8.8.8".parse().unwrap()) 356 .unwrap() 357 .as_str(), 358 "US" 359 ); 360 assert_eq!( 361 geolocation 362 .country_for_ip("2001:4860:4860::8888".parse().unwrap()) 363 .unwrap() 364 .as_str(), 365 "US" 366 ); 367 } 368 }