use std::fmt;
use std::future::Future;
use std::net::{IpAddr, SocketAddr};
use std::pin::Pin;
use std::sync::Arc;
use weida_core::{DEFAULT_PORT, Error};
use crate::exec::Exec;
pub type Resolved<'a> = Pin<Box<dyn Future<Output = Result<Vec<SocketAddr>, Error>> + Send + 'a>>;
pub trait Resolver: fmt::Debug + Send + Sync + 'static {
fn resolve<'a>(
&'a self,
exec: &'a Exec,
name: &'a str,
port: Option<u16>,
max_addresses: usize,
) -> Resolved<'a>;
}
#[derive(Clone, Copy, Debug, Default)]
pub struct SystemResolver;
impl Resolver for SystemResolver {
fn resolve<'a>(
&'a self,
exec: &'a Exec,
name: &'a str,
port: Option<u16>,
max_addresses: usize,
) -> Resolved<'a> {
Box::pin(async move {
let port = port.unwrap_or(DEFAULT_PORT);
if let Ok(ip) = name.parse::<IpAddr>() {
return Ok(vec![SocketAddr::new(ip, port)]);
}
let query = (name.to_owned(), port);
let looked_up = exec
.spawn(async move {
tokio::net::lookup_host(query)
.await
.map(|addrs| addrs.collect::<Vec<SocketAddr>>())
})
.await
.map_err(|e| Error::Runtime(format!("name resolution task failed: {e}")))?;
let addrs: Vec<SocketAddr> = looked_up
.map_err(|e| Error::InvalidAddress(format!("cannot resolve {name}:{port}: {e}")))?
.into_iter()
.take(max_addresses)
.collect();
if addrs.is_empty() {
return Err(Error::InvalidAddress(format!(
"{name}:{port} resolved to no addresses"
)));
}
Ok(addrs)
})
}
}
#[derive(Clone, Debug)]
pub struct SharedResolver(Arc<dyn Resolver>);
impl SharedResolver {
pub fn new(resolver: impl Resolver) -> SharedResolver {
SharedResolver(Arc::new(resolver))
}
pub fn resolve<'a>(
&'a self,
exec: &'a Exec,
name: &'a str,
port: Option<u16>,
max_addresses: usize,
) -> Resolved<'a> {
self.0.resolve(exec, name, port, max_addresses)
}
}
impl Default for SharedResolver {
fn default() -> SharedResolver {
SharedResolver::new(SystemResolver)
}
}
impl PartialEq for SharedResolver {
fn eq(&self, other: &SharedResolver) -> bool {
Arc::ptr_eq(&self.0, &other.0)
}
}
impl Eq for SharedResolver {}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Debug)]
struct Table(Vec<SocketAddr>);
impl Resolver for Table {
fn resolve<'a>(
&'a self,
_exec: &'a Exec,
_name: &'a str,
_port: Option<u16>,
max_addresses: usize,
) -> Resolved<'a> {
Box::pin(async move { Ok(self.0.iter().copied().take(max_addresses).collect()) })
}
}
#[tokio::test]
async fn a_literal_needs_no_resolver_and_takes_the_written_port() {
let exec = Exec::current().expect("ambient runtime");
let addrs = SystemResolver
.resolve(&exec, "127.0.0.1", Some(9000), 8)
.await
.expect("literal");
assert_eq!(addrs, vec!["127.0.0.1:9000".parse().expect("addr")]);
}
#[tokio::test]
async fn a_literal_without_a_written_port_takes_the_default() {
let exec = Exec::current().expect("ambient runtime");
let addrs = SystemResolver
.resolve(&exec, "127.0.0.1", None, 8)
.await
.expect("literal");
assert_eq!(addrs[0].port(), DEFAULT_PORT);
}
#[tokio::test]
async fn an_unresolvable_name_is_an_error_rather_than_an_empty_set() {
let exec = Exec::current().expect("ambient runtime");
let err = SystemResolver
.resolve(&exec, "no-such-host.invalid", Some(1), 8)
.await
.expect_err("`.invalid` never resolves");
assert!(
err.to_string().contains("no-such-host.invalid"),
"the error names what could not be resolved: {err}"
);
}
#[tokio::test]
async fn a_replaced_resolver_may_answer_several_ports_on_one_address() {
let exec = Exec::current().expect("ambient runtime");
let table = SharedResolver::new(Table(vec![
"203.0.113.7:7443".parse().expect("addr"),
"203.0.113.7:7444".parse().expect("addr"),
"203.0.113.7:7445".parse().expect("addr"),
]));
let addrs = table
.resolve(&exec, "lb.example", None, 8)
.await
.expect("table");
assert_eq!(addrs.len(), 3);
assert!(addrs.iter().all(|a| a.ip().to_string() == "203.0.113.7"));
assert_eq!(
addrs.iter().map(|a| a.port()).collect::<Vec<_>>(),
vec![7443, 7444, 7445],
"the order is the resolver's and is preserved"
);
}
#[tokio::test]
async fn the_cap_is_the_callers_and_the_resolver_honours_it() {
let exec = Exec::current().expect("ambient runtime");
let table = SharedResolver::new(Table(vec![
"203.0.113.7:7443".parse().expect("addr"),
"203.0.113.7:7444".parse().expect("addr"),
"203.0.113.7:7445".parse().expect("addr"),
]));
let addrs = table
.resolve(&exec, "lb.example", None, 2)
.await
.expect("table");
assert_eq!(addrs.len(), 2, "a resolver answer is remote input");
}
#[test]
fn two_resolvers_are_equal_only_when_they_are_the_same_one() {
let one = SharedResolver::default();
let same = one.clone();
let other = SharedResolver::default();
assert_eq!(one, same);
assert_ne!(
one, other,
"identical behaviour is not identity: the pool keys on this"
);
}
#[tokio::test]
async fn resolves_ip_literals_without_dns() {
let exec = Exec::current().expect("ambient runtime");
assert_eq!(
SystemResolver
.resolve(&exec, "127.0.0.1", Some(7443), 8)
.await
.expect("v4"),
vec![SocketAddr::from(([127, 0, 0, 1], 7443))]
);
let v6 = SystemResolver
.resolve(&exec, "::1", Some(7443), 8)
.await
.expect("v6");
assert_eq!(v6.len(), 1);
assert_eq!(v6[0].port(), 7443);
assert!(v6[0].is_ipv6());
}
#[tokio::test]
async fn a_hostname_resolves_to_every_address_up_to_the_cap() {
let exec = Exec::current().expect("ambient runtime");
let all = SystemResolver
.resolve(&exec, "localhost", Some(7443), 8)
.await
.expect("localhost");
assert!(!all.is_empty());
assert!(all.iter().all(|a| a.port() == 7443));
let capped = SystemResolver
.resolve(&exec, "localhost", Some(7443), 1)
.await
.expect("localhost");
assert_eq!(capped.len(), 1, "the cap must bound the answer");
assert_eq!(capped[0], all[0], "and it must keep the resolver's order");
}
}