use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::sync::Arc;
use reqwest::dns::{Addrs, Name, Resolve, Resolving};
use crate::upstream::UpstreamError;
type BoxError = Box<dyn std::error::Error + Send + Sync>;
pub struct GuardedResolver {
allow_private_addresses: bool,
}
impl GuardedResolver {
pub fn new(allow_private_addresses: bool) -> Arc<GuardedResolver> {
Arc::new(GuardedResolver {
allow_private_addresses,
})
}
}
impl Resolve for GuardedResolver {
fn resolve(&self, name: Name) -> Resolving {
let allow_private_addresses = self.allow_private_addresses;
let host = name.as_str().to_owned();
Box::pin(async move {
let resolved: Vec<SocketAddr> = tokio::net::lookup_host((host.as_str(), 0))
.await
.map_err(|source| {
let reason = format!("cannot resolve {host}: {source}");
Box::new(UpstreamError::Transport(reason)) as BoxError
})?
.collect();
if !allow_private_addresses {
if let Some(addr) = resolved.iter().find(|addr| !is_public(addr.ip())) {
let rejected = UpstreamError::RejectedAddress { addr: addr.ip() };
return Err(Box::new(rejected) as BoxError);
}
}
Ok(Box::new(resolved.into_iter()) as Addrs)
})
}
}
pub fn is_public(ip: IpAddr) -> bool {
match ip {
IpAddr::V4(ip) => is_public_v4(ip),
IpAddr::V6(ip) => is_public_v6(ip),
}
}
fn is_public_v4(ip: Ipv4Addr) -> bool {
let [a, b, ..] = ip.octets();
!(ip.is_loopback()
|| ip.is_private()
|| ip.is_link_local()
|| a == 0
|| ip.is_unspecified()
|| ip.is_multicast()
|| ip.is_broadcast()
|| ip.is_documentation()
|| (a == 100 && (64..128).contains(&b))
|| (a == 192 && b == 0 && ip.octets()[2] == 0)
|| (a == 198 && (b == 18 || b == 19))
|| a >= 240)
}
fn is_public_v6(ip: Ipv6Addr) -> bool {
if let Some(embedded) = embedded_v4(ip) {
return is_public_v4(embedded);
}
let segments = ip.segments();
if segments[0] & 0xe000 != 0x2000 {
return false;
}
let protocol_assignments = segments[0] == 0x2001 && segments[1] & 0xfe00 == 0;
let documentation = (segments[0] == 0x2001 && segments[1] == 0x0db8)
|| (segments[0] == 0x3fff && segments[1] & 0xf000 == 0);
let as112 = segments[0] == 0x2620 && segments[1] == 0x004f && segments[2] == 0x8000;
!(protocol_assignments || documentation || as112)
}
fn embedded_v4(ip: Ipv6Addr) -> Option<Ipv4Addr> {
let segments = ip.segments();
let octets = ip.octets();
#[allow(deprecated)]
if let Some(v4) = ip.to_ipv4() {
return Some(v4);
}
if segments[0] == 0x0064 && segments[1] == 0xff9b && segments[2..6] == [0, 0, 0, 0] {
return Some(Ipv4Addr::new(
octets[12], octets[13], octets[14], octets[15],
));
}
if segments[0] == 0x2002 {
return Some(Ipv4Addr::new(octets[2], octets[3], octets[4], octets[5]));
}
None
}