#![deny(missing_docs)]
use crate::ClientBuilder;
use std::{
collections::HashMap,
net::{IpAddr, SocketAddr},
str::FromStr,
sync::{
Arc, LazyLock,
atomic::{AtomicBool, Ordering::Relaxed},
},
time::Duration,
};
use arc_swap::ArcSwap;
use hickory_resolver::{
ConnectionProvider, Resolver,
config::{CLOUDFLARE, NameServerConfig, QUAD9, ResolverConfig, ResolverOpts},
net::{NetError, runtime::TokioRuntimeProvider},
};
use once_cell::sync::OnceCell;
use reqwest::dns::{Addrs, Name, Resolve, Resolving};
use tracing::*;
mod constants;
mod static_resolver;
mod trial;
pub(crate) use static_resolver::*;
pub(crate) const DEFAULT_POSITIVE_LOOKUP_CACHE_TTL: Duration = Duration::from_secs(1800);
pub(crate) const DEFAULT_OVERALL_LOOKUP_TIMEOUT: Duration = Duration::from_secs(10);
pub(crate) const DEFAULT_QUERY_TIMEOUT: Duration = Duration::from_secs(5);
impl ClientBuilder {
pub fn dns_resolver<R: Resolve + 'static>(mut self, resolver: Arc<R>) -> Self {
self = self.non_shared();
if let Some(rb) = self.reqwest_client_builder {
self.reqwest_client_builder = Some(rb.dns_resolver(resolver));
}
self.use_secure_dns = false;
self
}
pub fn no_hickory_dns(mut self) -> Self {
self = self.non_shared();
self.use_secure_dns = false;
self
}
}
static SHARED_RESOLVER: LazyLock<HickoryDnsResolver> = LazyLock::new(|| {
tracing::debug!("Initializing shared DNS resolver");
HickoryDnsResolver {
use_shared: false, ..Default::default()
}
});
pub trait SharedResolverState: ConnectionProvider + Default {
fn shared_resolver() -> Option<&'static HickoryDnsResolver<Self>> {
None
}
}
impl SharedResolverState for TokioRuntimeProvider {
fn shared_resolver() -> Option<&'static HickoryDnsResolver<Self>> {
Some(&SHARED_RESOLVER)
}
}
#[derive(Debug, thiserror::Error)]
#[allow(missing_docs)]
pub enum ResolveError {
#[error("invalid name: {0}")]
InvalidNameError(String),
#[error("hickory-dns resolver error: {0}")]
ResolveError(#[from] NetError),
#[error("high level lookup timed out")]
Timeout,
#[error("hostname not found in static lookup table")]
StaticLookupMiss,
}
impl ResolveError {
pub fn is_timeout(&self) -> bool {
matches!(
self,
ResolveError::Timeout | ResolveError::ResolveError(NetError::Timeout)
)
}
}
#[derive(Debug, Clone)]
pub struct HickoryDnsResolver<C: ConnectionProvider = TokioRuntimeProvider> {
state: Arc<ArcSwap<OnceCell<Resolver<C>>>>,
use_system: Arc<AtomicBool>,
system_resolver: Arc<OnceCell<Resolver<C>>>,
static_base: Option<Arc<OnceCell<StaticResolver>>>,
name_servers: Arc<ArcSwap<Vec<NameServerConfig>>>,
use_shared: bool,
overall_dns_timeout: Duration,
}
impl<C: ConnectionProvider> Default for HickoryDnsResolver<C> {
fn default() -> Self {
Self {
state: Default::default(),
use_system: Arc::new(AtomicBool::new(false)),
system_resolver: Default::default(),
static_base: Some(Default::default()),
name_servers: Arc::new(ArcSwap::from_pointee(default_nameserver_group_ipv4_only())),
use_shared: true,
overall_dns_timeout: DEFAULT_OVERALL_LOOKUP_TIMEOUT,
}
}
}
impl HickoryDnsResolver<TokioRuntimeProvider> {
pub fn new() -> Self {
Self::default()
}
}
impl<C: SharedResolverState> Resolve for HickoryDnsResolver<C> {
fn resolve(&self, name: Name) -> Resolving {
let use_system = self.use_system.load(std::sync::atomic::Ordering::Relaxed);
let use_shared = self.use_shared;
let result: Result<Resolver<C>, ResolveError> = if use_system {
self.system_resolver
.get_or_try_init(|| Self::new_resolver_system(use_shared))
.cloned()
} else {
self.build_configured_resolver()
};
let resolver = match result {
Ok(r) => r,
Err(err) => return Box::pin(return_err(err)),
};
let maybe_static = self.static_base.clone();
let overall_dns_timeout = self.overall_dns_timeout;
Box::pin(async move {
resolve(
name,
resolver,
maybe_static,
use_shared,
overall_dns_timeout,
)
.await
.map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)
})
}
}
async fn return_err(e: ResolveError) -> Result<Addrs, Box<dyn std::error::Error + Send + Sync>> {
Err(Box::new(e) as Box<dyn std::error::Error + Send + Sync>)
}
async fn resolve<C: SharedResolverState>(
name: Name,
resolver: Resolver<C>,
maybe_static: Option<Arc<OnceCell<StaticResolver>>>,
independent: bool,
overall_dns_timeout: Duration,
) -> Result<Addrs, ResolveError> {
if let Some(ref static_resolver) = maybe_static {
let resolver = static_resolver
.get_or_init(|| HickoryDnsResolver::<C>::new_static_fallback(independent));
if let Some(addrs) = resolver.pre_resolve(name.as_str()) {
let addrs: Addrs =
Box::new(addrs.into_iter().map(|ip_addr| SocketAddr::new(ip_addr, 0)));
return Ok(addrs);
}
}
let resolve_fut = tokio::time::timeout(overall_dns_timeout, resolver.lookup_ip(name.as_str()));
let primary_err = match resolve_fut.await {
Err(_) => ResolveError::Timeout,
Ok(Ok(lookup)) => {
let mut ips = Vec::from_iter(lookup.iter());
fastrand::shuffle(&mut ips);
let addrs: Addrs = Box::new(ips.into_iter().map(|ip| SocketAddr::new(ip, 0)));
return Ok(addrs);
}
Ok(Err(e)) => {
if !e.is_no_records_found() {
warn!("primary DNS failed w/ error: {e}");
}
e.into()
}
};
if let Some(ref static_resolver) = maybe_static {
debug!("checking static");
let resolver = static_resolver
.get_or_init(|| HickoryDnsResolver::<C>::new_static_fallback(independent));
if let Ok(addrs) = resolver.resolve(name).await {
return Ok(addrs);
}
}
Err(primary_err)
}
impl<C: SharedResolverState> HickoryDnsResolver<C> {
pub fn shared() -> Self {
C::shared_resolver().cloned().unwrap_or_default()
}
pub async fn resolve_str(
&self,
name: &str,
) -> Result<impl Iterator<Item = IpAddr> + use<C>, ResolveError> {
let n =
Name::from_str(name).map_err(|_| ResolveError::InvalidNameError(name.to_string()))?;
let use_system = self.use_system.load(std::sync::atomic::Ordering::Relaxed);
let resolver = if use_system {
self.system_resolver
.get_or_try_init(|| Self::new_resolver_system(self.use_shared))?
.clone()
} else {
self.build_configured_resolver()?
};
resolve(
n,
resolver,
self.static_base.clone(),
self.use_shared,
self.overall_dns_timeout,
)
.await
.map(|addrs| addrs.map(|socket_addr| socket_addr.ip()))
}
pub fn thread_resolver() -> Self {
Self {
use_shared: false,
..Default::default()
}
}
fn build_configured_resolver(&self) -> Result<Resolver<C>, ResolveError> {
match self.use_shared.then(C::shared_resolver).flatten() {
Some(shared) => shared
.state
.load()
.get_or_try_init(|| {
configure_and_build_resolver::<C>(
shared.name_servers.load_full().as_ref().clone(),
)
})
.cloned(),
None => self
.state
.load()
.get_or_try_init(|| {
configure_and_build_resolver::<C>(
self.name_servers.load_full().as_ref().clone(),
)
})
.cloned(),
}
}
fn new_resolver_system(use_shared: bool) -> Result<Resolver<C>, ResolveError> {
match use_shared.then(C::shared_resolver).flatten() {
Some(shared) => Ok(shared
.system_resolver
.get_or_try_init(new_resolver_system::<C>)?
.clone()),
None => new_resolver_system::<C>(),
}
}
fn new_static_fallback(use_shared: bool) -> StaticResolver {
match use_shared.then(C::shared_resolver).flatten() {
Some(shared) if shared.static_base.is_some() => shared
.static_base
.as_ref()
.unwrap()
.get_or_init(new_default_static_fallback)
.clone(),
_ => new_default_static_fallback(),
}
}
pub fn use_system_resolver(&self) {
self.use_system.store(true, Relaxed);
if let Some(shared) = self.use_shared.then(C::shared_resolver).flatten() {
shared.use_system_resolver();
}
}
pub fn use_configured_resolver(&self) {
self.use_system.store(false, Relaxed);
if let Some(shared) = self.use_shared.then(C::shared_resolver).flatten() {
shared.use_configured_resolver();
}
}
pub fn clear_preresolve(&self) {
debug!("clearing pre-resolve table");
if let Some(cell) = &self.static_base
&& let Some(static_base) = cell.get()
{
static_base.clear_preresolve()
}
}
pub fn get_static_fallbacks(&self) -> Option<HashMap<String, Vec<IpAddr>>> {
Some(self.static_base.as_ref()?.get()?.get_fallback_addrs())
}
pub fn set_fallback_addrs(&mut self, addrs: HashMap<String, Vec<IpAddr>>) {
debug!("setting fallback entries for {:?}", addrs.keys());
if self.static_base.is_none() {
let cell = OnceCell::new();
self.static_base = Some(Arc::new(cell));
}
self.static_base
.as_ref()
.unwrap()
.get_or_init(|| Self::new_static_fallback(self.use_shared))
.set_fallback(addrs);
}
pub fn get_static_preresolve(&self) -> Option<HashMap<String, Vec<IpAddr>>> {
Some(self.static_base.as_ref()?.get()?.get_preresolve_addrs())
}
pub fn set_static_preresolve(&mut self, addrs: HashMap<String, Vec<IpAddr>>) {
debug!("setting pre-resolve entries for {:?}", addrs.keys());
if self.static_base.is_none() {
let cell = OnceCell::new();
self.static_base = Some(Arc::new(cell));
}
self.static_base
.as_ref()
.unwrap()
.get_or_init(|| Self::new_static_fallback(self.use_shared))
.set_preresolve(addrs);
}
pub fn default_name_servers(&self) -> Vec<NameServerConfig> {
default_nameserver_group()
}
pub fn get_name_servers(&self) -> Vec<NameServerConfig> {
self.name_servers.load_full().as_ref().clone()
}
pub fn set_name_servers(&self, name_servers: Vec<NameServerConfig>) {
debug!("setting nameserver group to {name_servers:?}");
self.name_servers.store(Arc::new(name_servers));
self.state.store(Arc::new(OnceCell::new()));
if let Some(shared) = self.use_shared.then(C::shared_resolver).flatten() {
shared.name_servers.store(self.name_servers.load_full());
shared.state.store(Arc::new(OnceCell::new()));
}
}
}
fn default_options() -> ResolverOpts {
let mut opts = ResolverOpts::default();
opts.positive_min_ttl = Some(DEFAULT_POSITIVE_LOOKUP_CACHE_TTL);
opts.timeout = DEFAULT_QUERY_TIMEOUT;
opts.attempts = 0;
opts
}
fn configure_and_build_resolver<C: ConnectionProvider + Default>(
name_servers: Vec<NameServerConfig>,
) -> Result<Resolver<C>, ResolveError> {
let options = default_options();
info!("building new configured resolver");
debug!("configuring resolver with {options:?}, {name_servers:?}");
let config = ResolverConfig::from_parts(None, Vec::new(), name_servers);
let mut resolver_builder = Resolver::<C>::builder_with_config(config, C::default());
resolver_builder = resolver_builder.with_options(options);
Ok(resolver_builder.build()?)
}
fn filter_ipv4(nameservers: impl IntoIterator<Item = NameServerConfig>) -> Vec<NameServerConfig> {
nameservers
.into_iter()
.filter(|ns| ns.ip.is_ipv4())
.collect()
}
#[allow(unused)]
fn filter_ipv6(nameservers: impl IntoIterator<Item = NameServerConfig>) -> Vec<NameServerConfig> {
nameservers
.into_iter()
.filter(|ns| ns.ip.is_ipv6())
.collect()
}
fn default_nameserver_group() -> Vec<NameServerConfig> {
QUAD9
.tls()
.chain(QUAD9.https())
.chain(CLOUDFLARE.tls())
.chain(CLOUDFLARE.https())
.collect()
}
fn default_nameserver_group_ipv4_only() -> Vec<NameServerConfig> {
filter_ipv4(default_nameserver_group())
}
#[allow(unused)]
fn default_nameserver_group_ipv6_only() -> Vec<NameServerConfig> {
filter_ipv6(default_nameserver_group())
}
fn new_resolver_system<C: ConnectionProvider + Default>() -> Result<Resolver<C>, ResolveError> {
let mut resolver_builder = Resolver::<C>::builder(C::default())?;
let options = default_options();
info!("building new fallback system resolver");
debug!("fallback system resolver with {options:?}");
resolver_builder = resolver_builder.with_options(options);
Ok(resolver_builder.build()?)
}
fn new_default_static_fallback() -> StaticResolver {
StaticResolver::new().with_fallback(constants::default_static_addrs())
}
#[cfg(test)]
mod test;