Selaa lähdekoodia

net: Have Unix sockets work in the same way as the Tcp transport.

parazyd 3 vuotta sitten
vanhempi
sitoutus
6b563e63aa
5 muutettua tiedostoa jossa 98 lisäystä ja 35 poistoa
  1. 1 1
      src/net/transport/tcp.rs
  2. 65 20
      src/net/transport/unix.rs
  3. 2 7
      src/rpc/client.rs
  4. 2 6
      src/rpc/server.rs
  5. 28 1
      tests/network_transports.rs

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

@@ -63,7 +63,7 @@ impl TransportListener for (TlsAcceptor, TcpListener) {
         let url = socket_addr_to_url(peer_addr, "tcp+tls")?;
 
         if let Err(err) = stream {
-            error!("Error wraping the connection {} with tls: {}", url, err);
+            error!("Error wrapping the connection {} with tls: {}", url, err);
             return Err(Error::AcceptTlsConnectionFailed(self.1.local_addr()?.to_string()))
         }
 

+ 65 - 20
src/net/transport/unix.rs

@@ -16,19 +16,24 @@
  * along with this program.  If not, see <https://www.gnu.org/licenses/>.
  */
 
-use async_std::os::unix::net::{UnixListener, UnixStream};
+use std::{os::unix::net::SocketAddr, pin::Pin, time::Duration};
 
+use async_std::os::unix::net::{UnixListener, UnixStream};
 use async_trait::async_trait;
+use futures::prelude::*;
+use futures_rustls::{TlsAcceptor, TlsStream};
 use log::{debug, error};
 use url::Url;
 
-use super::{TransportListener, TransportStream};
+use super::{Transport, TransportListener, TransportStream};
 use crate::{Error, Result};
 
 fn unix_socket_addr_to_string(addr: std::os::unix::net::SocketAddr) -> String {
     addr.as_pathname().unwrap_or(&std::path::PathBuf::from("")).to_str().unwrap_or("").into()
 }
 
+impl TransportStream for UnixStream {}
+
 #[async_trait]
 impl TransportListener for UnixListener {
     async fn next(&self) -> Result<(Box<dyn TransportStream>, Url)> {
@@ -46,42 +51,82 @@ impl TransportListener for UnixListener {
     }
 }
 
-impl TransportStream for UnixStream {}
+#[async_trait]
+impl TransportListener for (TlsAcceptor, UnixListener) {
+    async fn next(&self) -> Result<(Box<dyn TransportStream>, Url)> {
+        unimplemented!("TLS not supported for Unix sockets");
+    }
+}
 
-#[derive(Default, Copy, Clone)]
+#[derive(Copy, Clone)]
 pub struct UnixTransport {}
 
-impl UnixTransport {
-    pub fn new() -> Self {
-        Self {}
-    }
-    pub async fn listen(self, url: Url) -> Result<UnixListener> {
+impl Transport for UnixTransport {
+    type Acceptor = UnixListener;
+    type Connector = UnixStream;
+
+    type Listener = Pin<Box<dyn Future<Output = Result<Self::Acceptor>> + Send>>;
+    type Dial = Pin<Box<dyn Future<Output = Result<Self::Connector>> + 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> {
         match url.scheme() {
             "unix" => {}
             x => return Err(Error::UnsupportedTransport(x.to_string())),
         }
 
-        if !cfg!(unix) {
-            return Err(Error::UnsupportedOS)
-        }
+        let socket_path = url.path();
+        let socket_addr = SocketAddr::from_pathname(&socket_path)?;
+        debug!(target: "net", "{} transport: listening on {}", url.scheme(), socket_path);
+        Ok(Box::pin(self.do_listen(socket_addr)))
+    }
 
-        let listener = UnixListener::bind(url.as_str()).await?;
-        debug!("{} transport: listening on {}", url.scheme(), url);
-        Ok(listener)
+    fn upgrade_listener(self, _acceptor: Self::Acceptor) -> Result<Self::TlsListener> {
+        unimplemented!("TLS not supported for Unix sockets");
     }
 
-    pub async fn dial(self, url: Url) -> Result<UnixStream> {
+    fn dial(self, url: Url, timeout: Option<Duration>) -> Result<Self::Dial> {
         match url.scheme() {
             "unix" => {}
             x => return Err(Error::UnsupportedTransport(x.to_string())),
         }
 
-        if !cfg!(unix) {
-            return Err(Error::UnsupportedOS)
+        let socket_path = url.path();
+        let socket_addr = SocketAddr::from_pathname(&socket_path)?;
+        debug!(target: "net", "{} transport: listening on {}", url.scheme(), socket_path);
+        Ok(Box::pin(self.do_dial(socket_addr, timeout)))
+    }
+
+    fn upgrade_dialer(self, _connector: Self::Connector) -> Result<Self::TlsDialer> {
+        unimplemented!("TLS not supported for Unix sockets");
+    }
+}
+
+impl UnixTransport {
+    pub fn new() -> Self {
+        Self {}
+    }
+
+    async fn do_listen(self, socket_addr: SocketAddr) -> Result<UnixListener> {
+        // We're a bit rough here and delete the socket.
+        let socket_path = socket_addr.as_pathname().unwrap();
+        if std::fs::metadata(socket_path).is_ok() {
+            std::fs::remove_file(socket_path)?;
         }
 
-        let stream = UnixStream::connect(url.as_str()).await?;
-        debug!("{} transport: dialing to {}", url.scheme(), url);
+        let socket = UnixListener::bind(socket_path).await?;
+        Ok(socket)
+    }
+
+    async fn do_dial(
+        self,
+        socket_addr: SocketAddr,
+        _timeout: Option<Duration>,
+    ) -> Result<UnixStream> {
+        let socket_path = socket_addr.as_pathname().unwrap();
+        let stream = UnixStream::connect(&socket_path).await?;
         Ok(stream)
     }
 }

+ 2 - 7
src/rpc/client.rs

@@ -216,13 +216,8 @@ impl RpcClient {
             }
             TransportName::Unix => {
                 let transport = UnixTransport::new();
-                let stream = transport.dial(uri.clone()).await;
-                if let Err(err) = stream {
-                    error!("JSON-RPC client connection to {} failed: {}", uri, err);
-                    return Err(Error::ConnectFailed)
-                }
-
-                smol::spawn(Self::reqrep_loop(stream?, result_send, data_recv, stop_recv)).detach();
+                let stream = transport.dial(uri.clone(), None);
+                reqrep!(stream, transport, None);
             }
             _ => unimplemented!(),
         }

+ 2 - 6
src/rpc/server.rs

@@ -192,12 +192,8 @@ pub async fn listen_and_serve(
         }
         TransportName::Unix => {
             let transport = UnixTransport::new();
-            let listener = transport.listen(accept_url.clone()).await;
-            if let Err(err) = listener {
-                error!("JSON-RPC Unix socket bind to {} failed: {}", accept_url, err);
-                return Err(Error::BindFailed(accept_url.as_str().into()))
-            }
-            run_accept_loop(Box::new(listener?), rh, ex.clone()).await?;
+            let listener = transport.listen_on(accept_url.clone());
+            accept!(listener, transport, None);
         }
         _ => unimplemented!(),
     }

+ 28 - 1
tests/network_transports.rs

@@ -26,7 +26,34 @@ use async_std::{
 };
 use url::Url;
 
-use darkfi::net::transport::{TcpTransport, TorTransport, Transport};
+use darkfi::net::transport::{TcpTransport, TorTransport, Transport, UnixTransport};
+
+#[async_std::test]
+async fn unix_transport() {
+    let unix = UnixTransport::new();
+    let url = Url::parse("unix:///tmp/darkfi_test.sock").unwrap();
+
+    let listener = unix.listen_on(url.clone()).unwrap().await.unwrap();
+
+    let _ = task::spawn(async move {
+        let mut incoming = listener.incoming();
+        while let Some(stream) = incoming.next().await {
+            let stream = stream.unwrap();
+            let (reader, writer) = &mut (&stream, &stream);
+            io::copy(reader, writer).await.unwrap();
+        }
+    });
+
+    let payload = b"ohai unix";
+
+    let mut client = unix.dial(url, None).unwrap().await.unwrap();
+    client.write_all(payload).await.unwrap();
+    let mut buf = vec![0_u8; 9];
+    client.read_exact(&mut buf).await.unwrap();
+
+    std::fs::remove_file("/tmp/darkfi_test.sock").unwrap();
+    assert_eq!(buf, payload);
+}
 
 #[async_std::test]
 async fn tcp_transport() {