use core::net::{IpAddr, Ipv4Addr, Ipv6Addr};
use std::io::{self, Read};
use std::path::Path;
use super::mmdb::{IpVersion, MmdbBuilder, MmdbWriteError};
use super::{AsOrg, Coordinates, GeoIpError, GeoLocation, MmdbReader, Subdivision, TimeZoneName};
use crate::asn::LossyAsn;
use rama_core::geo::Country;
use ipnet::{IpNet, Ipv4Net, Ipv6Net};
#[derive(Debug)]
#[non_exhaustive]
pub enum CsvError {
Io(io::Error),
Parse {
record: usize,
reason: Box<str>,
},
Write(MmdbWriteError),
Build(GeoIpError),
}
impl core::fmt::Display for CsvError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::Io(err) => write!(f, "csv: i/o error: {err}"),
Self::Parse { record, reason } => write!(f, "csv: record {record}: {reason}"),
Self::Write(err) => write!(f, "csv: {err}"),
Self::Build(err) => write!(f, "csv: {err}"),
}
}
}
impl core::error::Error for CsvError {
fn source(&self) -> Option<&(dyn core::error::Error + 'static)> {
match self {
Self::Io(err) => Some(err),
Self::Write(err) => Some(err),
Self::Build(err) => Some(err),
Self::Parse { .. } => None,
}
}
}
#[derive(Debug, Clone)]
pub struct CsvGeoRecord {
pub start: IpAddr,
pub end: IpAddr,
pub location: GeoLocation,
}
pub fn compile_csv_into<R, F>(
read: R,
builder: &mut MmdbBuilder,
mut map_row: F,
) -> Result<(), CsvError>
where
R: Read,
F: FnMut(&[&str]) -> Result<Option<CsvGeoRecord>, Box<str>>,
{
let mut reader = csv::ReaderBuilder::new()
.has_headers(false)
.flexible(true)
.from_reader(read);
let mut record = csv::StringRecord::new();
let mut record_num = 0usize;
let mut cidrs: Vec<(u128, u8)> = Vec::new();
loop {
match reader.read_record(&mut record) {
Ok(true) => {}
Ok(false) => break,
Err(err) => return Err(map_csv_err(err, record_num + 1)),
}
record_num += 1;
if record.iter().all(str::is_empty) {
continue;
}
let fields: Vec<&str> = record.iter().collect();
match map_row(&fields) {
Ok(Some(rec)) => {
insert_record(builder, &rec, &mut cidrs).map_err(|reason| CsvError::Parse {
record: record_num,
reason,
})?
}
Ok(None) => {}
Err(reason) => {
return Err(CsvError::Parse {
record: record_num,
reason,
});
}
}
}
Ok(())
}
pub fn compile_csv<R, F>(
read: R,
ip_version: IpVersion,
database_type: impl Into<String>,
languages: &[&str],
map_row: F,
) -> Result<MmdbReader, CsvError>
where
R: Read,
F: FnMut(&[&str]) -> Result<Option<CsvGeoRecord>, Box<str>>,
{
let mut builder = make_builder(ip_version, database_type, languages);
compile_csv_into(read, &mut builder, map_row)?;
let bytes = builder.build().map_err(CsvError::Write)?;
MmdbReader::from_bytes(bytes).map_err(CsvError::Build)
}
fn make_builder(
ip_version: IpVersion,
database_type: impl Into<String>,
languages: &[&str],
) -> MmdbBuilder {
let builder = MmdbBuilder::new(ip_version, database_type);
if languages.is_empty() {
builder
} else {
builder.with_languages(languages.iter().copied())
}
}
fn map_csv_err(err: csv::Error, record: usize) -> CsvError {
if err.is_io_error() {
match err.into_kind() {
csv::ErrorKind::Io(io) => CsvError::Io(io),
other => CsvError::Parse {
record,
reason: format!("{other:?}").into_boxed_str(),
},
}
} else {
CsvError::Parse {
record,
reason: err.to_string().into_boxed_str(),
}
}
}
fn insert_record(
builder: &mut MmdbBuilder,
record: &CsvGeoRecord,
cidrs: &mut Vec<(u128, u8)>,
) -> Result<(), Box<str>> {
let (start, end, bits) = match (record.start, record.end) {
(IpAddr::V4(a), IpAddr::V4(b)) => {
(u128::from(u32::from(a)), u128::from(u32::from(b)), 32u32)
}
(IpAddr::V6(a), IpAddr::V6(b)) => (u128::from(a), u128::from(b), 128u32),
_ => return Err("range endpoints mix IPv4 and IPv6".into()),
};
if start > end {
return Err("range start is greater than range end".into());
}
let to_box = |e: ipnet::PrefixLenError| e.to_string().into_boxed_str();
range_to_cidrs_into(start, end, bits, cidrs);
for &(addr, prefix) in cidrs.iter() {
let net = if bits == 32 {
IpNet::V4(Ipv4Net::new(Ipv4Addr::from(addr as u32), prefix).map_err(to_box)?)
} else {
IpNet::V6(Ipv6Net::new(Ipv6Addr::from(addr), prefix).map_err(to_box)?)
};
builder
.insert(net, &record.location)
.map_err(|e| e.to_string().into_boxed_str())?;
}
Ok(())
}
#[cfg(test)]
fn range_to_cidrs(start: u128, end: u128, bits: u32) -> Vec<(u128, u8)> {
let mut out = Vec::new();
range_to_cidrs_into(start, end, bits, &mut out);
out
}
fn range_to_cidrs_into(start: u128, end: u128, bits: u32, out: &mut Vec<(u128, u8)>) {
out.clear();
let mut cur = start;
loop {
let align = if cur == 0 {
bits
} else {
cur.trailing_zeros().min(bits)
};
let remaining = end - cur; let size_bits = if remaining == u128::MAX {
128
} else {
floor_log2(remaining + 1)
};
let n = align.min(size_bits).min(bits - 1);
out.push((cur, (bits - n) as u8));
let block = 1u128 << n;
match cur.checked_add(block) {
Some(next) if next <= end => cur = next,
_ => break,
}
}
}
fn floor_log2(x: u128) -> u32 {
127 - x.leading_zeros()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum Ip2LocationLite {
Country,
City,
Asn,
}
pub fn compile_ip2location_lite<R: Read>(
read: R,
ip_version: IpVersion,
kind: Ip2LocationLite,
) -> Result<MmdbReader, CsvError> {
let (database_type, languages) = ip2location_meta(kind);
let mut builder = make_builder(ip_version, database_type, languages);
fill_ip2location(read, ip_version, kind, &mut builder)?;
let bytes = builder.build().map_err(CsvError::Write)?;
MmdbReader::from_bytes(bytes).map_err(CsvError::Build)
}
pub fn compile_ip2location_lite_to_file<R: Read>(
read: R,
ip_version: IpVersion,
kind: Ip2LocationLite,
path: impl AsRef<Path>,
) -> Result<(), CsvError> {
let (database_type, languages) = ip2location_meta(kind);
let mut builder = make_builder(ip_version, database_type, languages);
fill_ip2location(read, ip_version, kind, &mut builder)?;
builder.write_to_file(path).map_err(CsvError::Write)
}
fn ip2location_meta(kind: Ip2LocationLite) -> (&'static str, &'static [&'static str]) {
match kind {
Ip2LocationLite::Country => ("IP2LOCATION-LITE-DB1", &["en"]),
Ip2LocationLite::City => ("IP2LOCATION-LITE-DB11", &["en"]),
Ip2LocationLite::Asn => ("IP2LOCATION-LITE-ASN", &[]),
}
}
fn fill_ip2location<R: Read>(
read: R,
ip_version: IpVersion,
kind: Ip2LocationLite,
builder: &mut MmdbBuilder,
) -> Result<(), CsvError> {
match kind {
Ip2LocationLite::Country => compile_csv_into(read, builder, |f| map_country(f, ip_version)),
Ip2LocationLite::City => compile_csv_into(read, builder, |f| map_city(f, ip_version)),
Ip2LocationLite::Asn => compile_csv_into(read, builder, |f| map_asn(f, ip_version)),
}
}
fn parse_ip(field: &str, ip_version: IpVersion) -> Result<IpAddr, Box<str>> {
let field = field.trim();
match ip_version {
IpVersion::V4 => field
.parse::<u32>()
.map(|n| IpAddr::V4(Ipv4Addr::from(n)))
.map_err(|e| format!("invalid ipv4 decimal {field:?}: {e}").into_boxed_str()),
IpVersion::V6 => field
.parse::<u128>()
.map(|n| IpAddr::V6(Ipv6Addr::from(n)))
.map_err(|e| format!("invalid ipv6 decimal {field:?}: {e}").into_boxed_str()),
}
}
fn cell<'a>(s: Option<&&'a str>) -> Option<&'a str> {
s.map(|s| s.trim()).filter(|s| !s.is_empty() && *s != "-")
}
fn record(start: IpAddr, end: IpAddr, location: GeoLocation) -> Option<CsvGeoRecord> {
(!location.is_empty()).then_some(CsvGeoRecord {
start,
end,
location,
})
}
fn map_country(fields: &[&str], ip_version: IpVersion) -> Result<Option<CsvGeoRecord>, Box<str>> {
if fields.len() < 4 {
return Err("expected at least 4 columns (ip_from, ip_to, code, name)".into());
}
let start = parse_ip(fields[0], ip_version)?;
let end = parse_ip(fields[1], ip_version)?;
let loc = GeoLocation {
country: cell(fields.get(2)).map(Country::from_code),
..Default::default()
};
Ok(record(start, end, loc))
}
fn map_city(fields: &[&str], ip_version: IpVersion) -> Result<Option<CsvGeoRecord>, Box<str>> {
if fields.len() < 6 {
return Err("expected at least 6 columns for a city layout".into());
}
let start = parse_ip(fields[0], ip_version)?;
let end = parse_ip(fields[1], ip_version)?;
let mut loc = GeoLocation {
country: cell(fields.get(2)).map(Country::from_code),
city: cell(fields.get(5)).map(Box::from),
postal_code: cell(fields.get(8)).map(Box::from),
..Default::default()
};
if let Some(region) = cell(fields.get(4)) {
loc.subdivisions.push(Subdivision {
iso_code: None,
name: Some(region.into()),
});
}
if let (Some(lat), Some(lon)) = (cell(fields.get(6)), cell(fields.get(7)))
&& let (Ok(latitude), Ok(longitude)) = (lat.parse::<f64>(), lon.parse::<f64>())
&& !(latitude == 0.0 && longitude == 0.0)
{
loc.location = Some(Coordinates {
latitude,
longitude,
accuracy_radius_km: None,
time_zone: cell(fields.get(9)).map(TimeZoneName::from),
});
}
Ok(record(start, end, loc))
}
fn map_asn(fields: &[&str], ip_version: IpVersion) -> Result<Option<CsvGeoRecord>, Box<str>> {
if fields.len() < 5 {
return Err("expected at least 5 columns (ip_from, ip_to, cidr, asn, name)".into());
}
let start = parse_ip(fields[0], ip_version)?;
let end = parse_ip(fields[1], ip_version)?;
let asn = cell(fields.get(3))
.and_then(|s| s.parse::<u32>().ok())
.filter(|n| *n != 0);
let organization = cell(fields.get(4)).map(Box::from);
let loc = GeoLocation {
autonomous_system: (asn.is_some() || organization.is_some()).then(|| AsOrg {
asn: asn.map(LossyAsn::from),
organization,
}),
..Default::default()
};
Ok(record(start, end, loc))
}
#[cfg(test)]
mod tests {
use super::*;
fn ip(s: &str) -> IpAddr {
s.parse().unwrap()
}
#[test]
fn range_to_cidrs_cases() {
assert_eq!(
range_to_cidrs(0x0102_0300, 0x0102_03ff, 32),
vec![(0x0102_0300, 24)]
);
assert_eq!(range_to_cidrs(0, 9, 32), vec![(0, 29), (8, 31)]);
assert_eq!(
range_to_cidrs(0, u128::from(u32::MAX), 32),
vec![(0, 1), (0x8000_0000, 1)]
);
}
#[test]
fn range_to_cidrs_covers_exactly() {
for (start, end) in [(0u128, 0u128), (5, 5), (10, 37), (256, 1000), (0, 255)] {
let blocks = range_to_cidrs(start, end, 32);
let mut next = start;
let mut count = 0u128;
for (addr, prefix) in &blocks {
assert_eq!(*addr, next, "blocks must be contiguous");
let size = 1u128 << (32 - u32::from(*prefix));
assert_eq!(addr % size, 0, "block must be aligned");
next += size;
count += size;
}
assert_eq!(count, end - start + 1);
assert_eq!(next, end + 1);
}
}
#[test]
fn compile_ip2location_country_db1() {
let csv = "\"16777216\",\"16777471\",\"US\",\"United States of America\"\n\
\"16777472\",\"16777727\",\"CN\",\"China\"\n\
\"16777728\",\"16777983\",\"-\",\"-\"\n";
let reader =
compile_ip2location_lite(csv.as_bytes(), IpVersion::V4, Ip2LocationLite::Country)
.unwrap();
let us = reader.lookup(ip("1.0.0.5")).unwrap();
assert_eq!(us.country().unwrap().to_owned(), Country::UnitedStates);
assert_eq!(us.country().unwrap().name(), Some("United States"));
let cn = reader.lookup(ip("1.0.1.5")).unwrap();
assert_eq!(cn.country().unwrap().to_owned(), Country::China);
assert!(reader.lookup(ip("1.0.2.5")).is_none());
}
#[test]
fn compile_ip2location_city_db11() {
let csv = "\"16777216\",\"16777471\",\"US\",\"United States\",\"New York\",\"Buffalo\",\"42.886\",\"-78.878\",\"14202\",\"-05:00\"\n";
let reader =
compile_ip2location_lite(csv.as_bytes(), IpVersion::V4, Ip2LocationLite::City).unwrap();
let loc = reader.lookup(ip("1.0.0.5")).unwrap();
assert_eq!(loc.country().unwrap().to_owned(), Country::UnitedStates);
assert_eq!(loc.city(), Some("Buffalo"));
assert_eq!(loc.postal_code(), Some("14202"));
assert_eq!(loc.latitude(), Some(42.886));
assert_eq!(loc.time_zone(), Some("-05:00"));
let owned = loc.to_owned();
assert_eq!(owned.subdivisions.len(), 1);
assert_eq!(owned.subdivisions[0].name.as_deref(), Some("New York"));
}
#[test]
fn compile_csv_into_streams_with_custom_mapper() {
let mut builder = MmdbBuilder::new(IpVersion::V4, "Test-Country");
let csv = "16777216,16777471,BE\n16777472,16777727,FR\n";
compile_csv_into(csv.as_bytes(), &mut builder, |f| {
let start = parse_ip(f[0], IpVersion::V4)?;
let end = parse_ip(f[1], IpVersion::V4)?;
let loc = GeoLocation {
country: Some(Country::from_code(f[2])),
..Default::default()
};
Ok(record(start, end, loc))
})
.unwrap();
let reader = MmdbReader::from_bytes(builder.build().unwrap()).unwrap();
assert_eq!(
reader
.lookup(ip("1.0.0.5"))
.unwrap()
.country()
.unwrap()
.code(),
"BE"
);
assert_eq!(
reader
.lookup(ip("1.0.1.5"))
.unwrap()
.country()
.unwrap()
.code(),
"FR"
);
}
#[test]
fn compile_city_skips_null_island_coords() {
let csv = "\"16777216\",\"16777471\",\"US\",\"United States\",\"\",\"\",\"0.000000\",\"0.000000\",\"\",\"\"\n\
\"16777472\",\"16777727\",\"US\",\"United States\",\"NY\",\"Buffalo\",\"42.886\",\"-78.878\",\"14202\",\"-05:00\"\n";
let reader =
compile_ip2location_lite(csv.as_bytes(), IpVersion::V4, Ip2LocationLite::City).unwrap();
let null = reader.lookup(ip("1.0.0.5")).unwrap();
assert_eq!(null.country().unwrap().code(), "US");
assert_eq!(null.latitude(), None);
let real = reader.lookup(ip("1.0.1.5")).unwrap();
assert_eq!(real.latitude(), Some(42.886));
}
#[test]
fn compile_ip2location_asn() {
let csv = "\"16777216\",\"16777471\",\"1.0.0.0/24\",\"13335\",\"Cloudflare, Inc.\"\n";
let reader =
compile_ip2location_lite(csv.as_bytes(), IpVersion::V4, Ip2LocationLite::Asn).unwrap();
let loc = reader.lookup(ip("1.0.0.5")).unwrap();
assert_eq!(loc.asn().map(|a| a.as_u32()), Some(13335));
assert_eq!(loc.as_organization(), Some("Cloudflare, Inc."));
}
#[test]
fn compile_to_file_streams_and_loads() {
let dir = tempfile::tempdir().expect("tempdir");
let path = dir.path().join("country.mmdb");
let csv = "\"16777216\",\"16777471\",\"BE\",\"Belgium\"\n";
compile_ip2location_lite_to_file(
csv.as_bytes(),
IpVersion::V4,
Ip2LocationLite::Country,
&path,
)
.unwrap();
let reader = MmdbReader::open(&path).unwrap();
assert_eq!(
reader
.lookup(ip("1.0.0.5"))
.unwrap()
.country()
.unwrap()
.to_owned(),
Country::Belgium
);
}
#[test]
fn generic_compile_csv_with_custom_mapper() {
let csv = "1.0.0.0,1.0.0.255,BE\n";
let reader = compile_csv(
csv.as_bytes(),
IpVersion::V4,
"Custom-Country",
&["en"],
|f| {
if f.len() < 3 {
return Err("need 3 cols".into());
}
let start: IpAddr = f[0]
.parse()
.map_err(|e| format!("bad ip {:?}: {e}", f[0]).into_boxed_str())?;
let end: IpAddr = f[1]
.parse()
.map_err(|e| format!("bad ip {:?}: {e}", f[1]).into_boxed_str())?;
let location = GeoLocation {
country: Some(Country::from_code(f[2])),
..Default::default()
};
Ok(Some(CsvGeoRecord {
start,
end,
location,
}))
},
)
.unwrap();
assert_eq!(
reader
.lookup(ip("1.0.0.42"))
.unwrap()
.country()
.unwrap()
.to_owned(),
Country::Belgium
);
}
#[quickcheck_macros::quickcheck]
fn prop_range_to_cidrs_covers_v4(a: u32, b: u32) -> bool {
let (start, end) = if a <= b { (a, b) } else { (b, a) };
let (start, end) = (u128::from(start), u128::from(end));
let blocks = range_to_cidrs(start, end, 32);
if blocks.is_empty() {
return false;
}
let mut next = start;
for (addr, prefix) in &blocks {
let p = u32::from(*prefix);
if *addr != next || !(1..=32).contains(&p) {
return false;
}
let size = 1u128 << (32 - p);
if addr % size != 0 {
return false; }
next += size;
}
next == end + 1
}
#[test]
fn compile_ipv6_country() {
let from: u128 = u128::from(ip6("2001:db8::"));
let to: u128 = u128::from(ip6("2001:db8:0:ffff:ffff:ffff:ffff:ffff"));
let csv = format!("\"{from}\",\"{to}\",\"BE\",\"Belgium\"\n");
let reader =
compile_ip2location_lite(csv.as_bytes(), IpVersion::V6, Ip2LocationLite::Country)
.unwrap();
assert_eq!(
reader
.lookup(ip("2001:db8:0:1234::1"))
.unwrap()
.country()
.unwrap()
.to_owned(),
Country::Belgium
);
assert!(reader.lookup(ip("2001:db9::1")).is_none());
}
#[test]
fn malformed_rows_report_line_and_range_errors() {
let err = compile_ip2location_lite(
"\"1\",\"2\",\"US\",\"x\"\n\"oops\"\n".as_bytes(),
IpVersion::V4,
Ip2LocationLite::Country,
)
.unwrap_err();
assert!(matches!(err, CsvError::Parse { record: 2, .. }));
let err = compile_ip2location_lite(
"\"100\",\"50\",\"US\",\"x\"\n".as_bytes(),
IpVersion::V4,
Ip2LocationLite::Country,
)
.unwrap_err();
assert!(
matches!(&err, CsvError::Parse { reason, .. } if reason.contains("greater than")),
"unexpected error: {err}"
);
}
fn ip6(s: &str) -> core::net::Ipv6Addr {
s.parse().unwrap()
}
}