use std::future::poll_fn;
use std::io;
use std::net::{SocketAddr, TcpStream as StdTcpStream};
#[cfg(unix)]
use std::os::unix::io::AsRawFd;
#[cfg(windows)]
use std::task::Poll;
use std::time::Duration;
use super::{AsyncTcpStream, SocketLease, poll_ready_op};
use crate::Interest;
#[cfg(windows)]
mod reprobe;
#[cfg(test)]
mod tests;
impl AsyncTcpStream {
pub async fn connect(addr: SocketAddr) -> io::Result<Self> {
let stream = Self::from_nonblocking(start_connect(addr)?);
let owner = SocketLease::from(&stream.inner);
let mut waiter = None;
#[cfg(windows)]
let mut reprobe = reprobe::Registration::new();
poll_fn(|cx| {
let polled = poll_ready_op(cx, owner.clone(), Interest::WRITABLE, &mut waiter, || {
connect_outcome(&stream.inner, Duration::ZERO)
});
#[cfg(windows)]
if polled.is_pending()
&& let Err(error) = reprobe.arm(cx.waker())
{
return Poll::Ready(Err(error));
}
polled
})
.await?;
Ok(stream)
}
}
fn start_connect(addr: SocketAddr) -> io::Result<StdTcpStream> {
use socket2::{Domain, Protocol, Socket, Type};
let socket = Socket::new(Domain::for_address(addr), Type::STREAM, Some(Protocol::TCP))?;
socket.set_nonblocking(true)?;
match socket.connect(&addr.into()) {
Ok(()) => {}
Err(error) if connect_in_progress(&error) => {}
Err(error) => return Err(error),
}
Ok(socket.into())
}
fn connect_in_progress(error: &io::Error) -> bool {
#[cfg(unix)]
{
error.raw_os_error() == Some(libc::EINPROGRESS)
}
#[cfg(windows)]
{
error.kind() == io::ErrorKind::WouldBlock
}
}
fn connect_outcome(stream: &StdTcpStream, wait: Duration) -> io::Result<()> {
if !connect_settled(stream, wait)? {
return Err(io::ErrorKind::WouldBlock.into());
}
if let Some(error) = stream.take_error()? {
return Err(error);
}
stream.peer_addr().map(drop)
}
#[cfg(unix)]
fn connect_settled(stream: &StdTcpStream, wait: Duration) -> io::Result<bool> {
let wait_ms = libc::c_int::try_from(wait.as_millis()).unwrap_or(libc::c_int::MAX);
let mut probe = libc::pollfd {
fd: stream.as_raw_fd(),
events: libc::POLLOUT,
revents: 0,
};
match unsafe { libc::poll(&raw mut probe, 1, wait_ms) } {
0 => Ok(false),
-1 => {
let error = io::Error::last_os_error();
if error.kind() == io::ErrorKind::Interrupted {
Ok(false)
} else {
Err(error)
}
}
_ => Ok(probe.revents & (libc::POLLOUT | libc::POLLERR | libc::POLLHUP) != 0),
}
}
#[cfg(windows)]
fn connect_settled(stream: &StdTcpStream, wait: Duration) -> io::Result<bool> {
use std::os::windows::io::AsRawSocket;
use windows::Win32::Networking::WinSock::{
FD_SET, SOCKET, SOCKET_ERROR, TIMEVAL, WSAGetLastError, select,
};
let socket = SOCKET(
usize::try_from(stream.as_raw_socket())
.map_err(|_| io::Error::other("socket handle exceeds the platform word"))?,
);
let single = || {
let mut fd_array = FD_SET::default().fd_array;
fd_array[0] = socket;
FD_SET {
fd_count: 1,
fd_array,
}
};
let mut writable = single();
let mut failed = single();
let timeout = TIMEVAL {
tv_sec: i32::try_from(wait.as_secs()).unwrap_or(i32::MAX),
tv_usec: i32::try_from(wait.subsec_micros()).unwrap_or(0),
};
let ready = unsafe {
select(
0,
None,
Some(&raw mut writable),
Some(&raw mut failed),
Some(&raw const timeout),
)
};
if ready == SOCKET_ERROR {
return Err(io::Error::from_raw_os_error(unsafe { WSAGetLastError() }.0));
}
Ok(writable.fd_count > 0 || failed.fd_count > 0)
}