use ipnet::{IpNet, Ipv4Net, Ipv4Subnets, Ipv6Net, Ipv6Subnets};
use serde::{
de::{self, DeserializeOwned, Visitor},
Deserialize, Deserializer,
};
use std::{
borrow::Borrow,
fmt::Debug,
io::Read,
marker::PhantomData,
net::{IpAddr, Ipv4Addr, Ipv6Addr},
};
use thiserror::Error;
#[doc(hidden)]
#[derive(Debug)]
pub struct Record<A> {
start: A,
end: A,
as_number: u32,
country: String,
owner: String,
}
impl From<Record<Ipv4Addr>> for Record<IpAddr> {
fn from(value: Record<Ipv4Addr>) -> Self {
Record {
start: value.start.into(),
end: value.end.into(),
as_number: value.as_number,
country: value.country,
owner: value.owner,
}
}
}
impl From<Record<Ipv6Addr>> for Record<IpAddr> {
fn from(value: Record<Ipv6Addr>) -> Self {
Record {
start: value.start.into(),
end: value.end.into(),
as_number: value.as_number,
country: value.country,
owner: value.owner,
}
}
}
pub type Ipv4Entry = Entry<Ipv4Net>;
pub type Ipv6Entry = Entry<Ipv6Net>;
#[derive(Debug, PartialEq, Eq, Clone)]
pub struct Entry<T> {
pub subnet: T,
pub as_number: u32,
pub country: String,
pub owner: String,
}
impl From<Entry<Ipv4Net>> for Entry<IpNet> {
fn from(value: Entry<Ipv4Net>) -> Self {
Entry {
subnet: value.subnet.into(),
as_number: value.as_number,
country: value.country,
owner: value.owner,
}
}
}
impl From<Entry<Ipv6Net>> for Entry<IpNet> {
fn from(value: Entry<Ipv6Net>) -> Self {
Entry {
subnet: value.subnet.into(),
as_number: value.as_number,
country: value.country,
owner: value.owner,
}
}
}
impl<T: PartialOrd> PartialOrd for Entry<T> {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
self.subnet.partial_cmp(&other.subnet)
}
}
impl<T: Ord> Ord for Entry<T> {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
self.subnet.cmp(&other.subnet)
}
}
impl<'de, A: Deserialize<'de>> Deserialize<'de> for Record<A> {
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
where
D: Deserializer<'de>,
{
deserializer.deserialize_seq(RecordVisitor {
_phantom: PhantomData::default(),
})
}
}
struct RecordVisitor<A> {
_phantom: PhantomData<A>,
}
impl<'de, A: Deserialize<'de>> Visitor<'de> for RecordVisitor<A> {
type Value = Record<A>;
fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
formatter.write_str("a valid ipv4 entry")
}
fn visit_seq<S>(self, mut seq: S) -> std::result::Result<Self::Value, S::Error>
where
S: serde::de::SeqAccess<'de>,
{
let start: A = seq
.next_element()?
.ok_or_else(|| de::Error::invalid_length(0, &self))?;
let end: A = seq
.next_element()?
.ok_or_else(|| de::Error::invalid_length(1, &self))?;
let as_number = seq
.next_element()?
.ok_or_else(|| de::Error::invalid_length(2, &self))?;
let country = seq
.next_element()?
.ok_or_else(|| de::Error::invalid_length(3, &self))?;
let owner = seq
.next_element()?
.ok_or_else(|| de::Error::invalid_length(4, &self))?;
Ok(Record {
start,
end,
as_number,
country,
owner,
})
}
}
#[derive(Debug, Error)]
pub enum Error {
#[error(transparent)]
Csv(#[from] csv::Error),
}
pub type Result<T> = std::result::Result<T, crate::Error>;
pub type Ipv4Database = IpDatabase<Ipv4Addr>;
pub type Ipv6Database = IpDatabase<Ipv6Addr>;
pub struct IpDatabase<A: MapSubnets> {
entries: Vec<Entry<A::Net>>,
}
impl<A: MapSubnets> Debug for IpDatabase<A> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"IpDatabase<{}>{{ {} entries }}",
std::any::type_name::<A>(),
self.entries.len()
)
}
}
#[doc(hidden)]
pub trait MapSubnets: Sized {
type Net: Ord + Debug + Copy;
fn map(record: Record<Self>) -> Box<dyn std::iter::Iterator<Item = Entry<Self::Net>>>;
fn contains(net: &Self::Net, address: &Self) -> bool;
}
impl MapSubnets for Ipv4Addr {
type Net = Ipv4Net;
fn map(record: Record<Self>) -> Box<dyn std::iter::Iterator<Item = Entry<Self::Net>>> {
let iter = Ipv4Subnets::new(record.start, record.end, 8).map(move |subnet| Ipv4Entry {
subnet,
as_number: record.as_number,
country: record.country.clone(),
owner: record.owner.clone(),
});
Box::new(iter)
}
fn contains(net: &Self::Net, address: &Self) -> bool {
net.contains(address)
}
}
impl MapSubnets for Ipv6Addr {
type Net = Ipv6Net;
fn map(record: Record<Self>) -> Box<dyn std::iter::Iterator<Item = Entry<Self::Net>>> {
let iter = Ipv6Subnets::new(record.start, record.end, 8).map(move |subnet| Ipv6Entry {
subnet,
as_number: record.as_number,
country: record.country.clone(),
owner: record.owner.clone(),
});
Box::new(iter)
}
fn contains(net: &Self::Net, address: &Self) -> bool {
net.contains(address)
}
}
macro_rules! read_entries {
($reader:expr, $type:tt) => {
$reader
.deserialize::<Record<$type>>()
.filter(|value| match value {
Ok(record) => record.owner != "Not routed",
Err(_) => true,
})
.map(|value| value.map(<$type>::map))
.flat_map(|subnets| {
let mut records = None;
let mut error = None;
match subnets {
Ok(subnets) => records = Some(subnets),
Err(err) => error = Some(err),
}
records
.into_iter()
.flatten()
.map(Ok)
.chain(error.into_iter().map(Err))
})
};
}
impl<A> IpDatabase<A>
where
A: DeserializeOwned + MapSubnets + Debug + Into<A::Net> + Copy,
{
pub fn from_reader(reader: impl Read) -> std::result::Result<Self, crate::Error> {
tracing::debug!("Parsing tsv data.");
let mut reader = create_reader(reader);
let entries = read_entries!(reader, A).collect::<std::result::Result<_, _>>()?;
Ok(Self { entries })
}
pub fn lookup(&self, address: A) -> Option<&Entry<A::Net>> {
tracing::debug!(?address, "Looking up ip address");
match self
.entries
.binary_search_by_key(&address.into(), |entry| entry.subnet)
{
Ok(index) => return Some(&self.entries[index]),
Err(index) => {
if index != 0 {
let entry = &self.entries[index - 1];
if A::contains(&entry.subnet, &address) {
return Some(entry);
}
}
}
}
None
}
}
fn create_reader(reader: impl Read) -> csv::Reader<impl Read> {
csv::ReaderBuilder::new()
.has_headers(false)
.delimiter(b'\t')
.from_reader(reader)
}
impl MapSubnets for IpAddr {
type Net = IpNet;
fn map(record: Record<Self>) -> Box<dyn std::iter::Iterator<Item = Entry<Self::Net>>> {
match (record.start, record.end) {
(IpAddr::V4(start), IpAddr::V4(end)) => Box::new(
Ipv4Addr::map(Record {
start,
end,
as_number: record.as_number,
country: record.country,
owner: record.owner,
})
.map(|entry| Entry {
subnet: IpNet::from(entry.subnet),
as_number: entry.as_number,
country: entry.country,
owner: entry.owner,
}),
),
(IpAddr::V6(start), IpAddr::V6(end)) => Box::new(
Ipv6Addr::map(Record {
start,
end,
as_number: record.as_number,
country: record.country,
owner: record.owner,
})
.map(|entry| Entry {
subnet: IpNet::from(entry.subnet),
as_number: entry.as_number,
country: entry.country,
owner: entry.owner,
}),
),
_ => panic!("invalid ip combination"),
}
}
fn contains(net: &Self::Net, address: &Self) -> bool {
net.contains(address)
}
}
#[derive(Debug)]
pub struct Database {
v4: Ipv4Database,
v6: Ipv6Database,
}
fn partition_map<I, A, B>(iter: I) -> (A, B)
where
I: Sized,
I: Iterator<Item = Entry<IpNet>>,
A: Default + Extend<Entry<Ipv4Net>>,
B: Default + Extend<Entry<Ipv6Net>>,
{
let mut left = A::default();
let mut right = B::default();
iter.for_each(|val| match val.subnet {
IpNet::V4(v4) => left.extend(Some(Entry {
subnet: v4,
as_number: val.as_number,
country: val.country,
owner: val.owner,
})),
IpNet::V6(v6) => right.extend(Some(Entry {
subnet: v6,
as_number: val.as_number,
country: val.country,
owner: val.owner,
})),
});
(left, right)
}
impl Database {
pub fn from_reader(reader: impl Read) -> std::result::Result<Self, crate::Error> {
let mut reader = create_reader(reader);
let entries = read_entries!(reader, IpAddr).collect::<std::result::Result<Vec<_>, _>>()?;
let (v4, v6): (Vec<_>, Vec<_>) = partition_map(entries.into_iter());
Ok(Self {
v4: Ipv4Database { entries: v4 },
v6: Ipv6Database { entries: v6 },
})
}
pub fn lookup(&self, address: IpAddr) -> Option<IpEntry<'_>> {
match address {
IpAddr::V4(v4) => self.v4.lookup(v4).map(IpEntry::V4),
IpAddr::V6(v6) => self.v6.lookup(v6).map(IpEntry::V6),
}
}
#[inline(always)]
pub fn lookup_v4(&self, address: Ipv4Addr) -> Option<&Ipv4Entry> {
self.v4.lookup(address)
}
#[inline(always)]
pub fn lookup_v6(&self, address: Ipv6Addr) -> Option<&Ipv6Entry> {
self.v6.lookup(address)
}
}
#[derive(Debug, PartialEq, Clone)]
pub enum IpEntry<'a> {
V4(&'a Ipv4Entry),
V6(&'a Ipv6Entry),
}
#[derive(Debug, PartialEq, Clone)]
pub enum IpEntryBuf {
V4(Ipv4Entry),
V6(Ipv6Entry),
}
impl IpEntry<'_> {
pub fn to_ip_entry_buf(self) -> IpEntryBuf {
match self {
IpEntry::V4(v4) => IpEntryBuf::V4(v4.clone()),
IpEntry::V6(v6) => IpEntryBuf::V6(v6.clone()),
}
}
}
#[cfg(test)]
mod tests {
use crate::{Ipv4Database, Ipv6Database};
use std::{
fs::File,
io::BufReader,
net::{Ipv4Addr, Ipv6Addr},
str::FromStr,
};
#[test]
fn test_database_v4() -> Result<(), Box<dyn std::error::Error>> {
let reader = BufReader::new(File::open("ip2asn-v4.tsv")?);
let db = Ipv4Database::from_reader(reader)?;
println!("{:#?}", db);
let record = db.lookup(Ipv4Addr::from_str("1.1.1.1")?).unwrap();
println!("{:#?}", record);
Ok(())
}
#[test]
fn test_database_v6() -> Result<(), Box<dyn std::error::Error>> {
let reader = BufReader::new(File::open("ip2asn-v6.tsv")?);
let db = Ipv6Database::from_reader(reader)?;
println!("{:#?}", db);
let record = db.lookup(Ipv6Addr::from_str("2a05:dfc2::")?).unwrap();
println!("{:#?}", record);
assert_eq!(record.as_number, 200242);
Ok(())
}
}