use std::collections::HashSet;
use std::fmt;
use std::future::Future;
use std::net::{IpAddr, SocketAddr};
use std::time::Duration;
use futures::StreamExt;
use futures::stream::FuturesUnordered;
pub(crate) const DEFAULT_DELAY: Duration = Duration::from_millis(250);
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Failure<E> {
pub addr: SocketAddr,
pub error: E,
}
impl<E: fmt::Display> fmt::Display for Failure<E> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}: {}", self.addr, self.error)
}
}
impl<E: std::error::Error + 'static> std::error::Error for Failure<E> {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
Some(&self.error)
}
}
pub(crate) trait Aggregate: Sized {
fn aggregate(failures: Vec<Failure<Self>>) -> Self;
}
pub(crate) fn interleave(addrs: impl IntoIterator<Item = SocketAddr>) -> Vec<SocketAddr> {
let (mut a, mut b): (Vec<SocketAddr>, Vec<SocketAddr>) = (Vec::new(), Vec::new());
for addr in addrs {
if a.is_empty() || a[0].is_ipv4() == addr.is_ipv4() {
a.push(addr);
} else {
b.push(addr);
}
}
let mut out = Vec::with_capacity(a.len() + b.len());
let (mut a, mut b) = (a.into_iter(), b.into_iter());
loop {
match (a.next(), b.next()) {
(Some(x), Some(y)) => {
out.push(x);
out.push(y);
}
(Some(x), None) => out.push(x),
(None, Some(y)) => out.push(y),
(None, None) => break,
}
}
out
}
pub(crate) fn match_local(
addrs: impl IntoIterator<Item = SocketAddr>,
local: SocketAddr,
dual_stack: bool,
) -> Vec<SocketAddr> {
let mut seen = HashSet::new();
let candidates: Vec<SocketAddr> = interleave(addrs)
.into_iter()
.map(|addr| normalize_family(addr, local))
.filter(|addr| seen.insert(*addr))
.collect();
let usable: Vec<SocketAddr> = candidates
.iter()
.copied()
.filter(|addr| addressable(*addr, local, dual_stack))
.collect();
if usable.is_empty() { candidates } else { usable }
}
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,
}
}
pub(crate) fn describe<E: fmt::Display>(failures: &[Failure<E>]) -> String {
failures.iter().map(|f| f.to_string()).collect::<Vec<_>>().join("; ")
}
pub(crate) async fn race<C, E, F, Fut>(candidates: Vec<SocketAddr>, delay: Duration, mut dial: F) -> Result<C, E>
where
F: FnMut(SocketAddr) -> Fut,
Fut: Future<Output = Result<C, E>>,
E: Aggregate + fmt::Display,
{
let mut remaining = candidates.into_iter();
let mut attempts = FuturesUnordered::new();
let mut failures: Vec<(usize, Failure<E>)> = Vec::new();
let mut next_index = 0;
let mut start = |addr: SocketAddr, attempts: &mut FuturesUnordered<_>| {
let index = next_index;
next_index += 1;
tracing::debug!(%addr, index, "dialing");
let attempt = dial(addr);
attempts.push(async move { (index, addr, attempt.await) });
};
let first = remaining.next().expect("no candidates to dial");
start(first, &mut attempts);
loop {
tokio::select! {
biased;
res = attempts.next() => {
let (index, addr, res) = res.expect("attempts can't be empty here");
match res {
Ok(conn) => {
tracing::debug!(%addr, index, "connected");
return Ok(conn);
}
Err(err) => {
tracing::debug!(%addr, index, %err, "connection attempt failed");
failures.push((index, Failure { addr, error: err }));
if let Some(addr) = remaining.next() {
start(addr, &mut attempts);
} else if attempts.is_empty() {
failures.sort_by_key(|(index, _)| *index);
return Err(collapse(failures.into_iter().map(|(_, failure)| failure).collect()));
}
}
}
}
_ = tokio::time::sleep(delay), if remaining.len() > 0 => {
let addr = remaining.next().expect("guarded by remaining.len()");
start(addr, &mut attempts);
}
}
}
}
fn collapse<E: Aggregate>(mut failures: Vec<Failure<E>>) -> E {
match failures.len() {
1 => failures.pop().expect("checked len").error,
_ => E::aggregate(failures),
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
fn v4(s: &str) -> SocketAddr {
s.parse().unwrap()
}
#[derive(Debug, PartialEq, Eq)]
enum TestError {
Dial(&'static str),
All(Vec<Failure<TestError>>),
}
impl fmt::Display for TestError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Dial(err) => write!(f, "{err}"),
Self::All(failures) => write!(f, "all {} attempts failed: {}", failures.len(), describe(failures)),
}
}
}
impl Aggregate for TestError {
fn aggregate(failures: Vec<Failure<Self>>) -> Self {
Self::All(failures)
}
}
fn failed(addr: &str, err: &'static str) -> Failure<TestError> {
Failure {
addr: v4(addr),
error: TestError::Dial(err),
}
}
#[test]
fn interleave_alternates_families() {
let addrs = [
v4("[2001:db8::1]:443"),
v4("[2001:db8::2]:443"),
v4("1.2.3.4:443"),
v4("5.6.7.8:443"),
];
assert_eq!(
interleave(addrs),
vec![
v4("[2001:db8::1]:443"),
v4("1.2.3.4:443"),
v4("[2001:db8::2]:443"),
v4("5.6.7.8:443"),
]
);
}
#[test]
fn interleave_keeps_the_resolver_preferred_family_first() {
let addrs = [v4("1.2.3.4:443"), v4("[2001:db8::1]:443")];
assert_eq!(interleave(addrs), vec![v4("1.2.3.4:443"), v4("[2001:db8::1]:443")]);
}
#[test]
fn interleave_single_family_passthrough() {
let addrs = [v4("1.2.3.4:443"), v4("5.6.7.8:443")];
assert_eq!(interleave(addrs), addrs.to_vec());
}
#[test]
fn match_local_prefers_matching_family() {
let a4 = v4("127.0.0.1:443");
let a6 = v4("[::1]:443");
assert_eq!(match_local([a6, a4], v4("0.0.0.0:0"), false), vec![a4]);
assert_eq!(
match_local([a4, a6], v4("[::]:0"), true),
vec![v4("[::ffff:127.0.0.1]:443"), a6]
);
}
#[test]
fn match_local_skips_mapped_ipv4_on_a_v6_only_socket() {
let a4 = v4("192.0.2.1:443");
let a6 = v4("[2001:db8::1]:443");
assert_eq!(match_local([a4, a6], v4("[::]:0"), false), vec![a6]);
}
#[test]
fn match_local_skips_ipv4_for_a_concrete_v6_bind() {
let a4 = v4("192.0.2.1:443");
let a6 = v4("[2001:db8::1]:443");
assert_eq!(match_local([a4, a6], v4("[2001:db8::5]:0"), true), vec![a6]);
}
#[test]
fn match_local_keeps_normalized_fallback_when_none_are_usable() {
let a4 = v4("192.0.2.1:443");
assert_eq!(
match_local([a4], v4("[::]:0"), false),
vec![v4("[::ffff:192.0.2.1]:443")]
);
}
#[test]
fn match_local_unwraps_v4_mapped_for_v4_socket() {
let mapped = v4("[::ffff:127.0.0.1]:443");
assert_eq!(match_local([mapped], v4("0.0.0.0:0"), false), vec![v4("127.0.0.1:443")]);
}
#[test]
fn match_local_falls_back_for_unmappable_v6() {
let a6 = v4("[2001:db8::1]:443");
assert_eq!(match_local([a6], v4("0.0.0.0:0"), false), vec![a6]);
}
#[test]
fn match_local_empty() {
assert!(match_local(std::iter::empty(), v4("0.0.0.0:0"), false).is_empty());
}
#[test]
fn match_local_dedups_across_the_interleave() {
let a4 = v4("1.2.3.4:443");
let a6 = v4("[2001:db8::1]:443");
assert_eq!(match_local([a4, a4, a6], v4("0.0.0.0:0"), false), vec![a4]);
assert_eq!(
match_local([a4, a4, a6], v4("[::]:0"), true),
vec![v4("[::ffff:1.2.3.4]:443"), a6]
);
}
#[test]
fn match_local_dedups_normalized_forms() {
let a4 = v4("1.2.3.4:443");
let mapped = v4("[::ffff:1.2.3.4]:443");
assert_eq!(match_local([a4, mapped], v4("[::]:0"), true), vec![mapped]);
assert_eq!(match_local([mapped, a4], v4("0.0.0.0:0"), false), vec![a4]);
}
#[tokio::test(start_paused = true)]
async fn first_success_returns_immediately() {
let dials = Arc::new(AtomicUsize::new(0));
let counter = dials.clone();
let res: Result<&str, TestError> = race(vec![v4("1.1.1.1:1"), v4("2.2.2.2:2")], DEFAULT_DELAY, move |_| {
counter.fetch_add(1, Ordering::SeqCst);
async { Ok("winner") }
})
.await;
assert_eq!(res, Ok("winner"));
assert_eq!(dials.load(Ordering::SeqCst), 1, "no second dial after a fast success");
}
#[tokio::test(start_paused = true)]
async fn second_wins_when_first_hangs() {
let start = tokio::time::Instant::now();
let res: Result<&str, TestError> = race(
vec![v4("1.1.1.1:1"), v4("2.2.2.2:2")],
DEFAULT_DELAY,
|addr| async move {
if addr == v4("1.1.1.1:1") {
std::future::pending().await
} else {
Ok("second")
}
},
)
.await;
assert_eq!(res, Ok("second"));
assert_eq!(start.elapsed(), DEFAULT_DELAY, "second dial waits out the stagger");
}
#[tokio::test(start_paused = true)]
async fn failure_starts_the_next_attempt_immediately() {
let start = tokio::time::Instant::now();
let res: Result<&str, TestError> = race(
vec![v4("1.1.1.1:1"), v4("2.2.2.2:2")],
DEFAULT_DELAY,
|addr| async move {
if addr == v4("1.1.1.1:1") {
Err(TestError::Dial("boom"))
} else {
Ok("second")
}
},
)
.await;
assert_eq!(res, Ok("second"));
assert_eq!(start.elapsed(), Duration::ZERO, "failure must not wait for the timer");
}
#[tokio::test(start_paused = true)]
async fn all_failures_are_reported_when_the_preferred_fails_first() {
let res: Result<&str, TestError> = race(
vec![v4("1.1.1.1:1"), v4("2.2.2.2:2")],
Duration::from_millis(10),
|addr| async move {
if addr == v4("1.1.1.1:1") {
Err(TestError::Dial("network unreachable"))
} else {
tokio::time::sleep(Duration::from_secs(1)).await;
Err(TestError::Dial("invalid peer certificate"))
}
},
)
.await;
assert_eq!(
res,
Err(TestError::All(vec![
failed("1.1.1.1:1", "network unreachable"),
failed("2.2.2.2:2", "invalid peer certificate"),
]))
);
}
#[tokio::test(start_paused = true)]
async fn all_failures_are_reported_when_the_preferred_times_out_last() {
let res: Result<&str, TestError> = race(
vec![v4("1.1.1.1:1"), v4("2.2.2.2:2")],
Duration::from_millis(10),
|addr| async move {
if addr == v4("1.1.1.1:1") {
tokio::time::sleep(Duration::from_secs(30)).await;
Err(TestError::Dial("timed out"))
} else {
Err(TestError::Dial("invalid peer certificate"))
}
},
)
.await;
assert_eq!(
res,
Err(TestError::All(vec![
failed("1.1.1.1:1", "timed out"),
failed("2.2.2.2:2", "invalid peer certificate"),
]))
);
}
#[tokio::test(start_paused = true)]
async fn a_lone_failure_is_returned_unwrapped() {
let res: Result<&str, TestError> = race(vec![v4("1.1.1.1:1")], DEFAULT_DELAY, |_| async {
Err(TestError::Dial("invalid peer certificate"))
})
.await;
assert_eq!(res, Err(TestError::Dial("invalid peer certificate")));
}
#[test]
fn describe_lists_every_attempt() {
let failures = [failed("1.1.1.1:1", "timed out"), failed("2.2.2.2:2", "bad cert")];
assert_eq!(describe(&failures), "1.1.1.1:1: timed out; 2.2.2.2:2: bad cert");
}
#[tokio::test(start_paused = true)]
async fn losers_are_dropped_on_success() {
struct Guard(Arc<AtomicUsize>);
impl Drop for Guard {
fn drop(&mut self) {
self.0.fetch_add(1, Ordering::SeqCst);
}
}
let dropped = Arc::new(AtomicUsize::new(0));
let count = dropped.clone();
let res: Result<&str, TestError> = race(vec![v4("1.1.1.1:1"), v4("2.2.2.2:2")], Duration::ZERO, move |addr| {
let guard = Guard(count.clone());
async move {
if addr == v4("1.1.1.1:1") {
let _guard = guard;
std::future::pending().await
} else {
drop(guard);
tokio::time::sleep(Duration::from_millis(1)).await;
Ok("second")
}
}
})
.await;
assert_eq!(res, Ok("second"));
assert_eq!(dropped.load(Ordering::SeqCst), 2, "the hung attempt was not aborted");
}
#[tokio::test(start_paused = true)]
async fn zero_delay_dials_all_at_once() {
let start = tokio::time::Instant::now();
let res: Result<&str, TestError> = race(
vec![v4("1.1.1.1:1"), v4("2.2.2.2:2")],
Duration::ZERO,
|addr| async move {
if addr == v4("1.1.1.1:1") {
std::future::pending().await
} else {
Ok("second")
}
},
)
.await;
assert_eq!(res, Ok("second"));
assert_eq!(start.elapsed(), Duration::ZERO);
}
}