use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, ToSocketAddrs};
use thiserror::Error;
use url::Url;
#[non_exhaustive]
#[derive(Debug, Error)]
pub enum DomainFilterError {
#[error("URL scheme '{0}' is not allowed; only http and https are permitted")]
InvalidScheme(String),
#[error("Domain '{0}' is on the deny list")]
DeniedDomain(String),
#[error("Domain '{0}' is not on the allow list")]
NotAllowlisted(String),
#[error("Address {0} is a private/internal IP and is blocked")]
PrivateIp(String),
#[error("Failed to parse URL: {0}")]
InvalidUrl(String),
#[error("DNS resolution failed for '{0}': {1}")]
DnsError(String, String),
}
#[non_exhaustive]
#[derive(Debug, Clone, Default)]
pub struct DomainFilter {
pub allowlist: Vec<String>,
pub denylist: Vec<String>,
pub block_private_ips: bool,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct ResolvedHost {
pub host: String,
pub addr: SocketAddr,
}
impl DomainFilter {
#[must_use]
pub fn blocking_private_ips() -> Self {
Self {
block_private_ips: true,
..Self::default()
}
}
pub fn is_allowed(&self, url: &Url) -> Result<(), DomainFilterError> {
self.validate_and_resolve(url).map(|_| ())
}
pub(crate) fn validate_and_resolve(
&self,
url: &Url,
) -> Result<Option<ResolvedHost>, DomainFilterError> {
let scheme = url.scheme();
if scheme != "http" && scheme != "https" {
return Err(DomainFilterError::InvalidScheme(scheme.to_string()));
}
let host = url
.host_str()
.ok_or_else(|| DomainFilterError::InvalidUrl("URL has no host".to_string()))?;
if !self.allowlist.is_empty()
&& !self.allowlist.iter().any(|a| a.eq_ignore_ascii_case(host))
{
return Err(DomainFilterError::NotAllowlisted(host.to_string()));
}
if self.denylist.iter().any(|d| d.eq_ignore_ascii_case(host)) {
return Err(DomainFilterError::DeniedDomain(host.to_string()));
}
if self.block_private_ips {
match url.host() {
Some(url::Host::Ipv4(ip)) => {
if is_private_ip(&IpAddr::V4(ip)) {
return Err(DomainFilterError::PrivateIp(ip.to_string()));
}
}
Some(url::Host::Ipv6(ip)) => {
if is_private_ip(&IpAddr::V6(ip)) {
return Err(DomainFilterError::PrivateIp(ip.to_string()));
}
}
Some(url::Host::Domain(domain)) => {
let port = url.port_or_known_default().unwrap_or(80);
let mut first_public_addr = None;
let addrs = (domain, port).to_socket_addrs().map_err(|e| {
DomainFilterError::DnsError(domain.to_string(), e.to_string())
})?;
for addr in addrs {
if is_private_ip(&addr.ip()) {
return Err(DomainFilterError::PrivateIp(addr.ip().to_string()));
}
if first_public_addr.is_none() {
first_public_addr = Some(addr);
}
}
let addr = first_public_addr.ok_or_else(|| {
DomainFilterError::DnsError(
domain.to_string(),
"no addresses found".to_string(),
)
})?;
return Ok(Some(ResolvedHost {
host: domain.to_string(),
addr,
}));
}
None => {
return Err(DomainFilterError::InvalidUrl("URL has no host".to_string()));
}
}
}
Ok(None)
}
}
fn is_private_ip(ip: &IpAddr) -> bool {
match ip {
IpAddr::V4(v4) => is_private_ipv4(v4),
IpAddr::V6(v6) => is_private_ipv6(v6),
}
}
fn is_private_ipv4(ip: &Ipv4Addr) -> bool {
let octets = ip.octets();
if octets[0] == 0 {
return true;
}
if octets[0] == 127 {
return true;
}
if octets[0] == 10 {
return true;
}
if octets[0] == 172 && (16..=31).contains(&octets[1]) {
return true;
}
if octets[0] == 192 && octets[1] == 168 {
return true;
}
if octets[0] == 169 && octets[1] == 254 {
return true;
}
if octets[0] == 100 && (64..=127).contains(&octets[1]) {
return true;
}
if octets[0] == 198 && (18..=19).contains(&octets[1]) {
return true;
}
if (octets[0] == 192 && octets[1] == 0 && octets[2] == 2)
|| (octets[0] == 198 && octets[1] == 51 && octets[2] == 100)
|| (octets[0] == 203 && octets[1] == 0 && octets[2] == 113)
{
return true;
}
if octets[0] >= 224 {
return true;
}
false
}
fn is_private_ipv6(ip: &Ipv6Addr) -> bool {
if let Some(mapped) = ip.to_ipv4_mapped() {
return is_private_ipv4(&mapped);
}
if ip.is_unspecified() {
return true;
}
if ip.is_loopback() {
return true;
}
let segments = ip.segments();
if segments[0] & 0xfe00 == 0xfc00 {
return true;
}
if segments[0] & 0xffc0 == 0xfe80 {
return true;
}
if segments[0] & 0xff00 == 0xff00 {
return true;
}
if segments[0] == 0x2001 && segments[1] == 0x0db8 {
return true;
}
false
}
#[cfg(test)]
mod tests {
use super::{DomainFilter, DomainFilterError};
use url::Url;
#[test]
fn rejects_invalid_schemes() {
let filter = DomainFilter::default();
let file = Url::parse("file:///etc/passwd").unwrap();
let ftp = Url::parse("ftp://example.com/pub").unwrap();
assert!(matches!(
filter.is_allowed(&file).unwrap_err(),
DomainFilterError::InvalidScheme(_)
));
assert!(matches!(
filter.is_allowed(&ftp).unwrap_err(),
DomainFilterError::InvalidScheme(_)
));
}
#[test]
fn allowlist_and_denylist_are_enforced() {
let allow_filter = DomainFilter {
allowlist: vec!["example.com".to_string()],
..Default::default()
};
let deny_filter = DomainFilter {
denylist: vec!["evil.com".to_string()],
..Default::default()
};
assert!(
allow_filter
.is_allowed(&Url::parse("https://example.com/page").unwrap())
.is_ok()
);
assert!(matches!(
allow_filter
.is_allowed(&Url::parse("https://evil.com").unwrap())
.unwrap_err(),
DomainFilterError::NotAllowlisted(_)
));
assert!(matches!(
deny_filter
.is_allowed(&Url::parse("https://evil.com/malware").unwrap())
.unwrap_err(),
DomainFilterError::DeniedDomain(_)
));
}
#[test]
fn private_ip_ranges_are_blocked() {
let filter = DomainFilter::blocking_private_ips();
for url in [
"http://0.0.0.0/admin",
"http://127.0.0.1/admin",
"http://10.0.0.1/internal",
"http://100.64.0.1/cgnat",
"http://172.16.0.1/secret",
"http://192.168.1.1/router",
"http://198.18.0.1/benchmark",
"http://192.0.2.1/docs",
"http://224.0.0.1/multicast",
] {
assert!(filter.is_allowed(&Url::parse(url).unwrap()).is_err());
}
}
#[test]
fn bracketed_ipv6_private_literals_are_blocked_as_private_not_dns_error() {
let filter = DomainFilter::blocking_private_ips();
for url in [
"http://[::1]/admin",
"http://[::]/admin",
"http://[fd00::1]/internal",
"http://[fe80::1]/link-local",
"http://[2001:db8::1]/docs",
] {
let err = filter.is_allowed(&Url::parse(url).unwrap()).unwrap_err();
assert!(
matches!(err, DomainFilterError::PrivateIp(_)),
"{url} should be PrivateIp, got {err:?}"
);
}
}
#[test]
fn bracketed_public_ipv6_literal_is_allowed_without_dns_resolution() {
let filter = DomainFilter::blocking_private_ips();
let url = Url::parse("http://[2606:4700:4700::1111]/").unwrap();
let resolved = filter
.validate_and_resolve(&url)
.expect("public IPv6 literal should pass the filter");
assert!(resolved.is_none());
}
#[test]
fn ipv6_non_routable_ranges_are_private() {
for ip in [
"::",
"::1",
"fc00::1",
"fd00::1",
"fe80::1",
"ff02::1",
"2001:db8::1",
"::ffff:0.0.0.0",
"::ffff:10.0.0.1",
"::ffff:127.0.0.1",
"::ffff:169.254.0.1",
"::ffff:172.16.0.1",
"::ffff:192.168.1.1",
] {
assert!(
super::is_private_ip(&ip.parse().unwrap()),
"{ip} should be blocked"
);
}
for ip in ["2606:4700:4700::1111", "::ffff:93.184.216.34"] {
assert!(
!super::is_private_ip(&ip.parse().unwrap()),
"{ip} should be allowed"
);
}
}
}