use std::fmt::Debug;
use std::future::Future;
use std::net::IpAddr;
use std::panic::{RefUnwindSafe, UnwindSafe};
use std::pin::Pin;
pub type DnsError = Box<dyn std::error::Error + Send + Sync>;
pub type DnsFuture = Pin<Box<dyn Future<Output = Result<Vec<IpAddr>, DnsError>> + Send>>;
pub trait DnsResolver: Debug + Send + Sync + UnwindSafe + RefUnwindSafe {
fn resolve(&self, host: &str) -> DnsFuture;
}
#[cfg(feature = "reqwest")]
mod reqwest_impl {
use super::{DnsError, DnsFuture, DnsResolver};
use rand::prelude::SliceRandom;
use std::net::{SocketAddr, ToSocketAddrs};
use std::sync::Arc;
use tokio::task::JoinSet;
pub(crate) struct ReqwestResolver(pub(crate) Arc<dyn DnsResolver>);
impl reqwest::dns::Resolve for ReqwestResolver {
fn resolve(&self, name: reqwest::dns::Name) -> reqwest::dns::Resolving {
let resolver = Arc::clone(&self.0);
let host = name.as_str().to_string();
Box::pin(async move {
let ips = resolver.resolve(&host).await?;
let addrs: reqwest::dns::Addrs =
Box::new(ips.into_iter().map(|ip| SocketAddr::new(ip, 0)));
Ok(addrs)
})
}
}
#[derive(Debug)]
pub(crate) struct ShuffleResolver;
impl DnsResolver for ShuffleResolver {
fn resolve(&self, host: &str) -> DnsFuture {
let host = host.to_string();
Box::pin(async move {
let mut tasks = JoinSet::new();
tasks.spawn_blocking(move || {
let it = (host.as_str(), 0).to_socket_addrs()?;
let mut addrs = it.map(|addr| addr.ip()).collect::<Vec<_>>();
addrs.shuffle(&mut rand::rng());
Ok(addrs)
});
tasks
.join_next()
.await
.expect("spawned one task")
.map_err(|err| Box::new(err) as DnsError)?
})
}
}
}
#[cfg(feature = "reqwest")]
pub(crate) use reqwest_impl::{ReqwestResolver, ShuffleResolver};
#[cfg(all(test, feature = "reqwest"))]
mod tests {
use super::*;
#[tokio::test]
async fn shuffle_resolver_resolves_localhost() {
let ips = ShuffleResolver.resolve("localhost").await.unwrap();
assert!(!ips.is_empty());
assert!(ips.iter().all(|ip| ip.is_loopback()));
}
#[derive(Debug)]
struct FailingResolver;
impl DnsResolver for FailingResolver {
fn resolve(&self, _host: &str) -> DnsFuture {
Box::pin(async { Err("boom".into()) })
}
}
#[tokio::test]
async fn adapter_propagates_errors() {
use reqwest::dns::Resolve;
use std::sync::Arc;
let adapter = ReqwestResolver(Arc::new(FailingResolver));
let err = match adapter.resolve("localhost".parse().unwrap()).await {
Ok(_) => panic!("expected resolution to fail"),
Err(e) => e,
};
assert!(err.to_string().contains("boom"));
}
}