소스 검색

net: Support TLS dialer/listener without requiring client certificates

x 2 달 전
부모
커밋
346e987f10
7개의 변경된 파일68개의 추가작업 그리고 45개의 파일을 삭제
  1. 1 1
      src/net/acceptor.rs
  2. 5 4
      src/net/connector.rs
  3. 31 22
      src/net/transport/mod.rs
  4. 16 7
      src/net/transport/tls.rs
  5. 2 2
      src/rpc/client.rs
  6. 1 1
      src/rpc/server.rs
  7. 12 8
      tests/network_transports.rs

+ 1 - 1
src/net/acceptor.rs

@@ -80,7 +80,7 @@ impl Acceptor {
             self.session.upgrade().unwrap().p2p().settings().read().await.p2p_datastore.clone();
 
         // Initialize listener
-        let listener = Listener::new(endpoint.clone(), datastore).await?;
+        let listener = Listener::new(endpoint.clone(), datastore, true).await?;
 
         // Open socket
         let ptlistener = listener.listen().await?;

+ 5 - 4
src/net/connector.rs

@@ -82,10 +82,11 @@ impl Connector {
         let outbound_connect_timeout = settings.outbound_connect_timeout(endpoint.scheme());
         drop(settings);
 
-        let dialer = match Dialer::new(endpoint.clone(), datastore, Some(i2p_socks5_proxy)).await {
-            Ok(dialer) => dialer,
-            Err(err) => return Err(Error::ConnectFailed(format!("[{endpoint}]: {err}"))),
-        };
+        let dialer =
+            match Dialer::new(endpoint.clone(), datastore, Some(i2p_socks5_proxy), true).await {
+                Ok(dialer) => dialer,
+                Err(err) => return Err(Error::ConnectFailed(format!("[{endpoint}]: {err}"))),
+            };
         let timeout = Duration::from_secs(outbound_connect_timeout);
 
         let stop_fut = async {

+ 31 - 22
src/net/transport/mod.rs

@@ -122,6 +122,8 @@ pub struct Dialer {
     endpoint: Url,
     /// The dialer variant (transport protocol)
     variant: DialerVariant,
+    /// Marker if TLS client certificate should be provided for this instance
+    provide_tls_client_cert: bool,
 }
 
 macro_rules! enforce_hostport {
@@ -151,6 +153,7 @@ impl Dialer {
         endpoint: Url,
         datastore: Option<String>,
         i2p_socks5_proxy: Option<Url>,
+        provide_tls_client_cert: bool,
     ) -> io::Result<Self> {
         match endpoint.scheme().to_lowercase().as_str() {
             "tcp" => {
@@ -158,7 +161,7 @@ impl Dialer {
                 enforce_hostport!(endpoint);
                 let variant = tcp::TcpDialer::new(None).await?;
                 let variant = DialerVariant::Tcp(variant);
-                Ok(Self { endpoint, variant })
+                Ok(Self { endpoint, variant, provide_tls_client_cert })
             }
 
             "tcp+tls" => {
@@ -166,7 +169,7 @@ impl Dialer {
                 enforce_hostport!(endpoint);
                 let variant = tcp::TcpDialer::new(None).await?;
                 let variant = DialerVariant::TcpTls(variant);
-                Ok(Self { endpoint, variant })
+                Ok(Self { endpoint, variant, provide_tls_client_cert })
             }
 
             #[cfg(feature = "p2p-tor")]
@@ -175,7 +178,7 @@ impl Dialer {
                 enforce_hostport!(endpoint);
                 let variant = tor::TorDialer::new(datastore).await?;
                 let variant = DialerVariant::Tor(variant);
-                Ok(Self { endpoint, variant })
+                Ok(Self { endpoint, variant, provide_tls_client_cert })
             }
 
             #[cfg(feature = "p2p-tor")]
@@ -184,7 +187,7 @@ impl Dialer {
                 enforce_hostport!(endpoint);
                 let variant = tor::TorDialer::new(datastore).await?;
                 let variant = DialerVariant::TorTls(variant);
-                Ok(Self { endpoint, variant })
+                Ok(Self { endpoint, variant, provide_tls_client_cert })
             }
 
             #[cfg(feature = "p2p-nym")]
@@ -193,7 +196,7 @@ impl Dialer {
                 enforce_hostport!(endpoint);
                 let variant = nym::NymDialer::new().await?;
                 let variant = DialerVariant::Nym(variant);
-                Ok(Self { endpoint, variant })
+                Ok(Self { endpoint, variant, provide_tls_client_cert })
             }
 
             #[cfg(feature = "p2p-nym")]
@@ -202,7 +205,7 @@ impl Dialer {
                 enforce_hostport!(endpoint);
                 let variant = nym::NymDialer::new().await?;
                 let variant = DialerVariant::NymTls(variant);
-                Ok(Self { endpoint, variant })
+                Ok(Self { endpoint, variant, provide_tls_client_cert })
             }
 
             #[cfg(feature = "p2p-unix")]
@@ -211,7 +214,7 @@ impl Dialer {
                 enforce_abspath!(endpoint);
                 let variant = unix::UnixDialer::new().await?;
                 let variant = DialerVariant::Unix(variant);
-                Ok(Self { endpoint, variant })
+                Ok(Self { endpoint, variant, provide_tls_client_cert })
             }
 
             #[cfg(feature = "p2p-socks5")]
@@ -220,7 +223,7 @@ impl Dialer {
                 enforce_hostport!(endpoint);
                 let variant = socks5::Socks5Dialer::new(&endpoint).await?;
                 let variant = DialerVariant::Socks5(variant);
-                Ok(Self { endpoint, variant })
+                Ok(Self { endpoint, variant, provide_tls_client_cert })
             }
 
             #[cfg(feature = "p2p-socks5")]
@@ -229,7 +232,7 @@ impl Dialer {
                 enforce_hostport!(endpoint);
                 let variant = socks5::Socks5Dialer::new(&endpoint).await?;
                 let variant = DialerVariant::Socks5Tls(variant);
-                Ok(Self { endpoint, variant })
+                Ok(Self { endpoint, variant, provide_tls_client_cert })
             }
 
             #[cfg(feature = "p2p-i2p")]
@@ -240,7 +243,7 @@ impl Dialer {
                 url.set_path(&format!("{}:{}", endpoint.host().unwrap(), endpoint.port().unwrap()));
                 let variant = socks5::Socks5Dialer::new(&url).await?;
                 let variant = DialerVariant::Socks5(variant);
-                Ok(Self { endpoint, variant })
+                Ok(Self { endpoint, variant, provide_tls_client_cert })
             }
 
             #[cfg(feature = "p2p-i2p")]
@@ -252,7 +255,7 @@ impl Dialer {
                 url.set_scheme("socks5+tls").unwrap();
                 let variant = socks5::Socks5Dialer::new(&url).await?;
                 let variant = DialerVariant::Socks5Tls(variant);
-                Ok(Self { endpoint, variant })
+                Ok(Self { endpoint, variant, provide_tls_client_cert })
             }
 
             #[cfg(feature = "p2p-quic")]
@@ -261,7 +264,7 @@ impl Dialer {
                 enforce_hostport!(endpoint);
                 let variant = quic::QuicDialer::new().await?;
                 let variant = DialerVariant::Quic(variant);
-                Ok(Self { endpoint, variant })
+                Ok(Self { endpoint, variant, provide_tls_client_cert })
             }
 
             x => {
@@ -288,7 +291,7 @@ impl Dialer {
             DialerVariant::TcpTls(dialer) => {
                 let sockaddr = self.endpoint.socket_addrs(|| None)?;
                 let stream = dialer.do_dial(sockaddr[0], timeout).await?;
-                let tlsupgrade = tls::TlsUpgrade::new().await?;
+                let tlsupgrade = tls::TlsUpgrade::new(self.provide_tls_client_cert).await?;
                 let stream = tlsupgrade.upgrade_dialer_tls(stream).await?;
                 Ok(Box::new(stream))
             }
@@ -306,7 +309,7 @@ impl Dialer {
                 let host = self.endpoint.host_str().unwrap();
                 let port = self.endpoint.port().unwrap();
                 let stream = dialer.do_dial(host, port, timeout).await?;
-                let tlsupgrade = tls::TlsUpgrade::new().await?;
+                let tlsupgrade = tls::TlsUpgrade::new(self.provide_tls_client_cert).await?;
                 let stream = tlsupgrade.upgrade_dialer_tls(stream).await?;
                 Ok(Box::new(stream))
             }
@@ -340,7 +343,7 @@ impl Dialer {
             #[cfg(feature = "p2p-socks5")]
             DialerVariant::Socks5Tls(dialer) => {
                 let stream = dialer.do_dial().await?;
-                let tlsupgrade = tls::TlsUpgrade::new().await?;
+                let tlsupgrade = tls::TlsUpgrade::new(self.provide_tls_client_cert).await?;
                 let stream = tlsupgrade.upgrade_dialer_tls(stream).await?;
                 Ok(Box::new(stream))
             }
@@ -366,19 +369,25 @@ pub struct Listener {
     endpoint: Url,
     /// The listener variant (transport protocol)
     variant: ListenerVariant,
+    /// Marker if TLS client cert should be required for this instance
+    require_tls_client_cert: bool,
 }
 
 impl Listener {
     /// Instantiate a new [`Listener`] with the given [`Url`] and datastore path.
     /// Must contain a scheme, host string, and a port.
-    pub async fn new(endpoint: Url, datastore: Option<String>) -> io::Result<Self> {
+    pub async fn new(
+        endpoint: Url,
+        datastore: Option<String>,
+        require_tls_client_cert: bool,
+    ) -> io::Result<Self> {
         match endpoint.scheme().to_lowercase().as_str() {
             "tcp" => {
                 // Build a TCP listener
                 enforce_hostport!(endpoint);
                 let variant = tcp::TcpListener::new(1024).await?;
                 let variant = ListenerVariant::Tcp(variant);
-                Ok(Self { endpoint, variant })
+                Ok(Self { endpoint, variant, require_tls_client_cert })
             }
 
             "tcp+tls" => {
@@ -386,7 +395,7 @@ impl Listener {
                 enforce_hostport!(endpoint);
                 let variant = tcp::TcpListener::new(1024).await?;
                 let variant = ListenerVariant::TcpTls(variant);
-                Ok(Self { endpoint, variant })
+                Ok(Self { endpoint, variant, require_tls_client_cert })
             }
 
             #[cfg(feature = "p2p-tor")]
@@ -395,7 +404,7 @@ impl Listener {
                 enforce_hostport!(endpoint);
                 let variant = tor::TorListener::new(datastore).await?;
                 let variant = ListenerVariant::Tor(variant);
-                Ok(Self { endpoint, variant })
+                Ok(Self { endpoint, variant, require_tls_client_cert })
             }
 
             #[cfg(feature = "p2p-unix")]
@@ -403,7 +412,7 @@ impl Listener {
                 enforce_abspath!(endpoint);
                 let variant = unix::UnixListener::new().await?;
                 let variant = ListenerVariant::Unix(variant);
-                Ok(Self { endpoint, variant })
+                Ok(Self { endpoint, variant, require_tls_client_cert })
             }
 
             #[cfg(feature = "p2p-quic")]
@@ -411,7 +420,7 @@ impl Listener {
                 enforce_hostport!(endpoint);
                 let variant = quic::QuicListener::new().await?;
                 let variant = ListenerVariant::Quic(variant);
-                Ok(Self { endpoint, variant })
+                Ok(Self { endpoint, variant, require_tls_client_cert })
             }
 
             x => {
@@ -434,7 +443,7 @@ impl Listener {
             ListenerVariant::TcpTls(listener) => {
                 let sockaddr = self.endpoint.socket_addrs(|| None)?;
                 let l = listener.do_listen(sockaddr[0]).await?;
-                let tlsupgrade = tls::TlsUpgrade::new().await?;
+                let tlsupgrade = tls::TlsUpgrade::new(self.require_tls_client_cert).await?;
                 let l = tlsupgrade.upgrade_listener_tcp_tls(l).await?;
                 Ok(Box::new(l))
             }

+ 16 - 7
src/net/transport/tls.rs

@@ -257,18 +257,27 @@ pub struct TlsUpgrade {
 }
 
 impl TlsUpgrade {
-    pub async fn new() -> io::Result<Self> {
+    pub async fn new(enable_tls_client_cert: bool) -> io::Result<Self> {
         // On each instantiation, generate a new keypair and certificate
         let (certificate, secret_key_der) = generate_certificate()?;
 
         // Server-side config
         let client_cert_verifier = Arc::new(ClientCertificateVerifier {});
-        let server_config = Arc::new(
-            ServerConfig::builder_with_protocol_versions(&[&TLS13])
-                .with_client_cert_verifier(client_cert_verifier)
-                .with_single_cert(vec![certificate.clone()], secret_key_der.clone_key())
-                .unwrap(),
-        );
+        let server_config = if enable_tls_client_cert {
+            Arc::new(
+                ServerConfig::builder_with_protocol_versions(&[&TLS13])
+                    .with_client_cert_verifier(client_cert_verifier)
+                    .with_single_cert(vec![certificate.clone()], secret_key_der.clone_key())
+                    .unwrap(),
+            )
+        } else {
+            Arc::new(
+                ServerConfig::builder_with_protocol_versions(&[&TLS13])
+                    .with_no_client_auth()
+                    .with_single_cert(vec![certificate.clone()], secret_key_der.clone_key())
+                    .unwrap(),
+            )
+        };
 
         // Client-side config
         let server_cert_verifier = Arc::new(ServerCertificateVerifier {});

+ 2 - 2
src/rpc/client.rs

@@ -72,7 +72,7 @@ impl RpcClient {
 
         // Instantiate Dialer and dial the server
         // TODO: Could add a timeout here
-        let dialer = Dialer::new(dialer_url, None, None).await?;
+        let dialer = Dialer::new(dialer_url, None, None, false).await?;
         let stream = dialer.dial(None).await?;
 
         // Create the StoppableTask running the request-reply loop.
@@ -325,7 +325,7 @@ impl RpcChadClient {
 
         // Instantiate Dialer and dial the server
         // TODO: Could add a timeout here
-        let dialer = Dialer::new(dialer_url, None, None).await?;
+        let dialer = Dialer::new(dialer_url, None, None, false).await?;
         let stream = dialer.dial(None).await?;
 
         // Create the StoppableTask running the request-reply loop.

+ 1 - 1
src/rpc/server.rs

@@ -511,7 +511,7 @@ pub async fn listen_and_serve<'a, T: 'a>(
         listen_url = url_str.parse()?;
     }
 
-    let listener = Listener::new(listen_url, None).await?.listen().await?;
+    let listener = Listener::new(listen_url, None, false).await?.listen().await?;
 
     run_accept_loop(listener, rh, conn_limit, settings, ex.clone()).await
 }

+ 12 - 8
tests/network_transports.rs

@@ -32,7 +32,8 @@ fn tcp_transport() {
         drop(listener);
         let url = Url::parse(&format!("tcp://127.0.0.1:{port}")).unwrap();
 
-        let listener = Listener::new(url.clone(), None).await.unwrap().listen().await.unwrap();
+        let listener =
+            Listener::new(url.clone(), None, true).await.unwrap().listen().await.unwrap();
         executor
             .spawn(async move {
                 let (stream, _) = listener.next().await.unwrap();
@@ -43,7 +44,7 @@ fn tcp_transport() {
 
         let payload = "ohai tcp";
 
-        let dialer = Dialer::new(url, None, None).await.unwrap();
+        let dialer = Dialer::new(url, None, None, true).await.unwrap();
         let mut client = dialer.dial(None).await.unwrap();
         payload.encode_async(&mut client).await.unwrap();
 
@@ -67,7 +68,8 @@ fn tcp_tls_transport() {
         drop(listener);
         let url = Url::parse(&format!("tcp://127.0.0.1:{port}")).unwrap();
 
-        let listener = Listener::new(url.clone(), None).await.unwrap().listen().await.unwrap();
+        let listener =
+            Listener::new(url.clone(), None, true).await.unwrap().listen().await.unwrap();
         executor
             .spawn(async move {
                 let (stream, _) = listener.next().await.unwrap();
@@ -78,7 +80,7 @@ fn tcp_tls_transport() {
 
         let payload = "ohai tls";
 
-        let dialer = Dialer::new(url, None, None).await.unwrap();
+        let dialer = Dialer::new(url, None, None, true).await.unwrap();
         let mut client = dialer.dial(None).await.unwrap();
         payload.encode_async(&mut client).await.unwrap();
 
@@ -98,7 +100,8 @@ fn quic_transport() {
         drop(listener);
         let url = Url::parse(&format!("quic://127.0.0.1:{port}")).unwrap();
 
-        let listener = Listener::new(url.clone(), None).await.unwrap().listen().await.unwrap();
+        let listener =
+            Listener::new(url.clone(), None, true).await.unwrap().listen().await.unwrap();
 
         executor
             .spawn(async move {
@@ -110,7 +113,7 @@ fn quic_transport() {
 
         let payload = "ohai quic";
 
-        let dialer = Dialer::new(url, None, None).await.unwrap();
+        let dialer = Dialer::new(url, None, None, true).await.unwrap();
         let mut client = dialer.dial(None).await.unwrap();
         payload.encode_async(&mut client).await.unwrap();
 
@@ -132,7 +135,8 @@ fn unix_transport() {
     .unwrap();
 
     smol::block_on(executor.run(async {
-        let listener = Listener::new(url.clone(), None).await.unwrap().listen().await.unwrap();
+        let listener =
+            Listener::new(url.clone(), None, true).await.unwrap().listen().await.unwrap();
         executor
             .spawn(async move {
                 let (stream, _) = listener.next().await.unwrap();
@@ -143,7 +147,7 @@ fn unix_transport() {
 
         let payload = "ohai unix";
 
-        let dialer = Dialer::new(url, None, None).await.unwrap();
+        let dialer = Dialer::new(url, None, None, true).await.unwrap();
         let mut client = dialer.dial(None).await.unwrap();
         payload.encode_async(&mut client).await.unwrap();