use std::borrow::Cow;
use std::error::Error as StdError;
use std::fmt;
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use reqwest::dns::{Addrs, Name, Resolve, Resolving};
pub type EgressReason = Cow<'static, str>;
pub trait EgressPolicy: Send + Sync + fmt::Debug {
fn check_ip(&self, ip: IpAddr) -> Result<(), EgressReason>;
}
#[derive(Debug, Default)]
pub struct AllowAll;
impl EgressPolicy for AllowAll {
fn check_ip(&self, _ip: IpAddr) -> Result<(), EgressReason> {
Ok(())
}
}
#[derive(Debug)]
pub struct EgressDenied {
pub host: String,
pub ip: IpAddr,
pub reason: EgressReason,
}
impl fmt::Display for EgressDenied {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"egress blocked: host '{}' resolves to {} address {}",
self.host, self.reason, self.ip
)
}
}
impl StdError for EgressDenied {}
pub(crate) fn check_addrs(
policy: &dyn EgressPolicy,
host: &str,
addrs: Vec<SocketAddr>,
) -> Result<Vec<SocketAddr>, EgressDenied> {
for addr in &addrs {
let ip = addr.ip().to_canonical();
if let Err(reason) = policy.check_ip(ip) {
return Err(EgressDenied {
host: host.to_string(),
ip,
reason,
});
}
}
Ok(addrs)
}
#[derive(Debug)]
pub struct PolicyDns {
policy: Arc<dyn EgressPolicy>,
}
impl PolicyDns {
pub fn new(policy: Arc<dyn EgressPolicy>) -> Self {
Self { policy }
}
}
impl Resolve for PolicyDns {
fn resolve(&self, name: Name) -> Resolving {
let policy = Arc::clone(&self.policy);
Box::pin(async move {
let host = name.as_str().to_string();
let addrs: Vec<SocketAddr> =
tokio::net::lookup_host((host.as_str(), 0)).await?.collect();
let validated = check_addrs(policy.as_ref(), &host, addrs)?;
Ok(Box::new(validated.into_iter()) as Addrs)
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::net::Ipv4Addr;
#[derive(Debug)]
struct DenyAll;
impl EgressPolicy for DenyAll {
fn check_ip(&self, _ip: IpAddr) -> Result<(), EgressReason> {
Err("test-denied".into())
}
}
#[derive(Debug)]
struct DenyV4(Ipv4Addr);
impl EgressPolicy for DenyV4 {
fn check_ip(&self, ip: IpAddr) -> Result<(), EgressReason> {
if ip == IpAddr::V4(self.0) {
Err("test-denied".into())
} else {
Ok(())
}
}
}
#[test]
fn allow_all_permits_every_address() {
let policy = AllowAll;
for ip in ["127.0.0.1", "169.254.169.254", "10.0.0.1", "1.1.1.1"] {
policy
.check_ip(ip.parse().unwrap())
.unwrap_or_else(|_| panic!("AllowAll must permit {ip}"));
}
}
#[test]
fn allow_all_check_addrs_returns_addrs_unchanged() {
let policy = AllowAll;
let addrs: Vec<SocketAddr> = vec!["10.0.0.5:0".parse().unwrap()];
assert_eq!(check_addrs(&policy, "any", addrs.clone()).unwrap(), addrs);
}
#[test]
fn check_addrs_surfaces_denied_address_with_host_and_reason() {
let policy = DenyAll;
let addrs = vec!["10.0.0.5:0".parse().unwrap()];
let err = check_addrs(&policy, "evil.example", addrs).unwrap_err();
assert_eq!(err.host, "evil.example");
assert_eq!(err.ip, "10.0.0.5".parse::<IpAddr>().unwrap());
assert_eq!(err.reason, "test-denied");
}
#[test]
fn check_addrs_canonicalizes_mapped_ipv6_before_policy() {
let policy = DenyV4("10.0.0.1".parse().unwrap());
let addrs: Vec<SocketAddr> = vec!["[::ffff:10.0.0.1]:80".parse().unwrap()];
assert!(check_addrs(&policy, "evil.example", addrs).is_err());
}
#[test]
fn egress_denied_display_names_host_reason_and_ip() {
let err = EgressDenied {
host: "evil.example".to_string(),
ip: "10.0.0.5".parse().unwrap(),
reason: "private".into(),
};
assert_eq!(
err.to_string(),
"egress blocked: host 'evil.example' resolves to private address 10.0.0.5"
);
}
}