use std::net::SocketAddr;
use std::time::Duration;
use futures::FutureExt;
use futures::future::BoxFuture;
use tokio::net::TcpStream;
use tokio::time::sleep;
use tracing::debug;
const DEFAULT_IPV6_HEAD_START: Duration = Duration::from_millis(250);
pub async fn connect_happy_eyeballs(addrs: Vec<SocketAddr>) -> std::io::Result<TcpStream> {
connect_with_head_start(addrs, DEFAULT_IPV6_HEAD_START).await
}
pub async fn connect_with_head_start(
addrs: Vec<SocketAddr>,
ipv6_head_start: Duration,
) -> std::io::Result<TcpStream> {
if addrs.is_empty() {
return Err(std::io::Error::new(
std::io::ErrorKind::AddrNotAvailable,
"no addresses to connect to",
));
}
let mut ipv6_addrs: Vec<SocketAddr> = Vec::new();
let mut ipv4_addrs: Vec<SocketAddr> = Vec::new();
for addr in addrs {
match addr {
SocketAddr::V6(_) => ipv6_addrs.push(addr),
SocketAddr::V4(_) => ipv4_addrs.push(addr),
}
}
if ipv6_addrs.is_empty() {
debug!("Happy Eyeballs: single-family (IPv4 only), skipping race");
return connect_first_available(ipv4_addrs).await;
}
if ipv4_addrs.is_empty() {
debug!("Happy Eyeballs: single-family (IPv6 only), skipping race");
return connect_first_available(ipv6_addrs).await;
}
debug!(
"Happy Eyeballs: racing {} IPv6 and {} IPv4 addresses (head start {:?})",
ipv6_addrs.len(),
ipv4_addrs.len(),
ipv6_head_start,
);
let mut ipv6_fut: BoxFuture<'static, std::io::Result<TcpStream>> =
connect_first_available(ipv6_addrs).boxed();
let mut ipv4_fut: BoxFuture<'static, std::io::Result<TcpStream>> =
connect_first_available(ipv4_addrs).boxed();
let mut head_start: BoxFuture<'static, ()> = sleep(ipv6_head_start).boxed();
let mut ipv4_started = false;
let mut ipv6_done = false;
let mut ipv4_done = false;
let mut ipv6_err: Option<std::io::Error> = None;
let mut ipv4_err: Option<std::io::Error> = None;
loop {
tokio::select! {
result = &mut ipv6_fut, if !ipv6_done => {
ipv6_done = true;
match result {
Ok(stream) => {
debug!("Happy Eyeballs: IPv6 connected first");
return Ok(stream);
}
Err(e) => {
debug!("Happy Eyeballs: IPv6 attempt failed: {e}");
ipv6_err = Some(e);
ipv4_started = true;
}
}
}
result = &mut ipv4_fut, if ipv4_started && !ipv4_done => {
ipv4_done = true;
match result {
Ok(stream) => {
debug!("Happy Eyeballs: IPv4 connected");
return Ok(stream);
}
Err(e) => {
debug!("Happy Eyeballs: IPv4 attempt failed: {e}");
ipv4_err = Some(e);
}
}
}
_ = &mut head_start, if !ipv4_started => {
ipv4_started = true;
debug!("Happy Eyeballs: IPv6 head start elapsed, starting IPv4");
}
}
if ipv6_done && ipv4_done {
return Err(ipv4_err
.or(ipv6_err)
.unwrap_or_else(|| std::io::Error::other("all connection attempts failed")));
}
}
}
async fn connect_first_available(addrs: Vec<SocketAddr>) -> std::io::Result<TcpStream> {
let mut last_err = std::io::Error::other("no addresses");
for addr in addrs {
match TcpStream::connect(addr).await {
Ok(stream) => {
debug!("Happy Eyeballs: connected to {addr}");
return Ok(stream);
}
Err(e) => {
debug!("Happy Eyeballs: connect to {addr} failed: {e}");
last_err = e;
}
}
}
Err(last_err)
}
#[cfg(test)]
mod tests {
use super::*;
use std::net::{Ipv4Addr, Ipv6Addr, SocketAddrV4, SocketAddrV6};
#[tokio::test]
async fn test_connect_empty_addresses_returns_error() {
let result = connect_happy_eyeballs(vec![]).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_connect_ipv4_only_closed_port_fails() {
let addr = SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 1));
let result = connect_happy_eyeballs(vec![addr]).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_connect_to_local_listener_succeeds() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let connect_task = tokio::spawn(async move { connect_happy_eyeballs(vec![addr]).await });
let _conn = listener.accept().await.unwrap();
let result = connect_task.await.unwrap();
assert!(result.is_ok());
}
#[tokio::test]
async fn test_happy_eyeballs_ipv6_unreachable_falls_back_to_ipv4() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let v4_addr = listener.local_addr().unwrap();
let v6_addr = SocketAddr::V6(SocketAddrV6::new(Ipv6Addr::LOCALHOST, 1, 0, 0));
let addrs = vec![v6_addr, v4_addr];
let start = std::time::Instant::now();
let connect_task = tokio::spawn(async move {
connect_with_head_start(addrs, Duration::from_millis(250)).await
});
let _conn = listener.accept().await.unwrap();
let result = connect_task.await.unwrap();
let elapsed = start.elapsed();
assert!(
result.is_ok(),
"should connect via IPv4 fallback: {:?}",
result.err()
);
assert!(
elapsed < Duration::from_secs(5),
"fallback took too long: {:?}",
elapsed
);
}
#[tokio::test]
async fn test_connect_first_available_returns_first_success() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let connect_task = tokio::spawn(async move { connect_first_available(vec![addr]).await });
let _conn = listener.accept().await.unwrap();
let result = connect_task.await.unwrap();
assert!(result.is_ok());
}
#[tokio::test]
async fn test_connect_first_available_all_fail_returns_error() {
let addrs = vec![
SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 1)),
SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 2)),
];
let result = connect_first_available(addrs).await;
assert!(result.is_err());
}
}