From 949f491bdca5f22b53a062b9f5d80ee8a4aa88b3 Mon Sep 17 00:00:00 2001 From: Dirkjan Ochtman Date: Thu, 3 May 2018 16:21:00 +0200 Subject: [PATCH] Directly use QuicClientTls/QuicServerTls from rustls --- examples/server.rs | 2 +- src/client.rs | 6 +- src/endpoint.rs | 19 +++--- src/server.rs | 6 +- src/tests.rs | 14 ++-- src/tls.rs | 156 +++++++++++++++++++++------------------------ 6 files changed, 98 insertions(+), 105 deletions(-) diff --git a/examples/server.rs b/examples/server.rs index 158d02cba..efad4aaff 100644 --- a/examples/server.rs +++ b/examples/server.rs @@ -23,6 +23,6 @@ fn main() { pemfile::rsa_private_keys(&mut reader).expect("cannot read private keys") }; - let tls_config = quinn::tls::ServerTls::build_config(certs, key[0].clone()); + let tls_config = quinn::tls::build_server_config(certs, key[0].clone()); quinn::Server::new("0.0.0.0", 4433, tls_config).run(); } diff --git a/src/client.rs b/src/client.rs index 48787a3b4..fd43c4e83 100644 --- a/src/client.rs +++ b/src/client.rs @@ -2,7 +2,7 @@ use futures::{Async, Future, Poll}; use endpoint::Endpoint; use packet::Packet; -use tls::ClientTls; +use tls; use types::Side; use std::io; @@ -14,7 +14,7 @@ pub struct QuicStream {} impl QuicStream { pub fn connect(server: &str, port: u16) -> ConnectFuture { - let mut endpoint = Endpoint::new(ClientTls::new(), Side::Client, None); + let mut endpoint = Endpoint::new(tls::client_session(None), Side::Client, None); let packet = endpoint.initial(server); let mut buf = Vec::with_capacity(65536); packet.encode(&endpoint.encode_key(&packet.header), &mut buf); @@ -33,7 +33,7 @@ impl QuicStream { #[must_use = "futures do nothing unless polled"] pub struct ConnectFuture { - endpoint: Endpoint, + endpoint: Endpoint, socket: UdpSocket, buf: Vec, state: ConnectionState, diff --git a/src/endpoint.rs b/src/endpoint.rs index 043aa2f58..723f445a3 100644 --- a/src/endpoint.rs +++ b/src/endpoint.rs @@ -1,13 +1,14 @@ use rand::{thread_rng, Rng}; use std::mem; +use std::ops::{Deref, DerefMut}; use codec::BufLen; use crypto::{PacketKey, Secret}; use frame::{Ack, AckFrame, Frame, PaddingFrame, StreamFrame}; use packet::{Header, LongType, Packet, ShortType}; -use tls::{ClientTls, QuicTls}; -use types::{ConnectionId, DRAFT_11, GENERATED_CID_LENGTH, Side}; +use tls; +use types::{ConnectionId, DRAFT_11, Side, GENERATED_CID_LENGTH}; pub struct Endpoint { side: Side, @@ -19,9 +20,10 @@ pub struct Endpoint { tls: T, } -impl Endpoint +impl Endpoint where - T: QuicTls, + T: DerefMut + Deref, + S: tls::Session, { pub fn new(tls: T, side: Side, secret: Option) -> Self { let mut rng = thread_rng(); @@ -144,9 +146,8 @@ where .next() .unwrap(); - let (handshake, new_secret) = self.tls - .process_handshake_messages(&tls_frame.data) - .unwrap(); + let (handshake, new_secret) = + tls::process_handshake_messages(&mut self.tls, Some(&tls_frame.data)).unwrap(); if let Some(secret) = new_secret { self.set_secret(secret); } @@ -178,9 +179,9 @@ where } } -impl Endpoint { +impl Endpoint { pub(crate) fn initial(&mut self, server: &str) -> Packet { - let (handshake, new_secret) = self.tls.get_handshake(server).unwrap(); + let (handshake, new_secret) = tls::start_handshake(&mut self.tls, server).unwrap(); if let Some(secret) = new_secret { self.set_secret(secret); } diff --git a/src/server.rs b/src/server.rs index a0a2e5350..49d57d048 100644 --- a/src/server.rs +++ b/src/server.rs @@ -4,7 +4,7 @@ use crypto::Secret; use endpoint::Endpoint; use packet::{LongType, Packet}; use types::{ConnectionId, Side}; -use tls::{self, ServerTls}; +use tls; use std::collections::{HashMap, hash_map::Entry}; use std::io; @@ -18,7 +18,7 @@ pub struct Server { tls_config: Arc, in_buf: Vec, out_buf: Vec, - connections: HashMap)>, + connections: HashMap)>, } impl Server { @@ -50,7 +50,7 @@ impl Future for Server { let cid = if partial.header.ptype() == Some(LongType::Initial) { let mut endpoint = Endpoint::new( - ServerTls::with_config(&self.tls_config), + tls::server_session(&self.tls_config), Side::Server, Some(Secret::Handshake(dst_cid)), ); diff --git a/src/tests.rs b/src/tests.rs index 414f2a965..2a8f3e623 100644 --- a/src/tests.rs +++ b/src/tests.rs @@ -8,7 +8,7 @@ use std::sync::Arc; use crypto::Secret; use endpoint::Endpoint; use packet::Packet; -use tls::{ClientTls, ServerTls}; +use tls; use types::{ConnectionId, Side}; use self::untrusted::Input; @@ -67,7 +67,7 @@ fn test_handshake() { assert!(c.handle_handshake(&server_hello).is_some()); } -fn server_endpoint(hs_cid: ConnectionId) -> Endpoint { +fn server_endpoint(hs_cid: ConnectionId) -> Endpoint { let certs = { let f = File::open("certs/server.chain").expect("cannot open 'certs/server.chain'"); let mut reader = BufReader::new(f); @@ -80,15 +80,15 @@ fn server_endpoint(hs_cid: ConnectionId) -> Endpoint { pemfile::rsa_private_keys(&mut reader).expect("cannot read private keys") }; - let tls_config = Arc::new(ServerTls::build_config(certs, keys[0].clone())); + let tls_config = Arc::new(tls::build_server_config(certs, keys[0].clone())); Endpoint::new( - ServerTls::with_config(&tls_config), + tls::server_session(&tls_config), Side::Server, Some(Secret::Handshake(hs_cid)), ) } -fn client_endpoint() -> Endpoint { +fn client_endpoint() -> Endpoint { let tls = { let mut f = File::open("certs/ca.der").expect("cannot open 'certs/ca.der'"); let mut bytes = Vec::new(); @@ -97,8 +97,8 @@ fn client_endpoint() -> Endpoint { let anchor = webpki::trust_anchor_util::cert_der_as_trust_anchor(Input::from(&bytes)).unwrap(); let anchor_vec = vec![anchor]; - let config = ClientTls::build_config(Some(&webpki::TLSServerTrustAnchors(&anchor_vec))); - ClientTls::with_config(config) + let config = tls::build_client_config(Some(&webpki::TLSServerTrustAnchors(&anchor_vec))); + tls::client_session(Some(config)) }; Endpoint::new(tls, Side::Client, None) diff --git a/src/tls.rs b/src/tls.rs index 8358e9ab2..f499bca44 100644 --- a/src/tls.rs +++ b/src/tls.rs @@ -1,7 +1,8 @@ use rustls::internal::msgs::codec::{self, Codec}; -use rustls::{ClientConfig, NoClientAuth, ProtocolVersion}; -use rustls::quic::{ClientSession, QuicSecret, ServerSession, TLSResult}; +use rustls::{ClientConfig, NoClientAuth, ProtocolVersion, TLSError}; +use std::io::Cursor; +use std::ops::{Deref, DerefMut}; use std::sync::Arc; use crypto::Secret; @@ -10,106 +11,97 @@ use types::{DRAFT_11, TransportParameters}; use webpki::{DNSNameRef, TLSServerTrustAnchors}; use webpki_roots; -pub use rustls::{Certificate, PrivateKey, ServerConfig, SupportedCipherSuite, TLSError}; +pub use rustls::quic::{QuicClientTls, QuicServerTls}; +pub use rustls::{Certificate, PrivateKey, ServerConfig, Session}; -pub struct ClientTls { - pub session: ClientSession, +pub fn client_session(config: Option) -> QuicClientTls { + QuicClientTls::new(&Arc::new(config.unwrap_or(build_client_config(None)))) } -impl ClientTls { - pub fn new() -> Self { - Self::with_config(Self::build_config(None)) - } +pub fn build_client_config(anchors: Option<&TLSServerTrustAnchors>) -> ClientConfig { + let mut config = ClientConfig::new(); + let anchors = anchors.unwrap_or(&webpki_roots::TLS_SERVER_ROOTS); + config.root_store.add_server_trust_anchors(anchors); + config.versions = vec![ProtocolVersion::TLSv1_3]; + config.alpn_protocols = vec![ALPN_PROTOCOL.into()]; + config +} - pub fn with_config(config: ClientConfig) -> Self { - Self { - session: ClientSession::new(&Arc::new(config)), - } - } +pub fn start_handshake(tls: &mut QuicClientTls, hostname: &str) -> Result { + let pki_server_name = DNSNameRef::try_from_ascii_str(hostname).unwrap(); + let params = ClientTransportParameters { + initial_version: 1, + parameters: TransportParameters::default(), + }; + tls.start_handshake(pki_server_name, to_vec(params)); + process_handshake_messages(tls, None) +} - pub fn build_config(anchors: Option<&TLSServerTrustAnchors>) -> ClientConfig { - let mut config = ClientConfig::new(); - let anchors = anchors.unwrap_or(&webpki_roots::TLS_SERVER_ROOTS); - config.root_store.add_server_trust_anchors(anchors); - config.versions = vec![ProtocolVersion::TLSv1_3]; - config.alpn_protocols = vec![ALPN_PROTOCOL.into()]; - config - } - - pub fn get_handshake(&mut self, hostname: &str) -> Result<(Vec, Option), TLSError> { - let pki_server_name = DNSNameRef::try_from_ascii_str(hostname).unwrap(); - let params = ClientTransportParameters { - initial_version: 1, +pub fn server_session(config: &Arc) -> QuicServerTls { + QuicServerTls::new( + config, + to_vec(ServerTransportParameters { + negotiated_version: DRAFT_11, + supported_versions: vec![DRAFT_11], parameters: TransportParameters::default(), - }; - Ok(process_tls_result(self.session.get_handshake(pki_server_name, to_vec(params))?)) + }), + ) +} + +pub fn build_server_config(cert_chain: Vec, key: PrivateKey) -> ServerConfig { + let mut config = ServerConfig::new(NoClientAuth::new()); + config.set_protocols(&[ALPN_PROTOCOL.into()]); + config.set_single_cert(cert_chain, key); + config +} + +pub fn process_handshake_messages( + session: &mut T, + msgs: Option<&[u8]>, +) -> Result +where + T: DerefMut + Deref, + S: Session, +{ + if let Some(data) = msgs { + let mut read = Cursor::new(data); + let did_read = session.read_tls(&mut read).unwrap(); + debug_assert_eq!(did_read, data.len()); + session.process_new_packets()?; } -} -impl QuicTls for ClientTls { - fn process_handshake_messages( - &mut self, - input: &[u8], - ) -> Result<(Vec, Option), TLSError> { - Ok(process_tls_result(self.session.process_handshake_messages(input)?)) - } -} + let key_ready = if !session.is_handshaking() { + let suite = session + .get_negotiated_ciphersuite() + .ok_or(TLSError::HandshakeNotComplete) + .unwrap(); -pub struct ServerTls { - session: ServerSession, -} + let mut secret = vec![0u8; suite.enc_key_len]; + session.export_keying_material(&mut secret, b"EXPORTER-QUIC client 1rtt", None)?; + Some((suite, secret)) + } else { + None + }; -impl ServerTls { - pub fn with_config(config: &Arc) -> Self { - Self { - session: ServerSession::new( - config, - to_vec(ServerTransportParameters { - negotiated_version: DRAFT_11, - supported_versions: vec![DRAFT_11], - parameters: TransportParameters::default(), - }), - ), + let mut messages = Vec::new(); + loop { + let size = session.write_tls(&mut messages).unwrap(); + if size == 0 { + break; } } - pub fn build_config(cert_chain: Vec, key: PrivateKey) -> ServerConfig { - let mut config = ServerConfig::new(NoClientAuth::new()); - config.set_protocols(&[ALPN_PROTOCOL.into()]); - config.set_single_cert(cert_chain, key); - config - } -} - -impl QuicTls for ServerTls { - fn process_handshake_messages( - &mut self, - input: &[u8], - ) -> Result<(Vec, Option), TLSError> { - Ok(process_tls_result(self.session.get_handshake(input)?)) - } -} - -fn process_tls_result(res: TLSResult) -> (Vec, Option) { - let TLSResult { - messages, - key_ready, - } = res; - let secret = if let Some((suite, QuicSecret::For1RTT(secret))) = key_ready { + let secret = if let Some((suite, secret)) = key_ready { let (aead_alg, hash_alg) = (suite.get_aead_alg(), suite.get_hash()); Some(Secret::For1Rtt(aead_alg, hash_alg, secret)) } else { None }; - (messages, secret) + + Ok((messages, secret)) } -pub trait QuicTls { - fn process_handshake_messages( - &mut self, - input: &[u8], - ) -> Result<(Vec, Option), TLSError>; -} +type TlsResult = (Vec, Option); macro_rules! try_ret( ($e:expr) => (match $e { Some(e) => e, None => return None })