Просмотр исходного кода

transport: non-blocking tcp connect

`socket2::connect_with_timeout()` is blocking, and so would wait up to
`outbound_connect_timeout` seconds before returning even when the
underlying connection has been stopped.

To make the connection non-blocking, we wrap it in an Async object and
wait for it to complete using the socket method `take_error()`. If
there's a timeout configured, we use futures::select to return whatever
completes first. Otherwise, we wait on the connection.

For more info, see: https://github.com/rust-lang/socket2/issues/466

Also helpful: https://stackoverflow.com/questions/73101446/non-blocking-tcp-socket-fails-to-connect-using-socket2
draoi 2 лет назад
Родитель
Сommit
a092f0d0eb
1 измененных файлов с 67 добавлено и 14 удалено
  1. 67 14
      src/net/transport/tcp.rs

+ 67 - 14
src/net/transport/tcp.rs

@@ -19,14 +19,21 @@
 use std::{io, time::Duration};
 
 use async_trait::async_trait;
+use futures::{
+    future::{select, Either},
+    pin_mut,
+};
 use futures_rustls::{TlsAcceptor, TlsStream};
 use log::debug;
-use smol::net::{SocketAddr, TcpListener as SmolTcpListener, TcpStream};
+use smol::{
+    net::{SocketAddr, TcpListener as SmolTcpListener, TcpStream},
+    Async, Timer,
+};
 use socket2::{Domain, Socket, TcpKeepalive, Type};
 use url::Url;
 
 use super::{PtListener, PtStream};
-use crate::Result;
+use crate::{Error, Result};
 
 /// TCP Dialer implementation
 #[derive(Debug, Clone)]
@@ -71,26 +78,72 @@ impl TcpDialer {
         debug!(target: "net::tcp::do_dial", "Dialing {} with TCP...", socket_addr);
         let socket = self.create_socket(socket_addr).await?;
 
-        let connection = if let Some(timeout) = timeout {
-            socket.connect_timeout(&socket_addr.into(), timeout)
-        } else {
-            socket.connect(&socket_addr.into())
-        };
+        socket.set_nonblocking(true)?;
 
-        match connection {
+        // Sync start socket connect. A WouldBlock error means this
+        // connection is in progress.
+        match socket.connect(&socket_addr.into()) {
             Ok(()) => {}
             Err(err) if err.raw_os_error() == Some(libc::EINPROGRESS) => {}
             Err(err) if err.kind() == io::ErrorKind::WouldBlock => {}
             Err(err) => return Err(err.into()),
-        }
+        };
 
-        socket.set_nonblocking(true)?;
+        // Wrap socket in an Async wrapper.
+        let async_socket = Async::new(socket)?;
 
-        let stream = std::net::TcpStream::from(socket);
-        let stream = smol::Async::<std::net::TcpStream>::try_from(stream)?;
-        let stream = TcpStream::from(stream);
+        // Wait until the async object becomes writable.
+        let connect = async {
+            match async_socket.get_ref().take_error()? {
+                Some(err) => Err(Error::Io(err.kind())),
+                None => Ok(()),
+            }
+        };
 
-        Ok(stream)
+        // If a timeout is configured, run both the connect and timeout
+        // futures and return whatever finishes first. Otherwise wait on
+        // the connect future.
+        match timeout {
+            Some(t) => {
+                let timeout = Timer::after(t);
+                pin_mut!(timeout);
+                pin_mut!(connect);
+
+                match select(connect, timeout).await {
+                    Either::Left((Ok(_), _)) => {
+                        debug!(target: "net::tcp::do_dial", "Connection successful!");
+                        let stream = {
+                            let socket = async_socket.into_inner()?;
+                            std::net::TcpStream::from(socket)
+                        };
+                        let stream = Async::<std::net::TcpStream>::try_from(stream)?;
+                        let stream = TcpStream::from(stream);
+
+                        Ok(stream)
+                    }
+                    Either::Left((Err(e), _)) => {
+                        debug!(target: "net::tcp::do_dial", "Connection error: {}", e);
+                        return Err(e.into());
+                    }
+
+                    Either::Right((_, _)) => {
+                        debug!(target: "net::tcp::do_dial", "Connection timeed out!");
+                        return Err(Error::ConnectTimeout)
+                    }
+                }
+            }
+            None => {
+                connect.await?;
+                let stream = {
+                    let socket = async_socket.into_inner()?;
+                    std::net::TcpStream::from(socket)
+                };
+                let stream = Async::<std::net::TcpStream>::try_from(stream)?;
+                let stream = TcpStream::from(stream);
+
+                Ok(stream)
+            }
+        }
     }
 }