Bläddra i källkod

net/transport/tcp: Restore usage of socket2 to create TCP sockets.

parazyd 2 år sedan
förälder
incheckning
33631ab318
4 ändrade filer med 88 tillägg och 17 borttagningar
  1. 12 1
      Cargo.lock
  2. 2 0
      Cargo.toml
  3. 2 2
      src/net/transport.rs
  4. 72 14
      src/net/transport/tcp.rs

+ 12 - 1
Cargo.lock

@@ -383,7 +383,7 @@ dependencies = [
  "polling",
  "rustix 0.37.23",
  "slab",
- "socket2",
+ "socket2 0.4.9",
  "waker-fn",
 ]
 
@@ -1434,6 +1434,7 @@ dependencies = [
  "sled",
  "sled-overlay",
  "smol",
+ "socket2 0.5.3",
  "structopt",
  "structopt-toml",
  "thiserror",
@@ -5161,6 +5162,16 @@ dependencies = [
  "winapi",
 ]
 
+[[package]]
+name = "socket2"
+version = "0.5.3"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "2538b18701741680e0322a2302176d3253a35388e2e62f172f64f4f16605f877"
+dependencies = [
+ "libc",
+ "windows-sys 0.48.0",
+]
+
 [[package]]
 name = "spin"
 version = "0.5.2"

+ 2 - 0
Cargo.toml

@@ -66,6 +66,7 @@ pin-project-lite = {version = "0.2.12", optional = true}
 async-rustls = {version = "0.4.0", features = ["dangerous_configuration"], optional = true}
 
 # Pluggable Transports
+socket2 = {version = "0.5.3", features = ["all"], optional = true}
 arti-client = {version = "0.10.0", default-features = false, features = ["async-std", "rustls", "onion-service-client"], optional = true}
 tor-hscrypto = {version = "0.3.1", optional = true}
 
@@ -192,6 +193,7 @@ net = [
     "rustls-pemfile",
     "semver",
     "smol",
+    "socket2",
     "structopt",
     "structopt-toml",
     "url",

+ 2 - 2
src/net/transport.rs

@@ -270,7 +270,7 @@ impl Listener {
             "tcp" => {
                 // Build a TCP listener
                 enforce_hostport!(endpoint);
-                let variant = tcp::TcpListener::new().await?;
+                let variant = tcp::TcpListener::new(1024).await?;
                 let variant = ListenerVariant::Tcp(variant);
                 Ok(Self { endpoint, variant })
             }
@@ -279,7 +279,7 @@ impl Listener {
             "tcp+tls" => {
                 // Build a TCP listener wrapped with TLS
                 enforce_hostport!(endpoint);
-                let variant = tcp::TcpListener::new().await?;
+                let variant = tcp::TcpListener::new(1024).await?;
                 let variant = ListenerVariant::TcpTls(variant);
                 Ok(Self { endpoint, variant })
             }

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

@@ -16,16 +16,17 @@
  * along with this program.  If not, see <https://www.gnu.org/licenses/>.
  */
 
-use std::time::Duration;
+use std::{io, time::Duration};
 
 use async_rustls::{TlsAcceptor, TlsStream};
 use async_trait::async_trait;
 use log::debug;
 use smol::net::{SocketAddr, TcpListener as SmolTcpListener, TcpStream};
+use socket2::{Domain, Socket, TcpKeepalive, Type};
 use url::Url;
 
 use super::{PtListener, PtStream};
-use crate::{system::io_timeout, Result};
+use crate::Result;
 
 /// TCP Dialer implementation
 #[derive(Debug, Clone)]
@@ -40,24 +41,54 @@ impl TcpDialer {
         Ok(Self { ttl })
     }
 
+    /// Internal helper function to create a TCP socket.
+    async fn create_socket(&self, socket_addr: SocketAddr) -> io::Result<Socket> {
+        let domain = if socket_addr.is_ipv4() { Domain::IPV4 } else { Domain::IPV6 };
+        let socket = Socket::new(domain, Type::STREAM, Some(socket2::Protocol::TCP))?;
+
+        if socket_addr.is_ipv6() {
+            socket.set_only_v6(true)?;
+        }
+
+        if let Some(ttl) = self.ttl {
+            socket.set_ttl(ttl)?;
+        }
+
+        socket.set_nodelay(true)?;
+        let keepalive = TcpKeepalive::new().with_time(Duration::from_secs(20));
+        socket.set_tcp_keepalive(&keepalive)?;
+        socket.set_reuse_port(true)?;
+
+        Ok(socket)
+    }
+
     /// Internal dial function
     pub(crate) async fn do_dial(
         &self,
         socket_addr: SocketAddr,
-        conn_timeout: Option<Duration>,
+        timeout: Option<Duration>,
     ) -> Result<TcpStream> {
         debug!(target: "net::tcp::do_dial", "Dialing {} with TCP...", socket_addr);
-        let stream = if let Some(conn_timeout) = conn_timeout {
-            io_timeout(conn_timeout, TcpStream::connect(socket_addr)).await?
+        let socket = self.create_socket(socket_addr).await?;
+
+        let connection = if let Some(timeout) = timeout {
+            socket.connect_timeout(&socket_addr.into(), timeout)
         } else {
-            TcpStream::connect(socket_addr).await?
+            socket.connect(&socket_addr.into())
         };
 
-        if let Some(ttl) = self.ttl {
-            stream.set_ttl(ttl)?;
+        match connection {
+            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()),
         }
 
-        stream.set_nodelay(true)?;
+        socket.set_nonblocking(true)?;
+
+        let stream = std::net::TcpStream::from(socket);
+        let stream = smol::Async::<std::net::TcpStream>::try_from(stream)?;
+        let stream = TcpStream::from(stream);
 
         Ok(stream)
     }
@@ -65,18 +96,45 @@ impl TcpDialer {
 
 /// TCP Listener implementation
 #[derive(Debug, Clone)]
-pub struct TcpListener;
+pub struct TcpListener {
+    /// Size of the listen backlog for listen sockets
+    backlog: i32,
+}
 
 impl TcpListener {
     /// Instantiate a new [`TcpListener`] with given backlog size.
-    pub async fn new() -> Result<Self> {
-        Ok(Self {})
+    pub async fn new(backlog: i32) -> Result<Self> {
+        Ok(Self { backlog })
+    }
+
+    /// Internal helper function to create a TCP socket.
+    async fn create_socket(&self, socket_addr: SocketAddr) -> io::Result<Socket> {
+        let domain = if socket_addr.is_ipv4() { Domain::IPV4 } else { Domain::IPV6 };
+        let socket = Socket::new(domain, Type::STREAM, Some(socket2::Protocol::TCP))?;
+
+        if socket_addr.is_ipv6() {
+            socket.set_only_v6(true)?;
+        }
+
+        socket.set_nodelay(true)?;
+        let keepalive = TcpKeepalive::new().with_time(Duration::from_secs(20));
+        socket.set_tcp_keepalive(&keepalive)?;
+        socket.set_reuse_port(true)?;
+
+        Ok(socket)
     }
 
     /// Internal listen function
     pub(crate) async fn do_listen(&self, socket_addr: SocketAddr) -> Result<SmolTcpListener> {
-        let listener = SmolTcpListener::bind(socket_addr).await?;
-        Ok(listener)
+        let socket = self.create_socket(socket_addr).await?;
+        socket.bind(&socket_addr.into())?;
+        socket.listen(self.backlog)?;
+        socket.set_nonblocking(true)?;
+
+        let listener = std::net::TcpListener::from(socket);
+        let listener = smol::Async::<std::net::TcpListener>::try_from(listener)?;
+
+        Ok(SmolTcpListener::from(listener))
     }
 }