use std::error::Error;
use std::net::IpAddr;
use std::str::FromStr;
use std::sync::Arc;
use ipnet::IpNet;
use reqwest::dns::{Addrs, Name, Resolve, Resolving};
use tokio::net::lookup_host;
use crate::dbs::capabilities::{NetTarget, Targets};
pub(crate) struct NetFilter {
pub(crate) allow: Targets<NetTarget>,
pub(crate) deny: Targets<NetTarget>,
}
fn is_private_ip(ip: IpAddr) -> bool {
match ip.to_canonical() {
IpAddr::V4(v4) => {
v4.is_loopback() || v4.is_private() || v4.is_link_local() || v4.is_broadcast() || v4.is_unspecified() || (u32::from(v4) & 0xFFC0_0000) == 0x6440_0000
}
IpAddr::V6(v6) => {
v6.is_loopback() || v6.is_unspecified() || (v6.segments()[0] & 0xFE00) == 0xFC00
|| (v6.segments()[0] & 0xFFC0) == 0xFE80
}
}
}
pub(crate) struct FilteringResolver {
pub(crate) filter: Arc<NetFilter>,
}
impl FilteringResolver {
pub(crate) fn from_net_filter(filter: Arc<NetFilter>) -> Self {
FilteringResolver {
filter,
}
}
}
impl Resolve for FilteringResolver {
fn resolve(&self, name: Name) -> Resolving {
let filter = Arc::clone(&self.filter);
let name_str = name.as_str().to_string();
Box::pin(async move {
let name_target = NetTarget::from_str(&name_str)
.map_err(|x| Box::new(x) as Box<dyn Error + Send + Sync>)?;
let name_is_allowed =
filter.allow.matches(&name_target) && !filter.deny.matches(&name_target);
if !name_is_allowed {
return Err(
Box::new(crate::err::Error::NetTargetNotAllowed(name_target.to_string()))
as Box<dyn Error + Send + Sync>,
);
}
let addrs: Vec<std::net::SocketAddr> = lookup_host((name_str, 0_u16))
.await
.map_err(|x| Box::new(x) as Box<dyn Error + Send + Sync>)?
.collect();
let mut allowed = Vec::new();
let mut first_denied = None;
for addr in addrs {
let target = IpNet::from(addr.ip());
let ip_target = NetTarget::IPNet(target);
if filter.deny.matches(&ip_target) {
if first_denied.is_none() {
first_denied = Some(target);
}
} else if !matches!(filter.allow, Targets::All)
&& is_private_ip(addr.ip())
&& !filter.allow.matches(&ip_target)
{
if first_denied.is_none() {
first_denied = Some(target);
}
} else {
allowed.push(addr);
}
}
if allowed.is_empty()
&& let Some(denied) = first_denied
{
return Err(Box::new(crate::err::Error::NetTargetNotAllowed(denied.to_string()))
as Box<dyn Error + Send + Sync>);
}
Ok(Box::new(allowed.into_iter()) as Addrs)
}) as Resolving
}
}
#[cfg(test)]
mod tests {
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
use std::str::FromStr;
use std::sync::Arc;
use reqwest::dns::{Name, Resolve};
use super::{FilteringResolver, NetFilter, is_private_ip};
use crate::dbs::capabilities::{NetTarget, Targets};
fn make_resolver(allow: Targets<NetTarget>, deny: Targets<NetTarget>) -> FilteringResolver {
FilteringResolver::from_net_filter(Arc::new(NetFilter {
allow,
deny,
}))
}
#[tokio::test]
async fn test_filtering_resolver_private_ip_via_hostname_blocked() {
let resolver = make_resolver(
Targets::Some([NetTarget::from_str("localhost").unwrap()].into()),
Targets::None,
);
let name = Name::from_str("localhost").unwrap();
let result = resolver.resolve(name).await;
match result {
Ok(_) => panic!(
"Expected FilteringResolver to block private IP resolved from an allowed hostname"
),
Err(e) => assert!(
e.to_string().contains("Access to network target"),
"Expected a NetTargetNotAllowed error, got: {e}"
),
}
}
#[tokio::test]
async fn test_filtering_resolver_private_ip_allowed_when_allow_all() {
let resolver = make_resolver(Targets::All, Targets::None);
let name = Name::from_str("localhost").unwrap();
let result = resolver.resolve(name).await;
assert!(result.is_ok(), "Expected resolution to succeed with allow_net = all");
}
#[test]
fn test_is_private_ip_loopback() {
assert!(is_private_ip(IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1))));
assert!(is_private_ip(IpAddr::V4(Ipv4Addr::new(127, 255, 255, 255))));
assert!(is_private_ip(IpAddr::V6(Ipv6Addr::LOCALHOST)));
}
#[test]
fn test_is_private_ip_rfc1918() {
assert!(is_private_ip(IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1))));
assert!(is_private_ip(IpAddr::V4(Ipv4Addr::new(172, 16, 0, 1))));
assert!(is_private_ip(IpAddr::V4(Ipv4Addr::new(172, 31, 255, 255))));
assert!(is_private_ip(IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1))));
}
#[test]
fn test_is_private_ip_link_local() {
assert!(is_private_ip(IpAddr::V4(Ipv4Addr::new(169, 254, 0, 1))));
assert!(is_private_ip(IpAddr::V4(Ipv4Addr::new(169, 254, 169, 254))));
assert!(is_private_ip(IpAddr::V6(Ipv6Addr::new(0xfe80, 0, 0, 0, 0, 0, 0, 1))));
}
#[test]
fn test_is_private_ip_shared_address_space() {
assert!(is_private_ip(IpAddr::V4(Ipv4Addr::new(100, 64, 0, 1))));
assert!(is_private_ip(IpAddr::V4(Ipv4Addr::new(100, 127, 255, 255))));
assert!(!is_private_ip(IpAddr::V4(Ipv4Addr::new(100, 128, 0, 0))));
}
#[test]
fn test_is_private_ip_unspecified() {
assert!(is_private_ip(IpAddr::V4(Ipv4Addr::UNSPECIFIED)));
assert!(is_private_ip(IpAddr::V6(Ipv6Addr::UNSPECIFIED)));
}
#[test]
fn test_is_private_ip_ipv6_unique_local() {
assert!(is_private_ip(IpAddr::V6(Ipv6Addr::new(0xfc00, 0, 0, 0, 0, 0, 0, 1))));
assert!(is_private_ip(IpAddr::V6(Ipv6Addr::new(0xfd00, 0, 0, 0, 0, 0, 0, 1))));
}
#[test]
fn test_is_private_ip_ipv4_mapped_ipv6() {
assert!(is_private_ip(IpAddr::V6(Ipv6Addr::new(0, 0, 0, 0, 0, 0xffff, 0x7f00, 0x0001)))); assert!(is_private_ip(IpAddr::V6(Ipv6Addr::new(0, 0, 0, 0, 0, 0xffff, 0xc0a8, 0x0101)))); assert!(is_private_ip(IpAddr::V6(Ipv6Addr::new(0, 0, 0, 0, 0, 0xffff, 0xac10, 0x0001)))); assert!(!is_private_ip(IpAddr::V6(Ipv6Addr::new(0, 0, 0, 0, 0, 0xffff, 0x0101, 0x0101)))); }
#[test]
fn test_is_private_ip_public_addresses() {
assert!(!is_private_ip(IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1))));
assert!(!is_private_ip(IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8))));
assert!(!is_private_ip(IpAddr::V4(Ipv4Addr::new(93, 184, 216, 34))));
assert!(!is_private_ip(IpAddr::V6(Ipv6Addr::new(
0x2001, 0x4860, 0x4860, 0, 0, 0, 0, 0x8888
))));
assert!(!is_private_ip(IpAddr::V4(Ipv4Addr::new(100, 128, 0, 1))));
}
}