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

net: Rework transports for protocol upgrades.

parazyd 4 лет назад
Родитель
Сommit
be3f87cfea

+ 7 - 11
src/error.rs

@@ -96,9 +96,6 @@ pub enum Error {
     #[error("Unsupported network transport: {0}")]
     UnsupportedTransport(String),
 
-    #[error("Transport error: {0}")]
-    TransportError(String),
-
     #[error("Connection failed")]
     ConnectFailed,
 
@@ -189,6 +186,9 @@ pub enum Error {
     #[error("Cashier error: {0}")]
     CashierError(String),
 
+    #[error("Tor error: {0}")]
+    TorError(String),
+
     // ===============
     // Database errors
     // ===============
@@ -273,6 +273,10 @@ pub enum Error {
     #[error("Failed decoding bincode: {0}")]
     ZkasDecoderError(&'static str),
 
+    #[cfg(feature = "regex")]
+    #[error(transparent)]
+    RegexError(#[from] regex::Error),
+
     // ==============================================
     // Wrappers for other error types in this library
     // ==============================================
@@ -357,14 +361,6 @@ impl From<VerifyFailed> for ClientFailed {
     }
 }
 
-// TEMP
-#[cfg(feature = "net3")]
-impl<T: std::fmt::Display> From<crate::net3::transport::TransportError<T>> for Error {
-    fn from(err: crate::net3::transport::TransportError<T>) -> Self {
-        Self::TransportError(err.to_string())
-    }
-}
-
 #[cfg(feature = "async-std")]
 impl From<async_std::future::TimeoutError> for Error {
     fn from(_err: async_std::future::TimeoutError) -> Self {

+ 5 - 5
src/net3/acceptor.rs

@@ -11,11 +11,11 @@ use crate::{
     Error, Result,
 };
 
-use super::{Channel, ChannelPtr, TcpTransport, TlsTransport, Transport};
+use super::{Channel, ChannelPtr, TcpTransport, Transport};
 
 /// A helper function to convert peer addr to Url and add scheme
 fn peer_addr_to_url(addr: SocketAddr, scheme: &str) -> Result<Url> {
-    let url = Url::parse(&format!("{}://{}", scheme, addr.to_string()))?;
+    let url = Url::parse(&format!("{}://{}", scheme, addr))?;
     Ok(url)
 }
 
@@ -105,8 +105,8 @@ impl Acceptor {
                     }
                 }
             }
-            "tls" => {
-                let transport = TlsTransport::new(None, 1024);
+            "tcp+tls" => {
+                let transport = TcpTransport::new(None, 1024);
 
                 let listener = transport.listen_on(accept_url);
 
@@ -122,7 +122,7 @@ impl Acceptor {
                     return Err(Error::OperationFailed)
                 }
 
-                let (acceptor, listener) = listener?;
+                let (acceptor, listener) = transport.upgrade_listener(listener?)?.await?;
 
                 let mut incoming = listener.incoming();
                 while let Some(stream) = incoming.next().await {

+ 5 - 3
src/net3/connector.rs

@@ -6,7 +6,7 @@ use url::Url;
 
 use crate::{Error, Result};
 
-use super::{Channel, ChannelPtr, SettingsPtr, TcpTransport, TlsTransport, Transport};
+use super::{Channel, ChannelPtr, SettingsPtr, TcpTransport, Transport};
 
 /// Create outbound socket connections.
 pub struct Connector {
@@ -42,8 +42,8 @@ impl Connector {
 
                         Ok(Channel::new(Box::new(stream?), connect_url).await)
                     }
-                    "tls" => {
-                        let transport = TlsTransport::new(None, 1024);
+                    "tcp+tls" => {
+                        let transport = TcpTransport::new(None, 1024);
                         let stream = transport.dial(connect_url.clone());
 
                         if let Err(err) = stream {
@@ -58,6 +58,8 @@ impl Connector {
                             return Err(Error::ConnectFailed)
                         }
 
+                        let stream = transport.upgrade_dialer(stream?)?.await;
+
                         Ok(Channel::new(Box::new(stream?), connect_url).await)
                     }
                     "tor" => todo!(),

+ 1 - 1
src/net3/mod.rs

@@ -98,4 +98,4 @@ pub use p2p::{P2p, P2pPtr};
 pub use protocol::{ProtocolBase, ProtocolBasePtr, ProtocolJobsManager, ProtocolJobsManagerPtr};
 pub use session::{SESSION_ALL, SESSION_INBOUND, SESSION_MANUAL, SESSION_OUTBOUND, SESSION_SEED};
 pub use settings::{Settings, SettingsPtr};
-pub use transport::{TcpTransport, TlsTransport, TorTransport, Transport};
+pub use transport::{TcpTransport, TorTransport, Transport};

+ 23 - 21
src/net3/transport.rs

@@ -1,44 +1,46 @@
-use std::error::Error;
-
 use async_trait::async_trait;
 use futures::prelude::*;
+use futures_rustls::{TlsAcceptor, TlsStream};
 use url::Url;
 
-mod tcp;
-mod tls;
-mod tor;
+use crate::Result;
+
+mod upgrade_tls;
+pub use upgrade_tls::TlsUpgrade;
 
+mod tcp;
 pub use tcp::TcpTransport;
-pub use tls::TlsTransport;
+
+mod tor;
 pub use tor::TorTransport;
 
+/// The `Transport` trait serves as a base for implementing transport protocols.
+/// Base transports can optionally be upgraded with TLS in order to support encryption.
+/// The implementation of our TLS authentication can be found in the [`upgrade_tls`] module.
 #[async_trait]
 pub trait Transport {
     type Acceptor;
     type Connector;
 
-    type Error: Error;
+    type Listener: Future<Output = Result<Self::Acceptor>>;
+    type Dial: Future<Output = Result<Self::Connector>>;
 
-    type Listener: Future<Output = Result<Self::Acceptor, Self::Error>>;
-    type Dial: Future<Output = Result<Self::Connector, Self::Error>>;
+    type TlsListener: Future<Output = Result<(TlsAcceptor, Self::Acceptor)>>;
+    type TlsDialer: Future<Output = Result<TlsStream<Self::Connector>>>;
 
-    fn listen_on(self, url: Url) -> Result<Self::Listener, TransportError<Self::Error>>
+    fn listen_on(self, url: Url) -> Result<Self::Listener>
     where
         Self: Sized;
 
-    fn dial(self, url: Url) -> Result<Self::Dial, TransportError<Self::Error>>
+    fn upgrade_listener(self, acceptor: Self::Acceptor) -> Result<Self::TlsListener>
     where
         Self: Sized;
-}
-
-#[derive(Debug, thiserror::Error)]
-pub enum TransportError<TErr> {
-    #[error("Address not supported: {0}")]
-    AddrNotSupported(Url),
 
-    #[error("Transport IO Error: {0}")]
-    IoError(#[from] std::io::Error),
+    fn dial(self, url: Url) -> Result<Self::Dial>
+    where
+        Self: Sized;
 
-    #[error("{0}")]
-    Other(TErr),
+    fn upgrade_dialer(self, stream: Self::Connector) -> Result<Self::TlsDialer>
+    where
+        Self: Sized;
 }

+ 31 - 16
src/net3/transport/tcp.rs

@@ -2,13 +2,15 @@ use async_std::net::{TcpListener, TcpStream};
 use std::{io, net::SocketAddr, pin::Pin};
 
 use futures::prelude::*;
+use futures_rustls::{TlsAcceptor, TlsStream};
 use log::debug;
 use socket2::{Domain, Socket, Type};
 use url::Url;
 
-use super::{Transport, TransportError};
+use super::{TlsUpgrade, Transport};
+use crate::{Error, Result};
 
-#[derive(Clone)]
+#[derive(Copy, Clone)]
 pub struct TcpTransport {
     /// TTL to set for opened sockets, or `None` for default
     ttl: Option<u32>,
@@ -20,30 +22,43 @@ impl Transport for TcpTransport {
     type Acceptor = TcpListener;
     type Connector = TcpStream;
 
-    type Error = io::Error;
+    type Listener = Pin<Box<dyn Future<Output = Result<Self::Acceptor>> + Send>>;
+    type Dial = Pin<Box<dyn Future<Output = Result<Self::Connector>> + Send>>;
 
-    type Listener = Pin<Box<dyn Future<Output = Result<Self::Acceptor, Self::Error>> + Send>>;
-    type Dial = Pin<Box<dyn Future<Output = Result<Self::Connector, Self::Error>> + Send>>;
+    type TlsListener = Pin<Box<dyn Future<Output = Result<(TlsAcceptor, Self::Acceptor)>> + Send>>;
+    type TlsDialer = Pin<Box<dyn Future<Output = Result<TlsStream<Self::Connector>>> + Send>>;
 
-    fn listen_on(self, url: Url) -> Result<Self::Listener, TransportError<Self::Error>> {
-        if url.scheme() != "tcp" {
-            return Err(TransportError::AddrNotSupported(url))
+    fn listen_on(self, url: Url) -> Result<Self::Listener> {
+        match url.scheme() {
+            "tcp" | "tcp+tls" => {}
+            x => return Err(Error::UnsupportedTransport(x.to_string())),
         }
 
         let socket_addr = url.socket_addrs(|| None)?[0];
-        debug!(target: "tcptransport", "listening on {}", socket_addr);
+        debug!("{} transport: listening on {}", url.scheme(), socket_addr);
         Ok(Box::pin(self.do_listen(socket_addr)))
     }
 
-    fn dial(self, url: Url) -> Result<Self::Dial, TransportError<Self::Error>> {
-        if url.scheme() != "tcp" {
-            return Err(TransportError::AddrNotSupported(url))
+    fn upgrade_listener(self, acceptor: Self::Acceptor) -> Result<Self::TlsListener> {
+        let tlsupgrade = TlsUpgrade::new();
+        Ok(Box::pin(tlsupgrade.upgrade_listener_tls(acceptor)))
+    }
+
+    fn dial(self, url: Url) -> Result<Self::Dial> {
+        match url.scheme() {
+            "tcp" | "tcp+tls" => {}
+            x => return Err(Error::UnsupportedTransport(x.to_string())),
         }
 
         let socket_addr = url.socket_addrs(|| None)?[0];
-        debug!(target: "tcptransport", "dialing {}", socket_addr);
+        debug!("{} transport: dialing {}", url.scheme(), socket_addr);
         Ok(Box::pin(self.do_dial(socket_addr)))
     }
+
+    fn upgrade_dialer(self, connector: Self::Connector) -> Result<Self::TlsDialer> {
+        let tlsupgrade = TlsUpgrade::new();
+        Ok(Box::pin(tlsupgrade.upgrade_dialer_tls(connector)))
+    }
 }
 
 impl TcpTransport {
@@ -66,7 +81,7 @@ impl TcpTransport {
         Ok(socket)
     }
 
-    async fn do_listen(self, socket_addr: SocketAddr) -> Result<TcpListener, io::Error> {
+    async fn do_listen(self, socket_addr: SocketAddr) -> Result<TcpListener> {
         let socket = self.create_socket(socket_addr)?;
         socket.bind(&socket_addr.into())?;
         socket.listen(self.backlog)?;
@@ -74,7 +89,7 @@ impl TcpTransport {
         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> {
         let socket = self.create_socket(socket_addr)?;
         socket.set_nonblocking(true)?;
 
@@ -82,7 +97,7 @@ impl TcpTransport {
             Ok(()) => {}
             Err(err) if err.raw_os_error() == Some(libc::EINPROGRESS) => {}
             Err(err) if err.kind() == io::ErrorKind::WouldBlock => {}
-            Err(err) => return Err(err),
+            Err(err) => return Err(err.into()),
         };
 
         let stream = TcpStream::from(std::net::TcpStream::from(socket));

+ 49 - 53
src/net3/transport/tor.rs

@@ -1,16 +1,3 @@
-use super::{Transport, TransportError};
-use async_std::{
-    net::{TcpListener, TcpStream},
-    sync::Arc,
-};
-use fast_socks5::{
-    client::{Config, Socks5Stream},
-    Result, SocksError,
-};
-use futures::prelude::*;
-
-use regex::Regex;
-use socket2::{Domain, Socket, Type};
 use std::{
     io,
     io::{BufRead, BufReader, Write},
@@ -18,8 +5,21 @@ use std::{
     pin::Pin,
     time::Duration,
 };
+
+use async_std::{
+    net::{TcpListener, TcpStream},
+    sync::Arc,
+};
+use fast_socks5::client::{Config, Socks5Stream};
+use futures::prelude::*;
+use futures_rustls::{TlsAcceptor, TlsStream};
+use regex::Regex;
+use socket2::{Domain, Socket, Type};
 use url::Url;
 
+use super::{TlsUpgrade, Transport};
+use crate::{Error, Result};
+
 /// Implements communication through the tor proxy service.
 ///
 /// ## Dialing
@@ -63,21 +63,6 @@ struct TorController {
     auth: String,
 }
 
-/// Wraps the errors, because dialing and listening use different communication
-#[derive(Debug, thiserror::Error)]
-pub enum TorError {
-    #[error("Transport IO Error: {0}")]
-    IoError(#[from] io::Error),
-    #[error("Socks: {0}")]
-    Socks5Error(#[from] SocksError),
-    #[error("Url parse error: {0}")]
-    UrlParseError(#[from] url::ParseError),
-    #[error("Regex parse error: {0}")]
-    RegexError(#[from] regex::Error),
-    #[error("Unexpected response from tor: {0}")]
-    GeneralError(String),
-}
-
 /// Contains the configuration to communicate with the Tor Controler
 ///
 /// When cloned, the socket is not reopened since we use reference count.
@@ -95,7 +80,7 @@ impl TorController {
     /// Cookie string: `assert_eq!(auth,"886b9177aec471965abd34b6a846dc32cf617dcff0625cba7a414e31dd4b75a0")`
     ///
     /// Password string: `assert_eq!(auth,"\"mypassword\"")`
-    pub fn new(url: Url, auth: String) -> Result<Self, io::Error> {
+    pub fn new(url: Url, auth: String) -> Result<Self> {
         let socket_addr = url.socket_addrs(|| None)?[0];
         let domain = if socket_addr.is_ipv4() { Domain::IPV4 } else { Domain::IPV6 };
         let socket = Socket::new(domain, Type::STREAM, Some(socket2::Protocol::TCP))?;
@@ -107,7 +92,7 @@ impl TorController {
             Ok(()) => {}
             Err(err) if err.raw_os_error() == Some(libc::EINPROGRESS) => {}
             Err(err) if err.kind() == io::ErrorKind::WouldBlock => {}
-            Err(err) => return Err(err),
+            Err(err) => return Err(err.into()),
         };
         Ok(Self { socket: Arc::new(socket), auth })
     }
@@ -116,17 +101,13 @@ impl TorController {
     /// # Arguments
     ///
     /// * `url` - url that the hidden service maps to.
-    pub fn create_ehs(&self, url: Url) -> Result<Url, TorError> {
+    pub fn create_ehs(&self, url: Url) -> Result<Url> {
         let local_socket = self.socket.try_clone()?;
         let mut stream = std::net::TcpStream::from(local_socket);
 
         stream.set_write_timeout(Some(Duration::from_secs(2)))?;
-        let host = url
-            .host()
-            .ok_or_else(|| TorError::GeneralError("No host on url for listening".to_string()))?;
-        let port = url
-            .port()
-            .ok_or_else(|| TorError::GeneralError("No port on url for listening".to_string()))?;
+        let host = url.host().unwrap();
+        let port = url.port().unwrap();
 
         let payload = format!(
             "AUTHENTICATE {a}\r\nADD_ONION NEW:BEST Flags=DiscardPK Port={p},{h}:{p}\r\n",
@@ -144,10 +125,9 @@ impl TorController {
             }
         }
         let re = Regex::new(r"250-ServiceID=(\w+*)")?;
-        let cap: Result<regex::Captures<'_>, TorError> =
-            re.captures(&repl).ok_or_else(|| TorError::GeneralError(repl.clone()));
-        let hurl =
-            cap?.get(1).map_or(Err(TorError::GeneralError(repl.clone())), |m| Ok(m.as_str()))?;
+        //let cap: Result<regex::Captures<'_>> =
+        let cap = re.captures(&repl).ok_or_else(|| Error::TorError(repl.clone()));
+        let hurl = cap?.get(1).map_or(Err(Error::TorError(repl.clone())), |m| Ok(m.as_str()))?;
         let hurl = format!("tcp://{}.onion:{}", &hurl, port);
         Ok(Url::parse(&hurl)?)
     }
@@ -164,7 +144,7 @@ impl TorTransport {
     /// services that live as long as the TorTransport.
     /// It is a tuple of the control socket url and authentication cookie as string
     /// represented in hex.
-    pub fn new(socks_url: Url, control_info: Option<(Url, String)>) -> Result<Self, TorError> {
+    pub fn new(socks_url: Url, control_info: Option<(Url, String)>) -> Result<Self> {
         match control_info {
             Some(info) => {
                 let (url, auth) = info;
@@ -181,16 +161,16 @@ impl TorTransport {
     /// # Arguments
     ///
     /// * `url` - url that the hidden service maps to.
-    pub fn create_ehs(&self, url: Url) -> Result<Url, TorError> {
+    pub fn create_ehs(&self, url: Url) -> Result<Url> {
         self.tor_controller
             .as_ref()
             .ok_or_else(|| {
-                TorError::GeneralError("No controller configured for this transport".to_string())
+                Error::TorError("No controller configured for this transport".to_string())
             })?
             .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>> {
         let socks_url_str = self.socks_url.socket_addrs(|| None)?[0].to_string();
         let host = url.host().unwrap().to_string();
         let port = url.port().unwrap_or(80);
@@ -222,7 +202,7 @@ impl TorTransport {
         Ok(socket)
     }
 
-    pub async fn do_listen(self, url: Url) -> Result<TcpListener, TorError> {
+    pub async fn do_listen(self, url: Url) -> Result<TcpListener> {
         let socket_addr = url.socket_addrs(|| None)?[0];
         let socket = self.create_socket(socket_addr)?;
         socket.bind(&socket_addr.into())?;
@@ -236,19 +216,35 @@ impl Transport for TorTransport {
     type Acceptor = TcpListener;
     type Connector = Socks5Stream<TcpStream>;
 
-    type Error = TorError;
+    type Listener = Pin<Box<dyn Future<Output = Result<Self::Acceptor>> + Send>>;
+    type Dial = Pin<Box<dyn Future<Output = Result<Self::Connector>> + Send>>;
 
-    type Listener = Pin<Box<dyn Future<Output = Result<Self::Acceptor, Self::Error>> + Send>>;
-    type Dial = Pin<Box<dyn Future<Output = Result<Self::Connector, Self::Error>> + Send>>;
+    type TlsListener = Pin<Box<dyn Future<Output = Result<(TlsAcceptor, Self::Acceptor)>> + Send>>;
+    type TlsDialer = Pin<Box<dyn Future<Output = Result<TlsStream<Self::Connector>>> + Send>>;
 
-    fn listen_on(self, url: Url) -> Result<Self::Listener, TransportError<Self::Error>> {
-        if url.scheme() != "tcp" {
-            return Err(TransportError::AddrNotSupported(url))
+    fn listen_on(self, url: Url) -> Result<Self::Listener> {
+        match url.scheme() {
+            "tor" | "tor+tls" => {}
+            x => return Err(Error::UnsupportedTransport(x.to_string())),
         }
         Ok(Box::pin(self.do_listen(url)))
     }
 
-    fn dial(self, url: Url) -> Result<Self::Dial, TransportError<Self::Error>> {
+    fn upgrade_listener(self, acceptor: Self::Acceptor) -> Result<Self::TlsListener> {
+        let tlsupgrade = TlsUpgrade::new();
+        Ok(Box::pin(tlsupgrade.upgrade_listener_tls(acceptor)))
+    }
+
+    fn dial(self, url: Url) -> Result<Self::Dial> {
+        match url.scheme() {
+            "tor" | "tor+tls" => {}
+            x => return Err(Error::UnsupportedTransport(x.to_string())),
+        }
         Ok(Box::pin(self.do_dial(url)))
     }
+
+    fn upgrade_dialer(self, connector: Self::Connector) -> Result<Self::TlsDialer> {
+        let tlsupgrade = TlsUpgrade::new();
+        Ok(Box::pin(tlsupgrade.upgrade_dialer_tls(connector)))
+    }
 }

+ 26 - 86
src/net3/transport/tls.rs → src/net3/transport/upgrade_tls.rs

@@ -1,6 +1,6 @@
-use std::{io, net::SocketAddr, pin::Pin, sync::Arc, time::SystemTime};
+use std::time::SystemTime;
 
-use async_std::net::{TcpListener, TcpStream};
+use async_std::{net::TcpListener, sync::Arc};
 use futures::prelude::*;
 use futures_rustls::{
     rustls,
@@ -13,12 +13,9 @@ use futures_rustls::{
     },
     TlsAcceptor, TlsConnector, TlsStream,
 };
-use log::debug;
 use rustls_pemfile::pkcs8_private_keys;
-use socket2::{Domain, Socket, Type};
-use url::Url;
 
-use super::{Transport, TransportError};
+use crate::Result;
 
 const CIPHER_SUITE: &str = "TLS13_CHACHA20_POLY1305_SHA256";
 
@@ -41,10 +38,10 @@ impl ServerCertVerifier for ServerCertificateVerifier {
         _end_entity: &Certificate,
         _intermediates: &[Certificate],
         _server_name: &ServerName,
-        _scts: &mut dyn Iterator<Item = &[u8]>,
+        _scrs: &mut dyn Iterator<Item = &[u8]>,
         _ocsp_response: &[u8],
         _now: SystemTime,
-    ) -> Result<ServerCertVerified, rustls::Error> {
+    ) -> std::result::Result<ServerCertVerified, rustls::Error> {
         // TODO: upsycle
         Ok(ServerCertVerified::assertion())
     }
@@ -61,55 +58,22 @@ impl ClientCertVerifier for ClientCertificateVerifier {
         _end_entity: &Certificate,
         _intermediates: &[Certificate],
         _now: SystemTime,
-    ) -> Result<ClientCertVerified, rustls::Error> {
+    ) -> std::result::Result<ClientCertVerified, rustls::Error> {
         // TODO: upsycle
         Ok(ClientCertVerified::assertion())
     }
 }
 
-#[derive(Clone)]
-pub struct TlsTransport {
-    /// TTL to set for opened sockets, or `None` for default
-    ttl: Option<u32>,
-    /// Size of the listen backlog for listen sockets
-    backlog: i32,
+pub struct TlsUpgrade {
     /// TLS server configuration
     server_config: Arc<ServerConfig>,
     /// TLS client configuration
     client_config: Arc<ClientConfig>,
 }
 
-impl Transport for TlsTransport {
-    type Acceptor = (TlsAcceptor, TcpListener);
-    type Connector = TlsStream<TcpStream>;
-
-    type Error = io::Error;
-
-    type Listener = Pin<Box<dyn Future<Output = Result<Self::Acceptor, Self::Error>> + Send>>;
-    type Dial = Pin<Box<dyn Future<Output = Result<Self::Connector, Self::Error>> + Send>>;
-
-    fn listen_on(self, url: Url) -> Result<Self::Listener, TransportError<Self::Error>> {
-        if url.scheme() != "tls" {
-            return Err(TransportError::AddrNotSupported(url))
-        }
-
-        debug!(target: "tlstransport", "listening on {}", url);
-        Ok(Box::pin(self.do_listen(url)))
-    }
-
-    fn dial(self, url: Url) -> Result<Self::Dial, TransportError<Self::Error>> {
-        if url.scheme() != "tls" {
-            return Err(TransportError::AddrNotSupported(url))
-        }
-
-        debug!(target: "tlstransport", "dialing {}", url);
-        Ok(Box::pin(self.do_dial(url)))
-    }
-}
-
-impl TlsTransport {
-    pub fn new(ttl: Option<u32>, backlog: i32) -> Self {
-        // On each instantiation, generate a new keypair and certificate
+impl TlsUpgrade {
+    pub fn new() -> Self {
+        // On each instantiation, generate a new keypair and certificate.
         let keypair_pem = ed25519_compact::KeyPair::generate().to_pem();
         let secret_key = pkcs8_private_keys(&mut keypair_pem.as_bytes()).unwrap();
         let secret_key = rustls::PrivateKey(secret_key[0].clone());
@@ -147,53 +111,29 @@ impl TlsTransport {
                 .unwrap(),
         );
 
-        Self { ttl, backlog, server_config, client_config }
-    }
-
-    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)?;
-        }
-
-        Ok(socket)
+        Self { server_config, client_config }
     }
 
-    async fn do_listen(self, url: Url) -> Result<(TlsAcceptor, TcpListener), io::Error> {
-        let socket_addr = url.socket_addrs(|| None)?[0];
-        let socket = self.create_socket(socket_addr)?;
-        socket.bind(&socket_addr.into())?;
-        socket.listen(self.backlog)?;
-        socket.set_nonblocking(true)?;
-
-        let listener = TcpListener::from(std::net::TcpListener::from(socket));
-        let acceptor = TlsAcceptor::from(self.server_config);
-        Ok((acceptor, listener))
+    pub async fn upgrade_listener_tls(
+        self,
+        listener: TcpListener,
+    ) -> Result<(TlsAcceptor, TcpListener)> {
+        Ok((TlsAcceptor::from(self.server_config), listener))
     }
 
-    async fn do_dial(self, url: Url) -> Result<TlsStream<TcpStream>, io::Error> {
-        let socket_addr = url.socket_addrs(|| None)?[0];
+    pub async fn upgrade_dialer_tls<IO>(self, stream: IO) -> Result<TlsStream<IO>>
+    where
+        IO: AsyncRead + AsyncWrite + Unpin,
+    {
         let server_name = ServerName::try_from("dark.fi").unwrap();
-        let socket = self.create_socket(socket_addr)?;
-        socket.set_nonblocking(true)?;
-
         let connector = TlsConnector::from(self.client_config);
-
-        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),
-        };
-
-        let stream = TcpStream::from(std::net::TcpStream::from(socket));
         let stream = connector.connect(server_name, stream).await?;
         Ok(TlsStream::Client(stream))
     }
 }
+
+impl Default for TlsUpgrade {
+    fn default() -> Self {
+        Self::new()
+    }
+}

+ 16 - 14
tests/network_transports.rs

@@ -1,15 +1,15 @@
+use std::{env::var, fs};
+
 use async_std::{
     io,
     io::{ReadExt, WriteExt},
     stream::StreamExt,
     task,
 };
-
-use darkfi::net::transport::{TcpTransport, TlsTransport, TorTransport, Transport};
-
-use std::{env::var, fs};
 use url::Url;
 
+use darkfi::net3::transport::{TcpTransport, TorTransport, Transport};
+
 #[async_std::test]
 async fn tcp_transport() {
     let tcp = TcpTransport::new(None, 1024);
@@ -37,11 +37,12 @@ async fn tcp_transport() {
 }
 
 #[async_std::test]
-async fn tls_transport() {
-    let tls = TlsTransport::new(None, 1024);
-    let url = Url::parse("tls://127.0.0.1:5433").unwrap();
+async fn tcp_tls_transport() {
+    let tcp = TcpTransport::new(None, 1024);
+    let url = Url::parse("tcp+tls://127.0.0.1:5433").unwrap();
 
-    let (acceptor, listener) = tls.clone().listen_on(url.clone()).unwrap().await.unwrap();
+    let listener = tcp.clone().listen_on(url.clone()).unwrap().await.unwrap();
+    let (acceptor, listener) = tcp.upgrade_listener(listener).unwrap().await.unwrap();
 
     let _ = task::spawn(async move {
         let mut incoming = listener.incoming();
@@ -55,7 +56,8 @@ async fn tls_transport() {
 
     let payload = b"ohai tls";
 
-    let mut client = tls.dial(url).unwrap().await.unwrap();
+    let client = tcp.dial(url).unwrap().await.unwrap();
+    let mut client = tcp.upgrade_dialer(client).unwrap().await.unwrap();
     client.write_all(payload).await.unwrap();
     let mut buf = vec![0_u8; 8];
     client.read_exact(&mut buf).await.unwrap();
@@ -68,13 +70,13 @@ async fn tls_transport() {
 async fn tor_transport_no_control() {
     let url = Url::parse("socks5://127.0.0.1:9050").unwrap();
     let hurl = var("DARKFI_TOR_LOCAL_ADDRESS")
-        .expect("Please set the env var DARKFI_TOR_LOCAL_ADDRESS to the configured local address in hidden service. \
-        For example: \'export DARKFI_TOR_LOCAL_ADDRESS=\"tcp://127.0.0.1:8080\"\'");
+.expect("Please set the env var DARKFI_TOR_LOCAL_ADDRESS to the configured local address in hidden service. \
+For example: \'export DARKFI_TOR_LOCAL_ADDRESS=\"tcp://127.0.0.1:8080\"\'");
     let hurl = Url::parse(&hurl).unwrap();
 
     let onion = var("DARKFI_TOR_PUBLIC_ADDRESS").expect(
         "Please set the env var DARKFI_TOR_PUBLIC_ADDRESS to the configured onion address. \
-        For example: \'export DARKFI_TOR_PUBLIC_ADDRESS=\"tor://abcdefghij234567.onion\"\'",
+For example: \'export DARKFI_TOR_PUBLIC_ADDRESS=\"tor://abcdefghij234567.onion\"\'",
     );
 
     let tor = TorTransport::new(url, None).unwrap();
@@ -103,7 +105,7 @@ async fn tor_transport_no_control() {
 async fn tor_transport_with_control() {
     let auth_cookie = var("DARKFI_TOR_COOKIE").expect(
         "Please set the env var DARKFI_TOR_COOKIE to the configured tor cookie file. \
-        For example: \'export DARKFI_TOR_COOKIE=\"/var/lib/tor/control_auth_cookie\"\'",
+For example: \'export DARKFI_TOR_COOKIE=\"/var/lib/tor/control_auth_cookie\"\'",
     );
     let auth_cookie = hex::encode(&fs::read(auth_cookie).unwrap());
     let socks_url = Url::parse("socks5://127.0.0.1:9050").unwrap();
@@ -140,7 +142,7 @@ async fn tor_transport_with_control() {
 async fn tor_transport_with_control_dropped() {
     let auth_cookie = var("DARKFI_TOR_COOKIE").expect(
         "Please set the env var DARKFI_TOR_COOKIE to the configured tor cookie file. \
-        For example: \'export DARKFI_TOR_COOKIE=\"/var/lib/tor/control_auth_cookie\"\'",
+For example: \'export DARKFI_TOR_COOKIE=\"/var/lib/tor/control_auth_cookie\"\'",
     );
     let auth_cookie = hex::encode(&fs::read(auth_cookie).unwrap());
     let socks_url = Url::parse("socks5://127.0.0.1:9050").unwrap();