use std::net::{IpAddr, Ipv4Addr, SocketAddr};
use std::time::Duration;
use socket2::{Domain, Protocol, Socket, Type};
use tokio::net::UdpSocket;
use crate::policy::dns;
pub const UDP_BUFFER_BYTES: usize = 4 * 1024 * 1024;
const SERVER_QUERY_ID: u16 = 0x5650;
const RESOLVE_BUF_BYTES: usize = 4096;
pub fn bind_udp(addr: SocketAddr) -> std::io::Result<UdpSocket> {
let domain = if addr.is_ipv4() {
Domain::IPV4
} else {
Domain::IPV6
};
let sock = Socket::new(domain, Type::DGRAM, Some(Protocol::UDP))?;
let _ = sock.set_recv_buffer_size(UDP_BUFFER_BYTES);
let _ = sock.set_send_buffer_size(UDP_BUFFER_BYTES);
sock.set_nonblocking(true)?;
sock.bind(&addr.into())?;
UdpSocket::from_std(sock.into())
}
pub async fn resolve_server(
server: &str,
upstreams: &[SocketAddr],
timeout: Duration,
) -> std::io::Result<SocketAddr> {
if let Ok(addr) = server.parse::<SocketAddr>() {
return Ok(addr);
}
let (host, port) = server
.rsplit_once(':')
.ok_or_else(|| resolve_err(format!("server `{server}` is missing a `:port`")))?;
let port: u16 = port
.parse()
.map_err(|_| resolve_err(format!("server `{server}` has an invalid port")))?;
if let Ok(ip) = host.parse::<IpAddr>() {
return Ok(SocketAddr::new(ip, port));
}
if upstreams.is_empty() {
return Err(resolve_err(
"no DNS upstreams configured to resolve the server".into(),
));
}
let query = dns::build_query(SERVER_QUERY_ID, host);
for &upstream in upstreams {
match query_a(upstream, &query, timeout).await {
Ok(Some(ip)) => return Ok(SocketAddr::new(IpAddr::V4(ip), port)),
Ok(None) => {
log::debug!("internal resolver: {upstream} returned no A record for {host}")
}
Err(e) => log::debug!("internal resolver: {upstream} failed for {host}: {e}"),
}
}
Err(resolve_err(format!(
"internal resolver could not resolve `{host}` via any upstream ({upstreams:?})"
)))
}
async fn query_a(
upstream: SocketAddr,
query: &[u8],
timeout: Duration,
) -> std::io::Result<Option<Ipv4Addr>> {
let bind: SocketAddr = if upstream.is_ipv4() {
(Ipv4Addr::UNSPECIFIED, 0).into()
} else {
(std::net::Ipv6Addr::UNSPECIFIED, 0).into()
};
let sock = UdpSocket::bind(bind).await?;
sock.connect(upstream).await?;
sock.send(query).await?;
let mut buf = vec![0u8; RESOLVE_BUF_BYTES];
let n = tokio::time::timeout(timeout, sock.recv(&mut buf))
.await
.map_err(|_| resolve_err(format!("DNS upstream {upstream} timed out")))??;
buf.truncate(n);
Ok(dns::a_records(&buf).into_iter().next())
}
fn resolve_err(msg: String) -> std::io::Error {
std::io::Error::other(msg)
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn literal_addr_short_circuits() {
let bogus = "192.0.2.1:9".parse().unwrap();
let got = resolve_server("157.245.227.200:443", &[bogus], Duration::from_millis(1))
.await
.unwrap();
assert_eq!(got, "157.245.227.200:443".parse().unwrap());
}
#[tokio::test]
async fn bare_ip_host_short_circuits() {
let bogus = "192.0.2.1:9".parse().unwrap();
let got = resolve_server("10.9.0.1:1234", &[bogus], Duration::from_millis(1))
.await
.unwrap();
assert_eq!(
got,
SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 9, 0, 1)), 1234)
);
}
#[tokio::test]
async fn missing_port_is_an_error() {
assert!(resolve_server("example.com", &[], Duration::from_millis(1))
.await
.is_err());
}
}