use std::net::{IpAddr, SocketAddr};
use async_trait::async_trait;
use hickory_resolver::{
TokioResolver,
config::{ConnectionConfig, NameServerConfig, ResolverConfig, ResolverOpts},
net::runtime::TokioRuntimeProvider,
proto::rr::{Name, RData, rdata::TXT},
};
use crate::config::DnsConfig;
#[async_trait]
pub trait Resolver: Send + Sync {
async fn reverse(&self, ip: IpAddr) -> Result<Vec<String>, String>;
async fn forward(&self, name: &str) -> Result<Vec<IpAddr>, String>;
async fn txt(&self, name: &str) -> Result<Vec<String>, String>;
}
pub struct HickoryResolver {
inner: TokioResolver,
}
impl HickoryResolver {
pub fn from_system() -> anyhow::Result<Self> {
Self::build(None, |_| {})
}
pub fn from_system_uncached() -> anyhow::Result<Self> {
Self::build(None, |options| options.cache_size = 0)
}
pub fn from_address(addr: SocketAddr) -> anyhow::Result<Self> {
Self::build(Some(addr), |_| {})
}
pub fn from_address_uncached(addr: SocketAddr) -> anyhow::Result<Self> {
Self::build(Some(addr), |options| options.cache_size = 0)
}
fn build(
addr: Option<SocketAddr>,
configure: impl FnOnce(&mut ResolverOpts),
) -> anyhow::Result<Self> {
let mut builder = match addr {
None => TokioResolver::builder_tokio()
.map_err(|error| anyhow::anyhow!("reading system resolver config: {error}"))?,
Some(addr) => {
let mut udp = ConnectionConfig::udp();
udp.port = addr.port();
let mut tcp = ConnectionConfig::tcp();
tcp.port = addr.port();
let name_server = NameServerConfig::new(addr.ip(), true, vec![udp, tcp]);
let config = ResolverConfig::from_parts(None, vec![], vec![name_server]);
TokioResolver::builder_with_config(config, TokioRuntimeProvider::default())
}
};
configure(builder.options_mut());
let inner = builder
.build()
.map_err(|error| anyhow::anyhow!("building resolver: {error}"))?;
Ok(Self { inner })
}
}
#[async_trait]
impl Resolver for HickoryResolver {
async fn reverse(&self, ip: IpAddr) -> Result<Vec<String>, String> {
let lookup = match self.inner.reverse_lookup(Name::from(ip)).await {
Ok(lookup) => lookup,
Err(error) if error.is_no_records_found() => return Ok(Vec::new()),
Err(error) => return Err(error.to_string()),
};
Ok(lookup
.answers()
.iter()
.filter_map(|record| match &record.data {
RData::PTR(ptr) => Some(strip_root(&ptr.to_string())),
_ => None,
})
.collect())
}
async fn forward(&self, name: &str) -> Result<Vec<IpAddr>, String> {
let lookup = match self.inner.lookup_ip(name).await {
Ok(lookup) => lookup,
Err(error) if error.is_no_records_found() => return Ok(Vec::new()),
Err(error) => return Err(error.to_string()),
};
Ok(lookup.iter().collect())
}
async fn txt(&self, name: &str) -> Result<Vec<String>, String> {
let lookup = match self.inner.txt_lookup(name).await {
Ok(lookup) => lookup,
Err(error) if error.is_no_records_found() => return Ok(Vec::new()),
Err(error) => return Err(error.to_string()),
};
Ok(lookup
.answers()
.iter()
.filter_map(|record| match &record.data {
RData::TXT(txt) => Some(join_character_strings(txt)),
_ => None,
})
.collect())
}
}
fn join_character_strings(txt: &TXT) -> String {
let bytes: Vec<u8> = txt
.txt_data
.iter()
.flat_map(|chunk| chunk.iter().copied())
.collect();
String::from_utf8_lossy(&bytes).into_owned()
}
pub(crate) fn strip_root(name: &str) -> String {
name.strip_suffix('.').unwrap_or(name).to_string()
}
pub fn resolver_addr(dns: &DnsConfig) -> anyhow::Result<Option<SocketAddr>> {
match dns.resolver.as_deref() {
None | Some("") => Ok(None),
Some(addr) => addr.parse::<SocketAddr>().map(Some).map_err(|error| {
anyhow::anyhow!("dns.resolver {addr:?} is not a valid address: {error}")
}),
}
}
pub(crate) async fn connect(
resolver: &dyn Resolver,
host: &str,
port: u16,
) -> Result<tokio::net::TcpStream, String> {
let ips = match host.parse::<IpAddr>() {
Ok(ip) => vec![ip],
Err(_) => resolver.forward(host).await?,
};
if ips.is_empty() {
return Err(format!("no address found for {host}"));
}
let mut last_error = None;
for ip in ips {
match tokio::net::TcpStream::connect((ip, port)).await {
Ok(stream) => return Ok(stream),
Err(error) => last_error = Some(format!("{ip}: {error}")),
}
}
Err(last_error.unwrap_or_else(|| format!("no address found for {host}")))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn strip_root_removes_the_trailing_dot() {
assert_eq!(strip_root("host.example.com."), "host.example.com");
assert_eq!(strip_root("host.example.com"), "host.example.com");
}
#[test]
fn txt_character_strings_are_concatenated() {
let single = TXT::new(vec!["one-piece".to_string()]);
assert_eq!(join_character_strings(&single), "one-piece");
let split = TXT::new(vec!["first".to_string(), "second".to_string()]);
assert_eq!(join_character_strings(&split), "firstsecond");
assert_eq!(join_character_strings(&TXT::new(vec![])), "");
}
#[test]
fn non_utf8_txt_data_is_lossy_rather_than_fatal() {
let raw = TXT::from_bytes(vec![&[0xff, 0xfe]]);
let joined = join_character_strings(&raw);
assert!(!joined.is_empty());
assert_ne!(joined, "expected-digest");
}
#[test]
fn both_constructors_read_the_same_system_configuration() {
assert_eq!(
HickoryResolver::from_system().is_ok(),
HickoryResolver::from_system_uncached().is_ok(),
);
}
#[test]
fn from_address_builds_without_reading_system_configuration() {
let addr: SocketAddr = "127.0.0.1:5300".parse().unwrap();
assert!(HickoryResolver::from_address(addr).is_ok());
assert!(HickoryResolver::from_address_uncached(addr).is_ok());
}
struct UnreachableResolver;
#[async_trait]
impl Resolver for UnreachableResolver {
async fn reverse(&self, _ip: IpAddr) -> Result<Vec<String>, String> {
unreachable!("connect never looks up PTR records")
}
async fn forward(&self, _name: &str) -> Result<Vec<IpAddr>, String> {
unreachable!("a literal IP must short-circuit before this is called")
}
async fn txt(&self, _name: &str) -> Result<Vec<String>, String> {
unreachable!("connect never looks up TXT records")
}
}
#[tokio::test]
async fn connect_short_circuits_a_literal_ip() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
tokio::spawn(async move {
let _ = listener.accept().await;
});
assert!(
connect(&UnreachableResolver, "127.0.0.1", port)
.await
.is_ok()
);
}
struct StubForward(Vec<IpAddr>);
#[async_trait]
impl Resolver for StubForward {
async fn reverse(&self, _ip: IpAddr) -> Result<Vec<String>, String> {
unreachable!()
}
async fn forward(&self, _name: &str) -> Result<Vec<IpAddr>, String> {
Ok(self.0.clone())
}
async fn txt(&self, _name: &str) -> Result<Vec<String>, String> {
unreachable!()
}
}
#[tokio::test]
async fn connect_errors_when_forward_is_empty() {
let error = connect(&StubForward(vec![]), "example.com", 1234)
.await
.unwrap_err();
assert!(error.contains("example.com"), "{error}");
}
#[tokio::test]
async fn connect_falls_back_past_an_unreachable_first_address() {
let listener = tokio::net::TcpListener::bind("127.0.0.2:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
tokio::spawn(async move {
let _ = listener.accept().await;
});
let unreachable_first = "127.0.0.1".parse().unwrap();
let reachable_second = "127.0.0.2".parse().unwrap();
let resolver = StubForward(vec![unreachable_first, reachable_second]);
assert!(connect(&resolver, "example.com", port).await.is_ok());
}
#[tokio::test]
async fn connect_errors_when_every_address_refuses() {
let port = {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
listener.local_addr().unwrap().port()
};
let resolver = StubForward(vec!["127.0.0.1".parse().unwrap()]);
let error = connect(&resolver, "example.com", port).await.unwrap_err();
assert!(error.contains("127.0.0.1"), "{error}");
}
#[test]
fn resolver_addr_is_none_when_unset() {
assert!(resolver_addr(&DnsConfig::default()).unwrap().is_none());
}
#[test]
fn resolver_addr_treats_an_empty_string_as_unset() {
let dns = DnsConfig {
resolver: Some(String::new()),
};
assert!(resolver_addr(&dns).unwrap().is_none());
}
#[test]
fn resolver_addr_parses_a_valid_socket_address() {
let dns = DnsConfig {
resolver: Some("10.60.0.2:53".to_string()),
};
assert_eq!(
resolver_addr(&dns).unwrap(),
Some("10.60.0.2:53".parse().unwrap())
);
}
#[test]
fn resolver_addr_rejects_a_hostname_without_a_port() {
let dns = DnsConfig {
resolver: Some("not-an-address".to_string()),
};
let error = resolver_addr(&dns).unwrap_err().to_string();
assert!(error.contains("not-an-address"), "{error}");
}
mod loopback {
use super::*;
use hickory_proto::op::{Message, MessageType, OpCode, ResponseCode};
use hickory_proto::rr::rdata::{A, PTR};
use hickory_proto::rr::{DNSClass, Record, RecordType};
use hickory_proto::serialize::binary::{BinDecodable, BinEncodable};
use tokio::net::UdpSocket;
async fn spawn(answers: Vec<(&'static str, RecordType, RData)>) -> SocketAddr {
let socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let addr = socket.local_addr().unwrap();
tokio::spawn(async move {
let mut buffer = vec![0u8; 4096];
loop {
let Ok((read, peer)) = socket.recv_from(&mut buffer).await else {
return;
};
let Ok(request) = Message::from_bytes(&buffer[..read]) else {
continue;
};
let query = request.queries.first().cloned();
let mut response = Message::response(request.id, OpCode::Query);
response.metadata.message_type = MessageType::Response;
response.metadata.authoritative = true;
response.metadata.response_code = ResponseCode::NoError;
if let Some(query) = &query {
response.queries.push(query.clone());
for (name, record_type, data) in &answers {
let name = Name::from_utf8(name).unwrap();
if query.name() == &name && query.query_type() == *record_type {
let mut record = Record::from_rdata(name, 60, data.clone());
record.dns_class = DNSClass::IN;
response.answers.push(record);
}
}
}
let bytes = response.to_bytes().unwrap();
let _ = socket.send_to(&bytes, peer).await;
}
});
addr
}
#[tokio::test]
async fn reverse_returns_the_ptr_name_without_its_root_dot() {
let addr = spawn(vec![(
"10.2.0.192.in-addr.arpa.",
RecordType::PTR,
RData::PTR(PTR(Name::from_utf8("host.example.com.").unwrap())),
)])
.await;
let names = HickoryResolver::from_address(addr)
.unwrap()
.reverse("192.0.2.10".parse().unwrap())
.await
.unwrap();
assert_eq!(names, vec!["host.example.com".to_string()]);
}
#[tokio::test]
async fn forward_returns_the_addresses() {
let addr = spawn(vec![(
"host.example.com.",
RecordType::A,
RData::A(A("192.0.2.10".parse().unwrap())),
)])
.await;
let addresses = HickoryResolver::from_address(addr)
.unwrap()
.forward("host.example.com.")
.await
.unwrap();
assert!(
addresses.contains(&"192.0.2.10".parse::<IpAddr>().unwrap()),
"{addresses:?}"
);
}
#[tokio::test]
async fn txt_returns_the_concatenated_value() {
let addr = spawn(vec![(
"_acme-challenge.example.com.",
RecordType::TXT,
RData::TXT(TXT::new(vec!["first".to_string(), "second".to_string()])),
)])
.await;
let values = HickoryResolver::from_address_uncached(addr)
.unwrap()
.txt("_acme-challenge.example.com.")
.await
.unwrap();
assert_eq!(values, vec!["firstsecond".to_string()]);
}
#[tokio::test]
async fn an_empty_answer_is_no_records_rather_than_an_error() {
let addr = spawn(Vec::new()).await;
let resolver = HickoryResolver::from_address_uncached(addr).unwrap();
assert_eq!(
resolver
.reverse("192.0.2.10".parse().unwrap())
.await
.unwrap(),
Vec::<String>::new()
);
assert_eq!(
resolver.forward("nothing.example.com.").await.unwrap(),
Vec::<IpAddr>::new()
);
assert_eq!(
resolver.txt("nothing.example.com.").await.unwrap(),
Vec::<String>::new()
);
}
}
}