use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr};
use std::sync::Arc;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, UdpSocket};
pub const DEFAULT_DNS_PORT: u16 = 15353;
const MAX_UDP_PAYLOAD: usize = 512;
const MAX_TCP_MESSAGE: usize = 65535;
const TTL: u32 = 60;
const ACCEPT_ERROR_BACKOFF: std::time::Duration = std::time::Duration::from_millis(100);
const TCP_IDLE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(10);
const REFUSAL_LOG_INTERVAL: std::time::Duration = std::time::Duration::from_secs(60);
static REFUSED_TCP: crate::proxy::LogThrottle = crate::proxy::LogThrottle::new();
const MAX_TCP_CONNECTIONS: usize = 64;
const TCP_CONNECTION_LIFETIME: std::time::Duration = std::time::Duration::from_secs(60);
const TYPE_A: u16 = 1;
const TYPE_AAAA: u16 = 28;
const CLASS_IN: u16 = 1;
const RCODE_NOERROR: u16 = 0;
const RCODE_FORMERR: u16 = 1;
const RCODE_NOTIMP: u16 = 4;
const RCODE_REFUSED: u16 = 5;
const FLAG_QR: u16 = 0x8000;
const FLAG_AA: u16 = 0x0400;
const FLAG_TC: u16 = 0x0200;
const FLAG_RD: u16 = 0x0100;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ResolverConfig {
pub tld: String,
pub ipv4: Option<Ipv4Addr>,
pub ipv6: Option<Ipv6Addr>,
}
impl ResolverConfig {
pub fn loopback(tld: impl Into<String>) -> Self {
Self {
tld: tld.into(),
ipv4: Some(Ipv4Addr::LOCALHOST),
ipv6: None,
}
}
pub fn lan(tld: impl Into<String>, ip: Ipv4Addr) -> Self {
Self {
tld: tld.into(),
ipv4: Some(ip),
ipv6: None,
}
}
pub fn for_bind(tld: impl Into<String>, bind_ip: std::net::IpAddr) -> Self {
let tld = tld.into();
match bind_ip {
std::net::IpAddr::V4(ip) if ip.is_unspecified() => Self {
tld,
ipv4: Some(Ipv4Addr::LOCALHOST),
ipv6: None,
},
std::net::IpAddr::V4(ip) => Self {
tld,
ipv4: Some(ip),
ipv6: None,
},
std::net::IpAddr::V6(ip) if ip.is_unspecified() => Self {
tld,
ipv4: Some(Ipv4Addr::LOCALHOST),
ipv6: Some(Ipv6Addr::LOCALHOST),
},
std::net::IpAddr::V6(ip) => Self {
tld,
ipv4: None,
ipv6: Some(ip),
},
}
}
fn owns(&self, name: &str) -> bool {
super::owns_name(&self.tld, name)
}
}
#[derive(Debug, PartialEq, Eq)]
struct Question {
name: String,
qtype: u16,
qclass: u16,
end: usize,
}
fn parse_question(msg: &[u8]) -> Option<Question> {
let mut pos = 12;
let mut name = String::new();
loop {
let len = *msg.get(pos)? as usize;
pos += 1;
if len == 0 {
break;
}
if len & 0xC0 != 0 {
return None;
}
let label = msg.get(pos..pos + len)?;
pos += len;
if !name.is_empty() {
name.push('.');
}
name.push_str(&String::from_utf8_lossy(label));
if name.len() > 255 {
return None;
}
}
let qtype = u16::from_be_bytes([*msg.get(pos)?, *msg.get(pos + 1)?]);
let qclass = u16::from_be_bytes([*msg.get(pos + 2)?, *msg.get(pos + 3)?]);
Some(Question {
name,
qtype,
qclass,
end: pos + 4,
})
}
fn header_only(id: u16, flags: u16, rcode: u16) -> Vec<u8> {
let mut out = Vec::with_capacity(12);
out.extend_from_slice(&id.to_be_bytes());
out.extend_from_slice(&(flags | rcode).to_be_bytes());
out.extend_from_slice(&0u16.to_be_bytes()); out.extend_from_slice(&0u16.to_be_bytes()); out.extend_from_slice(&0u16.to_be_bytes()); out.extend_from_slice(&0u16.to_be_bytes()); out
}
pub fn handle_query(query: &[u8], cfg: &ResolverConfig) -> Option<Vec<u8>> {
if query.len() < 12 {
return None;
}
let id = u16::from_be_bytes([query[0], query[1]]);
let req_flags = u16::from_be_bytes([query[2], query[3]]);
if req_flags & FLAG_QR != 0 {
return None;
}
let opcode = req_flags & 0x7800;
let qdcount = u16::from_be_bytes([query[4], query[5]]);
let base_flags = FLAG_QR | opcode | (req_flags & FLAG_RD);
if opcode != 0 {
return Some(header_only(id, base_flags, RCODE_NOTIMP));
}
if qdcount != 1 {
return Some(header_only(id, base_flags, RCODE_FORMERR));
}
let Some(q) = parse_question(query) else {
return Some(header_only(id, base_flags, RCODE_FORMERR));
};
let owned = q.qclass == CLASS_IN && cfg.owns(&q.name);
let answer = if !owned {
None
} else {
match q.qtype {
TYPE_A => cfg.ipv4.map(|ip| ip.octets().to_vec()),
TYPE_AAAA => cfg.ipv6.map(|ip| ip.octets().to_vec()),
_ => None,
}
};
let rcode = if owned { RCODE_NOERROR } else { RCODE_REFUSED };
let ancount: u16 = u16::from(answer.is_some());
let mut out = Vec::with_capacity(query.len() + 32);
out.extend_from_slice(&id.to_be_bytes());
let aa = if owned { FLAG_AA } else { 0 };
out.extend_from_slice(&(base_flags | aa | rcode).to_be_bytes());
out.extend_from_slice(&1u16.to_be_bytes()); out.extend_from_slice(&ancount.to_be_bytes());
out.extend_from_slice(&0u16.to_be_bytes()); out.extend_from_slice(&0u16.to_be_bytes()); out.extend_from_slice(&query[12..q.end]);
if let Some(rdata) = answer {
out.extend_from_slice(&[0xC0, 0x0C]);
out.extend_from_slice(&q.qtype.to_be_bytes());
out.extend_from_slice(&CLASS_IN.to_be_bytes());
out.extend_from_slice(&TTL.to_be_bytes());
out.extend_from_slice(&(rdata.len() as u16).to_be_bytes());
out.extend_from_slice(&rdata);
}
Some(out)
}
fn truncate_for_udp(mut resp: Vec<u8>) -> Vec<u8> {
if resp.len() <= MAX_UDP_PAYLOAD {
return resp;
}
let flags = u16::from_be_bytes([resp[2], resp[3]]) | FLAG_TC;
resp[2..4].copy_from_slice(&flags.to_be_bytes());
resp[6..8].copy_from_slice(&0u16.to_be_bytes());
resp.truncate(MAX_UDP_PAYLOAD);
resp
}
static ACTIVE_CONFIG: std::sync::RwLock<Option<Arc<std::sync::RwLock<ResolverConfig>>>> =
std::sync::RwLock::new(None);
pub fn update_lan_ip(ip: Ipv4Addr) {
let cfg = match ACTIVE_CONFIG.read() {
Ok(active) => active.clone(),
Err(e) => {
log::warn!("Could not read the active DNS resolver config: {e}");
return;
}
};
let Some(cfg) = cfg else {
return;
};
match cfg.write() {
Ok(mut cfg) if cfg.ipv4 != Some(ip) => {
log::info!("DNS resolver now answering *.{} with {ip}", cfg.tld);
cfg.ipv4 = Some(ip);
}
Ok(_) => {}
Err(e) => log::warn!("Could not update the DNS resolver address: {e}"),
}
}
pub async fn serve(
cfg: ResolverConfig,
addr: SocketAddr,
bind_tx: tokio::sync::oneshot::Sender<std::result::Result<(), String>>,
cancel: tokio_util::sync::CancellationToken,
) -> crate::Result<()> {
let udp = match UdpSocket::bind(addr).await {
Ok(s) => s,
Err(e) => {
let msg = format!("DNS resolver failed to bind UDP {addr}: {e}");
let _ = bind_tx.send(Err(msg.clone()));
miette::bail!("{msg}");
}
};
let tcp = match TcpListener::bind(addr).await {
Ok(l) => l,
Err(e) => {
let msg = format!("DNS resolver failed to bind TCP {addr}: {e}");
let _ = bind_tx.send(Err(msg.clone()));
miette::bail!("{msg}");
}
};
let _ = bind_tx.send(Ok(()));
{
let answers = [
cfg.ipv4.map(|ip| ip.to_string()),
cfg.ipv6.map(|ip| ip.to_string()),
]
.into_iter()
.flatten()
.collect::<Vec<_>>()
.join(", ");
log::info!(
"DNS resolver listening on {addr} (udp+tcp), answering *.{} with {answers}",
cfg.tld,
);
}
let cfg = Arc::new(std::sync::RwLock::new(cfg));
match ACTIVE_CONFIG.write() {
Ok(mut active) => *active = Some(Arc::clone(&cfg)),
Err(e) => log::warn!("Could not publish the DNS resolver config: {e}"),
}
fn answer(cfg: &std::sync::RwLock<ResolverConfig>, query: &[u8]) -> Option<Vec<u8>> {
match cfg.read() {
Ok(cfg) => handle_query(query, &cfg),
Err(e) => {
log::warn!("DNS resolver config lock poisoned: {e}");
None
}
}
}
let mut buf = vec![0u8; MAX_UDP_PAYLOAD];
let mut conns: tokio::task::JoinSet<()> = tokio::task::JoinSet::new();
loop {
while conns.try_join_next().is_some() {}
tokio::select! {
recv = udp.recv_from(&mut buf) => {
let (len, peer) = match recv {
Ok(v) => v,
Err(e) => {
log::debug!("DNS UDP receive error: {e}");
tokio::select! {
_ = tokio::time::sleep(ACCEPT_ERROR_BACKOFF) => continue,
_ = cancel.cancelled() => {
log::info!("DNS resolver shutting down");
break;
}
}
}
};
if let Some(resp) = answer(&cfg, &buf[..len])
&& let Err(e) = udp.send_to(&truncate_for_udp(resp), peer).await
{
log::debug!("DNS UDP send error to {peer}: {e}");
}
}
accept = tcp.accept() => {
let (stream, peer) = match accept {
Ok(v) => v,
Err(e) => {
log::debug!("DNS TCP accept error: {e}");
tokio::select! {
_ = tokio::time::sleep(ACCEPT_ERROR_BACKOFF) => continue,
_ = cancel.cancelled() => {
log::info!("DNS resolver shutting down");
break;
}
}
}
};
while conns.try_join_next().is_some() {}
if conns.len() >= MAX_TCP_CONNECTIONS {
if let Some(suppressed) = REFUSED_TCP.allow(REFUSAL_LOG_INTERVAL) {
log::warn!(
"DNS resolver refused a TCP connection from {peer}: \
{MAX_TCP_CONNECTIONS} already in flight \
({suppressed} similar refusals since the last message)"
);
}
drop(stream);
continue;
}
let cfg = Arc::clone(&cfg);
conns.spawn(async move {
match tokio::time::timeout(
TCP_CONNECTION_LIFETIME,
serve_tcp_conn(stream, &cfg, TCP_IDLE_TIMEOUT),
)
.await
{
Ok(Ok(())) => {}
Ok(Err(e)) => log::debug!("DNS TCP connection from {peer} ended: {e}"),
Err(_) => log::debug!(
"DNS TCP connection from {peer} closed after \
{TCP_CONNECTION_LIFETIME:?}"
),
}
});
}
_ = cancel.cancelled() => {
log::info!("DNS resolver shutting down");
break;
}
}
}
conns.abort_all();
if let Ok(mut active) = ACTIVE_CONFIG.write()
&& active.as_ref().is_some_and(|c| Arc::ptr_eq(c, &cfg))
{
*active = None;
}
Ok(())
}
async fn serve_tcp_conn(
mut stream: tokio::net::TcpStream,
cfg: &std::sync::RwLock<ResolverConfig>,
idle: std::time::Duration,
) -> std::io::Result<()> {
async fn read_exact_timeout(
stream: &mut tokio::net::TcpStream,
buf: &mut [u8],
idle: std::time::Duration,
) -> std::io::Result<()> {
tokio::time::timeout(idle, stream.read_exact(buf))
.await
.map_err(|_| {
std::io::Error::new(std::io::ErrorKind::TimedOut, "idle DNS connection")
})??;
Ok(())
}
loop {
let mut len_buf = [0u8; 2];
match read_exact_timeout(&mut stream, &mut len_buf, idle).await {
Ok(()) => {}
Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => return Ok(()),
Err(e) => return Err(e),
}
let len = usize::from(u16::from_be_bytes(len_buf));
if len == 0 || len > MAX_TCP_MESSAGE {
return Ok(());
}
let mut msg = vec![0u8; len];
read_exact_timeout(&mut stream, &mut msg, idle).await?;
let Some(resp) = (match cfg.read() {
Ok(cfg) => handle_query(&msg, &cfg),
Err(_) => None,
}) else {
continue;
};
let Ok(resp_len) = u16::try_from(resp.len()) else {
log::debug!(
"DNS reply of {} bytes cannot be framed over TCP; closing the connection",
resp.len()
);
return Ok(());
};
tokio::time::timeout(idle, async {
stream.write_all(&resp_len.to_be_bytes()).await?;
stream.write_all(&resp).await?;
stream.flush().await
})
.await
.map_err(|_| {
std::io::Error::new(
std::io::ErrorKind::TimedOut,
"DNS client not reading replies",
)
})??;
}
}
pub fn config_from_settings(
s: &crate::settings::Settings,
lan_ip: Option<Ipv4Addr>,
) -> ResolverConfig {
let lan_enabled = s.proxy.lan || !s.proxy.lan_ip.is_empty();
let tld = crate::proxy::effective_tld(s).to_string();
if lan_enabled {
return match lan_ip {
Some(ip) => ResolverConfig::lan(tld, ip),
None => ResolverConfig::loopback(tld),
};
}
let bind_ip = s
.proxy
.host
.parse()
.unwrap_or(std::net::IpAddr::V4(Ipv4Addr::LOCALHOST));
ResolverConfig::for_bind(tld, bind_ip)
}
pub fn dns_port(s: &crate::settings::Settings) -> u16 {
u16::try_from(s.proxy.dns_port)
.ok()
.filter(|&p| p > 0)
.unwrap_or_else(|| {
log::warn!(
"proxy.dns_port {} is out of valid port range (1-65535), using {DEFAULT_DNS_PORT}",
s.proxy.dns_port
);
DEFAULT_DNS_PORT
})
}
#[cfg(test)]
mod tests {
use super::*;
fn query(id: u16, name: &str, qtype: u16) -> Vec<u8> {
let mut out = Vec::new();
out.extend_from_slice(&id.to_be_bytes());
out.extend_from_slice(&FLAG_RD.to_be_bytes());
out.extend_from_slice(&1u16.to_be_bytes());
out.extend_from_slice(&0u16.to_be_bytes());
out.extend_from_slice(&0u16.to_be_bytes());
out.extend_from_slice(&0u16.to_be_bytes());
for label in name.split('.') {
out.push(label.len() as u8);
out.extend_from_slice(label.as_bytes());
}
out.push(0);
out.extend_from_slice(&qtype.to_be_bytes());
out.extend_from_slice(&CLASS_IN.to_be_bytes());
out
}
fn rcode(resp: &[u8]) -> u16 {
u16::from_be_bytes([resp[2], resp[3]]) & 0x000F
}
fn ancount(resp: &[u8]) -> u16 {
u16::from_be_bytes([resp[6], resp[7]])
}
fn rdata(resp: &[u8]) -> Vec<u8> {
let q = parse_question(resp).expect("response echoes the question");
let rdlen = usize::from(u16::from_be_bytes([resp[q.end + 10], resp[q.end + 11]]));
resp[q.end + 12..q.end + 12 + rdlen].to_vec()
}
fn cfg() -> ResolverConfig {
ResolverConfig::for_bind("localhost", "::".parse().unwrap())
}
#[test]
fn a_query_under_tld_answers_loopback() {
let resp = handle_query(&query(0x1234, "myapp.localhost", TYPE_A), &cfg()).unwrap();
assert_eq!(&resp[0..2], &0x1234u16.to_be_bytes());
assert_eq!(rcode(&resp), RCODE_NOERROR);
assert_eq!(ancount(&resp), 1);
assert_eq!(rdata(&resp), vec![127, 0, 0, 1]);
let flags = u16::from_be_bytes([resp[2], resp[3]]);
assert_eq!(flags & FLAG_QR, FLAG_QR);
assert_eq!(flags & FLAG_AA, FLAG_AA);
assert_eq!(flags & FLAG_RD, FLAG_RD);
}
#[test]
fn a_query_answers_multi_level_names() {
let resp = handle_query(
&query(1, "core.fix-refs.entiredb.localhost", TYPE_A),
&cfg(),
)
.unwrap();
assert_eq!(rcode(&resp), RCODE_NOERROR);
assert_eq!(rdata(&resp), vec![127, 0, 0, 1]);
}
#[test]
fn tld_apex_resolves() {
let resp = handle_query(&query(1, "localhost", TYPE_A), &cfg()).unwrap();
assert_eq!(rcode(&resp), RCODE_NOERROR);
assert_eq!(ancount(&resp), 1);
}
#[test]
fn matching_is_case_insensitive() {
let resp = handle_query(&query(1, "MyApp.LOCALHOST", TYPE_A), &cfg()).unwrap();
assert_eq!(rcode(&resp), RCODE_NOERROR);
assert_eq!(ancount(&resp), 1);
}
#[test]
fn aaaa_query_answers_ipv6_loopback() {
let resp = handle_query(&query(1, "myapp.localhost", TYPE_AAAA), &cfg()).unwrap();
assert_eq!(rcode(&resp), RCODE_NOERROR);
assert_eq!(ancount(&resp), 1);
assert_eq!(rdata(&resp), Ipv6Addr::LOCALHOST.octets().to_vec());
}
#[test]
fn lan_mode_answers_lan_ip_and_nodata_for_aaaa() {
let cfg = ResolverConfig::lan("local", Ipv4Addr::new(192, 168, 1, 42));
let a = handle_query(&query(1, "myapp.local", TYPE_A), &cfg).unwrap();
assert_eq!(rdata(&a), vec![192, 168, 1, 42]);
let aaaa = handle_query(&query(1, "myapp.local", TYPE_AAAA), &cfg).unwrap();
assert_eq!(rcode(&aaaa), RCODE_NOERROR);
assert_eq!(ancount(&aaaa), 0);
}
#[test]
fn name_outside_tld_is_refused_not_nxdomain() {
let resp = handle_query(&query(1, "example.com", TYPE_A), &cfg()).unwrap();
assert_eq!(rcode(&resp), RCODE_REFUSED);
assert_eq!(ancount(&resp), 0);
assert_eq!(u16::from_be_bytes([resp[2], resp[3]]) & FLAG_AA, 0);
}
#[test]
fn tld_suffix_without_label_boundary_is_refused() {
let resp = handle_query(&query(1, "notlocalhost", TYPE_A), &cfg()).unwrap();
assert_eq!(rcode(&resp), RCODE_REFUSED);
}
#[test]
fn lan_mode_serves_ipv4_whatever_proxy_host_says() {
let cfg = ResolverConfig::lan("local", Ipv4Addr::new(192, 168, 1, 42));
assert_eq!(cfg.ipv4, Some(Ipv4Addr::new(192, 168, 1, 42)));
assert_eq!(cfg.ipv6, None);
let fallback = ResolverConfig::loopback("local");
assert_eq!(fallback.ipv4, Some(Ipv4Addr::LOCALHOST));
assert_eq!(fallback.ipv6, None);
}
#[test]
fn answers_name_only_addresses_the_proxy_listens_on() {
use std::net::IpAddr;
let v6 = ResolverConfig::for_bind("test", IpAddr::V6(Ipv6Addr::LOCALHOST));
assert_eq!(v6.ipv4, None);
assert_eq!(v6.ipv6, Some(Ipv6Addr::LOCALHOST));
let a = handle_query(&query(1, "x.test", TYPE_A), &v6).unwrap();
assert_eq!(rcode(&a), RCODE_NOERROR);
assert_eq!(ancount(&a), 0);
let aaaa = handle_query(&query(1, "x.test", TYPE_AAAA), &v6).unwrap();
assert_eq!(rdata(&aaaa), Ipv6Addr::LOCALHOST.octets().to_vec());
let specific_v4 = ResolverConfig::for_bind("test", "192.168.1.5".parse().unwrap());
assert_eq!(specific_v4.ipv4, Some(Ipv4Addr::new(192, 168, 1, 5)));
assert_eq!(specific_v4.ipv6, None);
let specific_v6 = ResolverConfig::for_bind("test", "fd00::1".parse().unwrap());
assert_eq!(specific_v6.ipv4, None);
assert_eq!(specific_v6.ipv6, Some("fd00::1".parse().unwrap()));
let any_v4 = ResolverConfig::for_bind("test", "0.0.0.0".parse().unwrap());
assert_eq!(any_v4.ipv4, Some(Ipv4Addr::LOCALHOST));
assert_eq!(any_v4.ipv6, None);
let any_v6 = ResolverConfig::for_bind("test", "::".parse().unwrap());
assert_eq!(any_v6.ipv4, Some(Ipv4Addr::LOCALHOST));
assert_eq!(any_v6.ipv6, Some(Ipv6Addr::LOCALHOST));
}
#[test]
fn aaaa_is_nodata_unless_the_proxy_listens_on_ipv6() {
let v4_only = ResolverConfig::loopback("localhost");
let resp = handle_query(&query(1, "myapp.localhost", TYPE_AAAA), &v4_only).unwrap();
assert_eq!(rcode(&resp), RCODE_NOERROR);
assert_eq!(ancount(&resp), 0);
let a = handle_query(&query(1, "myapp.localhost", TYPE_A), &v4_only).unwrap();
assert_eq!(ancount(&a), 1);
}
#[test]
fn unsupported_record_type_under_tld_is_nodata() {
const TYPE_MX: u16 = 15;
let resp = handle_query(&query(1, "myapp.localhost", TYPE_MX), &cfg()).unwrap();
assert_eq!(rcode(&resp), RCODE_NOERROR);
assert_eq!(ancount(&resp), 0);
}
#[test]
fn non_internet_class_is_refused() {
let mut q = query(1, "myapp.localhost", TYPE_A);
let len = q.len();
q[len - 2..].copy_from_slice(&3u16.to_be_bytes()); let resp = handle_query(&q, &cfg()).unwrap();
assert_eq!(rcode(&resp), RCODE_REFUSED);
}
#[test]
fn malformed_and_unsupported_messages() {
assert!(handle_query(&[0u8; 4], &cfg()).is_none());
let mut resp_msg = query(1, "myapp.localhost", TYPE_A);
resp_msg[2] |= 0x80;
assert!(handle_query(&resp_msg, &cfg()).is_none());
let q = query(1, "myapp.localhost", TYPE_A);
let resp = handle_query(&q[..16], &cfg()).unwrap();
assert_eq!(rcode(&resp), RCODE_FORMERR);
let mut upd = query(1, "myapp.localhost", TYPE_A);
upd[2] |= 5 << 3;
let resp = handle_query(&upd, &cfg()).unwrap();
assert_eq!(rcode(&resp), RCODE_NOTIMP);
}
#[test]
fn compression_pointer_in_question_is_rejected() {
let mut q = query(1, "myapp.localhost", TYPE_A);
q[12] = 0xC0;
let resp = handle_query(&q, &cfg()).unwrap();
assert_eq!(rcode(&resp), RCODE_FORMERR);
}
#[tokio::test]
async fn an_idle_tcp_client_is_dropped_rather_than_held() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let idle = std::time::Duration::from_millis(50);
let server = tokio::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
serve_tcp_conn(stream, &std::sync::RwLock::new(cfg()), idle).await
});
let mut client = tokio::net::TcpStream::connect(addr).await.unwrap();
client.write_all(&16u16.to_be_bytes()).await.unwrap();
let err = tokio::time::timeout(std::time::Duration::from_secs(5), server)
.await
.expect("the handler should give up on its own")
.unwrap()
.expect_err("an idle connection is an error, not a clean close");
assert_eq!(err.kind(), std::io::ErrorKind::TimedOut);
let mut buf = [0u8; 1];
let n = tokio::time::timeout(std::time::Duration::from_secs(5), client.read(&mut buf))
.await
.expect("the connection is already closed")
.unwrap();
assert_eq!(n, 0);
}
#[test]
fn the_tcp_connection_cap_is_bounded() {
assert!((8..=1024).contains(&MAX_TCP_CONNECTIONS));
}
#[tokio::test]
async fn serves_over_udp_and_tcp() {
let cancel = tokio_util::sync::CancellationToken::new();
let (tx, rx) = tokio::sync::oneshot::channel();
let probe = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = probe.local_addr().unwrap();
drop(probe);
let task = tokio::spawn({
let cancel = cancel.clone();
async move { serve(cfg(), addr, tx, cancel).await }
});
rx.await.unwrap().expect("resolver binds");
let q = query(0x4242, "deep.nested.myapp.localhost", TYPE_A);
let sock = UdpSocket::bind("127.0.0.1:0").await.unwrap();
sock.send_to(&q, addr).await.unwrap();
let mut buf = [0u8; 512];
let (n, _) =
tokio::time::timeout(std::time::Duration::from_secs(5), sock.recv_from(&mut buf))
.await
.expect("udp reply arrives")
.unwrap();
assert_eq!(rdata(&buf[..n]), vec![127, 0, 0, 1]);
let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
stream
.write_all(&(q.len() as u16).to_be_bytes())
.await
.unwrap();
stream.write_all(&q).await.unwrap();
let mut len_buf = [0u8; 2];
stream.read_exact(&mut len_buf).await.unwrap();
let mut resp = vec![0u8; usize::from(u16::from_be_bytes(len_buf))];
stream.read_exact(&mut resp).await.unwrap();
assert_eq!(rdata(&resp), vec![127, 0, 0, 1]);
stream
.write_all(&(q.len() as u16).to_be_bytes())
.await
.unwrap();
stream.write_all(&q).await.unwrap();
stream.read_exact(&mut len_buf).await.unwrap();
let mut resp2 = vec![0u8; usize::from(u16::from_be_bytes(len_buf))];
stream.read_exact(&mut resp2).await.unwrap();
assert_eq!(rcode(&resp2), RCODE_NOERROR);
drop(stream);
cancel.cancel();
tokio::time::timeout(std::time::Duration::from_secs(5), task)
.await
.expect("the resolver did not stop within 5s of cancellation")
.expect("the resolver task panicked")
.expect("the resolver returned an error");
}
}