iuna

iuna

iuna - experimental mainnet-candidate protocol
git clone https://getiuna.org/git/iuna.git
Log | Files | Refs | README | LICENSE

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 }