use std::{
error, fmt,
net::{IpAddr, SocketAddr},
};
use hickory_resolver::{ResolveError, TokioResolver, name_server::TokioConnectionProvider};
pub use hickory_resolver::{
IntoName, Name,
config::{
LookupIpStrategy, NameServerConfig, NameServerConfigGroup, ResolveHosts, ResolverConfig,
ResolverOpts, ServerOrderingStrategy,
},
lookup::{Lookup, ReverseLookup},
lookup_ip::LookupIp,
proto::{
rr::{RData, Record, RecordType},
xfer::Protocol,
},
};
pub type DnsResult<T> = std::result::Result<T, DnsError>;
#[derive(Debug, Clone)]
#[non_exhaustive]
pub enum DnsError {
NoNameServers,
Resolve(ResolveError),
}
impl DnsError {
#[must_use]
pub fn is_nx_domain(&self) -> bool {
match self {
Self::NoNameServers => false,
Self::Resolve(error) => error.is_nx_domain(),
}
}
#[must_use]
pub fn is_no_records_found(&self) -> bool {
match self {
Self::NoNameServers => false,
Self::Resolve(error) => error.is_no_records_found(),
}
}
#[must_use]
pub fn resolve_error(&self) -> Option<&ResolveError> {
match self {
Self::NoNameServers => None,
Self::Resolve(error) => Some(error),
}
}
}
impl fmt::Display for DnsError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::NoNameServers => formatter.write_str("at least one DNS name server is required"),
Self::Resolve(error) => write!(formatter, "DNS resolution failed: {error}"),
}
}
}
impl error::Error for DnsError {
fn source(&self) -> Option<&(dyn error::Error + 'static)> {
match self {
Self::NoNameServers => None,
Self::Resolve(error) => Some(error),
}
}
}
impl From<ResolveError> for DnsError {
fn from(error: ResolveError) -> Self {
Self::Resolve(error)
}
}
#[derive(Clone)]
pub struct DnsResolver {
inner: TokioResolver,
}
impl DnsResolver {
pub fn from_system() -> DnsResult<Self> {
let inner = TokioResolver::builder_tokio()?.build();
Ok(Self { inner })
}
#[must_use]
pub fn from_config(config: ResolverConfig, options: ResolverOpts) -> Self {
let mut builder =
TokioResolver::builder_with_config(config, TokioConnectionProvider::default());
*builder.options_mut() = options;
Self {
inner: builder.build(),
}
}
pub fn from_name_servers(
name_servers: impl IntoIterator<Item = SocketAddr>,
options: ResolverOpts,
) -> DnsResult<Self> {
let mut group = NameServerConfigGroup::new();
for socket_addr in name_servers {
group.push(NameServerConfig::new(socket_addr, Protocol::Udp));
group.push(NameServerConfig::new(socket_addr, Protocol::Tcp));
}
if group.is_empty() {
return Err(DnsError::NoNameServers);
}
Ok(Self::from_config(
ResolverConfig::from_parts(None, Vec::new(), group),
options,
))
}
#[must_use]
pub fn config(&self) -> &ResolverConfig {
self.inner.config()
}
#[must_use]
pub fn options(&self) -> &ResolverOpts {
self.inner.options()
}
pub fn clear_cache(&self) {
self.inner.clear_cache();
}
pub async fn lookup(&self, name: impl IntoName, record_type: RecordType) -> DnsResult<Lookup> {
self.inner
.lookup(name, record_type)
.await
.map_err(Into::into)
}
pub async fn lookup_ip(&self, host: impl IntoName) -> DnsResult<LookupIp> {
self.inner.lookup_ip(host).await.map_err(Into::into)
}
pub async fn reverse_lookup(&self, address: IpAddr) -> DnsResult<ReverseLookup> {
self.inner.reverse_lookup(address).await.map_err(Into::into)
}
}
impl fmt::Debug for DnsResolver {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("DnsResolver")
.field("config", self.config())
.field("options", self.options())
.finish_non_exhaustive()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn empty_name_server_list_is_rejected() {
let error = DnsResolver::from_name_servers([], ResolverOpts::default()).unwrap_err();
assert!(matches!(error, DnsError::NoNameServers));
assert!(error.resolve_error().is_none());
assert!(!error.is_nx_domain());
assert!(!error.is_no_records_found());
}
#[test]
fn explicit_configuration_is_observable() {
let mut options = ResolverOpts::default();
options.attempts = 1;
options.cache_size = 64;
let resolver =
DnsResolver::from_name_servers(["127.0.0.1:5353".parse().unwrap()], options).unwrap();
assert_eq!(resolver.config().name_servers().len(), 2);
assert_eq!(resolver.options().attempts, 1);
assert_eq!(resolver.options().cache_size, 64);
resolver.clear_cache();
}
}