From 4096cdc19c47933647d846825bc1bb5dc7deaa65 Mon Sep 17 00:00:00 2001 From: Benjamin Saunders Date: Fri, 21 Dec 2018 13:42:59 -0800 Subject: [PATCH] Issue new connection IDs on connection established --- quinn-proto/src/connection.rs | 71 +++++++++++++++++++++++++---------- quinn-proto/src/endpoint.rs | 45 +++++++++++++--------- quinn-proto/src/frame.rs | 10 +++++ quinn/examples/server.rs | 2 +- quinn/src/lib.rs | 11 ++++-- 5 files changed, 97 insertions(+), 42 deletions(-) diff --git a/quinn-proto/src/connection.rs b/quinn-proto/src/connection.rs index d37179950..f385fd10f 100644 --- a/quinn-proto/src/connection.rs +++ b/quinn-proto/src/connection.rs @@ -1,4 +1,4 @@ -use std::collections::{hash_map, BTreeMap, VecDeque}; +use std::collections::{hash_map, BTreeMap, HashMap, VecDeque}; use std::net::SocketAddrV6; use std::sync::Arc; use std::{cmp, io, mem}; @@ -8,7 +8,7 @@ use fnv::{FnvHashMap, FnvHashSet}; use slog::Logger; use crate::coding::{BufExt, BufMutExt}; -use crate::crypto::{self, Crypto, HeaderCrypto, TlsSession, ACK_DELAY_EXPONENT}; +use crate::crypto::{self, reset_token_for, Crypto, HeaderCrypto, TlsSession, ACK_DELAY_EXPONENT}; use crate::dedup::Dedup; use crate::endpoint::{Config, Event, Timer}; use crate::frame::FrameStruct; @@ -32,7 +32,7 @@ pub struct Connection { app_closed: bool, /// DCID of Initial packet pub(crate) init_cid: ConnectionId, - loc_cid: ConnectionId, + loc_cids: HashMap, rem_cid: ConnectionId, pub(crate) remote: SocketAddrV6, state: State, @@ -68,6 +68,8 @@ pub struct Connection { lost_packets: u64, io: IoQueue, events: VecDeque, + /// Number of local connection IDs that have been issued in NEW_CONNECTION_ID frames. + cids_issued: u64, // // Loss Detection @@ -202,7 +204,8 @@ impl Connection { Stream::new_bi(config.stream_receive_window as u64), ); } - + let mut loc_cids = HashMap::new(); + loc_cids.insert(0, loc_cid); let state = State::Handshake(state::Handshake { rem_cid_set: side.is_server(), token: None, @@ -212,7 +215,7 @@ impl Connection { tls, app_closed: false, init_cid, - loc_cid, + loc_cids, rem_cid, remote, side, @@ -238,6 +241,7 @@ impl Connection { lost_packets: 0, io: IoQueue::new(), events: VecDeque::new(), + cids_issued: 0, handshake_count: 0, tlp_count: 0, @@ -1314,6 +1318,17 @@ impl Connection { } } + pub fn issue_cid(&mut self, cid: ConnectionId) { + let token = reset_token_for(&self.config.reset_key, &cid); + self.cids_issued += 1; + self.pending.new_cids.push(frame::NewConnectionId { + id: cid, + sequence: self.cids_issued, + reset_token: token, + }); + self.loc_cids.insert(self.cids_issued, cid); + } + fn process_payload( &mut self, now: u64, @@ -1540,8 +1555,10 @@ impl Connection { } Frame::NewConnectionId { .. } => { if self.rem_cid.is_empty() { - debug!(self.log, "got NEW_CONNECTION_ID for connection {connection} with empty remote ID", - connection=self.loc_cid); + debug!( + self.log, + "got NEW_CONNECTION_ID when remote isn't using connection IDs" + ); return Err(TransportError::PROTOCOL_VIOLATION); } trace!(self.log, "ignoring NEW_CONNECTION_ID (unimplemented)"); @@ -1585,7 +1602,7 @@ impl Connection { { trace!(self.log, "sending initial packet"; "pn" => number); Header::Initial { - src_cid: self.loc_cid, + src_cid: *self.loc_cids.values().next().unwrap(), dst_cid: self.rem_cid, token: match self.state { State::Handshake(ref state) => { @@ -1599,7 +1616,7 @@ impl Connection { trace!(self.log, "sending handshake packet"; "pn" => number); Header::Long { ty: LongType::Handshake, - src_cid: self.loc_cid, + src_cid: *self.loc_cids.values().next().unwrap(), dst_cid: self.rem_cid, number: PacketNumber::new(number, self.largest_acked_packet), } @@ -1817,6 +1834,22 @@ impl Connection { buf.write_var(self.streams.max_remote_bi); } + // NEW_CONNECTION_ID + while buf.len() + 44 < max_size { + let frame = if let Some(x) = pending.new_cids.pop() { + x + } else { + break; + }; + trace!( + self.log, + "NEW_CONNECTION_ID {sequence}", + sequence = frame.sequence + ); + frame.encode(&mut buf); + sent.new_cids.push(frame); + } + // STREAM while buf.len() + frame::Stream::::SIZE_BOUND < max_size { let mut stream = if let Some(x) = pending.stream.pop_front() { @@ -1948,7 +1981,7 @@ impl Connection { CryptoLevel::Initial => Header::Long { ty: LongType::Handshake, dst_cid: self.rem_cid, - src_cid: self.loc_cid, + src_cid: *self.loc_cids.values().next().unwrap(), number, }, }; @@ -2350,12 +2383,12 @@ impl Connection { self.side } - /// The `ConnectionId` used for this Connection locally - pub fn loc_cid(&self) -> ConnectionId { - self.loc_cid + /// The `ConnectionId`s defined for this Connection locally. + pub fn loc_cids(&self) -> impl Iterator { + self.loc_cids.values() } - /// The `ConnectionId` used for this Connection by the peer + /// The `ConnectionId` defined for this Connection by the peer. pub fn rem_cid(&self) -> ConnectionId { self.rem_cid } @@ -2529,7 +2562,6 @@ pub struct Retransmits { max_uni_stream_id: bool, max_bi_stream_id: bool, ping: bool, - new_connection_id: Option, stream: VecDeque, /// packet number, token path_response: Option<(u64, u64)>, @@ -2537,6 +2569,7 @@ pub struct Retransmits { stop_sending: Vec<(StreamId, u16)>, max_stream_data: FnvHashSet, crypto: VecDeque, + new_cids: Vec, } impl Retransmits { @@ -2545,13 +2578,13 @@ impl Retransmits { && !self.max_uni_stream_id && !self.max_bi_stream_id && !self.ping - && self.new_connection_id.is_none() && self.stream.is_empty() && self.path_response.is_none() && self.rst_stream.is_empty() && self.stop_sending.is_empty() && self.max_stream_data.is_empty() && self.crypto.is_empty() + && self.new_cids.is_empty() } pub fn path_challenge(&mut self, packet: u64, token: u64) { @@ -2574,13 +2607,13 @@ impl Default for Retransmits { max_uni_stream_id: false, max_bi_stream_id: false, ping: false, - new_connection_id: None, stream: VecDeque::new(), path_response: None, rst_stream: Vec::new(), stop_sending: Vec::new(), max_stream_data: FnvHashSet::default(), crypto: VecDeque::new(), + new_cids: Vec::new(), } } } @@ -2591,9 +2624,6 @@ impl ::std::ops::AddAssign for Retransmits { self.ping |= rhs.ping; self.max_uni_stream_id |= rhs.max_uni_stream_id; self.max_bi_stream_id |= rhs.max_bi_stream_id; - if let Some(x) = rhs.new_connection_id { - self.new_connection_id = Some(x); - } self.stream.extend(rhs.stream.into_iter()); if let Some((packet, token)) = rhs.path_response { self.path_challenge(packet, token); @@ -2602,6 +2632,7 @@ impl ::std::ops::AddAssign for Retransmits { self.stop_sending.extend_from_slice(&rhs.stop_sending); self.max_stream_data.extend(&rhs.max_stream_data); self.crypto.extend(rhs.crypto.into_iter()); + self.new_cids.extend(&rhs.new_cids); } } diff --git a/quinn-proto/src/endpoint.rs b/quinn-proto/src/endpoint.rs index dff8159e6..6bf447501 100644 --- a/quinn-proto/src/endpoint.rs +++ b/quinn-proto/src/endpoint.rs @@ -203,12 +203,24 @@ impl Endpoint { .cloned() }; if let Some(conn_id) = conn { - let conn = &mut self.connections[conn_id.0]; - let was_handshake = conn.is_handshaking(); - let remaining = conn.handle_decode(now, ecn, partial_decode); - if was_handshake && !conn.is_handshaking() && conn.side().is_server() { - self.incoming_handshakes -= 1; - self.incoming.push_back(conn_id); + let was_handshake = self.connections[conn_id.0].is_handshaking(); + let remaining = self.connections[conn_id.0].handle_decode(now, ecn, partial_decode); + if was_handshake && !self.connections[conn_id.0].is_handshaking() { + // Newly established connection + if self.connections[conn_id.0].side().is_server() { + self.incoming_handshakes -= 1; + self.incoming.push_back(conn_id); + } + if self.config.local_cid_len != 0 { + /// Draft 17 ยง5.1.1: endpoints SHOULD provide and maintain at least eight + /// connection IDs + const LOCAL_CID_COUNT: usize = 8; + // We've already issued one CID as part of the normal handshake process. + for _ in 1..LOCAL_CID_COUNT { + let cid = self.new_cid(); + self.connections[conn_id.0].issue_cid(cid); + } + } } self.dirty_conns.insert(conn_id); self.eventful_conns.insert(conn_id); @@ -321,12 +333,10 @@ impl Endpoint { config: &Arc, server_name: &str, ) -> Result { - let local_id = self.new_cid(); let remote_id = ConnectionId::random(&mut self.rng, MAX_CID_SIZE); trace!(self.log, "initial dcid"; "value" => %remote_id); let conn = self.add_connection( remote_id, - local_id, remote_id, remote, ConnectionOpts::Client(ClientConfig { @@ -351,12 +361,11 @@ impl Endpoint { fn add_connection( &mut self, initial_id: ConnectionId, - local_id: ConnectionId, remote_id: ConnectionId, remote: SocketAddrV6, opts: ConnectionOpts, ) -> Result { - debug_assert!(!local_id.is_empty()); + let local_id = self.new_cid(); let (tls, client_config) = match opts { ConnectionOpts::Client(config) => ( TlsSession::new_client( @@ -432,7 +441,9 @@ impl Endpoint { debug!(self.log, "failed to authenticate initial packet"; "pn" => packet_number); return; }; - let loc_cid = self.new_cid(); + + // Local CID used for stateless packets + let temp_loc_cid = ConnectionId::random(&mut self.rng, self.config.local_cid_len); if self.incoming.len() + self.incoming_handshakes == self.server_config.as_ref().unwrap().accept_buffer as usize @@ -445,7 +456,7 @@ impl Endpoint { crypto, header_crypto, &src_cid, - &loc_cid, + &temp_loc_cid, 0, TransportError::SERVER_BUSY, ), @@ -482,7 +493,7 @@ impl Endpoint { ); let mut buf = Vec::new(); let header = Header::Retry { - src_cid: loc_cid, + src_cid: temp_loc_cid, dst_cid: src_cid, orig_dst_cid: dst_cid, }; @@ -503,7 +514,6 @@ impl Endpoint { let conn = self .add_connection( dst_cid, - loc_cid, src_cid, remote, ConnectionOpts::Server { @@ -528,7 +538,7 @@ impl Endpoint { self.io.push_back(Io::Transmit { destination: remote, ecn: None, - packet: handshake_close(crypto, header_crypto, &src_cid, &loc_cid, 0, e), + packet: handshake_close(crypto, header_crypto, &src_cid, &temp_loc_cid, 0, e), }); } } @@ -540,8 +550,9 @@ impl Endpoint { .remove(&self.connections[conn.0].init_cid); } if self.config.local_cid_len > 0 { - self.connection_ids - .remove(&self.connections[conn.0].loc_cid()); + for cid in self.connections[conn.0].loc_cids() { + self.connection_ids.remove(cid); + } } self.connection_remotes .remove(&self.connections[conn.0].remote); diff --git a/quinn-proto/src/frame.rs b/quinn-proto/src/frame.rs index 1f4aef86e..b375a5914 100644 --- a/quinn-proto/src/frame.rs +++ b/quinn-proto/src/frame.rs @@ -727,6 +727,16 @@ pub struct NewConnectionId { pub reset_token: [u8; 16], } +impl NewConnectionId { + pub fn encode(&self, out: &mut W) { + out.write(Type::NEW_CONNECTION_ID); + out.write_var(self.sequence); + out.write(self.id.len() as u8); + out.put_slice(&self.id); + out.put_slice(&self.reset_token); + } +} + #[cfg(test)] mod test { use super::*; diff --git a/quinn/examples/server.rs b/quinn/examples/server.rs index 213dc781b..64b5932fd 100644 --- a/quinn/examples/server.rs +++ b/quinn/examples/server.rs @@ -138,7 +138,7 @@ fn handle_connection(root: &PathBuf, log: &Logger, conn: quinn::NewConnection) { incoming, connection, } = conn; - let log = log.new(o!("local_id" => format!("{}", connection.local_id()))); + let log = log.clone(); info!(log, "got connection"; "remote_id" => %connection.remote_id(), "address" => %connection.remote_address(), diff --git a/quinn/src/lib.rs b/quinn/src/lib.rs index 93d0c437b..25e4c23bb 100644 --- a/quinn/src/lib.rs +++ b/quinn/src/lib.rs @@ -1030,16 +1030,19 @@ impl Connection { .into() } - /// The `ConnectionId` used for `conn` locally. - pub fn local_id(&self) -> ConnectionId { + /// The `ConnectionId`s defined for `conn` locally. + pub fn local_ids(&self) -> impl Iterator { self.0 .endpoint .borrow() .inner .connection(self.0.conn) - .loc_cid() + .loc_cids() + .cloned() + .collect::>() + .into_iter() } - /// The `ConnectionId` used for `conn` by the peer. + /// The `ConnectionId` defined for `conn` by the peer. pub fn remote_id(&self) -> ConnectionId { self.0 .endpoint