From b8adc7a8d9bbf292057fbb972fb59df698bdbe00 Mon Sep 17 00:00:00 2001 From: Dirkjan Ochtman Date: Thu, 17 Jan 2019 16:51:36 +0100 Subject: [PATCH] Always pass all datagram contents to Connection in one go --- quinn-proto/src/connection.rs | 43 +++++++++++++++-- quinn-proto/src/endpoint.rs | 87 ++++++++++++++++------------------- 2 files changed, 80 insertions(+), 50 deletions(-) diff --git a/quinn-proto/src/connection.rs b/quinn-proto/src/connection.rs index f00fa7278..c9bf6185d 100644 --- a/quinn-proto/src/connection.rs +++ b/quinn-proto/src/connection.rs @@ -3,7 +3,7 @@ use std::net::SocketAddr; use std::sync::Arc; use std::{cmp, io, mem}; -use bytes::{Buf, Bytes}; +use bytes::{Buf, Bytes, BytesMut}; use err_derive::Error; use fnv::{FnvHashMap, FnvHashSet}; use rand::{rngs::OsRng, Rng}; @@ -816,9 +816,11 @@ impl Connection { pub fn handle_initial( &mut self, now: u64, + remote: SocketAddr, ecn: Option, packet_number: u64, packet: Packet, + remaining: Option, ) -> Result<(), TransportError> { let len = packet.header_data.len() + packet.payload.len(); self.on_packet_authenticated(now, SpaceId::Initial, ecn, Some(packet_number), false, len); @@ -833,6 +835,7 @@ impl Connection { self.set_params(params)?; self.write_tls(); self.init_0rtt(); + self.handle_coalesced(now, remote, ecn, remaining); Ok(()) } @@ -963,12 +966,13 @@ impl Connection { } } - pub fn handle_decode( + pub fn handle_dgram( &mut self, now: u64, remote: SocketAddr, ecn: Option, - partial_decode: PartialDecode, + first_decode: PartialDecode, + remaining: Option, ) { if remote != self.remote && self.side.is_client() { trace!( @@ -978,6 +982,39 @@ impl Connection { ); return; } + + self.handle_decode(now, remote, ecn, first_decode); + self.handle_coalesced(now, remote, ecn, remaining); + } + + fn handle_coalesced( + &mut self, + now: u64, + remote: SocketAddr, + ecn: Option, + mut remaining: Option, + ) { + while let Some(data) = remaining { + match PartialDecode::new(data, self.config.local_cid_len) { + Ok((partial_decode, rest)) => { + remaining = rest; + self.handle_decode(now, remote, ecn, partial_decode); + } + Err(e) => { + trace!(self.log, "malformed header"; "reason" => %e); + return; + } + } + } + } + + fn handle_decode( + &mut self, + now: u64, + remote: SocketAddr, + ecn: Option, + partial_decode: PartialDecode, + ) { let header_crypto = if partial_decode.is_0rtt() { if let Some(ref crypto) = self.zero_rtt_crypto { Some(&crypto.header) diff --git a/quinn-proto/src/endpoint.rs b/quinn-proto/src/endpoint.rs index afec98cc8..6c541f9cd 100644 --- a/quinn-proto/src/endpoint.rs +++ b/quinn-proto/src/endpoint.rs @@ -145,55 +145,40 @@ impl Endpoint { data: BytesMut, ) { let datagram_len = data.len(); - let mut remaining = Some(data); - while let Some(data) = remaining { - match PartialDecode::new(data, self.config.local_cid_len) { - Ok((partial_decode, rest)) => { - remaining = rest; - self.handle_decode(now, remote, ecn, partial_decode, datagram_len); - } - Err(PacketDecodeError::UnsupportedVersion { - source, - destination, - }) => { - if !self.is_server() { - debug!(self.log, "dropping packet with unsupported version"); - return; - } - trace!(self.log, "sending version negotiation"); - // Negotiate versions - let mut buf = Vec::::new(); - Header::VersionNegotiate { - random: self.rng.gen(), - src_cid: destination, - dst_cid: source, - } - .encode(&mut buf); - buf.write::(0x0a1a_2a3a); // reserved version - buf.write(VERSION); // supported version - self.io.push_back(Io::Transmit { - destination: remote, - ecn: None, - packet: buf.into(), - }); + let (partial_decode, rest) = match PartialDecode::new(data, self.config.local_cid_len) { + Ok(x) => x, + Err(PacketDecodeError::UnsupportedVersion { + source, + destination, + }) => { + if !self.is_server() { + debug!(self.log, "dropping packet with unsupported version"); return; } - Err(e) => { - trace!(self.log, "malformed header"; "reason" => %e); - return; + trace!(self.log, "sending version negotiation"); + // Negotiate versions + let mut buf = Vec::::new(); + Header::VersionNegotiate { + random: self.rng.gen(), + src_cid: destination, + dst_cid: source, } + .encode(&mut buf); + buf.write::(0x0a1a_2a3a); // reserved version + buf.write(VERSION); // supported version + self.io.push_back(Io::Transmit { + destination: remote, + ecn: None, + packet: buf.into(), + }); + return; } - } - } + Err(e) => { + trace!(self.log, "malformed header"; "reason" => %e); + return; + } + }; - fn handle_decode( - &mut self, - now: u64, - remote: SocketAddr, - ecn: Option, - partial_decode: PartialDecode, - datagram_len: usize, - ) { // // Handle packet on existing connection, if any // @@ -219,7 +204,7 @@ impl Endpoint { }; if let Some(ch) = known_ch { let had_1rtt = self.connections[ch].has_1rtt(); - self.connections[ch].handle_decode(now, remote, ecn, partial_decode); + self.connections[ch].handle_dgram(now, remote, ecn, partial_decode, rest); if !had_1rtt && (self.connections[ch].has_1rtt() || !self.connections[ch].is_handshaking()) { @@ -259,7 +244,7 @@ impl Endpoint { let header_crypto = crypto.header_crypto(); match partial_decode.finish(Some(&header_crypto)) { Ok(packet) => { - self.handle_initial(now, remote, ecn, packet, &crypto, &header_crypto) + self.handle_initial(now, remote, ecn, packet, rest, &crypto, &header_crypto) } Err(e) => { trace!(self.log, "unable to decode packet"; "reason" => %e); @@ -423,6 +408,7 @@ impl Endpoint { remote: SocketAddr, ecn: Option, mut packet: Packet, + rest: Option, crypto: &Crypto, header_crypto: &HeaderCrypto, ) { @@ -554,7 +540,14 @@ impl Endpoint { if dst_cid.len() != 0 { self.connection_ids_initial.insert(dst_cid, ch); } - match self.connections[ch].handle_initial(now, ecn, packet_number as u64, packet) { + match self.connections[ch].handle_initial( + now, + remote, + ecn, + packet_number as u64, + packet, + rest, + ) { Ok(()) => { self.incoming_handshakes += 1; self.dirty_conns.insert(ch);