use super::{asn::Asn, asn::lookup_ip};
use futures::future::join_all;
use hickory_resolver::{Resolver, name_server::ConnectionProvider, proto::rr::RecordType};
use ip2asn::IpAsnMap;
use serde::Serialize;
use std::net::IpAddr;
use std::sync::Arc;
#[derive(Debug, Serialize, Clone)]
pub struct NameServer {
pub names: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub ips: Option<Vec<IpAddr>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub asn: Option<Vec<Asn>>,
}
pub async fn query_ns<T: ConnectionProvider>(
target: &str,
resolver: &Resolver<T>,
ip2asn_map: &Arc<IpAsnMap>,
) -> Option<NameServer> {
let lookup_ns_future = resolver.lookup(target, RecordType::NS);
match lookup_ns_future.await {
Ok(response_ns) => {
let ns_records = response_ns
.into_iter()
.filter_map(|r| r.into_ns().ok())
.map(|name| name.to_string())
.collect::<Vec<_>>();
let futures = ns_records.iter().map(|ns| query_ipv4_ipv6(ns, resolver));
let parallel_results = join_all(futures).await;
let ns_ips = parallel_results
.into_iter()
.flatten()
.flatten()
.collect::<Vec<_>>();
let asn = lookup_ip(&ns_ips, ip2asn_map);
let ip_records = match ns_ips.is_empty() {
true => None,
false => Some(ns_ips),
};
Some(NameServer {
names: ns_records,
ips: ip_records,
asn,
})
}
Err(_) => None,
}
}
pub async fn query_cname<T: ConnectionProvider>(
target: &str,
resolver: &Resolver<T>,
) -> Option<Vec<String>> {
let lookup_cname_future = resolver.lookup(target, RecordType::CNAME);
match lookup_cname_future.await {
Ok(response_cname) => {
let cnames = response_cname
.into_iter()
.filter_map(|r| r.into_cname().ok())
.map(|name| name.to_string())
.collect::<Vec<_>>();
if cnames.is_empty() {
None
} else {
Some(cnames)
}
}
Err(_) => None,
}
}
pub async fn query_ipv6<T: ConnectionProvider>(
target: &str,
resolver: &Resolver<T>,
) -> Option<Vec<IpAddr>> {
let lookup_aaaa_future = resolver.ipv6_lookup(target);
match lookup_aaaa_future.await {
Ok(response_aaaa) => {
let ipv6_addrs = response_aaaa
.into_iter()
.map(|addr| IpAddr::from(addr.0))
.collect::<Vec<_>>();
Some(ipv6_addrs)
}
Err(_) => None,
}
}
pub async fn query_ipv4<T: ConnectionProvider>(
target: &str,
resolver: &Resolver<T>,
) -> Option<Vec<IpAddr>> {
let lookup_a_future = resolver.ipv4_lookup(target);
match lookup_a_future.await {
Ok(response_a) => {
let ipv4_addrs = response_a
.into_iter()
.map(|addr| IpAddr::from(addr.0))
.collect::<Vec<_>>();
Some(ipv4_addrs)
}
Err(_) => None,
}
}
pub async fn query_ipv4_ipv6<T: ConnectionProvider>(
target: &str,
resolver: &Resolver<T>,
) -> Option<Vec<IpAddr>> {
let ipv4 = query_ipv4(target, resolver);
let ipv6 = query_ipv6(target, resolver);
let mut ip: Vec<IpAddr> = Vec::new();
let (ipv4, ipv6) = tokio::join!(ipv4, ipv6);
if let Some(v4) = ipv4 {
ip.extend(v4);
}
if let Some(v6) = ipv6 {
ip.extend(v6);
}
if ip.is_empty() { None } else { Some(ip) }
}
#[cfg(test)]
mod tests {
use super::*;
use hickory_resolver::Resolver;
use ip2asn::Builder;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
#[tokio::test]
async fn test_query_ipv4_some() {
let target = "localhost";
let resolver = Resolver::builder_tokio().unwrap().build();
let response = query_ipv4(target, &resolver).await;
assert!(response.is_some());
let mut response = response.unwrap();
let expected = [IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1))];
for ip in &mut response {
assert!(expected.contains(ip));
}
}
#[tokio::test]
async fn test_query_ipv6_some() {
let target = "localhost";
let resolver = Resolver::builder_tokio().unwrap().build();
let response = query_ipv6(target, &resolver).await;
assert!(response.is_some());
let mut response = response.unwrap();
let expected = [IpAddr::V6(Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 1))];
for ip in &mut response {
assert!(expected.contains(ip));
}
}
#[tokio::test]
async fn test_query_ipv4_ipv6_some() {
let target = "localhost";
let resolver = Resolver::builder_tokio().unwrap().build();
let response = query_ipv4_ipv6(target, &resolver).await;
assert!(response.is_some());
let mut response = response.unwrap();
let expected = [
IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)),
IpAddr::V6(Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 1)),
];
for ip in &mut response {
assert!(expected.contains(ip));
}
}
#[tokio::test]
async fn test_query_cname_some() {
let target = "www.example.com";
let resolver = Resolver::builder_tokio().unwrap().build();
let response = query_cname(target, &resolver).await;
assert!(response.is_some());
let mut response = response.unwrap();
let expected = ["www.example.com-v4.edgesuite.net.".to_string()];
for cname in &mut response {
assert!(expected.contains(cname));
}
}
#[tokio::test]
async fn test_query_ns_some() {
let target = "facebook.com";
let resolver = Resolver::builder_tokio().unwrap().build();
let data = "129.134.0.0\t129.134.255.255\t32934\tUS\tFACEBOOK-AS";
let ip2asn_map = Builder::new()
.with_source(data.as_bytes())
.unwrap()
.build()
.unwrap();
let ip2asn_map = Arc::new(ip2asn_map);
let response = query_ns(target, &resolver, &ip2asn_map).await;
assert!(response.is_some());
let response = response.unwrap();
let expected_names = [
"a.ns.facebook.com.".to_string(),
"b.ns.facebook.com.".to_string(),
"c.ns.facebook.com.".to_string(),
"d.ns.facebook.com.".to_string(),
];
for name in &response.names {
assert!(expected_names.contains(name));
}
assert!(response.ips.is_some());
let ips = response.ips.unwrap();
assert_eq!(ips.len(), 8);
}
}