use std::future::Future;
use std::net::SocketAddr;
use std::pin::Pin;
use std::sync::Arc;
pub type DnsError = Box<dyn std::error::Error + Send + Sync>;
pub type Resolving = Pin<Box<dyn Future<Output = Result<Vec<SocketAddr>, DnsError>> + Send>>;
pub trait DnsResolver: Send + Sync {
fn resolve(&self, host: &str) -> Resolving;
}
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 = self.0.clone();
Box::pin(async move {
let addrs = resolver.resolve(name.as_str()).await?;
Ok(Box::new(addrs.into_iter()) as reqwest::dns::Addrs)
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::net::fetcher::{Fetcher, FetcherConfig};
use crate::net::fetcher_context::NullContext;
use crate::net::test_support::{RouteConfig, TestServer};
use crate::net::types::{FetchRequest, FetchResult};
use http::Method;
use parking_lot::Mutex;
use std::collections::HashMap;
use std::time::Duration;
use tokio_util::sync::CancellationToken;
use url::Url;
struct MapResolver {
map: HashMap<String, SocketAddr>,
seen: Mutex<Vec<String>>,
}
impl MapResolver {
fn new(entries: &[(&str, SocketAddr)]) -> Arc<Self> {
Arc::new(Self {
map: entries.iter().map(|(h, a)| (h.to_string(), *a)).collect(),
seen: Mutex::new(Vec::new()),
})
}
fn seen(&self) -> Vec<String> {
self.seen.lock().clone()
}
}
impl DnsResolver for MapResolver {
fn resolve(&self, host: &str) -> Resolving {
self.seen.lock().push(host.to_string());
let result = match self.map.get(host) {
Some(addr) => Ok(vec![*addr]),
None => Err(format!("resolver policy refuses host {host}").into()),
};
Box::pin(async move { result })
}
}
fn config_with(resolver: Arc<MapResolver>) -> FetcherConfig {
FetcherConfig {
connect_timeout: Duration::from_secs(2),
req_timeout: Duration::from_secs(5),
dns_resolver: Some(resolver),
..FetcherConfig::default()
}
}
fn spawn_fetcher(cfg: FetcherConfig) -> (Arc<Fetcher>, CancellationToken) {
let fetcher = Arc::new(Fetcher::new(cfg, Arc::new(NullContext)).unwrap());
let shutdown = CancellationToken::new();
let f = fetcher.clone();
let s = shutdown.clone();
tokio::spawn(async move { f.run(s).await });
(fetcher, shutdown)
}
async fn fetch(fetcher: &Fetcher, url: Url) -> FetchResult {
let req = FetchRequest::builder(Method::GET, url).build();
tokio::time::timeout(Duration::from_secs(5), fetcher.fetch(req))
.await
.unwrap()
}
#[tokio::test(flavor = "current_thread")]
async fn custom_resolver_handles_the_lookup() {
let srv = TestServer::new()
.route("/fast", RouteConfig::ok(b"x"))
.start()
.await;
let addr = srv.socket_addr();
let resolver = MapResolver::new(&[("sonar-dns.test", addr)]);
let (fetcher, shutdown) = spawn_fetcher(config_with(resolver.clone()));
let url = Url::parse(&format!("http://sonar-dns.test:{}/fast", addr.port())).unwrap();
match fetch(&fetcher, url).await {
FetchResult::Buffered { meta, body } => {
assert_eq!(meta.status, 200);
assert_eq!(&body[..], b"x");
}
other => panic!("expected Buffered, got {other:?}"),
}
assert_eq!(resolver.seen(), vec!["sonar-dns.test"]);
shutdown.cancel();
}
#[tokio::test(flavor = "current_thread")]
async fn resolver_rejection_blocks_the_request() {
let resolver = MapResolver::new(&[]);
let (fetcher, shutdown) = spawn_fetcher(config_with(resolver.clone()));
let url = Url::parse("http://sonar-refused.test/").unwrap();
let result = fetch(&fetcher, url).await;
assert!(result.is_error(), "expected an error, got {result:?}");
assert_eq!(resolver.seen(), vec!["sonar-refused.test"]);
shutdown.cancel();
}
#[tokio::test(flavor = "current_thread")]
async fn resolver_covers_redirect_hops() {
let target_srv = TestServer::new()
.route("/fast", RouteConfig::ok(b"x"))
.start()
.await;
let target_addr = target_srv.socket_addr();
let hop_srv = TestServer::new()
.route(
"/hop",
RouteConfig::RedirectAbsolute(format!(
"http://sonar-hop.test:{}/fast",
target_addr.port()
)),
)
.start()
.await;
let hop_addr = hop_srv.socket_addr();
let resolver = MapResolver::new(&[
("sonar-dns.test", hop_addr),
("sonar-hop.test", target_addr),
]);
let (fetcher, shutdown) = spawn_fetcher(config_with(resolver.clone()));
let url = Url::parse(&format!("http://sonar-dns.test:{}/hop", hop_addr.port())).unwrap();
match fetch(&fetcher, url).await {
FetchResult::Buffered { meta, body } => {
assert_eq!(meta.status, 200);
assert_eq!(&body[..], b"x");
}
other => panic!("expected Buffered, got {other:?}"),
}
assert_eq!(resolver.seen(), vec!["sonar-dns.test", "sonar-hop.test"]);
shutdown.cancel();
}
#[tokio::test(flavor = "current_thread")]
async fn resolver_rejection_blocks_a_redirect_hop() {
let srv = TestServer::new()
.route(
"/hop",
RouteConfig::RedirectAbsolute("http://sonar-internal.test/".to_string()),
)
.start()
.await;
let addr = srv.socket_addr();
let resolver = MapResolver::new(&[("sonar-dns.test", addr)]);
let (fetcher, shutdown) = spawn_fetcher(config_with(resolver.clone()));
let url = Url::parse(&format!("http://sonar-dns.test:{}/hop", addr.port())).unwrap();
let result = fetch(&fetcher, url).await;
assert!(result.is_error(), "expected an error, got {result:?}");
assert_eq!(
resolver.seen(),
vec!["sonar-dns.test", "sonar-internal.test"]
);
shutdown.cancel();
}
}