use std::net::{IpAddr, Ipv6Addr, SocketAddr};
use tokio::net::TcpStream;
use crate::{BoxStream, ConnectError, TargetAddr, TargetHost};
pub fn is_reserved_or_private_ip(ip: &IpAddr) -> bool {
match ip {
IpAddr::V4(v4) => {
v4.is_loopback()
|| v4.is_link_local()
|| v4.is_private()
|| v4.is_unspecified()
|| v4.is_multicast()
|| v4.is_broadcast()
|| is_v4_documentation(v4)
|| is_v4_benchmarking(v4)
|| is_v4_reserved(v4)
|| is_v4_this_network(v4)
}
IpAddr::V6(v6) => {
if let Some(v4) = v6.to_ipv4_mapped() {
return is_reserved_or_private_ip(&IpAddr::V4(v4));
}
v6.is_loopback()
|| v6.is_unspecified()
|| v6.is_multicast()
|| is_v6_documentation(v6)
|| is_unicast_link_local_v6(v6)
|| is_unique_local_v6(v6)
|| is_v6_discard_prefix(v6)
}
}
}
fn is_unique_local_v6(ip: &Ipv6Addr) -> bool {
let octets = ip.octets();
(octets[0] & 0xfe) == 0xfc
}
fn is_unicast_link_local_v6(ip: &Ipv6Addr) -> bool {
let octets = ip.octets();
octets[0] == 0xfe && (octets[1] & 0xc0) == 0x80
}
fn is_v6_discard_prefix(ip: &Ipv6Addr) -> bool {
let octets = ip.octets();
octets[0] == 0x01 && octets[1..8].iter().all(|b| *b == 0)
}
fn is_v4_this_network(ip: &std::net::Ipv4Addr) -> bool {
ip.octets()[0] == 0
}
fn is_v4_documentation(ip: &std::net::Ipv4Addr) -> bool {
let octets = ip.octets();
matches!(
octets,
[192, 0, 2, _] | [198, 51, 100, _] | [203, 0, 113, _] | [192, 88, 99, _]
)
}
fn is_v4_benchmarking(ip: &std::net::Ipv4Addr) -> bool {
let octets = ip.octets();
octets[0] == 198 && (octets[1] == 18 || octets[1] == 19)
}
fn is_v4_reserved(ip: &std::net::Ipv4Addr) -> bool {
ip.octets()[0] >= 240
}
fn is_v6_documentation(ip: &Ipv6Addr) -> bool {
let octets = ip.octets();
octets[0] == 0x20 && octets[1] == 0x01 && octets[2] == 0x0d && octets[3] == 0xb8
}
pub fn is_dns_rebinding_risk(ip: &IpAddr) -> bool {
is_reserved_or_private_ip(ip)
}
#[trait_variant::make(Connector: Send)]
pub trait LocalConnector {
async fn connect(&self, target: &TargetAddr) -> Result<BoxStream, ConnectError>;
}
pub struct DirectConnector;
#[derive(Debug, Clone)]
pub struct ConnectOptions {
pub local_bind: Option<SocketAddr>,
pub enforce_dns_rebinding_check: bool,
pub enforce_literal_ip_check: bool,
}
impl Default for ConnectOptions {
fn default() -> Self {
Self {
local_bind: None,
enforce_dns_rebinding_check: true,
enforce_literal_ip_check: false,
}
}
}
impl DirectConnector {
pub async fn connect_with_options(
&self,
target: &TargetAddr,
options: &ConnectOptions,
) -> Result<BoxStream, ConnectError> {
let addrs = resolve_target(
target,
options.enforce_dns_rebinding_check,
options.enforce_literal_ip_check,
)
.await?;
connect_to_addrs(&addrs, options.local_bind).await
}
}
async fn connect_to_addrs(
addrs: &[SocketAddr],
local_bind: Option<SocketAddr>,
) -> Result<BoxStream, ConnectError> {
let mut last_error = None;
for &addr in addrs {
let result = if let Some(local) = local_bind {
let local = match local {
SocketAddr::V6(local) => local
.ip()
.to_ipv4_mapped()
.map(|ip| SocketAddr::new(ip.into(), local.port()))
.unwrap_or(local.into()),
local => local,
};
let socket = if local.is_ipv4() {
tokio::net::TcpSocket::new_v4()
} else {
tokio::net::TcpSocket::new_v6()
}
.map_err(ConnectError::Io)?;
socket.bind(local).map_err(ConnectError::Io)?;
socket.connect(addr).await.map_err(ConnectError::Io)
} else {
TcpStream::connect(addr).await.map_err(ConnectError::Io)
};
match result {
Ok(stream) => return Ok(Box::new(stream)),
Err(error) => last_error = Some(error),
}
}
Err(last_error.unwrap_or_else(|| ConnectError::DnsResolution("no addresses found".to_string())))
}
async fn resolve_target(
target: &TargetAddr,
enforce_dns_rebinding_check: bool,
enforce_literal_ip_check: bool,
) -> Result<Vec<SocketAddr>, ConnectError> {
match &target.host {
TargetHost::Ip(ip) => {
if enforce_literal_ip_check && is_dns_rebinding_risk(ip) {
return Err(ConnectError::ReservedTarget(*ip));
}
Ok(vec![SocketAddr::new(*ip, target.port)])
}
TargetHost::Domain(domain) => {
let lookup = format!("{}:{}", domain, target.port);
let addrs: Vec<_> = tokio::net::lookup_host(&lookup)
.await
.map_err(|e| ConnectError::DnsResolution(e.to_string()))?
.collect();
if addrs.is_empty() {
return Err(ConnectError::DnsResolution(
"no addresses found".to_string(),
));
}
if enforce_dns_rebinding_check {
if let Some(reserved) = addrs.iter().find(|addr| is_dns_rebinding_risk(&addr.ip()))
{
return Err(ConnectError::ReservedTarget(reserved.ip()));
}
}
Ok(addrs)
}
}
}
impl Connector for DirectConnector {
async fn connect(&self, target: &TargetAddr) -> Result<BoxStream, ConnectError> {
self.connect_with_options(target, &ConnectOptions::default())
.await
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::net::Ipv4Addr;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
#[tokio::test]
async fn test_direct_connect_echo() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let jh = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut buf = [0u8; 1024];
let n = stream.read(&mut buf).await.unwrap();
stream.write_all(&buf[..n]).await.unwrap();
});
let target = TargetAddr {
host: TargetHost::Ip(addr.ip()),
port: addr.port(),
};
let connector = DirectConnector;
let mut stream = Connector::connect(&connector, &target).await.unwrap();
stream.write_all(b"ping").await.unwrap();
let mut buf = [0u8; 4];
stream.read_exact(&mut buf).await.unwrap();
assert_eq!(&buf, b"ping");
jh.await.unwrap();
}
#[tokio::test]
async fn dns_rebinding_policy_applies_consistently_to_domains() {
let target = TargetAddr {
host: TargetHost::Domain("localhost".to_string()),
port: 80,
};
assert!(resolve_target(&target, false, false).await.is_ok());
assert!(ConnectOptions::default().enforce_dns_rebinding_check);
assert!(matches!(
resolve_target(
&target,
ConnectOptions::default().enforce_dns_rebinding_check,
ConnectOptions::default().enforce_literal_ip_check,
)
.await,
Err(ConnectError::ReservedTarget(_))
));
}
#[test]
fn reserved_ipv4_loopback() {
assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
127, 0, 0, 1
))));
}
#[test]
fn reserved_ipv4_private_10() {
assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
10, 0, 0, 1
))));
}
#[test]
fn reserved_ipv4_private_172() {
assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
172, 16, 0, 1
))));
}
#[test]
fn reserved_ipv4_private_192() {
assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
192, 168, 1, 1
))));
}
#[test]
fn reserved_ipv4_link_local() {
assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
169, 254, 1, 1
))));
}
#[test]
fn reserved_ipv4_unspecified() {
assert!(is_reserved_or_private_ip(&IpAddr::V4(
Ipv4Addr::UNSPECIFIED
)));
}
#[test]
fn not_reserved_ipv4_public() {
assert!(!is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
8, 8, 8, 8
))));
}
#[test]
fn reserved_ipv6_loopback() {
assert!(is_reserved_or_private_ip(&IpAddr::V6(Ipv6Addr::LOCALHOST)));
}
#[test]
fn reserved_ipv6_link_local() {
let ip = "fe80::1".parse::<Ipv6Addr>().unwrap();
assert!(is_reserved_or_private_ip(&IpAddr::V6(ip)));
}
#[test]
fn reserved_ipv4_mapped_ipv6() {
let ip = "::ffff:127.0.0.1".parse::<Ipv6Addr>().unwrap();
assert!(is_reserved_or_private_ip(&IpAddr::V6(ip)));
}
#[test]
fn reserved_ipv6_unique_local() {
let ip = "fd00::1".parse::<Ipv6Addr>().unwrap();
assert!(is_reserved_or_private_ip(&IpAddr::V6(ip)));
}
#[test]
fn reserved_ipv6_unspecified() {
assert!(is_reserved_or_private_ip(&IpAddr::V6(
Ipv6Addr::UNSPECIFIED
)));
}
#[test]
fn not_reserved_ipv6_public() {
let ip = "2606:4700:4700::1111".parse::<Ipv6Addr>().unwrap();
assert!(!is_reserved_or_private_ip(&IpAddr::V6(ip)));
}
#[test]
fn reserved_ipv4_multicast() {
assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
224, 0, 0, 1
))));
}
#[test]
fn reserved_ipv4_broadcast() {
assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::BROADCAST)));
}
#[test]
fn reserved_ipv4_documentation() {
assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
192, 0, 2, 1
))));
assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
198, 51, 100, 1
))));
assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
203, 0, 113, 1
))));
}
#[test]
fn reserved_ipv4_benchmarking() {
assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
198, 18, 0, 1
))));
}
#[test]
fn reserved_ipv4_reserved_future() {
assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
240, 0, 0, 1
))));
}
#[test]
fn reserved_ipv4_this_network() {
assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
0, 1, 2, 3
))));
}
#[test]
fn reserved_ipv6_multicast() {
let ip = "ff02::1".parse::<Ipv6Addr>().unwrap();
assert!(is_reserved_or_private_ip(&IpAddr::V6(ip)));
}
#[test]
fn reserved_ipv6_documentation() {
let ip = "2001:db8::1".parse::<Ipv6Addr>().unwrap();
assert!(is_reserved_or_private_ip(&IpAddr::V6(ip)));
}
#[test]
fn reserved_ipv6_discard_prefix() {
let ip = "0100::1".parse::<Ipv6Addr>().unwrap();
assert!(is_reserved_or_private_ip(&IpAddr::V6(ip)));
}
#[tokio::test]
async fn reject_domain_resolving_to_loopback() {
let connector = DirectConnector;
let target = TargetAddr {
host: TargetHost::Domain("localhost".to_string()),
port: 1,
};
let result = connector
.connect_with_options(
&target,
&ConnectOptions {
enforce_dns_rebinding_check: true,
..Default::default()
},
)
.await;
assert!(matches!(result, Err(ConnectError::ReservedTarget(_))));
}
#[tokio::test]
async fn direct_connect_falls_back_to_next_resolved_address() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let good_addr = listener.local_addr().unwrap();
let bad_addr = SocketAddr::new(good_addr.ip(), good_addr.port() + 1);
let accept = tokio::spawn(async move { listener.accept().await.unwrap() });
let stream = connect_to_addrs(&[bad_addr, good_addr], None)
.await
.expect("second resolved address should be attempted");
drop(stream);
accept.await.unwrap();
}
#[tokio::test]
async fn mapped_ipv6_local_bind_uses_ipv4_socket() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let accept = tokio::spawn(async move { listener.accept().await.unwrap() });
let mapped = SocketAddr::new("::ffff:127.0.0.1".parse().unwrap(), 0);
let stream = connect_to_addrs(&[addr], Some(mapped))
.await
.expect("mapped IPv6 local bind should connect to IPv4");
drop(stream);
accept.await.unwrap();
}
}