1#![forbid(unsafe_code)]
4use std::io;
16use std::net::{IpAddr, SocketAddr};
17use std::time::Duration;
18
19use tokio::net::TcpStream;
20use tokio::task::JoinSet;
21
22use crate::constants::HAPPY_EYEBALLS_ATTEMPT_DELAY_MS;
23
24pub async fn dial_tcp(host: &str, port: u16) -> io::Result<TcpStream> {
31 let addrs: Vec<SocketAddr> = tokio::net::lookup_host((host, port)).await?.collect();
32 if addrs.is_empty() {
33 return Err(io::Error::new(
34 io::ErrorKind::AddrNotAvailable,
35 format!("no addresses resolved for {host}:{port}"),
36 ));
37 }
38 let ordered = interleave_address_families(addrs);
39 tracing::debug!(
40 host,
41 port,
42 candidates = ordered.len(),
43 "dialing TCP with Happy Eyeballs ordering"
44 );
45 race_connect(ordered).await
46}
47
48#[must_use]
53pub fn interleave_address_families(addrs: Vec<SocketAddr>) -> Vec<SocketAddr> {
54 let mut v6 = Vec::new();
55 let mut v4 = Vec::new();
56 for addr in addrs {
57 match addr.ip() {
58 IpAddr::V6(_) => v6.push(addr),
59 IpAddr::V4(_) => v4.push(addr),
60 }
61 }
62 let mut out = Vec::with_capacity(v6.len() + v4.len());
63 let mut i6 = v6.into_iter();
64 let mut i4 = v4.into_iter();
65 loop {
66 match (i6.next(), i4.next()) {
67 (Some(a), Some(b)) => {
68 out.push(a);
69 out.push(b);
70 }
71 (Some(a), None) => {
72 out.push(a);
73 out.extend(i6);
74 break;
75 }
76 (None, Some(b)) => {
77 out.push(b);
78 out.extend(i4);
79 break;
80 }
81 (None, None) => break,
82 }
83 }
84 out
85}
86
87async fn race_connect(addrs: Vec<SocketAddr>) -> io::Result<TcpStream> {
90 debug_assert!(!addrs.is_empty());
91
92 let mut pending = addrs.into_iter();
93 let mut set: JoinSet<(SocketAddr, io::Result<TcpStream>)> = JoinSet::new();
94 let mut last_err: Option<io::Error> = None;
95 let delay = Duration::from_millis(HAPPY_EYEBALLS_ATTEMPT_DELAY_MS);
96
97 if let Some(addr) = pending.next() {
99 set.spawn(async move { (addr, TcpStream::connect(addr).await) });
100 }
101
102 let stagger = tokio::time::sleep(delay);
103 tokio::pin!(stagger);
104 let mut stagger_armed = true;
105
106 loop {
107 if set.is_empty() {
108 if let Some(addr) = pending.next() {
110 set.spawn(async move { (addr, TcpStream::connect(addr).await) });
111 stagger.as_mut().reset(tokio::time::Instant::now() + delay);
112 stagger_armed = true;
113 continue;
114 }
115 break;
116 }
117
118 tokio::select! {
119 biased;
120 Some(joined) = set.join_next() => {
121 match joined {
122 Ok((addr, Ok(stream))) => {
123 tracing::debug!(%addr, "TCP connect succeeded");
124 set.abort_all();
125 while set.join_next().await.is_some() {}
126 return Ok(stream);
127 }
128 Ok((addr, Err(e))) => {
129 tracing::debug!(%addr, err = %e, "TCP connect attempt failed");
130 last_err = Some(e);
131 if let Some(next) = pending.next() {
133 set.spawn(async move { (next, TcpStream::connect(next).await) });
134 stagger.as_mut().reset(tokio::time::Instant::now() + delay);
135 stagger_armed = true;
136 }
137 }
138 Err(join_err) if join_err.is_cancelled() => {
139 }
141 Err(join_err) => {
142 tracing::debug!(err = %join_err, "dial task join error");
143 last_err = Some(io::Error::other(join_err));
144 }
145 }
146 }
147 _ = &mut stagger, if stagger_armed => {
148 stagger_armed = false;
149 if let Some(addr) = pending.next() {
150 set.spawn(async move { (addr, TcpStream::connect(addr).await) });
151 stagger.as_mut().reset(tokio::time::Instant::now() + delay);
152 stagger_armed = true;
153 }
154 }
155 }
156 }
157
158 Err(last_err.unwrap_or_else(|| {
159 io::Error::new(io::ErrorKind::ConnectionRefused, "all dial attempts failed")
160 }))
161}
162
163#[cfg(test)]
164mod tests {
165 use super::*;
166 use std::net::{Ipv4Addr, Ipv6Addr, SocketAddrV4, SocketAddrV6};
167
168 #[test]
169 fn interleave_prefers_v6_then_pairs_with_v4() {
170 let addrs = vec![
171 SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 22)),
172 SocketAddr::V6(SocketAddrV6::new(Ipv6Addr::LOCALHOST, 22, 0, 0)),
173 SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::new(127, 0, 0, 2), 22)),
174 SocketAddr::V6(SocketAddrV6::new(Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 2), 22, 0, 0)),
175 ];
176 let ordered = interleave_address_families(addrs);
177 assert_eq!(ordered.len(), 4);
178 assert!(ordered[0].is_ipv6());
179 assert!(ordered[1].is_ipv4());
180 assert!(ordered[2].is_ipv6());
181 assert!(ordered[3].is_ipv4());
182 }
183
184 #[test]
185 fn interleave_v4_only_preserves_all() {
186 let addrs = vec![
187 SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 1)),
188 SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 2)),
189 ];
190 let ordered = interleave_address_families(addrs);
191 assert_eq!(ordered.len(), 2);
192 assert!(ordered.iter().all(|a| a.is_ipv4()));
193 }
194
195 #[tokio::test(flavor = "current_thread")]
196 async fn dial_unresolvable_host_errors() {
197 let err = dial_tcp("no-such-host.invalid", 22)
199 .await
200 .expect_err("must fail DNS or dial");
201 let kind = err.kind();
203 let msg = err.to_string();
204 assert!(
205 matches!(
206 kind,
207 io::ErrorKind::AddrNotAvailable
208 | io::ErrorKind::NotFound
209 | io::ErrorKind::Other
210 | io::ErrorKind::InvalidInput
211 | io::ErrorKind::TimedOut
212 | io::ErrorKind::ConnectionRefused
213 ) || msg.contains("failed to lookup")
214 || msg.contains("Name or service")
215 || msg.contains("no addresses")
216 || msg.contains("Temporary failure")
217 || msg.contains("No address"),
218 "unexpected dial error: {err:?}"
219 );
220 }
221}