use std::{
collections::VecDeque,
fmt,
future::Future,
net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr},
pin::Pin,
sync::Arc,
};
use arc_swap::ArcSwap;
use iroh_base::EndpointId;
use n0_error::{AnyError, StackError, e, stack_error};
use n0_future::{
Either, MaybeFuture, Stream, StreamExt,
boxed::BoxFuture,
stream,
time::{self, Duration},
};
use tokio::sync::Notify;
use url::Url;
use crate::{ParseError, endpoint_info::EndpointInfo};
pub const DNS_TIMEOUT: Duration = Duration::from_secs(3);
pub const N0_DNS_ENDPOINT_ORIGIN_PROD: &str = "dns.iroh.link.";
pub const N0_DNS_ENDPOINT_ORIGIN_STAGING: &str = "staging-dns.iroh.link.";
#[cfg(any(target_os = "android", doc))]
pub unsafe fn install_android_jni_context(
java_vm: *mut std::ffi::c_void,
application_context: *mut std::ffi::c_void,
) {
#[cfg(target_os = "android")]
unsafe {
n0_dns_resolver::install_android_jni_context(java_vm, application_context);
}
#[cfg(not(target_os = "android"))]
let _ = (java_vm, application_context);
}
const MAX_JITTER_PERCENT: u64 = 20;
pub trait Resolver: fmt::Debug + Send + Sync + 'static {
fn lookup_ipv4(&self, host: String) -> BoxFuture<Result<BoxIter<Ipv4Addr>, DnsError>>;
fn lookup_ipv6(&self, host: String) -> BoxFuture<Result<BoxIter<Ipv6Addr>, DnsError>>;
fn lookup_txt(&self, host: String) -> BoxFuture<Result<BoxIter<TxtRecordData>, DnsError>>;
fn clear_cache(&self);
fn reset(&self) -> Box<dyn Resolver>;
}
pub type BoxIter<T> = Box<dyn Iterator<Item = T> + Send + 'static>;
#[allow(missing_docs)]
#[stack_error(derive, add_meta, std_sources)]
#[non_exhaustive]
pub enum DnsError {
#[error("Request timed out")]
Timeout {},
#[error("No response")]
NoResponse {},
#[error("Resolve failed, IPv4: {ipv4}, IPv6: {ipv6}")]
ResolveBoth {
ipv4: Box<DnsError>,
ipv6: Box<DnsError>,
},
#[error("Missing host")]
MissingHost {},
#[error("Failed to resolve")]
Resolve {
#[error(from)]
source: AnyError,
},
#[error("Invalid DNS response: not a query for _iroh.z32encodedpubkey")]
InvalidResponse {},
#[error("Domain name does not exist (NXDOMAIN)")]
NxDomain {},
}
#[allow(missing_docs)]
#[stack_error(derive, add_meta, from_sources)]
#[non_exhaustive]
pub enum LookupError {
#[error("Malformed txt from lookup")]
ParseError { source: ParseError },
#[error("Failed to resolve TXT record")]
LookupFailed { source: DnsError },
}
#[stack_error(derive, add_meta)]
#[error("no calls succeeded: [{}]", errors.iter().map(|e| e.to_string()).collect::<Vec<_>>().join(""))]
pub struct StaggeredError<E: n0_error::StackError + 'static> {
errors: Vec<E>,
}
impl<E: StackError + 'static> StaggeredError<E> {
pub fn iter(&self) -> impl Iterator<Item = &E> {
self.errors.iter()
}
}
#[derive(Debug, Clone, Default)]
pub struct Builder {
use_system_defaults: bool,
nameservers: Vec<NameserverConfig>,
fallback_mode: FallbackMode,
fallback_nameservers: Vec<NameserverConfig>,
#[cfg(not(wasm_browser))]
tls_client_config: Option<rustls::ClientConfig>,
}
#[derive(Debug, Default, Copy, Clone, Eq, PartialEq)]
#[non_exhaustive]
pub enum FallbackMode {
Never,
Eager,
IfSystemEmpty,
#[default]
Deferred,
}
impl FallbackMode {
fn to_resolver_mode(self) -> Option<n0_dns_resolver::FallbackMode> {
match self {
FallbackMode::Never => None,
FallbackMode::Eager => Some(n0_dns_resolver::FallbackMode::Eager),
FallbackMode::IfSystemEmpty => Some(n0_dns_resolver::FallbackMode::IfSystemEmpty),
FallbackMode::Deferred => Some(n0_dns_resolver::FallbackMode::Deferred),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct NameserverConfig {
addr: SocketAddr,
protocol: DnsProtocol,
server_name: Option<String>,
}
impl NameserverConfig {
fn new(addr: IpAddr, port: u16, protocol: DnsProtocol) -> Self {
Self {
addr: SocketAddr::new(addr, port),
protocol,
server_name: None,
}
}
pub fn udp(addr: IpAddr) -> Self {
Self::new(addr, 53, DnsProtocol::Udp)
}
pub fn tcp(addr: IpAddr) -> Self {
Self::new(addr, 53, DnsProtocol::Tcp)
}
pub fn tls(addr: IpAddr) -> Self {
Self::new(addr, 853, DnsProtocol::Tls)
}
pub fn https(addr: IpAddr) -> Self {
Self::new(addr, 443, DnsProtocol::Https)
}
pub fn with_port(self, port: u16) -> Self {
Self {
addr: SocketAddr::new(self.addr.ip(), port),
..self
}
}
pub fn with_tls_server_name(self, server_name: impl Into<String>) -> Self {
Self {
server_name: Some(server_name.into()),
..self
}
}
fn into_resolver_nameserver(self) -> n0_dns_resolver::Nameserver {
if let Some(server_name) = self.server_name {
return n0_dns_resolver::Nameserver::with_server_name(
self.addr,
self.protocol.to_resolver_protocol(),
server_name,
);
}
n0_dns_resolver::Nameserver::new(self.addr, self.protocol.to_resolver_protocol())
}
}
#[derive(Debug, Default, Copy, Clone, Eq, PartialEq)]
#[non_exhaustive]
pub enum DnsProtocol {
#[default]
Udp,
Tcp,
Tls,
Https,
}
impl DnsProtocol {
fn to_resolver_protocol(self) -> n0_dns_resolver::DnsProtocol {
match self {
DnsProtocol::Udp => n0_dns_resolver::DnsProtocol::Udp,
DnsProtocol::Tcp => n0_dns_resolver::DnsProtocol::Tcp,
DnsProtocol::Tls => n0_dns_resolver::DnsProtocol::Tls,
DnsProtocol::Https => n0_dns_resolver::DnsProtocol::Https,
}
}
}
impl Builder {
pub fn with_system_defaults(mut self) -> Self {
self.use_system_defaults = true;
self
}
#[deprecated(since = "1.2.0", note = "use add_nameserer_config()")]
pub fn with_nameserver(mut self, addr: SocketAddr, protocol: DnsProtocol) -> Self {
self.nameservers.push(NameserverConfig {
addr,
protocol,
server_name: None,
});
self
}
#[deprecated(since = "1.2.0", note = "use add_nameserer_configs()")]
pub fn with_nameservers(
mut self,
nameservers: impl IntoIterator<Item = (SocketAddr, DnsProtocol)>,
) -> Self {
self.nameservers.extend(
nameservers
.into_iter()
.map(|(addr, protocol)| NameserverConfig {
addr,
protocol,
server_name: None,
}),
);
self
}
pub fn add_nameserver_config(mut self, nameserver: NameserverConfig) -> Self {
self.nameservers.push(nameserver);
self
}
pub fn add_nameserver_configs(
mut self,
nameservers: impl IntoIterator<Item = NameserverConfig>,
) -> Self {
self.nameservers.extend(nameservers);
self
}
#[cfg(not(wasm_browser))]
pub fn tls_client_config(mut self, client_config: rustls::ClientConfig) -> Self {
self.tls_client_config = Some(client_config);
self
}
pub fn with_fallback_mode(mut self, mode: FallbackMode) -> Self {
self.fallback_mode = mode;
self
}
pub fn disable_fallback(self) -> Self {
self.with_fallback_mode(FallbackMode::Never)
}
pub fn fallback_nameserver_configs(
mut self,
nameservers: impl IntoIterator<Item = NameserverConfig>,
) -> Self {
self.fallback_nameservers.extend(nameservers);
self
}
pub fn build(self) -> DnsResolver {
DnsResolver::custom(DefaultResolver(Arc::new(
self.into_resolver_builder().build(),
)))
}
fn into_resolver_builder(self) -> n0_dns_resolver::Builder {
let mut builder = n0_dns_resolver::DnsResolver::builder().nameservers(
self.nameservers
.into_iter()
.map(NameserverConfig::into_resolver_nameserver),
);
if self.use_system_defaults {
builder = builder.use_system_config();
}
if let Some(mode) = self.fallback_mode.to_resolver_mode() {
builder = builder.fallback_mode(mode);
builder = if self.fallback_nameservers.is_empty() {
builder.default_fallback_nameservers()
} else {
builder.fallback_nameservers(
self.fallback_nameservers
.into_iter()
.map(NameserverConfig::into_resolver_nameserver),
)
};
}
#[cfg(not(wasm_browser))]
if let Some(tls_client_config) = self.tls_client_config {
builder = builder.tls_client_config(tls_client_config);
}
builder
}
}
#[derive(Debug)]
struct DefaultResolver(Arc<n0_dns_resolver::DnsResolver>);
impl Resolver for DefaultResolver {
fn lookup_ipv4(&self, host: String) -> BoxFuture<Result<BoxIter<Ipv4Addr>, DnsError>> {
let this = self.0.clone();
Box::pin(async move {
let list = this.lookup_ipv4(host).await.map_err(map_resolve_error)?;
let iter: BoxIter<_> = Box::new(list.into_iter());
Ok(iter)
})
}
fn lookup_ipv6(&self, host: String) -> BoxFuture<Result<BoxIter<Ipv6Addr>, DnsError>> {
let this = self.0.clone();
Box::pin(async move {
let list = this.lookup_ipv6(host).await.map_err(map_resolve_error)?;
let iter: BoxIter<_> = Box::new(list.into_iter());
Ok(iter)
})
}
fn lookup_txt(&self, host: String) -> BoxFuture<Result<BoxIter<TxtRecordData>, DnsError>> {
let this = self.0.clone();
Box::pin(async move {
let list = this.lookup_txt(host).await.map_err(map_resolve_error)?;
let iter: BoxIter<TxtRecordData> = Box::new(list.into_iter().map(convert_txt));
Ok(iter)
})
}
fn clear_cache(&self) {
self.0.clear_cache();
}
fn reset(&self) -> Box<dyn Resolver> {
Box::new(DefaultResolver(Arc::new(self.0.reset())))
}
}
fn map_resolve_error(err: n0_dns_resolver::Error) -> DnsError {
use n0_dns_resolver::Error as E;
match err {
E::Timeout { .. } => e!(DnsError::Timeout),
E::NoResponse { .. } => e!(DnsError::NoResponse),
E::NxDomain { .. } => e!(DnsError::NxDomain),
E::InvalidResponse { .. } => e!(DnsError::InvalidResponse),
other => e!(DnsError::Resolve, AnyError::from_stack(other)),
}
}
fn convert_txt(txt: n0_dns_resolver::TxtRecordData) -> TxtRecordData {
TxtRecordData(txt.into_boxed_slices())
}
#[derive(Debug, Clone)]
pub struct DnsResolver {
inner: Arc<Inner>,
}
#[derive(Debug)]
struct Inner {
notify_reset: Notify,
resolver: ArcSwap<Box<dyn Resolver>>,
}
impl Inner {
fn new(inner: Box<dyn Resolver>) -> Self {
Self {
notify_reset: Notify::new(),
resolver: ArcSwap::from_pointee(inner),
}
}
fn reset(&self) {
let current = self.resolver.load();
let new = Arc::new(current.reset());
let prev = self.resolver.compare_and_swap(¤t, new);
if Arc::ptr_eq(¤t, &prev) {
self.notify_reset.notify_waiters();
}
}
fn clear_cache(&self) {
self.resolver.load().clear_cache();
}
async fn op<F, Fut, R, E>(&self, timeout: Duration, f: F) -> Result<R, DnsError>
where
E: 'static + Send + Into<DnsError>,
R: 'static + Send,
F: 'static + Send + Fn(Arc<Box<dyn Resolver>>) -> Fut,
Fut: 'static + Send + Future<Output = Result<R, E>>,
{
loop {
let notified = self.notify_reset.notified();
tokio::pin!(notified);
notified.as_mut().enable();
let timeout = n0_future::time::sleep(timeout);
tokio::pin!(timeout);
let resolver = self.resolver.load_full();
let fut = f(resolver);
tokio::pin!(fut);
tokio::select! {
biased;
res = fut => return res.map_err(Into::into),
_ = notified => continue,
_ = timeout => return Err(e!(DnsError::Timeout)),
}
}
}
}
impl DnsResolver {
pub fn new() -> Self {
Builder::default().with_system_defaults().build()
}
pub fn with_nameserver(nameserver: SocketAddr) -> Self {
Builder::default()
.add_nameserver_config(
NameserverConfig::udp(nameserver.ip()).with_port(nameserver.port()),
)
.build()
}
pub fn builder() -> Builder {
Builder::default()
}
pub fn custom(resolver: impl Resolver) -> Self {
Self {
inner: Arc::new(Inner::new(Box::new(resolver))),
}
}
pub fn clear_cache(&self) {
self.inner.clear_cache();
}
pub fn reset(&self) {
self.inner.reset();
}
pub async fn lookup_txt<T: ToString>(
&self,
host: T,
timeout: Duration,
) -> Result<impl Iterator<Item = TxtRecordData>, DnsError> {
let host = host.to_string();
let res = self
.inner
.op(timeout, move |resolver| resolver.lookup_txt(host.clone()))
.await?;
Ok(res)
}
pub async fn lookup_ipv4<T: ToString>(
&self,
host: T,
timeout: Duration,
) -> Result<impl Iterator<Item = IpAddr> + use<T>, DnsError> {
let host = host.to_string();
let addrs = self
.inner
.op(timeout, move |resolver| resolver.lookup_ipv4(host.clone()))
.await?;
Ok(addrs.into_iter().map(IpAddr::V4))
}
pub async fn lookup_ipv6<T: ToString>(
&self,
host: T,
timeout: Duration,
) -> Result<impl Iterator<Item = IpAddr> + use<T>, DnsError> {
let host = host.to_string();
let addrs = self
.inner
.op(timeout, move |resolver| resolver.lookup_ipv6(host.clone()))
.await?;
Ok(addrs.into_iter().map(IpAddr::V6))
}
pub async fn lookup_ipv4_ipv6<T: ToString>(
&self,
host: T,
timeout: Duration,
) -> Result<impl Iterator<Item = IpAddr> + use<T>, DnsError> {
let host = host.to_string();
let res = tokio::join!(
self.lookup_ipv4(host.clone(), timeout),
self.lookup_ipv6(host, timeout)
);
match res {
(Ok(ipv4), Ok(ipv6)) => Ok(LookupIter::Both(ipv4.chain(ipv6))),
(Ok(ipv4), Err(_)) => Ok(LookupIter::Ipv4(ipv4)),
(Err(_), Ok(ipv6)) => Ok(LookupIter::Ipv6(ipv6)),
(Err(ipv4_err), Err(ipv6_err)) => Err(e!(DnsError::ResolveBoth {
ipv4: Box::new(ipv4_err),
ipv6: Box::new(ipv6_err)
})),
}
}
pub async fn resolve_host(
&self,
url: &Url,
prefer_ipv6: bool,
timeout: Duration,
) -> Result<IpAddr, DnsError> {
let host = url.host().ok_or_else(|| e!(DnsError::MissingHost))?;
match host {
url::Host::Domain(domain) => {
let lookup = tokio::join!(
self.lookup_ipv4(domain, timeout),
self.lookup_ipv6(domain, timeout)
);
let (v4, v6) = match lookup {
(Err(ipv4_err), Err(ipv6_err)) => {
return Err(e!(DnsError::ResolveBoth {
ipv4: Box::new(ipv4_err),
ipv6: Box::new(ipv6_err)
}));
}
(Err(_), Ok(mut v6)) => (None, v6.next()),
(Ok(mut v4), Err(_)) => (v4.next(), None),
(Ok(mut v4), Ok(mut v6)) => (v4.next(), v6.next()),
};
if prefer_ipv6 {
v6.or(v4).ok_or_else(|| e!(DnsError::NoResponse))
} else {
v4.or(v6).ok_or_else(|| e!(DnsError::NoResponse))
}
}
url::Host::Ipv4(ip) => Ok(IpAddr::V4(ip)),
url::Host::Ipv6(ip) => Ok(IpAddr::V6(ip)),
}
}
pub fn resolve_host_all<'a>(
&'a self,
url: &Url,
timeout: Duration,
) -> impl Stream<Item = Result<IpAddr, DnsError>> + Send + 'a {
let host = match url.host() {
None => {
return Either::Left(stream::once(Err(e!(DnsError::MissingHost))));
}
Some(url::Host::Ipv4(ip)) => {
return Either::Left(stream::once(Ok(IpAddr::V4(ip))));
}
Some(url::Host::Ipv6(ip)) => {
return Either::Left(stream::once(Ok(IpAddr::V6(ip))));
}
Some(url::Host::Domain(domain)) => domain.to_string(),
};
type Lookup<'a, A> =
Pin<Box<dyn Future<Output = Result<BoxIter<A>, DnsError>> + Send + 'a>>;
struct State<'a> {
v4_fut: MaybeFuture<Lookup<'a, Ipv4Addr>>,
v6_fut: MaybeFuture<Lookup<'a, Ipv6Addr>>,
v4_err: Option<DnsError>,
v6_err: Option<DnsError>,
queue: VecDeque<IpAddr>,
closed: bool,
yielded: bool,
}
let state = State {
v4_fut: MaybeFuture::Some(Box::pin({
let host = host.clone();
self.inner.op(timeout, move |r| r.lookup_ipv4(host.clone()))
})),
v6_fut: MaybeFuture::Some(Box::pin({
let host = host.clone();
self.inner.op(timeout, move |r| r.lookup_ipv6(host.clone()))
})),
v4_err: None,
v6_err: None,
queue: VecDeque::new(),
closed: false,
yielded: false,
};
Either::Right(stream::unfold(state, async |mut state| {
loop {
if state.closed {
return None;
}
if let Some(item) = state.queue.pop_front() {
state.yielded = true;
return Some((Ok(item), state));
}
if state.v4_fut.is_none() && state.v6_fut.is_none() {
state.closed = true;
if let (Some(v4), Some(v6)) = (state.v4_err.take(), state.v6_err.take()) {
let error = e!(DnsError::ResolveBoth {
ipv4: Box::new(v4),
ipv6: Box::new(v6),
});
return Some((Err(error), state));
} else if !state.yielded {
return Some((Err(e!(DnsError::NoResponse)), state));
} else {
return None;
}
}
tokio::select! {
biased;
res = &mut state.v4_fut => {
match res {
Ok(items) => state.queue.extend(items.map(IpAddr::V4)),
Err(err) => state.v4_err = Some(err),
}
}
res = &mut state.v6_fut => {
match res {
Ok(items) => state.queue.extend(items.map(IpAddr::V6)),
Err(err) => state.v6_err = Some(err),
}
}
}
}
}))
}
pub async fn lookup_ipv4_staggered(
&self,
host: impl ToString,
timeout: Duration,
delays_ms: &[u64],
) -> Result<impl Iterator<Item = IpAddr>, StaggeredError<DnsError>> {
let host = host.to_string();
let f = || self.lookup_ipv4(host.clone(), timeout);
stagger_call(f, delays_ms).await
}
pub async fn lookup_ipv6_staggered(
&self,
host: impl ToString,
timeout: Duration,
delays_ms: &[u64],
) -> Result<impl Iterator<Item = IpAddr>, StaggeredError<DnsError>> {
let host = host.to_string();
let f = || self.lookup_ipv6(host.clone(), timeout);
stagger_call(f, delays_ms).await
}
pub async fn lookup_ipv4_ipv6_staggered(
&self,
host: impl ToString,
timeout: Duration,
delays_ms: &[u64],
) -> Result<impl Iterator<Item = IpAddr>, StaggeredError<DnsError>> {
let host = host.to_string();
let f = || self.lookup_ipv4_ipv6(host.clone(), timeout);
stagger_call(f, delays_ms).await
}
pub async fn lookup_endpoint_by_id(
&self,
endpoint_id: &EndpointId,
origin: &str,
) -> Result<EndpointInfo, LookupError> {
let name = format!("_iroh.{}.{}", endpoint_id.to_z32(), origin);
let lookup = self.lookup_txt(name.clone(), DNS_TIMEOUT).await?;
let info = EndpointInfo::from_txt_lookup(name, lookup)?;
Ok(info)
}
pub async fn lookup_endpoint_by_domain_name(
&self,
name: &str,
) -> Result<EndpointInfo, LookupError> {
let name = if name.starts_with("_iroh.") {
name.to_string()
} else {
format!("_iroh.{name}")
};
let lookup = self.lookup_txt(name.clone(), DNS_TIMEOUT).await?;
let info = EndpointInfo::from_txt_lookup(name, lookup)?;
Ok(info)
}
pub async fn lookup_endpoint_by_domain_name_staggered(
&self,
name: &str,
delays_ms: &[u64],
) -> Result<EndpointInfo, StaggeredError<LookupError>> {
let f = || self.lookup_endpoint_by_domain_name(name);
stagger_call(f, delays_ms).await
}
pub async fn lookup_endpoint_by_id_staggered(
&self,
endpoint_id: &EndpointId,
origin: &str,
delays_ms: &[u64],
) -> Result<EndpointInfo, StaggeredError<LookupError>> {
let f = || self.lookup_endpoint_by_id(endpoint_id, origin);
stagger_call(f, delays_ms).await
}
}
impl Default for DnsResolver {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone)]
pub struct TxtRecordData(Box<[Box<[u8]>]>);
impl TxtRecordData {
pub fn iter(&self) -> impl Iterator<Item = &[u8]> {
self.0.iter().map(|x| x.as_ref())
}
}
impl fmt::Display for TxtRecordData {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
for s in self.iter() {
write!(f, "{}", String::from_utf8_lossy(s))?
}
Ok(())
}
}
impl FromIterator<Box<[u8]>> for TxtRecordData {
fn from_iter<T: IntoIterator<Item = Box<[u8]>>>(iter: T) -> Self {
Self(iter.into_iter().collect())
}
}
impl From<Vec<Box<[u8]>>> for TxtRecordData {
fn from(value: Vec<Box<[u8]>>) -> Self {
Self(value.into_boxed_slice())
}
}
impl From<Vec<String>> for TxtRecordData {
fn from(value: Vec<String>) -> Self {
Self(
value
.into_iter()
.map(|s| s.into_bytes().into_boxed_slice())
.collect(),
)
}
}
enum LookupIter<A, B> {
Ipv4(A),
Ipv6(B),
Both(std::iter::Chain<A, B>),
}
impl<A: Iterator<Item = IpAddr>, B: Iterator<Item = IpAddr>> Iterator for LookupIter<A, B> {
type Item = IpAddr;
fn next(&mut self) -> Option<Self::Item> {
match self {
LookupIter::Ipv4(iter) => iter.next(),
LookupIter::Ipv6(iter) => iter.next(),
LookupIter::Both(iter) => iter.next(),
}
}
}
async fn stagger_call<
T,
E: StackError + 'static,
F: Fn() -> Fut,
Fut: Future<Output = Result<T, E>>,
>(
f: F,
delays_ms: &[u64],
) -> Result<T, StaggeredError<E>> {
let mut calls = n0_future::FuturesUnorderedBounded::new(delays_ms.len() + 1);
for delay in std::iter::once(&0u64).chain(delays_ms) {
let delay = add_jitter(delay);
let fut = f();
let staggered_fut = async move {
time::sleep(delay).await;
fut.await
};
calls.push(staggered_fut)
}
let mut errors = vec![];
while let Some(call_result) = calls.next().await {
match call_result {
Ok(t) => return Ok(t),
Err(e) => errors.push(e),
}
}
Err(e!(StaggeredError { errors }))
}
fn add_jitter(delay: &u64) -> Duration {
if *delay == 0 {
return Duration::ZERO;
}
let max_jitter = delay.saturating_mul(MAX_JITTER_PERCENT * 2) / 100;
let jitter = rand::random::<u64>() % max_jitter;
Duration::from_millis(delay.saturating_sub(max_jitter / 2).saturating_add(jitter))
}
#[cfg(test)]
pub(crate) mod tests {
use std::sync::atomic::AtomicUsize;
use n0_tracing_test::traced_test;
use super::*;
#[test]
fn builder_named_nameservers_carry_server_name() {
let addr = SocketAddr::new(std::net::Ipv4Addr::new(1, 1, 1, 1).into(), 443);
let builder = Builder::default()
.add_nameserver_config(
NameserverConfig::https(addr.ip()).with_tls_server_name("cloudflare-dns.com"),
)
.add_nameserver_config(
NameserverConfig::tls(addr.ip())
.with_port(addr.port())
.with_tls_server_name("cloudflare-dns.com"),
)
.fallback_nameserver_configs([
NameserverConfig::https(addr.ip()).with_tls_server_name("fallback.example.com")
]);
let ns = &builder.nameservers;
assert_eq!(ns[0].protocol, DnsProtocol::Https);
assert_eq!(ns[0].addr.port(), 443);
assert_eq!(ns[0].server_name.as_deref(), Some("cloudflare-dns.com"));
assert_eq!(ns[1].protocol, DnsProtocol::Tls);
assert_eq!(ns[1].addr.port(), 443);
assert_eq!(ns[1].server_name.as_deref(), Some("cloudflare-dns.com"));
assert_eq!(
builder.fallback_nameservers[0].server_name.as_deref(),
Some("fallback.example.com")
);
}
fn resolver_nameservers(builder: Builder) -> Vec<SocketAddr> {
builder
.into_resolver_builder()
.build()
.configured_nameservers()
.iter()
.map(|ns| ns.addr())
.collect()
}
#[test]
fn empty_fallback_list_uses_public_resolvers() {
let addrs = resolver_nameservers(Builder::default());
assert!(!addrs.is_empty());
assert!(addrs.iter().any(|addr| addr.ip() == CLOUDFLARE_IP));
}
#[test]
fn disable_fallback_drops_the_public_resolvers() {
let addrs = resolver_nameservers(Builder::default().disable_fallback());
assert!(addrs.is_empty(), "{addrs:?}");
}
#[test]
fn explicit_fallback_list_replaces_the_defaults() {
let custom = SocketAddr::new(std::net::Ipv4Addr::new(192, 0, 2, 1).into(), 53);
let addrs = resolver_nameservers(
Builder::default().fallback_nameserver_configs([NameserverConfig::udp(custom.ip())]),
);
assert_eq!(addrs, vec![custom]);
}
const CLOUDFLARE_IP: IpAddr = IpAddr::V4(std::net::Ipv4Addr::new(1, 1, 1, 1));
#[tokio::test]
#[traced_test]
async fn stagger_basic() {
const CALL_RESULTS: &[Result<u8, u8>] = &[Err(2), Ok(3), Ok(5), Ok(7)];
static DONE_CALL: AtomicUsize = AtomicUsize::new(0);
let f = || {
let r_pos = DONE_CALL.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
async move {
tracing::info!(r_pos, "call");
CALL_RESULTS[r_pos].map_err(|_| e!(DnsError::InvalidResponse))
}
};
let delays = [1000, 15];
let result = stagger_call(f, &delays).await.unwrap();
assert_eq!(result, 5)
}
#[test]
#[traced_test]
fn jitter_test_zero() {
let jittered_delay = add_jitter(&0);
assert_eq!(jittered_delay, Duration::from_secs(0));
}
#[test]
#[traced_test]
fn jitter_test_nonzero_lower_bound() {
let delay: u64 = 300;
for _ in 0..100 {
assert!(add_jitter(&delay) >= Duration::from_millis(delay * 8 / 10));
}
}
#[test]
#[traced_test]
fn jitter_test_nonzero_upper_bound() {
let delay: u64 = 300;
for _ in 0..100 {
assert!(add_jitter(&delay) < Duration::from_millis(delay * 12 / 10));
}
}
#[tokio::test]
#[traced_test]
async fn custom_resolver() {
#[derive(Debug)]
struct MyResolver;
impl Resolver for MyResolver {
fn lookup_ipv4(&self, host: String) -> BoxFuture<Result<BoxIter<Ipv4Addr>, DnsError>> {
Box::pin(async move {
let addr = if host == "foo.example" {
Ipv4Addr::new(1, 1, 1, 1)
} else {
return Err(e!(DnsError::NoResponse));
};
let iter: BoxIter<Ipv4Addr> = Box::new(vec![addr].into_iter());
Ok(iter)
})
}
fn lookup_ipv6(&self, _host: String) -> BoxFuture<Result<BoxIter<Ipv6Addr>, DnsError>> {
todo!()
}
fn lookup_txt(
&self,
_host: String,
) -> BoxFuture<Result<BoxIter<TxtRecordData>, DnsError>> {
todo!()
}
fn clear_cache(&self) {
todo!()
}
fn reset(&self) -> Box<dyn Resolver> {
todo!()
}
}
let resolver = DnsResolver::custom(MyResolver);
let mut iter = resolver
.lookup_ipv4("foo.example", Duration::from_secs(1))
.await
.expect("not to fail");
let addr = iter.next().expect("one result");
assert_eq!(addr, "1.1.1.1".parse::<IpAddr>().unwrap());
let res = resolver
.lookup_ipv4("bar.example", Duration::from_secs(1))
.await;
assert!(matches!(res, Err(DnsError::NoResponse { .. })))
}
}