use std::error::Error;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
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) => is_private_ipv4(v4),
IpAddr::V6(v6) => {
v6.is_loopback() || v6.is_unspecified() || (v6.segments()[0] & 0xFE00) == 0xFC00
|| (v6.segments()[0] & 0xFFC0) == 0xFE80
|| embedded_ipv4(&v6).is_some_and(is_private_ipv4)
}
}
}
fn is_private_ipv4(v4: Ipv4Addr) -> bool {
v4.is_loopback() || v4.is_private() || v4.is_link_local() || v4.is_broadcast() || v4.is_unspecified() || (u32::from(v4) & 0xFFC0_0000) == 0x6440_0000
}
fn embedded_ipv4(v6: &Ipv6Addr) -> Option<Ipv4Addr> {
let s = v6.segments();
let v4 = |hi: u16, lo: u16| Ipv4Addr::from(((hi as u32) << 16) | lo as u32);
match s {
[0x0064, 0xff9b, 0, 0, 0, 0, hi, lo] => Some(v4(hi, lo)),
[0x2002, hi, lo, ..] => Some(v4(hi, lo)),
[0x2001, 0x0000, .., hi, lo] => Some(v4(hi ^ 0xffff, lo ^ 0xffff)),
_ => None,
}
}
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::dbs::capabilities::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::dbs::capabilities::Error::NetTargetNotAllowed(
denied.to_string(),
)) as Box<dyn Error + Send + Sync>);
}
Ok(Box::new(allowed.into_iter()) as Addrs)
}) as Resolving
}
}
#[cfg(feature = "http")]
pub(crate) async fn resolve_net_target(
target: &NetTarget,
) -> Result<Vec<NetTarget>, std::io::Error> {
match target {
NetTarget::Host(h, p) => {
let port = p.unwrap_or(80);
let mut out = Vec::new();
for a in tokio::net::lookup_host((h.to_string(), port)).await? {
let ip = a.ip();
out.push(NetTarget::IPNet(ip.into()));
out.push(NetTarget::Host(net_target_host_from_ip(ip), Some(port)));
}
Ok(out)
}
NetTarget::IPNet(_) => Ok(vec![]),
}
}
#[cfg(feature = "http")]
fn net_target_host_from_ip(ip: IpAddr) -> url::Host<String> {
match ip.to_canonical() {
IpAddr::V4(v4) => url::Host::Ipv4(v4),
IpAddr::V6(v6) => url::Host::Ipv6(v6),
}
}
#[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, Target, 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_nat64_embedded_ipv4() {
assert!(is_private_ip(IpAddr::V6("64:ff9b::7f00:1".parse().unwrap())));
assert!(is_private_ip(IpAddr::V6("64:ff9b::a9fe:a9fe".parse().unwrap())));
assert!(is_private_ip(IpAddr::V6("64:ff9b::a00:1".parse().unwrap())));
assert!(!is_private_ip(IpAddr::V6("64:ff9b::101:101".parse().unwrap())));
}
#[test]
fn test_is_private_ip_6to4_embedded_ipv4() {
assert!(is_private_ip(IpAddr::V6("2002:a9fe:a9fe::".parse().unwrap())));
assert!(is_private_ip(IpAddr::V6("2002:7f00:1::".parse().unwrap())));
assert!(is_private_ip(IpAddr::V6("2002:c0a8:101::".parse().unwrap())));
assert!(!is_private_ip(IpAddr::V6("2002:101:101::".parse().unwrap())));
}
#[test]
fn test_is_private_ip_teredo_embedded_ipv4() {
assert!(is_private_ip(IpAddr::V6("2001:0:4136:e378:8000:63bf:f5ff:fffe".parse().unwrap())));
assert!(is_private_ip(IpAddr::V6("2001:0:4136:e378:8000:63bf:80ff:fffe".parse().unwrap())));
assert!(!is_private_ip(IpAddr::V6(
"2001:0:4136:e378:8000:63bf:fefe:fefe".parse().unwrap()
)));
}
#[test]
fn test_is_private_ip_non_transition_ipv6_unaffected() {
assert!(!is_private_ip(IpAddr::V6("2001:db8::1".parse().unwrap())));
assert!(!is_private_ip(IpAddr::V6("2003:7f00:1::".parse().unwrap())));
assert!(!is_private_ip(IpAddr::V6("64:ff9c::7f00:1".parse().unwrap())));
}
#[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))));
}
#[tokio::test]
#[cfg(feature = "http")]
async fn test_net_target_resolve_emits_port_bearing_form() {
let via_name = super::resolve_net_target(&NetTarget::from_str("localhost:9999").unwrap())
.await
.unwrap();
let ipv4_rule = NetTarget::from_str("127.0.0.1:9999").unwrap();
let ipv6_rule = NetTarget::from_str("[::1]:9999").unwrap();
assert!(
via_name.iter().any(|t| ipv4_rule.matches(t))
|| via_name.iter().any(|t| ipv6_rule.matches(t)),
"a port-bearing IP rule must match a hostname that resolves to it, got: {via_name:?}"
);
let via_ip = super::resolve_net_target(&NetTarget::from_str("127.0.0.1:9999").unwrap())
.await
.unwrap();
assert!(
via_ip.iter().any(|t| ipv4_rule.matches(t)),
"a port-bearing rule must match the resolved form of its own IP literal, got: {via_ip:?}"
);
let hostname_rule = NetTarget::from_str("example.com:9999").unwrap();
assert!(
!via_ip.iter().any(|t| hostname_rule.matches(t)),
"a hostname rule matches only the literal request host, not resolved IPs: {via_ip:?}"
);
let other_port = NetTarget::from_str("127.0.0.1:1234").unwrap();
assert!(
!via_ip.iter().any(|t| other_port.matches(t)),
"a port-bearing rule must not match a different port, got: {via_ip:?}"
);
let bare = NetTarget::from_str("127.0.0.1").unwrap();
assert!(
via_ip.iter().any(|t| bare.matches(t)),
"a bare-IP rule must still match, got: {via_ip:?}"
);
}
#[tokio::test]
#[cfg(feature = "http")]
async fn test_net_target_resolve_async() {
let r =
super::resolve_net_target(&NetTarget::from_str("localhost").unwrap()).await.unwrap();
let has_ipv4 = r.contains(&NetTarget::from_str("127.0.0.1").unwrap());
let has_ipv6 = r.contains(&NetTarget::from_str("::1/128").unwrap());
assert!(
has_ipv4 || has_ipv6,
"Expected localhost to resolve to at least 127.0.0.1 or ::1, got: {:?}",
r
);
}
}