use crate::{
traits::AsyncToSocketAddrs,
util::{self, SocketAddrsFromIpAddrs},
};
use hickory_resolver::{TokioResolver, proto::rr::IntoName};
use std::{
io,
net::{IpAddr, SocketAddr, ToSocketAddrs},
str::FromStr,
sync::OnceLock,
vec,
};
static RESOLVER: OnceLock<TokioResolver> = OnceLock::new();
fn new_resolver() -> io::Result<TokioResolver> {
TokioResolver::builder_tokio()
.map_err(io::Error::other)?
.build()
.map_err(io::Error::other)
}
fn get_or_init_resolver() -> io::Result<&'static TokioResolver> {
if let Some(r) = RESOLVER.get() {
return Ok(r);
}
let resolver = new_resolver()?;
Ok(RESOLVER.get_or_init(|| resolver))
}
#[derive(Debug, Clone)]
pub struct HickoryToSocketAddrs<T: IntoName + Send + 'static> {
host: T,
port: u16,
}
impl<H: IntoName + Send + 'static> HickoryToSocketAddrs<H> {
pub fn new(host: H, port: u16) -> Self {
Self { host, port }
}
async fn lookup(self) -> io::Result<SocketAddrsFromIpAddrs<vec::IntoIter<IpAddr>>> {
if !util::inside_tokio() {
return Err(io::Error::other(
"hickory-dns is only supported in a tokio context",
));
}
self.lookup_with(get_or_init_resolver()?).await
}
async fn lookup_with(
self,
resolver: &TokioResolver,
) -> io::Result<SocketAddrsFromIpAddrs<vec::IntoIter<IpAddr>>> {
Ok(SocketAddrsFromIpAddrs(
resolver
.lookup_ip(self.host)
.await
.map_err(io::Error::other)?
.iter()
.collect::<Vec<_>>() .into_iter(),
self.port,
))
}
}
impl FromStr for HickoryToSocketAddrs<String> {
type Err = io::Error;
fn from_str(s: &str) -> io::Result<Self> {
fn invalid(msg: &'static str) -> io::Error {
io::Error::new(io::ErrorKind::InvalidInput, msg)
}
if let Ok(addr) = s.parse::<SocketAddr>() {
if matches!(addr, SocketAddr::V6(addr) if addr.scope_id() != 0) {
return Err(invalid("IPv6 scope ids are not supported"));
}
return Ok(Self::new(addr.ip().to_string(), addr.port()));
}
if s.starts_with('[') {
return Err(invalid("bracketed host is not an IP address"));
}
let (host, port_str) = s
.rsplit_once(':')
.ok_or_else(|| invalid("invalid socket address"))?;
if host.is_empty() {
return Err(invalid("empty host"));
}
if host.contains(':') {
return Err(invalid("IPv6 literals must be bracketed"));
}
let port = port_str
.parse()
.map_err(|_| invalid("invalid port value"))?;
Ok(Self::new(host.to_owned(), port))
}
}
impl<T: IntoName + Clone + Send + 'static> ToSocketAddrs for HickoryToSocketAddrs<T> {
type Iter = SocketAddrsFromIpAddrs<vec::IntoIter<IpAddr>>;
fn to_socket_addrs(&self) -> io::Result<Self::Iter> {
if util::inside_tokio() {
return util::block_on_tokio(self.clone().lookup());
}
let this = self.clone();
util::block_on_tokio(async move { this.lookup_with(&new_resolver()?).await })
}
}
impl<T: IntoName + Send + 'static> AsyncToSocketAddrs for HickoryToSocketAddrs<T> {
fn to_socket_addrs(
self,
) -> impl Future<Output = io::Result<impl Iterator<Item = SocketAddr> + Send + 'static>>
+ Send
+ 'static {
self.lookup()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn parse(s: &str) -> (String, u16) {
let addrs: HickoryToSocketAddrs<String> = s.parse().expect("parse");
(addrs.host, addrs.port)
}
#[test]
fn from_str_splits_host_and_port() {
assert_eq!(parse("example.com:80"), ("example.com".to_owned(), 80));
}
#[test]
fn from_str_keeps_ip_literals_parseable_as_ip() {
for (input, host) in [("127.0.0.1:80", "127.0.0.1"), ("[::1]:80", "::1")] {
let (parsed, port) = parse(input);
assert_eq!((parsed.as_str(), port), (host, 80));
assert!(parsed.parse::<IpAddr>().is_ok(), "{input}");
}
}
#[test]
fn from_str_rejects_garbage() {
for input in [
"example.com",
"example.com:http",
"2001:db8::1",
"::1",
"[example.com]:80",
"[fe80::1%eth0]:80",
"[fe80::1%1]:80",
":80",
"[::1:80",
"[::1]",
"[::1]80",
] {
assert!(
input.parse::<HickoryToSocketAddrs<String>>().is_err(),
"{input}"
);
}
}
}