use std::net::SocketAddr;
use std::time::{Duration, Instant};
use tokio::net::{self, TcpStream};
use tokio::time::timeout;
use tracing::{debug, info, trace, warn};
use crate::core::TdsResult;
use crate::error::Error;
pub const MAX_PARALLEL_IPS: usize = 64;
pub const DEFAULT_PARALLEL_TIMEOUT_MS: u64 = 15000;
#[derive(Clone, Debug)]
pub struct ParallelConnectConfig {
pub timeout_ms: u64,
pub keep_alive_in_ms: u32,
pub keep_alive_interval_in_ms: u32,
}
impl Default for ParallelConnectConfig {
fn default() -> Self {
Self {
timeout_ms: DEFAULT_PARALLEL_TIMEOUT_MS,
keep_alive_in_ms: 30_000,
keep_alive_interval_in_ms: 1_000,
}
}
}
#[derive(Debug)]
pub struct ParallelConnectResult {
pub stream: TcpStream,
pub connected_address: SocketAddr,
pub total_addresses: usize,
pub failed_attempts: usize,
}
pub async fn parallel_connect(
host: &str,
port: u16,
config: &ParallelConnectConfig,
) -> TdsResult<ParallelConnectResult> {
info!(
"Starting parallel connection to {}:{} with timeout {}ms",
host, port, config.timeout_ms
);
let deadline = Instant::now() + Duration::from_millis(config.timeout_ms);
let addresses: Vec<SocketAddr> = match timeout(
deadline.saturating_duration_since(Instant::now()),
tokio::net::lookup_host((host, port)),
)
.await
{
Ok(Ok(resolved)) => resolved.take(MAX_PARALLEL_IPS).collect(),
Ok(Err(e)) => return Err(Error::from(e)),
Err(_elapsed) => {
warn!(
"DNS resolution for {}:{} timed out after {}ms",
host, port, config.timeout_ms
);
return Err(Error::TimeoutError(crate::error::TimeoutErrorType::String(
format!("Connection timeout: DNS resolution for {host}:{port} timed out"),
)));
}
};
if addresses.is_empty() {
return Err(Error::ConnectionError(format!(
"DNS resolution returned no addresses for {}:{}",
host, port
)));
}
let total_addresses = addresses.len();
info!(
"Resolved {} addresses for {}:{}: {:?}",
total_addresses, host, port, addresses
);
if total_addresses > MAX_PARALLEL_IPS {
warn!(
"Resolved {} addresses, but only {} will be used",
total_addresses, MAX_PARALLEL_IPS
);
}
let connect_futures: Vec<_> = addresses
.iter()
.enumerate()
.map(|(idx, addr)| {
let addr = *addr;
let keep_alive_in_ms = config.keep_alive_in_ms;
let keep_alive_interval_in_ms = config.keep_alive_interval_in_ms;
async move {
trace!("Attempting connection {} to {}", idx, addr);
match connect_with_keepalive(addr, keep_alive_in_ms, keep_alive_interval_in_ms)
.await
{
Ok(stream) => {
info!("Connection {} to {} succeeded", idx, addr);
Ok((stream, addr, idx))
}
Err(e) => {
debug!("Connection {} to {} failed: {}", idx, addr, e);
Err((e, addr, idx))
}
}
}
})
.collect();
let result = timeout(
deadline.saturating_duration_since(Instant::now()),
race_connections(connect_futures),
)
.await;
match result {
Ok(Ok((stream, addr, _idx, failed_attempts))) => {
info!(
"Parallel connection succeeded to {} after {} failed attempts",
addr, failed_attempts
);
Ok(ParallelConnectResult {
stream,
connected_address: addr,
total_addresses,
failed_attempts,
})
}
Ok(Err(last_error)) => {
warn!(
"All {} parallel connections failed. Last error: {}",
total_addresses, last_error
);
Err(Error::ConnectionError(format!(
"All parallel connection attempts failed to {}:{}. Last error: {}",
host, port, last_error
)))
}
Err(_) => {
warn!("Parallel connection timeout after {}ms", config.timeout_ms);
Err(Error::TimeoutError(crate::error::TimeoutErrorType::String(
"Connection timeout: Connection attempt timed out".to_string(),
)))
}
}
}
pub async fn parallel_connect_to_addresses(
addresses: Vec<SocketAddr>,
config: &ParallelConnectConfig,
) -> TdsResult<ParallelConnectResult> {
if addresses.is_empty() {
return Err(Error::ConnectionError(
"No addresses provided for parallel connection".to_string(),
));
}
let total_addresses = addresses.len();
info!(
"Starting parallel connection to {} addresses: {:?}",
total_addresses, addresses
);
if total_addresses > MAX_PARALLEL_IPS {
warn!(
"Provided {} addresses, but only {} will be used",
total_addresses, MAX_PARALLEL_IPS
);
}
let addresses: Vec<SocketAddr> = addresses.into_iter().take(MAX_PARALLEL_IPS).collect();
let total_addresses = addresses.len();
let connect_futures: Vec<_> = addresses
.iter()
.enumerate()
.map(|(idx, addr)| {
let addr = *addr;
let keep_alive_in_ms = config.keep_alive_in_ms;
let keep_alive_interval_in_ms = config.keep_alive_interval_in_ms;
async move {
trace!("Attempting connection {} to {}", idx, addr);
match connect_with_keepalive(addr, keep_alive_in_ms, keep_alive_interval_in_ms)
.await
{
Ok(stream) => {
info!("Connection {} to {} succeeded", idx, addr);
Ok((stream, addr, idx))
}
Err(e) => {
debug!("Connection {} to {} failed: {}", idx, addr, e);
Err((e, addr, idx))
}
}
}
})
.collect();
let result = timeout(
Duration::from_millis(config.timeout_ms),
race_connections(connect_futures),
)
.await;
match result {
Ok(Ok((stream, addr, _idx, failed_attempts))) => {
info!(
"Parallel connection succeeded to {} after {} failed attempts",
addr, failed_attempts
);
Ok(ParallelConnectResult {
stream,
connected_address: addr,
total_addresses,
failed_attempts,
})
}
Ok(Err(last_error)) => {
warn!(
"All {} parallel connections failed. Last error: {}",
total_addresses, last_error
);
Err(Error::ConnectionError(format!(
"All parallel connection attempts failed. Last error: {}",
last_error
)))
}
Err(_) => {
warn!("Parallel connection timeout after {}ms", config.timeout_ms);
Err(Error::TimeoutError(crate::error::TimeoutErrorType::String(
"Connection timeout: Connection attempt timed out".to_string(),
)))
}
}
}
async fn connect_with_keepalive(
addr: SocketAddr,
keep_alive_in_ms: u32,
keep_alive_interval_in_ms: u32,
) -> Result<TcpStream, std::io::Error> {
let socket = if addr.is_ipv6() {
net::TcpSocket::new_v6()?
} else {
net::TcpSocket::new_v4()?
};
let keep_alive_settings = socket2::TcpKeepalive::new()
.with_time(Duration::from_millis(keep_alive_in_ms as u64))
.with_interval(Duration::from_millis(keep_alive_interval_in_ms as u64));
let socket2_socket = socket2::SockRef::from(&socket);
socket2_socket.set_tcp_keepalive(&keep_alive_settings)?;
socket2_socket.set_nodelay(true)?;
socket.connect(addr).await
}
async fn race_connections(
connect_futures: Vec<
impl std::future::Future<
Output = Result<(TcpStream, SocketAddr, usize), (std::io::Error, SocketAddr, usize)>,
> + Send
+ 'static,
>,
) -> Result<(TcpStream, SocketAddr, usize, usize), std::io::Error> {
use tokio::sync::mpsc;
let (tx, mut rx) = mpsc::channel::<Result<(TcpStream, SocketAddr, usize), std::io::Error>>(1);
let total = connect_futures.len();
let handles: Vec<_> = connect_futures
.into_iter()
.map(|fut| {
let tx = tx.clone();
tokio::spawn(async move {
let result = fut.await;
match result {
Ok((stream, addr, idx)) => {
let _ = tx.send(Ok((stream, addr, idx))).await;
}
Err((e, _addr, _idx)) => {
let _ = tx.send(Err(e)).await;
}
}
})
})
.collect();
drop(tx);
let mut failed_attempts = 0;
let mut last_error = std::io::Error::new(
std::io::ErrorKind::NotConnected,
"No connection attempts made",
);
while let Some(result) = rx.recv().await {
match result {
Ok((stream, addr, idx)) => {
for handle in handles {
handle.abort();
}
return Ok((stream, addr, idx, failed_attempts));
}
Err(e) => {
failed_attempts += 1;
last_error = e;
if failed_attempts >= total {
break;
}
}
}
}
Err(last_error)
}
#[cfg(test)]
mod tests {
use super::*;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
use std::sync::atomic::{AtomicUsize, Ordering};
#[tokio::test(flavor = "current_thread")]
async fn parallel_connect_resolution_yields_to_the_executor() {
let heartbeats = std::sync::Arc::new(AtomicUsize::new(0));
let heartbeats_task = heartbeats.clone();
let heartbeat = tokio::spawn(async move {
loop {
heartbeats_task.fetch_add(1, Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(1)).await;
}
});
let config = ParallelConnectConfig {
timeout_ms: 5000,
..Default::default()
};
let result =
parallel_connect("invalid.host.that.does.not.exist.local", 1433, &config).await;
heartbeat.abort();
assert!(
result.is_err(),
"resolution must fail for a nonexistent host"
);
assert!(
heartbeats.load(Ordering::SeqCst) > 0,
"the heartbeat task never ran while resolution was in flight — \
resolution is blocking the executor instead of awaiting it"
);
}
#[test]
fn test_parallel_connect_config_default() {
let config = ParallelConnectConfig::default();
assert_eq!(config.timeout_ms, DEFAULT_PARALLEL_TIMEOUT_MS);
assert_eq!(config.keep_alive_in_ms, 30_000);
assert_eq!(config.keep_alive_interval_in_ms, 1_000);
}
#[test]
fn test_max_parallel_ips() {
assert_eq!(MAX_PARALLEL_IPS, 64);
}
#[tokio::test]
async fn test_parallel_connect_invalid_host() {
let config = ParallelConnectConfig {
timeout_ms: 1000,
..Default::default()
};
let result =
parallel_connect("invalid.host.that.does.not.exist.local", 1433, &config).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_parallel_connect_connection_refused() {
let config = ParallelConnectConfig {
timeout_ms: 1000,
..Default::default()
};
let result = parallel_connect("127.0.0.1", 59999, &config).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_connect_with_keepalive_v4() {
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)), 59999);
let result = connect_with_keepalive(addr, 30_000, 1_000).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_connect_with_keepalive_v6() {
let addr = SocketAddr::new(IpAddr::V6(Ipv6Addr::LOCALHOST), 59999);
let result = connect_with_keepalive(addr, 30_000, 1_000).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_parallel_connect_to_addresses_empty() {
let config = ParallelConnectConfig {
timeout_ms: 1000,
..Default::default()
};
let result = parallel_connect_to_addresses(vec![], &config).await;
assert!(result.is_err());
match result {
Err(Error::ConnectionError(msg)) => {
assert!(msg.contains("No addresses provided"));
}
_ => panic!("Expected ConnectionError for empty addresses"),
}
}
#[tokio::test]
async fn test_parallel_connect_to_addresses_single_failure() {
let config = ParallelConnectConfig {
timeout_ms: 1000,
..Default::default()
};
let addresses = vec![SocketAddr::new(
IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)),
59999,
)];
let result = parallel_connect_to_addresses(addresses, &config).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_parallel_connect_to_addresses_multiple_failures() {
let config = ParallelConnectConfig {
timeout_ms: 5000, ..Default::default()
};
let addresses = vec![
SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)), 59997),
SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)), 59998),
SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)), 59999),
];
let result = parallel_connect_to_addresses(addresses, &config).await;
assert!(result.is_err());
match result {
Err(Error::ConnectionError(_)) | Err(Error::TimeoutError(_)) => {
}
Err(e) => panic!("Unexpected error type: {:?}", e),
Ok(_) => panic!("Expected error but got success"),
}
}
#[tokio::test]
async fn test_parallel_connect_timeout() {
let config = ParallelConnectConfig {
timeout_ms: 1, ..Default::default()
};
let result = parallel_connect("10.255.255.1", 1433, &config).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_parallel_connect_to_addresses_timeout() {
let config = ParallelConnectConfig {
timeout_ms: 1, ..Default::default()
};
let addresses = vec![SocketAddr::new(
IpAddr::V4(Ipv4Addr::new(10, 255, 255, 1)),
1433,
)];
let result = parallel_connect_to_addresses(addresses, &config).await;
assert!(result.is_err());
}
#[test]
fn test_parallel_connect_result_fields() {
let config = ParallelConnectConfig {
timeout_ms: 5000,
keep_alive_in_ms: 10_000,
keep_alive_interval_in_ms: 2_000,
};
assert_eq!(config.timeout_ms, 5000);
assert_eq!(config.keep_alive_in_ms, 10_000);
assert_eq!(config.keep_alive_interval_in_ms, 2_000);
let debug_str = format!("{:?}", config);
assert!(debug_str.contains("ParallelConnectConfig"));
assert!(debug_str.contains("5000"));
}
}