From 232b62728ba2dc37cadfc2eaf34fb8a9d9ec5fcd Mon Sep 17 00:00:00 2001 From: Dirkjan Ochtman Date: Wed, 2 May 2018 12:01:17 +0200 Subject: [PATCH] Move Endpoint into new endpoint module --- src/client.rs | 3 +- src/endpoint.rs | 160 ++++++++++++++++++++++++++++++++++++++++++++++++ src/lib.rs | 1 + src/server.rs | 3 +- src/types.rs | 160 +----------------------------------------------- 5 files changed, 167 insertions(+), 160 deletions(-) create mode 100644 src/endpoint.rs diff --git a/src/client.rs b/src/client.rs index 77fbaa128..d893c36aa 100644 --- a/src/client.rs +++ b/src/client.rs @@ -1,8 +1,9 @@ use futures::{Async, Future, Poll}; +use endpoint::Endpoint; use packet::Packet; use tls::ClientTls; -use types::{Endpoint, Side}; +use types::Side; use std::io; use std::net::ToSocketAddrs; diff --git a/src/endpoint.rs b/src/endpoint.rs new file mode 100644 index 000000000..52f56bb20 --- /dev/null +++ b/src/endpoint.rs @@ -0,0 +1,160 @@ +use rand::{thread_rng, Rng}; + +use std::mem; + +use codec::BufLen; +use crypto::{PacketKey, Secret}; +use frame::{Ack, AckFrame, Frame, PaddingFrame, StreamFrame}; +use packet::{Header, LongType, Packet}; +use tls::{ClientTls, QuicTls}; +use types::{ConnectionId, DRAFT_11, Side}; + +pub struct Endpoint { + side: Side, + pub dst_cid: ConnectionId, + pub src_cid: ConnectionId, + pub src_pn: u32, + secret: Secret, + prev_secret: Option, + tls: T, +} + +impl Endpoint +where + T: QuicTls, +{ + pub fn new(tls: T, side: Side, secret: Option) -> Self { + let mut rng = thread_rng(); + let dst_cid = rng.gen(); + + let secret = if side == Side::Client { + debug_assert!(secret.is_none()); + Secret::Handshake(dst_cid) + } else if let Some(secret) = secret { + secret + } else { + panic!("need secret for client endpoint"); + }; + + Endpoint { + tls, + side, + dst_cid, + src_cid: rng.gen(), + src_pn: rng.gen(), + secret, + prev_secret: None, + } + } + + pub(crate) fn encode_key(&self, h: &Header) -> PacketKey { + if let Some(LongType::Handshake) = h.ptype() { + if let Some(Secret::Handshake(_)) = self.prev_secret { + return self.prev_secret.as_ref().unwrap().build_key(Side::Client); + } + } + self.secret.build_key(self.side) + } + + pub(crate) fn decode_key(&self, _: &Header) -> PacketKey { + self.secret.build_key(self.side.other()) + } + + pub(crate) fn set_secret(&mut self, secret: Secret) { + let old = mem::replace(&mut self.secret, secret); + self.prev_secret = Some(old); + } + + pub fn build_initial_packet(&mut self, mut payload: Vec) -> Packet { + let number = self.src_pn; + self.src_pn += 1; + + let mut payload_len = payload.buf_len() + self.secret.tag_len(); + if payload_len < 1200 { + payload.push(Frame::Padding(PaddingFrame(1200 - payload_len))); + payload_len = 1200; + } + + Packet { + header: Header::Long { + ptype: LongType::Initial, + version: DRAFT_11, + dst_cid: self.dst_cid, + src_cid: self.src_cid, + len: payload_len as u64, + number, + }, + payload, + } + } + + pub fn build_handshake_packet(&mut self, payload: Vec) -> Packet { + let number = self.src_pn; + self.src_pn += 1; + Packet { + header: Header::Long { + ptype: LongType::Handshake, + version: DRAFT_11, + dst_cid: self.dst_cid, + src_cid: self.src_cid, + len: (payload.buf_len() + self.secret.tag_len()) as u64, + number, + }, + payload, + } + } + + pub(crate) fn handle_handshake(&mut self, rsp: &Packet) -> Option { + self.dst_cid = rsp.dst_cid(); + let tls_frame = rsp.payload + .iter() + .filter_map(|f| match *f { + Frame::Stream(ref f) => Some(f), + _ => None, + }) + .next() + .unwrap(); + + let (handshake, new_secret) = self.tls + .process_handshake_messages(&tls_frame.data) + .unwrap(); + if let Some(secret) = new_secret { + self.set_secret(secret); + } + + Some(self.build_handshake_packet(vec![ + Frame::Ack(AckFrame { + largest: rsp.number(), + ack_delay: 0, + blocks: vec![Ack::Ack(0)], + }), + Frame::Stream(StreamFrame { + id: 0, + fin: false, + offset: 0, + len: Some(handshake.len() as u64), + data: handshake, + }), + ])) + } +} + +impl Endpoint { + pub(crate) fn initial(&mut self, server: &str) -> Packet { + let (handshake, new_secret) = self.tls.get_handshake(server).unwrap(); + if let Some(secret) = new_secret { + self.set_secret(secret); + } + + self.build_initial_packet(vec![ + Frame::Stream(StreamFrame { + id: 0, + fin: false, + offset: 0, + len: Some(handshake.len() as u64), + data: handshake, + }), + ]) + } +} + diff --git a/src/lib.rs b/src/lib.rs index 4dfc8e9c9..511071e64 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -15,6 +15,7 @@ pub use server::Server; mod client; mod codec; mod crypto; +mod endpoint; mod frame; mod packet; mod server; diff --git a/src/server.rs b/src/server.rs index 70da846e6..7bf01ba22 100644 --- a/src/server.rs +++ b/src/server.rs @@ -1,8 +1,9 @@ use futures::{Future, Poll}; use crypto::Secret; +use endpoint::Endpoint; use packet::{LongType, Packet}; -use types::{ConnectionId, Endpoint, Side}; +use types::{ConnectionId, Side}; use tls::{self, ServerTls}; use std::collections::{HashMap, hash_map::Entry}; diff --git a/src/types.rs b/src/types.rs index b9b8c2c7d..3abe75ba4 100644 --- a/src/types.rs +++ b/src/types.rs @@ -1,163 +1,7 @@ -use rand::{thread_rng, Rand, Rng}; +use rand::{Rand, Rng}; -use std::mem; use std::ops::Deref; -use codec::BufLen; -use crypto::{PacketKey, Secret}; -use frame::{Ack, AckFrame, Frame, PaddingFrame, StreamFrame}; -use packet::{Header, LongType, Packet}; -use tls::{ClientTls, QuicTls}; - -pub struct Endpoint { - side: Side, - pub dst_cid: ConnectionId, - pub src_cid: ConnectionId, - pub src_pn: u32, - secret: Secret, - prev_secret: Option, - tls: T, -} - -impl Endpoint -where - T: QuicTls, -{ - pub fn new(tls: T, side: Side, secret: Option) -> Self { - let mut rng = thread_rng(); - let dst_cid = rng.gen(); - - let secret = if side == Side::Client { - debug_assert!(secret.is_none()); - Secret::Handshake(dst_cid) - } else if let Some(secret) = secret { - secret - } else { - panic!("need secret for client endpoint"); - }; - - Endpoint { - tls, - side, - dst_cid, - src_cid: rng.gen(), - src_pn: rng.gen(), - secret, - prev_secret: None, - } - } - - pub(crate) fn encode_key(&self, h: &Header) -> PacketKey { - if let Some(LongType::Handshake) = h.ptype() { - if let Some(Secret::Handshake(_)) = self.prev_secret { - return self.prev_secret.as_ref().unwrap().build_key(Side::Client); - } - } - self.secret.build_key(self.side) - } - - pub(crate) fn decode_key(&self, _: &Header) -> PacketKey { - self.secret.build_key(self.side.other()) - } - - pub(crate) fn set_secret(&mut self, secret: Secret) { - let old = mem::replace(&mut self.secret, secret); - self.prev_secret = Some(old); - } - - pub fn build_initial_packet(&mut self, mut payload: Vec) -> Packet { - let number = self.src_pn; - self.src_pn += 1; - - let mut payload_len = payload.buf_len() + self.secret.tag_len(); - if payload_len < 1200 { - payload.push(Frame::Padding(PaddingFrame(1200 - payload_len))); - payload_len = 1200; - } - - Packet { - header: Header::Long { - ptype: LongType::Initial, - version: DRAFT_11, - dst_cid: self.dst_cid, - src_cid: self.src_cid, - len: payload_len as u64, - number, - }, - payload, - } - } - - pub fn build_handshake_packet(&mut self, payload: Vec) -> Packet { - let number = self.src_pn; - self.src_pn += 1; - Packet { - header: Header::Long { - ptype: LongType::Handshake, - version: DRAFT_11, - dst_cid: self.dst_cid, - src_cid: self.src_cid, - len: (payload.buf_len() + self.secret.tag_len()) as u64, - number, - }, - payload, - } - } - - pub(crate) fn handle_handshake(&mut self, rsp: &Packet) -> Option { - self.dst_cid = rsp.dst_cid(); - let tls_frame = rsp.payload - .iter() - .filter_map(|f| match *f { - Frame::Stream(ref f) => Some(f), - _ => None, - }) - .next() - .unwrap(); - - let (handshake, new_secret) = self.tls - .process_handshake_messages(&tls_frame.data) - .unwrap(); - if let Some(secret) = new_secret { - self.set_secret(secret); - } - - Some(self.build_handshake_packet(vec![ - Frame::Ack(AckFrame { - largest: rsp.number(), - ack_delay: 0, - blocks: vec![Ack::Ack(0)], - }), - Frame::Stream(StreamFrame { - id: 0, - fin: false, - offset: 0, - len: Some(handshake.len() as u64), - data: handshake, - }), - ])) - } -} - -impl Endpoint { - pub(crate) fn initial(&mut self, server: &str) -> Packet { - let (handshake, new_secret) = self.tls.get_handshake(server).unwrap(); - if let Some(secret) = new_secret { - self.set_secret(secret); - } - - self.build_initial_packet(vec![ - Frame::Stream(StreamFrame { - id: 0, - fin: false, - offset: 0, - len: Some(handshake.len() as u64), - data: handshake, - }), - ]) - } -} - #[derive(Clone, Debug, Eq, Hash, PartialEq)] pub struct ConnectionId { pub len: u8, @@ -224,7 +68,7 @@ pub enum Side { } impl Side { - fn other(&self) -> Side { + pub fn other(&self) -> Side { match *self { Side::Client => Side::Server, Side::Server => Side::Client,