Skip to main content

ssh_cli/
net.rs

1// SPDX-License-Identifier: MIT OR Apache-2.0
2// G-SECDEV-05: pure module — no `unsafe` permitted (crate root allows only OS FFI / test env).
3#![forbid(unsafe_code)]
4//! TCP dial helpers (Rules Rust — rede).
5//!
6//! Product surface is SSH client + optional local tunnel listener. This module
7//! owns dual-stack DNS resolution and Happy Eyeballs-style connect so callers
8//! never stick on a blackholed first address (`addrs.next().unwrap()` antipattern).
9//!
10//! - **DNS:** async via [`tokio::net::lookup_host`] (cancelable with outer timeouts).
11//! - **Dial:** try all resolved [`SocketAddr`]s; race families with a short delay
12//!   (RFC 8305-inspired) and abort losers on first success.
13//! - **Workload:** pure I/O; no CPU fan-out, no Rayon.
14
15use 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
24/// Resolve `host:port` and dial with multi-address Happy Eyeballs racing.
25///
26/// # Errors
27///
28/// Returns the last connect error when every candidate fails, or
29/// [`io::ErrorKind::AddrNotAvailable`] when DNS yields no addresses.
30pub 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/// Interleave IPv6 and IPv4 candidates (IPv6 first within each pair).
49///
50/// Prefer dual-stack interleave over "all AAAA then all A" so a dead IPv6 path
51/// does not delay every IPv4 attempt until the full v6 list is exhausted.
52#[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
87/// Race connect attempts: first address immediately; further addresses start
88/// after [`HAPPY_EYEBALLS_ATTEMPT_DELAY_MS`] or sooner if the previous attempt fails.
89async 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    // Kick the first candidate immediately.
98    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            // Start next if any remain (previous wave fully failed).
109            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                        // On failure, start the next candidate without waiting out the stagger.
132                        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                        // Expected after abort_all on success path; ignore if we continue.
140                    }
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        // RFC 6761 `.invalid` TLD must not resolve.
198        let err = dial_tcp("no-such-host.invalid", 22)
199            .await
200            .expect_err("must fail DNS or dial");
201        // Either DNS failure or empty-result style errors are acceptable.
202        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}