use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::time::Duration;
use async_trait::async_trait;
use dns_lattice_core::{Error, Result};
use dns_lattice_model::Message;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpStream, UdpSocket};
use tokio::time::timeout;
#[cfg(feature = "dot")]
mod dot;
#[cfg(feature = "dot")]
pub use dot::{DotBackend, DotBackendConfig};
#[cfg(feature = "doh")]
mod doh;
#[cfg(feature = "doh")]
pub use doh::{Doh3Backend, Doh3BackendConfig, DohBackend, DohBackendConfig, DohMethod};
#[cfg(feature = "doq")]
mod doq;
#[cfg(feature = "doq")]
pub(crate) use doq::QuicStream;
#[cfg(feature = "doq")]
pub use doq::{DoqBackend, DoqBackendConfig};
pub(crate) const UDP_MAX_RESPONSE_LEN: usize = 512;
#[async_trait]
pub trait UpstreamBackend: Send + Sync {
async fn resolve(&self, query: &Message) -> Result<Message>;
}
#[derive(Debug, Clone)]
pub struct UdpBackendConfig {
pub server: SocketAddr,
pub timeout: Duration,
pub bind_addr: Option<SocketAddr>,
}
pub struct UdpBackend {
config: UdpBackendConfig,
}
impl UdpBackend {
pub fn new(config: UdpBackendConfig) -> Self {
Self { config }
}
}
#[async_trait]
impl UpstreamBackend for UdpBackend {
async fn resolve(&self, query: &Message) -> Result<Message> {
let bind_addr = self
.config
.bind_addr
.unwrap_or_else(|| unspecified_like(self.config.server));
let socket = bind_udp(bind_addr, self.config.timeout).await?;
connect_udp(&socket, self.config.server, self.config.timeout).await?;
let payload = query.encode()?;
send_udp(&socket, &payload, self.config.timeout).await?;
let mut buf = [0u8; UDP_MAX_RESPONSE_LEN];
let len = recv_udp(&socket, &mut buf, self.config.timeout).await?;
let response = Message::decode(&buf[..len])?;
if response.header.truncated {
return tcp_query(
self.config.server,
self.config.timeout,
self.config.timeout,
query,
)
.await;
}
Ok(response)
}
}
#[derive(Debug, Clone)]
pub struct TcpBackendConfig {
pub server: SocketAddr,
pub connect_timeout: Duration,
pub read_timeout: Duration,
}
pub struct TcpBackend {
config: TcpBackendConfig,
}
impl TcpBackend {
pub fn new(config: TcpBackendConfig) -> Self {
Self { config }
}
}
#[async_trait]
impl UpstreamBackend for TcpBackend {
async fn resolve(&self, query: &Message) -> Result<Message> {
tcp_query(
self.config.server,
self.config.connect_timeout,
self.config.read_timeout,
query,
)
.await
}
}
fn unspecified_like(addr: SocketAddr) -> SocketAddr {
match addr {
SocketAddr::V4(_) => SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), 0),
SocketAddr::V6(_) => SocketAddr::new(IpAddr::V6(Ipv6Addr::UNSPECIFIED), 0),
}
}
async fn bind_udp(bind_addr: SocketAddr, budget: Duration) -> Result<UdpSocket> {
timeout(budget, UdpSocket::bind(bind_addr))
.await
.map_err(|_| Error::Timeout)?
.map_err(|err| Error::Transport(err.to_string()))
}
async fn connect_udp(socket: &UdpSocket, server: SocketAddr, budget: Duration) -> Result<()> {
timeout(budget, socket.connect(server))
.await
.map_err(|_| Error::Timeout)?
.map_err(|err| Error::Transport(err.to_string()))
}
async fn send_udp(socket: &UdpSocket, payload: &[u8], budget: Duration) -> Result<()> {
timeout(budget, socket.send(payload))
.await
.map_err(|_| Error::Timeout)?
.map_err(|err| Error::Transport(err.to_string()))?;
Ok(())
}
async fn recv_udp(socket: &UdpSocket, buf: &mut [u8], budget: Duration) -> Result<usize> {
timeout(budget, socket.recv(buf))
.await
.map_err(|_| Error::Timeout)?
.map_err(|err| Error::Transport(err.to_string()))
}
async fn tcp_query(
server: SocketAddr,
connect_timeout: Duration,
read_timeout: Duration,
query: &Message,
) -> Result<Message> {
let mut stream = timeout(connect_timeout, TcpStream::connect(server))
.await
.map_err(|_| Error::Timeout)?
.map_err(|err| Error::Transport(err.to_string()))?;
framed_query(&mut stream, read_timeout, query).await
}
pub(crate) async fn framed_query<S>(
stream: &mut S,
budget: Duration,
query: &Message,
) -> Result<Message>
where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
{
write_framed(stream, budget, query).await?;
read_framed(stream, budget).await
}
pub(crate) async fn write_framed<S>(
stream: &mut S,
budget: Duration,
message: &Message,
) -> Result<()>
where
S: tokio::io::AsyncWrite + Unpin,
{
let payload = message.encode()?;
let len: u16 = payload
.len()
.try_into()
.map_err(|_| Error::MessageTooLong)?;
let mut framed = Vec::with_capacity(payload.len() + 2);
framed.extend_from_slice(&len.to_be_bytes());
framed.extend_from_slice(&payload);
timeout(budget, stream.write_all(&framed))
.await
.map_err(|_| Error::Timeout)?
.map_err(|err| Error::Transport(err.to_string()))?;
Ok(())
}
pub(crate) async fn read_framed<S>(stream: &mut S, budget: Duration) -> Result<Message>
where
S: tokio::io::AsyncRead + Unpin,
{
let mut len_buf = [0u8; 2];
timeout(budget, stream.read_exact(&mut len_buf))
.await
.map_err(|_| Error::Timeout)?
.map_err(|err| Error::Transport(err.to_string()))?;
let response_len = u16::from_be_bytes(len_buf) as usize;
let mut response_buf = vec![0u8; response_len];
timeout(budget, stream.read_exact(&mut response_buf))
.await
.map_err(|_| Error::Timeout)?
.map_err(|err| Error::Transport(err.to_string()))?;
Message::decode(&response_buf)
}
#[cfg(test)]
mod tests {
use super::*;
use dns_lattice_model::{Class, Header, Name, Opcode, Question, Rcode, RecordType};
use tokio::net::TcpListener;
fn query_for(name: &str) -> Message {
Message {
header: Header {
id: 11,
qr: false,
opcode: Opcode::Query,
authoritative: false,
truncated: false,
recursion_desired: true,
recursion_available: false,
rcode: Rcode::NoError,
},
questions: vec![Question {
name: Name::from_ascii(name).unwrap(),
qtype: RecordType::A,
qclass: Class::In,
}],
answers: vec![],
authorities: vec![],
additionals: vec![],
}
}
fn answer_for(name: &str, id: u16) -> Message {
let mut msg = query_for(name);
msg.header.id = id;
msg.header.qr = true;
msg
}
#[tokio::test]
async fn udp_backend_resolves_against_a_loopback_server() {
let server = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let server_addr = server.local_addr().unwrap();
let responder = tokio::spawn(async move {
let mut buf = [0u8; 512];
let (len, from) = server.recv_from(&mut buf).await.unwrap();
let query = Message::decode(&buf[..len]).unwrap();
let response = answer_for("example.com", query.header.id);
server
.send_to(&response.encode().unwrap(), from)
.await
.unwrap();
});
let backend = UdpBackend::new(UdpBackendConfig {
server: server_addr,
timeout: Duration::from_secs(2),
bind_addr: None,
});
let answer = backend
.resolve(&query_for("example.com"))
.await
.expect("udp backend resolves");
assert!(answer.header.qr);
responder.await.unwrap();
}
#[tokio::test]
async fn udp_backend_falls_back_to_tcp_on_truncated_response() {
let udp_server = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let udp_addr = udp_server.local_addr().unwrap();
let tcp_listener = TcpListener::bind(udp_addr).await.unwrap();
let udp_responder = tokio::spawn(async move {
let mut buf = [0u8; 512];
let (len, from) = udp_server.recv_from(&mut buf).await.unwrap();
let query = Message::decode(&buf[..len]).unwrap();
let mut truncated = answer_for("example.com", query.header.id);
truncated.header.truncated = true;
udp_server
.send_to(&truncated.encode().unwrap(), from)
.await
.unwrap();
});
let tcp_responder = tokio::spawn(async move {
let (mut stream, _) = tcp_listener.accept().await.unwrap();
let mut len_buf = [0u8; 2];
stream.read_exact(&mut len_buf).await.unwrap();
let len = u16::from_be_bytes(len_buf) as usize;
let mut payload = vec![0u8; len];
stream.read_exact(&mut payload).await.unwrap();
let query = Message::decode(&payload).unwrap();
let response = answer_for("example.com", query.header.id);
let bytes = response.encode().unwrap();
let framed_len: u16 = bytes.len().try_into().unwrap();
let mut framed = Vec::new();
framed.extend_from_slice(&framed_len.to_be_bytes());
framed.extend_from_slice(&bytes);
stream.write_all(&framed).await.unwrap();
});
let backend = UdpBackend::new(UdpBackendConfig {
server: udp_addr,
timeout: Duration::from_secs(2),
bind_addr: None,
});
let answer = backend
.resolve(&query_for("example.com"))
.await
.expect("udp backend falls back to tcp on truncation");
assert!(!answer.header.truncated);
udp_responder.await.unwrap();
tcp_responder.await.unwrap();
}
#[tokio::test]
async fn tcp_backend_resolves_against_a_loopback_server() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let responder = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut len_buf = [0u8; 2];
stream.read_exact(&mut len_buf).await.unwrap();
let len = u16::from_be_bytes(len_buf) as usize;
let mut payload = vec![0u8; len];
stream.read_exact(&mut payload).await.unwrap();
let query = Message::decode(&payload).unwrap();
let response = answer_for("example.com", query.header.id);
let bytes = response.encode().unwrap();
let framed_len: u16 = bytes.len().try_into().unwrap();
let mut framed = Vec::new();
framed.extend_from_slice(&framed_len.to_be_bytes());
framed.extend_from_slice(&bytes);
stream.write_all(&framed).await.unwrap();
});
let backend = TcpBackend::new(TcpBackendConfig {
server: addr,
connect_timeout: Duration::from_secs(2),
read_timeout: Duration::from_secs(2),
});
let answer = backend
.resolve(&query_for("example.com"))
.await
.expect("tcp backend resolves");
assert!(answer.header.qr);
responder.await.unwrap();
}
#[tokio::test]
async fn udp_backend_times_out_when_server_never_responds() {
let server = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let server_addr = server.local_addr().unwrap();
let backend = UdpBackend::new(UdpBackendConfig {
server: server_addr,
timeout: Duration::from_millis(50),
bind_addr: None,
});
let err = backend
.resolve(&query_for("example.com"))
.await
.expect_err("no response within the timeout budget");
assert_eq!(err, Error::Timeout);
}
#[tokio::test]
async fn tcp_backend_returns_transport_when_peer_closes_connection() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let responder = tokio::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
drop(stream);
});
let backend = TcpBackend::new(TcpBackendConfig {
server: addr,
connect_timeout: Duration::from_secs(2),
read_timeout: Duration::from_secs(2),
});
let err = backend
.resolve(&query_for("example.com"))
.await
.expect_err("a peer that closes before a DNS response is transport failure");
assert!(matches!(err, Error::Transport(_)));
responder.await.unwrap();
}
}