use std::net::{SocketAddr, TcpStream, ToSocketAddrs};
use std::num::NonZeroU32;
use std::time::{Duration, Instant};
const MAX_PROBE_ADDRESSES: usize = 64;
pub fn probe(addr: impl ToSocketAddrs, timeout: Duration) -> bool {
let Ok(addrs) = addr.to_socket_addrs() else {
return false;
};
probe_candidates(addrs.take(MAX_PROBE_ADDRESSES), timeout)
}
#[must_use]
pub fn probe_resolved(addrs: &[SocketAddr], timeout: Duration) -> bool {
probe_candidates(addrs.iter().copied().take(MAX_PROBE_ADDRESSES), timeout)
}
fn probe_candidates(addrs: impl Iterator<Item = SocketAddr>, timeout: Duration) -> bool {
let started = Instant::now();
probe_candidates_with(
addrs,
timeout,
|| started.elapsed(),
|candidate, remaining| TcpStream::connect_timeout(candidate, remaining).is_ok(),
)
}
fn capped_candidate_count(count: usize) -> u32 {
let bytes = count.to_le_bytes();
u32::from_le_bytes([bytes[0], bytes[1], bytes[2], bytes[3]])
}
fn probe_candidates_with(
addrs: impl Iterator<Item = SocketAddr>,
timeout: Duration,
mut elapsed: impl FnMut() -> Duration,
mut connect: impl FnMut(&SocketAddr, Duration) -> bool,
) -> bool {
let candidates: Vec<SocketAddr> = addrs.take(MAX_PROBE_ADDRESSES).collect();
for (index, candidate) in candidates.iter().enumerate() {
let remaining = timeout.saturating_sub(elapsed());
if remaining.is_zero() {
return false;
}
let Some(left) = NonZeroU32::new(capped_candidate_count(
candidates.len().saturating_sub(index),
)) else {
return false;
};
let share = match remaining.checked_div(left.get()) {
Some(share) => share,
None => remaining,
};
if connect(candidate, share) {
return true;
}
}
false
}
#[must_use]
pub fn is_online(timeout: Duration) -> bool {
let addresses = [
SocketAddr::from(([1, 1, 1, 1], 443)),
SocketAddr::from(([8, 8, 8, 8], 53)),
];
probe_resolved(&addresses, timeout)
}
#[cfg(test)]
mod tests {
use super::*;
use std::net::{Ipv4Addr, TcpListener};
const WAIT: Duration = Duration::from_secs(2);
#[test]
fn open_port_probes_true() -> Result<(), std::io::Error> {
let listener = TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
assert!(probe(format!("127.0.0.1:{port}"), WAIT));
Ok(())
}
#[test]
fn closed_port_probes_false() -> Result<(), std::io::Error> {
let listener = TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
drop(listener);
assert!(!probe(format!("127.0.0.1:{port}"), WAIT));
Ok(())
}
#[test]
fn unresolvable_host_probes_false() {
assert!(!probe("nonexistent.invalid:80", WAIT));
}
#[test]
fn resolved_probe_reaches_a_loopback_listener() -> Result<(), std::io::Error> {
let listener = TcpListener::bind("127.0.0.1:0")?;
let address = listener.local_addr()?;
assert!(probe_resolved(&[address], WAIT));
Ok(())
}
const CANDIDATES: [SocketAddr; 2] = [
SocketAddr::new(std::net::IpAddr::V4(Ipv4Addr::new(192, 0, 2, 1)), 443),
SocketAddr::new(std::net::IpAddr::V4(Ipv4Addr::new(198, 51, 100, 1)), 443),
];
#[test]
fn address_candidates_share_one_remaining_budget() {
let clock = std::cell::Cell::new(Duration::ZERO);
let mut budgets = Vec::new();
let result = probe_candidates_with(
CANDIDATES.into_iter(),
Duration::from_millis(100),
|| clock.get(),
|_, remaining| {
budgets.push(remaining);
if budgets.len() == 1 {
clock.set(clock.get().saturating_add(Duration::from_millis(60)));
}
false
},
);
assert!(!result, "all candidate failures report unreachable");
assert_eq!(
budgets,
[Duration::from_millis(50), Duration::from_millis(40)],
"the first is offered half, the second what is left after the first spent 60 ms"
);
}
#[test]
fn a_spent_budget_dials_no_further_candidate() {
let clock = std::cell::Cell::new(Duration::ZERO);
let mut dialed = 0u32;
let result = probe_candidates_with(
CANDIDATES.into_iter(),
Duration::from_millis(100),
|| clock.get(),
|_, _| {
dialed = dialed.saturating_add(1);
clock.set(Duration::from_millis(100));
false
},
);
assert!(!result);
assert_eq!(
dialed, 1,
"the second candidate has no time left, so it is not dialed"
);
}
#[test]
fn a_blackholed_candidate_does_not_starve_the_next() {
let clock = std::cell::Cell::new(Duration::ZERO);
let timeout = Duration::from_millis(100);
let mut budgets = Vec::new();
let result = probe_candidates_with(
CANDIDATES.into_iter(),
timeout,
|| clock.get(),
|_, share| {
budgets.push(share);
if budgets.len() == 1 {
clock.set(clock.get().saturating_add(share));
return false;
}
true
},
);
assert!(result, "the second candidate answers, so the probe is true");
assert_eq!(
budgets,
[timeout / 2, timeout / 2],
"the first is offered its share, not the whole budget, and the second the rest"
);
}
#[test]
fn resolver_delay_is_outside_the_connection_budget() {
struct SlowResolver;
impl ToSocketAddrs for SlowResolver {
type Iter = std::vec::IntoIter<SocketAddr>;
fn to_socket_addrs(&self) -> std::io::Result<Self::Iter> {
std::thread::park_timeout(Duration::from_millis(80));
Ok(vec![SocketAddr::from((Ipv4Addr::LOCALHOST, 9))].into_iter())
}
}
let started = Instant::now();
assert!(!probe(SlowResolver, Duration::from_millis(10)));
assert!(started.elapsed() >= Duration::from_millis(80));
}
#[test]
fn a_whole_probe_fits_one_budget() {
let candidates = [
SocketAddr::from(([192, 0, 2, 1], 443)),
SocketAddr::from(([198, 51, 100, 1], 443)),
SocketAddr::from(([203, 0, 113, 1], 443)),
SocketAddr::from(([192, 0, 2, 2], 443)),
];
let timeout = Duration::from_millis(120);
let clock = std::cell::Cell::new(Duration::ZERO);
let mut shares = Vec::new();
let result = probe_candidates_with(
candidates.into_iter(),
timeout,
|| clock.get(),
|_, share| {
shares.push(share);
clock.set(clock.get().saturating_add(share));
false
},
);
assert!(
!result,
"every candidate fails, so the probe reports unreachable"
);
assert_eq!(
shares,
[Duration::from_millis(30); 4],
"each candidate is offered an equal share of what is left"
);
assert_eq!(
clock.get(),
timeout,
"the four candidates share one budget of {timeout:?}, not four"
);
}
}