Browse Source

src/net2: WIP clean up

ghassmo 4 years ago
parent
commit
c9c3c9881d

+ 11 - 0
src/error.rs

@@ -229,6 +229,10 @@ pub enum Error {
 
 
     #[error("Unsupported network transport")]
     #[error("Unsupported network transport")]
     UnsupportedTransport,
     UnsupportedTransport,
+
+    #[cfg(feature = "net2")]
+    #[error("TransportError: {0}")]
+    TransportError(String),
 }
 }
 
 
 #[cfg(feature = "node")]
 #[cfg(feature = "node")]
@@ -353,6 +357,13 @@ impl From<wasmer::ExportError> for Error {
     }
     }
 }
 }
 
 
+#[cfg(feature = "net2")]
+impl<T: std::fmt::Display> From<crate::net2::transport::TransportError<T>> for Error {
+    fn from(err: crate::net2::transport::TransportError<T>) -> Error {
+        Error::TransportError(err.to_string())
+    }
+}
+
 #[cfg(feature = "wasm-runtime")]
 #[cfg(feature = "wasm-runtime")]
 impl From<wasmer::RuntimeError> for Error {
 impl From<wasmer::RuntimeError> for Error {
     fn from(err: wasmer::RuntimeError) -> Error {
     fn from(err: wasmer::RuntimeError) -> Error {

+ 2 - 2
src/net2/acceptor.rs

@@ -61,9 +61,9 @@ impl<T: Transport> Acceptor<T> {
     /// Run the accept loop.
     /// Run the accept loop.
     async fn run_accept_loop(self: Arc<Self>, url: url::Url) -> Result<()> {
     async fn run_accept_loop(self: Arc<Self>, url: url::Url) -> Result<()> {
         let transport = T::new(None, 1024);
         let transport = T::new(None, 1024);
-        let listener = Arc::new(transport.listen_on(url.clone()).unwrap().await.unwrap());
+        let listener = Arc::new(transport.listen_on(url.clone())?.await?);
         loop {
         loop {
-            let stream = T::accept(listener.clone()).await;
+            let stream = T::accept(listener.clone()).await?;
             let channel = Channel::<T>::new(stream, url.clone()).await;
             let channel = Channel::<T>::new(stream, url.clone()).await;
             self.channel_subscriber.notify(Ok(channel)).await;
             self.channel_subscriber.notify(Ok(channel)).await;
         }
         }

+ 1 - 1
src/net2/connector.rs

@@ -23,7 +23,7 @@ impl Connector {
         let stream_result =
         let stream_result =
             timeout(Duration::from_secs(self.settings.connect_timeout_seconds.into()), async {
             timeout(Duration::from_secs(self.settings.connect_timeout_seconds.into()), async {
                 let transport = T::new(None, 1024);
                 let transport = T::new(None, 1024);
-                let connect_stream = transport.dial(hostaddr.clone()).unwrap().await.unwrap();
+                let connect_stream = transport.dial(hostaddr.clone())?.await?;
                 let channel = Channel::<T>::new(connect_stream, hostaddr).await;
                 let channel = Channel::<T>::new(connect_stream, hostaddr).await;
                 Ok(channel)
                 Ok(channel)
             })
             })

+ 7 - 3
src/net2/transport.rs

@@ -20,8 +20,10 @@ pub trait Transport: Sync + Send + 'static + Clone {
 
 
     type Error: Error;
     type Error: Error;
 
 
-    type Listener: Future<Output = Result<Self::Acceptor, Self::Error>> + Sync + Send;
-    type Dial: Future<Output = Result<Self::Connector, Self::Error>> + Sync + Send;
+    type Listener: Future<Output = Result<Self::Acceptor, TransportError<Self::Error>>>
+        + Sync
+        + Send;
+    type Dial: Future<Output = Result<Self::Connector, TransportError<Self::Error>>> + Sync + Send;
 
 
     fn listen_on(self, url: Url) -> Result<Self::Listener, TransportError<Self::Error>>
     fn listen_on(self, url: Url) -> Result<Self::Listener, TransportError<Self::Error>>
     where
     where
@@ -33,7 +35,9 @@ pub trait Transport: Sync + Send + 'static + Clone {
 
 
     fn new(ttl: Option<u32>, backlog: i32) -> Self;
     fn new(ttl: Option<u32>, backlog: i32) -> Self;
 
 
-    async fn accept(listener: Arc<Self::Acceptor>) -> Self::Connector;
+    async fn accept(
+        listener: Arc<Self::Acceptor>,
+    ) -> Result<Self::Connector, TransportError<Self::Error>>;
 }
 }
 
 
 #[derive(Debug, thiserror::Error)]
 #[derive(Debug, thiserror::Error)]

+ 21 - 8
src/net2/transport/tcp.rs

@@ -27,9 +27,14 @@ impl Transport for TcpTransport {
 
 
     type Error = io::Error;
     type Error = io::Error;
 
 
-    type Listener =
-        Pin<Box<dyn Future<Output = Result<Self::Acceptor, Self::Error>> + Send + Sync>>;
-    type Dial = Pin<Box<dyn Future<Output = Result<Self::Connector, Self::Error>> + Send + Sync>>;
+    type Listener = Pin<
+        Box<dyn Future<Output = Result<Self::Acceptor, TransportError<Self::Error>>> + Send + Sync>,
+    >;
+    type Dial = Pin<
+        Box<
+            dyn Future<Output = Result<Self::Connector, TransportError<Self::Error>>> + Send + Sync,
+        >,
+    >;
 
 
     fn listen_on(self, url: Url) -> Result<Self::Listener, TransportError<Self::Error>> {
     fn listen_on(self, url: Url) -> Result<Self::Listener, TransportError<Self::Error>> {
         if url.scheme() != "tcp" {
         if url.scheme() != "tcp" {
@@ -55,8 +60,10 @@ impl Transport for TcpTransport {
         Self { ttl, backlog }
         Self { ttl, backlog }
     }
     }
 
 
-    async fn accept(listener: Arc<Self::Acceptor>) -> Self::Connector {
-        listener.accept().await.unwrap().0
+    async fn accept(
+        listener: Arc<Self::Acceptor>,
+    ) -> Result<Self::Connector, TransportError<Self::Error>> {
+        Ok(listener.accept().await?.0)
     }
     }
 }
 }
 
 
@@ -76,7 +83,10 @@ impl TcpTransport {
         Ok(socket)
         Ok(socket)
     }
     }
 
 
-    async fn do_listen(self, socket_addr: SocketAddr) -> Result<TcpListener, io::Error> {
+    async fn do_listen(
+        self,
+        socket_addr: SocketAddr,
+    ) -> Result<TcpListener, TransportError<io::Error>> {
         let socket = self.create_socket(socket_addr)?;
         let socket = self.create_socket(socket_addr)?;
         socket.bind(&socket_addr.into())?;
         socket.bind(&socket_addr.into())?;
         socket.listen(self.backlog)?;
         socket.listen(self.backlog)?;
@@ -84,7 +94,10 @@ impl TcpTransport {
         Ok(TcpListener::from(std::net::TcpListener::from(socket)))
         Ok(TcpListener::from(std::net::TcpListener::from(socket)))
     }
     }
 
 
-    async fn do_dial(self, socket_addr: SocketAddr) -> Result<TcpStream, io::Error> {
+    async fn do_dial(
+        self,
+        socket_addr: SocketAddr,
+    ) -> Result<TcpStream, TransportError<io::Error>> {
         let socket = self.create_socket(socket_addr)?;
         let socket = self.create_socket(socket_addr)?;
         socket.set_nonblocking(true)?;
         socket.set_nonblocking(true)?;
 
 
@@ -92,7 +105,7 @@ impl TcpTransport {
             Ok(()) => {}
             Ok(()) => {}
             Err(err) if err.raw_os_error() == Some(libc::EINPROGRESS) => {}
             Err(err) if err.raw_os_error() == Some(libc::EINPROGRESS) => {}
             Err(err) if err.kind() == io::ErrorKind::WouldBlock => {}
             Err(err) if err.kind() == io::ErrorKind::WouldBlock => {}
-            Err(err) => return Err(err),
+            Err(err) => return Err(TransportError::Other(err)),
         };
         };
 
 
         let stream = TcpStream::from(std::net::TcpStream::from(socket));
         let stream = TcpStream::from(std::net::TcpStream::from(socket));

+ 19 - 9
src/net2/transport/tls.rs

@@ -87,9 +87,14 @@ impl Transport for TlsTransport {
 
 
     type Error = io::Error;
     type Error = io::Error;
 
 
-    type Listener =
-        Pin<Box<dyn Future<Output = Result<Self::Acceptor, Self::Error>> + Send + Sync>>;
-    type Dial = Pin<Box<dyn Future<Output = Result<Self::Connector, Self::Error>> + Send + Sync>>;
+    type Listener = Pin<
+        Box<dyn Future<Output = Result<Self::Acceptor, TransportError<Self::Error>>> + Send + Sync>,
+    >;
+    type Dial = Pin<
+        Box<
+            dyn Future<Output = Result<Self::Connector, TransportError<Self::Error>>> + Send + Sync,
+        >,
+    >;
 
 
     fn listen_on(self, url: Url) -> Result<Self::Listener, TransportError<Self::Error>> {
     fn listen_on(self, url: Url) -> Result<Self::Listener, TransportError<Self::Error>> {
         if url.scheme() != "tls" {
         if url.scheme() != "tls" {
@@ -151,9 +156,11 @@ impl Transport for TlsTransport {
         Self { ttl, backlog, server_config, client_config }
         Self { ttl, backlog, server_config, client_config }
     }
     }
 
 
-    async fn accept(listener: Arc<Self::Acceptor>) -> Self::Connector {
-        let stream = listener.1.accept().await.unwrap().0;
-        listener.0.accept(stream).await.unwrap().into()
+    async fn accept(
+        listener: Arc<Self::Acceptor>,
+    ) -> Result<Self::Connector, TransportError<Self::Error>> {
+        let stream = listener.1.accept().await?.0;
+        Ok(listener.0.accept(stream).await?.into())
     }
     }
 }
 }
 
 
@@ -173,7 +180,10 @@ impl TlsTransport {
         Ok(socket)
         Ok(socket)
     }
     }
 
 
-    async fn do_listen(self, url: Url) -> Result<(TlsAcceptor, TcpListener), io::Error> {
+    async fn do_listen(
+        self,
+        url: Url,
+    ) -> Result<(TlsAcceptor, TcpListener), TransportError<io::Error>> {
         let socket_addr = url.socket_addrs(|| None)?[0];
         let socket_addr = url.socket_addrs(|| None)?[0];
         let socket = self.create_socket(socket_addr)?;
         let socket = self.create_socket(socket_addr)?;
         socket.bind(&socket_addr.into())?;
         socket.bind(&socket_addr.into())?;
@@ -185,7 +195,7 @@ impl TlsTransport {
         Ok((acceptor, listener))
         Ok((acceptor, listener))
     }
     }
 
 
-    async fn do_dial(self, url: Url) -> Result<TlsStream<TcpStream>, io::Error> {
+    async fn do_dial(self, url: Url) -> Result<TlsStream<TcpStream>, TransportError<io::Error>> {
         let socket_addr = url.socket_addrs(|| None)?[0];
         let socket_addr = url.socket_addrs(|| None)?[0];
         let server_name = ServerName::try_from("dark.fi").unwrap();
         let server_name = ServerName::try_from("dark.fi").unwrap();
         let socket = self.create_socket(socket_addr)?;
         let socket = self.create_socket(socket_addr)?;
@@ -197,7 +207,7 @@ impl TlsTransport {
             Ok(()) => {}
             Ok(()) => {}
             Err(err) if err.raw_os_error() == Some(libc::EINPROGRESS) => {}
             Err(err) if err.raw_os_error() == Some(libc::EINPROGRESS) => {}
             Err(err) if err.kind() == io::ErrorKind::WouldBlock => {}
             Err(err) if err.kind() == io::ErrorKind::WouldBlock => {}
-            Err(err) => return Err(err),
+            Err(err) => return Err(TransportError::Other(err)),
         };
         };
 
 
         let stream = TcpStream::from(std::net::TcpStream::from(socket));
         let stream = TcpStream::from(std::net::TcpStream::from(socket));

+ 20 - 9
src/net2/transport/tor.rs

@@ -194,7 +194,10 @@ impl TorTransport {
             .create_ehs(url)
             .create_ehs(url)
     }
     }
 
 
-    pub async fn do_dial(self, url: Url) -> Result<Socks5Stream<TcpStream>, TorError> {
+    pub async fn do_dial(
+        self,
+        url: Url,
+    ) -> Result<Socks5Stream<TcpStream>, TransportError<TorError>> {
         let socks_url_str = self.socks_url.socket_addrs(|| None)?[0].to_string();
         let socks_url_str = self.socks_url.socket_addrs(|| None)?[0].to_string();
         let host = url.host().unwrap().to_string();
         let host = url.host().unwrap().to_string();
         let port = url.port().unwrap_or(80);
         let port = url.port().unwrap_or(80);
@@ -209,11 +212,12 @@ impl TorTransport {
                 self.socks_url.password().unwrap().to_string(),
                 self.socks_url.password().unwrap().to_string(),
                 config,
                 config,
             )
             )
-            .await?
+            .await
         } else {
         } else {
-            Socks5Stream::connect(socks_url_str, host, port, config).await?
+            Socks5Stream::connect(socks_url_str, host, port, config).await
         };
         };
-        Ok(stream)
+        // FIXME
+        Ok(stream.unwrap())
     }
     }
 
 
     fn create_socket(&self, socket_addr: SocketAddr) -> io::Result<Socket> {
     fn create_socket(&self, socket_addr: SocketAddr) -> io::Result<Socket> {
@@ -226,7 +230,7 @@ impl TorTransport {
         Ok(socket)
         Ok(socket)
     }
     }
 
 
-    pub async fn do_listen(self, url: Url) -> Result<TcpListener, TorError> {
+    pub async fn do_listen(self, url: Url) -> Result<TcpListener, TransportError<TorError>> {
         let socket_addr = url.socket_addrs(|| None)?[0];
         let socket_addr = url.socket_addrs(|| None)?[0];
         let socket = self.create_socket(socket_addr)?;
         let socket = self.create_socket(socket_addr)?;
         socket.bind(&socket_addr.into())?;
         socket.bind(&socket_addr.into())?;
@@ -243,9 +247,14 @@ impl Transport for TorTransport {
 
 
     type Error = TorError;
     type Error = TorError;
 
 
-    type Listener =
-        Pin<Box<dyn Future<Output = Result<Self::Acceptor, Self::Error>> + Send + Sync>>;
-    type Dial = Pin<Box<dyn Future<Output = Result<Self::Connector, Self::Error>> + Send + Sync>>;
+    type Listener = Pin<
+        Box<dyn Future<Output = Result<Self::Acceptor, TransportError<Self::Error>>> + Send + Sync>,
+    >;
+    type Dial = Pin<
+        Box<
+            dyn Future<Output = Result<Self::Connector, TransportError<Self::Error>>> + Send + Sync,
+        >,
+    >;
 
 
     fn listen_on(self, url: Url) -> Result<Self::Listener, TransportError<Self::Error>> {
     fn listen_on(self, url: Url) -> Result<Self::Listener, TransportError<Self::Error>> {
         if url.scheme() != "tcp" {
         if url.scheme() != "tcp" {
@@ -261,7 +270,9 @@ impl Transport for TorTransport {
         unimplemented!()
         unimplemented!()
     }
     }
 
 
-    async fn accept(_listener: Arc<Self::Acceptor>) -> Self::Connector {
+    async fn accept(
+        _listener: Arc<Self::Acceptor>,
+    ) -> Result<Self::Connector, TransportError<Self::Error>> {
         unimplemented!()
         unimplemented!()
     }
     }
 }
 }