Преглед изворни кода

net/transport: Implement TLS client protocol.

parazyd пре 4 година
родитељ
комит
41ecea27a7
5 измењених фајлова са 157 додато и 7 уклоњено
  1. 12 0
      Cargo.lock
  2. 2 0
      Cargo.toml
  3. 19 6
      example/net2.rs
  4. 4 1
      src/net/transport.rs
  5. 120 0
      src/net/transport/tls.rs

+ 12 - 0
Cargo.lock

@@ -1566,6 +1566,7 @@ dependencies = [
  "drk-sdk",
  "fast-socks5",
  "futures",
+ "futures-rustls",
  "fxhash",
  "group",
  "halo2_gadgets",
@@ -2440,6 +2441,17 @@ dependencies = [
  "syn 1.0.89",
 ]
 
+[[package]]
+name = "futures-rustls"
+version = "0.22.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "e01fe9932a224b72b45336d96040aa86386d674a31d0af27d800ea7bc8ca97fe"
+dependencies = [
+ "futures-io",
+ "rustls 0.20.4",
+ "webpki 0.22.0",
+]
+
 [[package]]
 name = "futures-sink"
 version = "0.3.21"

+ 2 - 0
Cargo.toml

@@ -51,6 +51,7 @@ native-tls = {version = "0.2.8", optional = true}
 
 # Networking
 socket2 = {version = "0.4.4", optional = true}
+futures-rustls = {version = "0.22.1", features = ["dangerous_configuration"], optional = true}
 
 # Encoding
 hex = {version = "0.4.3", optional = true}
@@ -201,6 +202,7 @@ system = [
 net = [
     "fxhash",
     "socket2",
+    "futures-rustls",
 
     "util",
     "system",

+ 19 - 6
example/net2.rs

@@ -1,13 +1,26 @@
-use async_std::io::WriteExt;
-use darkfi::net::transport::{TcpTransport, Transport};
+use async_std::io::{ReadExt, WriteExt};
+use darkfi::net::transport::{TcpTransport, TlsTransport, Transport};
+use std::{fs::File, io::Write};
 use url::Url;
 
 #[async_std::main]
 async fn main() {
-    let tcp = TcpTransport { ttl: None };
-    let url = Url::parse("tcp://127.0.0.1:5432").unwrap();
+    // nc -l 127.0.0.1 5432
+    // let tcp = TcpTransport { ttl: None };
+    // let url = Url::parse("tcp://127.0.0.1:5432").unwrap();
 
-    let mut socket = tcp.dial(url).unwrap().await.unwrap();
-    socket.write_all(b"ohai\n").await.unwrap();
+    // let mut socket = tcp.dial(url).unwrap().await.unwrap();
+    // socket.write_all(b"ohai tcp\n").await.unwrap();
     // socket.flush().await;
+
+    let tls = TlsTransport { ttl: None };
+    let url = Url::parse("tls://parazyd.org:70").unwrap();
+    let mut socket = tls.dial(url).unwrap().await.unwrap();
+    socket.write_all(b"/rms.png\r\n").await.unwrap();
+
+    let mut buf = vec![];
+    socket.read_to_end(&mut buf).await.unwrap();
+
+    let mut file = File::create("rms.png").unwrap();
+    file.write_all(&buf).unwrap();
 }

+ 4 - 1
src/net/transport.rs

@@ -3,8 +3,11 @@ use std::error::Error;
 use futures::prelude::*;
 use url::Url;
 
-pub mod tcp;
+mod tcp;
+mod tls;
+
 pub use tcp::TcpTransport;
+pub use tls::TlsTransport;
 
 pub trait Transport {
     type Output;

+ 120 - 0
src/net/transport/tls.rs

@@ -0,0 +1,120 @@
+use std::{io, net::SocketAddr, pin::Pin, sync::Arc, time::SystemTime};
+
+use async_std::net::TcpStream;
+use futures::prelude::*;
+use futures_rustls::{
+    rustls,
+    rustls::{
+        client::{ServerCertVerified, ServerCertVerifier},
+        kx_group::X25519,
+        version::TLS13,
+        Certificate, ClientConfig, RootCertStore, ServerName,
+    },
+    TlsConnector, TlsStream,
+};
+use log::debug;
+use socket2::{Domain, Socket, Type};
+use url::Url;
+
+use super::{Transport, TransportError};
+
+const CIPHER_SUITE: &str = "TLS13_CHACHA20_POLY1305_SHA256";
+
+fn cipher_suite() -> rustls::SupportedCipherSuite {
+    for suite in rustls::ALL_CIPHER_SUITES {
+        let sname = format!("{:?}", suite.suite()).to_lowercase();
+
+        if sname == CIPHER_SUITE.to_string().to_lowercase() {
+            return *suite
+        }
+    }
+
+    unreachable!()
+}
+
+struct ServerCertificateVerifier;
+
+impl ServerCertVerifier for ServerCertificateVerifier {
+    fn verify_server_cert(
+        &self,
+        _end_entity: &Certificate,
+        _intermediates: &[Certificate],
+        _server_name: &ServerName,
+        _scts: &mut dyn Iterator<Item = &[u8]>,
+        _ocsp_response: &[u8],
+        _now: SystemTime,
+    ) -> Result<ServerCertVerified, rustls::Error> {
+        // TODO: upsycle
+        Ok(ServerCertVerified::assertion())
+    }
+}
+
+pub struct TlsTransport {
+    pub ttl: Option<u32>,
+}
+
+impl Transport for TlsTransport {
+    type Output = TlsStream<TcpStream>;
+    type Error = io::Error;
+    type Dial = Pin<Box<dyn Future<Output = Result<Self::Output, Self::Error>> + Send>>;
+
+    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 {
+    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)
+    }
+
+    async fn do_dial(self, url: Url) -> Result<TlsStream<TcpStream>, io::Error> {
+        let socket_addr = url.socket_addrs(|| None)?[0];
+        // TODO: Handle host
+        let server_name = ServerName::try_from("example.com").unwrap();
+
+        let socket = self.create_socket(socket_addr)?;
+        socket.set_nonblocking(true)?;
+
+        // TODO: This should be in the struct
+        // TODO: Client auth (see upsycle)
+        let root_store = RootCertStore::empty();
+        let server_cert_verifier = Arc::new(ServerCertificateVerifier {});
+        let config = ClientConfig::builder()
+            .with_cipher_suites(&[cipher_suite()])
+            .with_kx_groups(&[&X25519])
+            .with_protocol_versions(&[&TLS13])
+            .unwrap()
+            .with_custom_certificate_verifier(server_cert_verifier)
+            .with_no_client_auth();
+
+        let connector = TlsConnector::from(Arc::new(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))
+    }
+}