use std::collections::{HashSet, VecDeque};
use std::io;
use std::net::{IpAddr, SocketAddr};
use std::time::Duration;
use futures::FutureExt;
use crate::client::DEFAULT_RESOLUTION_DELAY;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Lookup {
Full,
Ipv4,
}
impl Lookup {
fn hints(self) -> dns_lookup::AddrInfoHints {
dns_lookup::AddrInfoHints {
address: match self {
Self::Full => 0,
Self::Ipv4 => dns_lookup::AddrFamily::Inet.into(),
},
socktype: dns_lookup::SockType::Stream.into(),
..Default::default()
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Family {
V6,
V4,
}
impl Family {
fn of(addr: SocketAddr) -> Self {
match addr.is_ipv4() {
true => Self::V4,
false => Self::V6,
}
}
fn other(self) -> Self {
match self {
Self::V6 => Self::V4,
Self::V4 => Self::V6,
}
}
}
#[derive(Default)]
struct Query {
call: Option<tokio::task::JoinHandle<io::Result<Vec<SocketAddr>>>>,
error: Option<io::Error>,
}
impl Query {
fn start(host: &str, port: u16, lookup: Lookup) -> Self {
let host = host.to_owned();
let call = tokio::task::spawn_blocking(move || resolve(&host, port, lookup));
Self {
call: Some(call),
..Default::default()
}
}
fn pending(&self) -> bool {
self.call.is_some()
}
async fn answer(&mut self, lookup: Lookup) -> Vec<SocketAddr> {
let Some(call) = self.call.as_mut() else {
return std::future::pending().await;
};
let res = call.await;
self.record(lookup, res)
}
fn ready(&mut self, lookup: Lookup) -> Option<Vec<SocketAddr>> {
let res = self.call.as_mut()?.now_or_never()?;
Some(self.record(lookup, res))
}
fn record(
&mut self,
lookup: Lookup,
res: Result<io::Result<Vec<SocketAddr>>, tokio::task::JoinError>,
) -> Vec<SocketAddr> {
let res = res.unwrap_or_else(|err| Err(io::Error::other(err)));
self.call = None;
match res {
Ok(addrs) => {
tracing::debug!(?lookup, count = addrs.len(), "resolved");
addrs
}
Err(err) => {
tracing::debug!(?lookup, %err, "lookup failed");
self.error = Some(err);
Vec::new()
}
}
}
}
pub(crate) struct Candidates {
full: Query,
ipv4: Query,
v6: VecDeque<SocketAddr>,
v4: VecDeque<SocketAddr>,
next: Family,
delay: Duration,
delayed: bool,
local: Option<(SocketAddr, bool)>,
unreachable: VecDeque<SocketAddr>,
seen: HashSet<SocketAddr>,
usable: bool,
count: usize,
limit: usize,
}
impl Default for Candidates {
fn default() -> Self {
Self {
full: Query::default(),
ipv4: Query::default(),
v6: VecDeque::new(),
v4: VecDeque::new(),
next: Family::V6,
delay: DEFAULT_RESOLUTION_DELAY,
delayed: false,
local: None,
unreachable: VecDeque::new(),
seen: HashSet::new(),
usable: false,
count: 0,
limit: usize::MAX,
}
}
}
impl Candidates {
pub(crate) fn resolve(host: url::Host<&str>, port: u16, delay: Duration) -> Self {
let domain = match host {
url::Host::Ipv4(ip) => return Self::fixed([SocketAddr::new(ip.into(), port)]),
url::Host::Ipv6(ip) => return Self::fixed([SocketAddr::new(ip.into(), port)]),
url::Host::Domain(domain) => match domain.parse::<IpAddr>() {
Ok(ip) => return Self::fixed([SocketAddr::new(ip, port)]),
Err(_) => domain,
},
};
Self {
full: Query::start(domain, port, Lookup::Full),
ipv4: Query::start(domain, port, Lookup::Ipv4),
delay,
..Default::default()
}
}
pub(crate) fn fixed(addrs: impl IntoIterator<Item = SocketAddr>) -> Self {
let mut this = Self::default();
for addr in addrs {
if this.v6.is_empty() && this.v4.is_empty() {
this.next = Family::of(addr);
}
this.queue(addr);
}
this
}
#[cfg(any(feature = "noq", feature = "quinn", feature = "quiche", test))]
pub(crate) fn with_local(mut self, local: SocketAddr, dual_stack: bool) -> Self {
self.local = Some((local, dual_stack));
self
}
#[cfg(any(feature = "quiche", test))]
pub(crate) fn with_limit(mut self, max: usize) -> Self {
self.limit = max;
self
}
pub(crate) async fn next(&mut self) -> Option<SocketAddr> {
if self.count >= self.limit {
return None;
}
loop {
self.poll_answers();
if let Some(addr) = self.take(self.next) {
return Some(addr);
}
if !self.queued(self.next.other()).is_empty() {
if self.count == 0 && !self.delayed && self.full.pending() {
self.delayed = true;
let _ = tokio::time::timeout(self.delay, self.answer_full()).await;
continue;
}
if let Some(addr) = self.take(self.next.other()) {
return Some(addr);
}
continue;
}
if !self.full.pending() && !self.ipv4.pending() {
if !self.usable
&& let Some(addr) = self.unreachable.pop_front()
{
self.count += 1;
return Some(addr);
}
return None;
}
self.answer_any().await;
}
}
pub(crate) fn failure(&mut self) -> Option<io::Error> {
self.full.error.take().or_else(|| self.ipv4.error.take())
}
fn take(&mut self, family: Family) -> Option<SocketAddr> {
loop {
let addr = self.queued(family).pop_front()?;
let addr = match self.local {
Some((local, _)) => normalize_family(addr, local),
None => addr,
};
if !self.seen.insert(addr) {
continue;
}
if let Some((local, dual_stack)) = self.local
&& !addressable(addr, local, dual_stack)
{
self.unreachable.push_back(addr);
continue;
}
self.usable = true;
self.count += 1;
self.next = family.other();
return Some(addr);
}
}
fn queued(&mut self, family: Family) -> &mut VecDeque<SocketAddr> {
match family {
Family::V6 => &mut self.v6,
Family::V4 => &mut self.v4,
}
}
fn queue(&mut self, addr: SocketAddr) {
self.queued(Family::of(addr)).push_back(addr);
}
fn poll_answers(&mut self) {
if let Some(addrs) = self.full.ready(Lookup::Full) {
self.accept(Lookup::Full, addrs);
}
if let Some(addrs) = self.ipv4.ready(Lookup::Ipv4) {
self.accept(Lookup::Ipv4, addrs);
}
}
async fn answer_full(&mut self) {
let addrs = self.full.answer(Lookup::Full).await;
self.accept(Lookup::Full, addrs);
}
async fn answer_any(&mut self) {
let (full, ipv4) = (&mut self.full, &mut self.ipv4);
let (lookup, addrs) = tokio::select! {
addrs = full.answer(Lookup::Full) => (Lookup::Full, addrs),
addrs = ipv4.answer(Lookup::Ipv4) => (Lookup::Ipv4, addrs),
};
self.accept(lookup, addrs);
}
fn accept(&mut self, lookup: Lookup, addrs: Vec<SocketAddr>) {
if lookup == Lookup::Ipv4 || addrs.is_empty() {
for addr in addrs {
self.queue(addr);
}
return;
}
self.ipv4.call = None;
self.v6.clear();
self.v4.clear();
if self.count == 0
&& let Some(first) = addrs.first()
{
self.next = Family::of(*first);
}
for addr in addrs {
self.queue(addr);
}
}
}
fn resolve(host: &str, port: u16, lookup: Lookup) -> io::Result<Vec<SocketAddr>> {
let answers = dns_lookup::getaddrinfo(Some(host), None, Some(lookup.hints())).map_err(io::Error::from)?;
answers.map(|answer| Ok(with_port(answer?.sockaddr, port))).collect()
}
fn with_port(mut addr: SocketAddr, port: u16) -> SocketAddr {
addr.set_port(port);
addr
}
fn addressable(dest: SocketAddr, local: SocketAddr, dual_stack: bool) -> bool {
let (SocketAddr::V6(dest), SocketAddr::V6(local)) = (dest, local) else {
return dest.is_ipv4() == local.is_ipv4();
};
match (dest.ip().to_ipv4_mapped(), local.ip().to_ipv4_mapped()) {
(Some(_), None) => dual_stack && local.ip().is_unspecified(),
(None, Some(_)) => false,
_ => true,
}
}
fn normalize_family(addr: SocketAddr, local: SocketAddr) -> SocketAddr {
match (addr, local.is_ipv4()) {
(SocketAddr::V6(v6), true) => match v6.ip().to_ipv4_mapped() {
Some(v4) => SocketAddr::new(IpAddr::V4(v4), v6.port()),
None => addr,
},
(SocketAddr::V4(v4), false) => SocketAddr::new(IpAddr::V6(v4.ip().to_ipv6_mapped()), v4.port()),
_ => addr,
}
}
#[cfg(test)]
impl Candidates {
pub(crate) fn slow(full: (&[SocketAddr], Duration), ipv4: (&[SocketAddr], Duration)) -> Self {
Self {
full: Query::slow(full.0, full.1),
ipv4: Query::slow(ipv4.0, ipv4.1),
delay: Duration::ZERO,
..Default::default()
}
}
}
#[cfg(test)]
impl Query {
fn slow(addrs: &[SocketAddr], delay: Duration) -> Self {
let addrs = addrs.to_vec();
Self {
call: Some(tokio::spawn(async move {
tokio::time::sleep(delay).await;
Ok(addrs)
})),
..Default::default()
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn addr(s: &str) -> SocketAddr {
s.parse().unwrap()
}
fn addrs(list: &[&str]) -> Vec<SocketAddr> {
list.iter().map(|s| addr(s)).collect()
}
async fn drain(mut candidates: Candidates) -> Vec<SocketAddr> {
let mut out = Vec::new();
while let Some(addr) = candidates.next().await {
out.push(addr);
}
out
}
fn answered(list: &[&str]) -> Candidates {
let mut candidates = Candidates::default();
candidates.accept(Lookup::Full, addrs(list));
candidates
}
#[tokio::test]
async fn alternates_families() {
let candidates = answered(&["[2001:db8::1]:443", "[2001:db8::2]:443", "1.2.3.4:443", "5.6.7.8:443"]);
assert_eq!(
drain(candidates).await,
addrs(&["[2001:db8::1]:443", "1.2.3.4:443", "[2001:db8::2]:443", "5.6.7.8:443",])
);
}
#[tokio::test]
async fn the_resolver_picks_the_leading_family() {
let candidates = answered(&["1.2.3.4:443", "[2001:db8::1]:443"]);
assert_eq!(drain(candidates).await, addrs(&["1.2.3.4:443", "[2001:db8::1]:443"]));
let candidates = answered(&["[2001:db8::1]:443", "1.2.3.4:443"]);
assert_eq!(drain(candidates).await, addrs(&["[2001:db8::1]:443", "1.2.3.4:443"]));
}
#[tokio::test]
async fn a_single_attempt_follows_the_resolver() {
let candidates = answered(&["1.2.3.4:443", "[2001:db8::1]:443"]).with_limit(1);
assert_eq!(drain(candidates).await, addrs(&["1.2.3.4:443"]));
}
#[tokio::test]
async fn single_family_passthrough() {
let list = ["1.2.3.4:443", "5.6.7.8:443"];
assert_eq!(drain(answered(&list)).await, addrs(&list));
}
#[tokio::test]
async fn fixed_keeps_the_leading_family() {
let list = addrs(&["1.2.3.4:443", "[2001:db8::1]:443"]);
assert_eq!(drain(Candidates::fixed(list.clone())).await, list);
}
#[tokio::test]
async fn local_prefers_matching_family() {
let candidates = answered(&["[::1]:443", "127.0.0.1:443"]).with_local(addr("0.0.0.0:0"), false);
assert_eq!(drain(candidates).await, addrs(&["127.0.0.1:443"]));
let candidates = answered(&["[::1]:443", "127.0.0.1:443"]).with_local(addr("[::]:0"), true);
assert_eq!(drain(candidates).await, addrs(&["[::1]:443", "[::ffff:127.0.0.1]:443"]));
}
#[tokio::test]
async fn local_skips_mapped_ipv4_on_a_v6_only_socket() {
let candidates = answered(&["[2001:db8::1]:443", "192.0.2.1:443"]).with_local(addr("[::]:0"), false);
assert_eq!(drain(candidates).await, addrs(&["[2001:db8::1]:443"]));
}
#[tokio::test]
async fn local_skips_ipv4_for_a_concrete_v6_bind() {
let candidates = answered(&["[2001:db8::1]:443", "192.0.2.1:443"]).with_local(addr("[2001:db8::5]:0"), true);
assert_eq!(drain(candidates).await, addrs(&["[2001:db8::1]:443"]));
}
#[tokio::test]
async fn local_keeps_the_normalized_fallback_when_none_are_usable() {
let candidates = answered(&["192.0.2.1:443"]).with_local(addr("[::]:0"), false);
assert_eq!(drain(candidates).await, addrs(&["[::ffff:192.0.2.1]:443"]));
}
#[tokio::test]
async fn local_unwraps_v4_mapped_for_a_v4_socket() {
let candidates = answered(&["[::ffff:127.0.0.1]:443"]).with_local(addr("0.0.0.0:0"), false);
assert_eq!(drain(candidates).await, addrs(&["127.0.0.1:443"]));
}
#[tokio::test]
async fn local_falls_back_for_an_unmappable_v6() {
let candidates = answered(&["[2001:db8::1]:443"]).with_local(addr("0.0.0.0:0"), false);
assert_eq!(drain(candidates).await, addrs(&["[2001:db8::1]:443"]));
}
#[tokio::test]
async fn empty_yields_nothing() {
assert!(drain(Candidates::fixed([])).await.is_empty());
assert!(Candidates::fixed([]).failure().is_none());
}
#[tokio::test]
async fn dedups_the_answer() {
let candidates = answered(&["[2001:db8::1]:443", "1.2.3.4:443", "1.2.3.4:443"]);
assert_eq!(drain(candidates).await, addrs(&["[2001:db8::1]:443", "1.2.3.4:443"]));
}
#[tokio::test]
async fn dedups_normalized_forms() {
let candidates = answered(&["[::ffff:1.2.3.4]:443", "1.2.3.4:443"]).with_local(addr("[::]:0"), true);
assert_eq!(drain(candidates).await, addrs(&["[::ffff:1.2.3.4]:443"]));
let candidates = answered(&["[::ffff:1.2.3.4]:443", "1.2.3.4:443"]).with_local(addr("0.0.0.0:0"), false);
assert_eq!(drain(candidates).await, addrs(&["1.2.3.4:443"]));
}
#[tokio::test]
async fn ip_literals_skip_the_resolver() {
let cases = [
("https://192.0.2.1", "192.0.2.1:443"),
("moqt://192.0.2.1", "192.0.2.1:443"),
("https://[2001:db8::1]", "[2001:db8::1]:443"),
("moqt://[2001:db8::1]", "[2001:db8::1]:443"),
];
for (url, want) in cases {
let url = url::Url::parse(url).unwrap();
let candidates = Candidates::resolve(url.host().unwrap(), 443, DEFAULT_RESOLUTION_DELAY);
assert_eq!(drain(candidates).await, vec![addr(want)], "{url}");
}
}
#[tokio::test(start_paused = true)]
async fn ipv4_waits_out_the_resolution_delay_for_the_full_answer() {
let mut candidates = Candidates::slow(
(
&addrs(&["[2001:db8::1]:443", "1.2.3.4:443"]),
DEFAULT_RESOLUTION_DELAY / 2,
),
(&addrs(&["1.2.3.4:443"]), Duration::ZERO),
);
candidates.delay = DEFAULT_RESOLUTION_DELAY;
let start = tokio::time::Instant::now();
assert_eq!(candidates.next().await, Some(addr("[2001:db8::1]:443")));
assert_eq!(
start.elapsed(),
DEFAULT_RESOLUTION_DELAY / 2,
"waited longer than the answer took"
);
assert_eq!(candidates.next().await, Some(addr("1.2.3.4:443")));
}
#[tokio::test(start_paused = true)]
async fn ipv4_proceeds_once_the_resolution_delay_expires() {
let mut candidates = Candidates::slow(
(&[], Duration::from_secs(30)),
(&addrs(&["1.2.3.4:443"]), Duration::ZERO),
);
candidates.delay = DEFAULT_RESOLUTION_DELAY;
let start = tokio::time::Instant::now();
assert_eq!(candidates.next().await, Some(addr("1.2.3.4:443")));
assert_eq!(start.elapsed(), DEFAULT_RESOLUTION_DELAY);
}
#[tokio::test(start_paused = true)]
async fn ipv4_does_not_wait_for_a_failed_full_lookup() {
let mut candidates = Candidates::slow((&[], Duration::ZERO), (&addrs(&["1.2.3.4:443"]), Duration::ZERO));
candidates.delay = DEFAULT_RESOLUTION_DELAY;
candidates.full.call = Some(tokio::spawn(async { Err(io::Error::other("nope")) }));
let start = tokio::time::Instant::now();
assert_eq!(candidates.next().await, Some(addr("1.2.3.4:443")));
assert_eq!(start.elapsed(), Duration::ZERO);
}
#[tokio::test(start_paused = true)]
async fn only_the_first_candidate_waits() {
let mut candidates = Candidates::slow(
(&[], Duration::from_secs(30)),
(&addrs(&["1.2.3.4:443", "5.6.7.8:443"]), Duration::ZERO),
);
candidates.delay = DEFAULT_RESOLUTION_DELAY;
assert_eq!(candidates.next().await, Some(addr("1.2.3.4:443")));
let start = tokio::time::Instant::now();
assert_eq!(candidates.next().await, Some(addr("5.6.7.8:443")));
assert_eq!(start.elapsed(), Duration::ZERO);
}
#[tokio::test(start_paused = true)]
async fn the_full_answer_supersedes_the_ipv4_one() {
let mut candidates = Candidates::slow(
(
&addrs(&["1.2.3.4:443", "[2001:db8::1]:443", "5.6.7.8:443"]),
Duration::from_secs(1),
),
(&addrs(&["1.2.3.4:443", "5.6.7.8:443"]), Duration::ZERO),
);
assert_eq!(candidates.next().await, Some(addr("1.2.3.4:443")));
assert_eq!(candidates.next().await, Some(addr("5.6.7.8:443")));
let start = tokio::time::Instant::now();
assert_eq!(candidates.next().await, Some(addr("[2001:db8::1]:443")));
assert_eq!(start.elapsed(), Duration::from_secs(1), "waited on the full answer");
assert_eq!(candidates.next().await, None, "an address was dialed twice");
}
#[tokio::test(start_paused = true)]
async fn a_queued_address_yields_to_the_full_answer() {
let mut candidates = Candidates::slow(
(
&addrs(&["[2001:db8::1]:443", "1.2.3.4:443", "5.6.7.8:443"]),
Duration::from_millis(10),
),
(&addrs(&["1.2.3.4:443", "5.6.7.8:443"]), Duration::ZERO),
);
assert_eq!(candidates.next().await, Some(addr("1.2.3.4:443")));
tokio::time::sleep(Duration::from_millis(20)).await;
assert_eq!(candidates.next().await, Some(addr("[2001:db8::1]:443")));
assert_eq!(candidates.next().await, Some(addr("5.6.7.8:443")));
assert_eq!(candidates.next().await, None);
}
#[tokio::test]
async fn failure_reports_the_lookup_error() {
let mut candidates = Candidates::default();
candidates.full.error = Some(io::Error::other("no such host"));
candidates.ipv4.error = Some(io::Error::other("no A record"));
let err = candidates.failure().expect("no failure reported");
assert_eq!(err.to_string(), "no such host");
assert_eq!(
candidates.failure().map(|err| err.to_string()).as_deref(),
Some("no A record")
);
assert!(candidates.failure().is_none());
}
#[test]
fn the_port_is_stamped_without_losing_the_scope() {
use std::net::{Ipv6Addr, SocketAddrV6};
let answer = SocketAddrV6::new("fe80::1".parse::<Ipv6Addr>().unwrap(), 0, 7, 3);
let dialed = with_port(SocketAddr::V6(answer), 443);
let SocketAddr::V6(dialed) = dialed else {
panic!("family changed: {dialed}");
};
assert_eq!(dialed.port(), 443);
assert_eq!(dialed.scope_id(), 3, "dropped the interface scope");
assert_eq!(dialed.flowinfo(), 7);
}
#[tokio::test]
async fn resolves_localhost() {
let url = url::Url::parse("https://localhost").unwrap();
let candidates = Candidates::resolve(url.host().unwrap(), 443, DEFAULT_RESOLUTION_DELAY);
let addrs = drain(candidates).await;
assert!(!addrs.is_empty(), "localhost resolved to nothing");
assert!(addrs.iter().all(|addr| addr.ip().is_loopback() && addr.port() == 443));
}
#[tokio::test]
async fn a_rejected_host_reports_a_failure() {
let mut candidates = Candidates {
full: Query::start("no\0such\0host", 443, Lookup::Full),
ipv4: Query::start("no\0such\0host", 443, Lookup::Ipv4),
..Default::default()
};
assert_eq!(candidates.next().await, None);
assert!(candidates.failure().is_some());
}
}