Directly use QuicClientTls/QuicServerTls from rustls

This commit is contained in:
Dirkjan Ochtman
2018-05-03 16:21:00 +02:00
parent 010800dda6
commit 949f491bdc
6 changed files with 98 additions and 105 deletions
+1 -1
View File
@@ -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();
}
+3 -3
View File
@@ -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<ClientTls>,
endpoint: Endpoint<tls::QuicClientTls>,
socket: UdpSocket,
buf: Vec<u8>,
state: ConnectionState,
+10 -9
View File
@@ -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<T> {
side: Side,
@@ -19,9 +20,10 @@ pub struct Endpoint<T> {
tls: T,
}
impl<T> Endpoint<T>
impl<T, S> Endpoint<T>
where
T: QuicTls,
T: DerefMut + Deref<Target = S>,
S: tls::Session,
{
pub fn new(tls: T, side: Side, secret: Option<Secret>) -> 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<ClientTls> {
impl Endpoint<tls::QuicClientTls> {
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);
}
+3 -3
View File
@@ -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<tls::ServerConfig>,
in_buf: Vec<u8>,
out_buf: Vec<u8>,
connections: HashMap<ConnectionId, (SocketAddr, Endpoint<ServerTls>)>,
connections: HashMap<ConnectionId, (SocketAddr, Endpoint<tls::QuicServerTls>)>,
}
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)),
);
+7 -7
View File
@@ -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<ServerTls> {
fn server_endpoint(hs_cid: ConnectionId) -> Endpoint<tls::QuicServerTls> {
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<ServerTls> {
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<ClientTls> {
fn client_endpoint() -> Endpoint<tls::QuicClientTls> {
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<ClientTls> {
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)
+74 -82
View File
@@ -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<ClientConfig>) -> 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<TlsResult, TLSError> {
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<u8>, Option<Secret>), TLSError> {
let pki_server_name = DNSNameRef::try_from_ascii_str(hostname).unwrap();
let params = ClientTransportParameters {
initial_version: 1,
pub fn server_session(config: &Arc<ServerConfig>) -> 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<Certificate>, 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<T, S>(
session: &mut T,
msgs: Option<&[u8]>,
) -> Result<TlsResult, TLSError>
where
T: DerefMut + Deref<Target = S>,
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<u8>, Option<Secret>), 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<ServerConfig>) -> 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<Certificate>, 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<u8>, Option<Secret>), TLSError> {
Ok(process_tls_result(self.session.get_handshake(input)?))
}
}
fn process_tls_result(res: TLSResult) -> (Vec<u8>, Option<Secret>) {
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<u8>, Option<Secret>), TLSError>;
}
type TlsResult = (Vec<u8>, Option<Secret>);
macro_rules! try_ret(
($e:expr) => (match $e { Some(e) => e, None => return None })