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

net/transport: add timeout option for dial function

ghassmo 4 лет назад
Родитель
Сommit
f20dd33cb8
5 измененных файлов с 32 добавлено и 21 удалено
  1. 10 9
      src/net/connector.rs
  2. 2 2
      src/net/transport.rs
  3. 15 5
      src/net/transport/tcp.rs
  4. 1 1
      src/net/transport/tor.rs
  5. 4 4
      src/rpc/jsonrpc.rs

+ 10 - 9
src/net/connector.rs

@@ -1,4 +1,4 @@
-use async_std::{future::timeout, sync::Arc};
+use async_std::sync::Arc;
 use std::{env, time::Duration};
 
 use log::error;
@@ -24,18 +24,19 @@ impl Connector {
     /// Establish an outbound connection.
     pub async fn connect(&self, connect_url: Url) -> Result<ChannelPtr> {
         let transport_name = TransportName::try_from(connect_url.clone())?;
-        let result =
-            timeout(Duration::from_secs(self.settings.connect_timeout_seconds.into()), async {
-                self.connect_channel(connect_url, transport_name).await
-            })
-            .await?;
-        result
+        self.connect_channel(
+            connect_url,
+            transport_name,
+            Duration::from_secs(self.settings.connect_timeout_seconds.into()),
+        )
+        .await
     }
 
     async fn connect_channel(
         &self,
         connect_url: Url,
         transport_name: TransportName,
+        timeout: Duration,
     ) -> Result<Arc<Channel>> {
         macro_rules! connect {
             ($stream:expr, $transport:expr, $upgrade:expr) => {{
@@ -67,7 +68,7 @@ impl Connector {
         match transport_name {
             TransportName::Tcp(upgrade) => {
                 let transport = TcpTransport::new(None, 1024);
-                let stream = transport.dial(connect_url.clone());
+                let stream = transport.dial(connect_url.clone(), Some(timeout));
                 connect!(stream, transport, upgrade)
             }
             TransportName::Tor(upgrade) => {
@@ -78,7 +79,7 @@ impl Connector {
 
                 let transport = TorTransport::new(socks5_url, None)?;
 
-                let stream = transport.clone().dial(connect_url.clone());
+                let stream = transport.clone().dial(connect_url.clone(), None);
 
                 connect!(stream, transport, upgrade)
             }

+ 2 - 2
src/net/transport.rs

@@ -1,4 +1,4 @@
-use std::net::SocketAddr;
+use std::{net::SocketAddr, time::Duration};
 
 use async_trait::async_trait;
 // TODO remove *
@@ -82,7 +82,7 @@ pub trait Transport {
     where
         Self: Sized;
 
-    fn dial(self, url: Url) -> Result<Self::Dial>
+    fn dial(self, url: Url, timeout: Option<Duration>) -> Result<Self::Dial>
     where
         Self: Sized;
 

+ 15 - 5
src/net/transport/tcp.rs

@@ -1,5 +1,5 @@
 use async_std::net::{TcpListener, TcpStream};
-use std::{io, net::SocketAddr, pin::Pin};
+use std::{io, net::SocketAddr, pin::Pin, time::Duration};
 
 use async_trait::async_trait;
 use futures::prelude::*;
@@ -87,7 +87,7 @@ impl Transport for TcpTransport {
         Ok(Box::pin(tlsupgrade.upgrade_listener_tls(acceptor)))
     }
 
-    fn dial(self, url: Url) -> Result<Self::Dial> {
+    fn dial(self, url: Url, timeout: Option<Duration>) -> Result<Self::Dial> {
         match url.scheme() {
             "tcp" | "tcp+tls" | "tls" => {}
             x => return Err(Error::UnsupportedTransport(x.to_string())),
@@ -95,7 +95,7 @@ impl Transport for TcpTransport {
 
         let socket_addr = url.socket_addrs(|| None)?[0];
         debug!("{} transport: dialing {}", url.scheme(), socket_addr);
-        Ok(Box::pin(self.do_dial(socket_addr)))
+        Ok(Box::pin(self.do_dial(socket_addr, timeout)))
     }
 
     fn upgrade_dialer(self, connector: Self::Connector) -> Result<Self::TlsDialer> {
@@ -132,10 +132,20 @@ impl TcpTransport {
         Ok(TcpListener::from(std::net::TcpListener::from(socket)))
     }
 
-    async fn do_dial(self, socket_addr: SocketAddr) -> Result<TcpStream> {
+    async fn do_dial(
+        self,
+        socket_addr: SocketAddr,
+        timeout: Option<Duration>,
+    ) -> Result<TcpStream> {
         let socket = self.create_socket(socket_addr)?;
 
-        match socket.connect(&socket_addr.into()) {
+        let connection = if timeout.is_some() {
+            socket.connect_timeout(&socket_addr.into(), timeout.unwrap())
+        } else {
+            socket.connect(&socket_addr.into())
+        };
+
+        match connection {
             Ok(()) => {}
             Err(err) if err.raw_os_error() == Some(libc::EINPROGRESS) => {}
             Err(err) if err.kind() == io::ErrorKind::WouldBlock => {}

+ 1 - 1
src/net/transport/tor.rs

@@ -239,7 +239,7 @@ impl Transport for TorTransport {
         Ok(Box::pin(tlsupgrade.upgrade_listener_tls(acceptor)))
     }
 
-    fn dial(self, url: Url) -> Result<Self::Dial> {
+    fn dial(self, url: Url, _timeout: Option<Duration>) -> Result<Self::Dial> {
         match url.scheme() {
             "tor" | "tor+tls" => {}
             x => return Err(Error::UnsupportedTransport(x.to_string())),

+ 4 - 4
src/rpc/jsonrpc.rs

@@ -241,7 +241,7 @@ pub async fn open_channels(
     match transport_name {
         TransportName::Tcp(upgrade) => {
             let transport = TcpTransport::new(None, 1024);
-            let stream = transport.dial(uri.clone());
+            let stream = transport.dial(uri.clone(), None);
 
             reqrep!(stream, transport, upgrade);
         }
@@ -253,7 +253,7 @@ pub async fn open_channels(
 
             let transport = TorTransport::new(socks5_url, None)?;
 
-            let stream = transport.clone().dial(uri.clone());
+            let stream = transport.clone().dial(uri.clone(), None);
 
             reqrep!(stream, transport, upgrade);
         }
@@ -310,7 +310,7 @@ pub async fn send_request(uri: &Url, data: Value) -> Result<JsonResult> {
     match transport_name {
         TransportName::Tcp(upgrade) => {
             let transport = TcpTransport::new(None, 1024);
-            let stream = transport.dial(uri.clone());
+            let stream = transport.dial(uri.clone(), None);
 
             reply!(stream, transport, upgrade)
         }
@@ -322,7 +322,7 @@ pub async fn send_request(uri: &Url, data: Value) -> Result<JsonResult> {
 
             let transport = TorTransport::new(socks5_url, None)?;
 
-            let stream = transport.clone().dial(uri.clone());
+            let stream = transport.clone().dial(uri.clone(), None);
 
             reply!(stream, transport, upgrade)
         }