diff --git a/quinn-proto/src/connection.rs b/quinn-proto/src/connection.rs index 8c67d662f..3f14cfb8d 100644 --- a/quinn-proto/src/connection.rs +++ b/quinn-proto/src/connection.rs @@ -36,9 +36,6 @@ pub struct Connection { config: Arc, rng: OsRng, tls: TlsSession, - app_closed: bool, - /// DCID of Initial packet - pub(crate) init_cid: ConnectionId, loc_cids: HashMap, /// The CID we initially chose, for use during the handshake handshake_cid: ConnectionId, @@ -197,8 +194,6 @@ impl Connection { endpoint_config, rng, tls, - app_closed: false, - init_cid, loc_cids, handshake_cid: loc_cid, rem_cid, @@ -522,18 +517,16 @@ impl Connection { self.space_mut(space).pending_acks.subtract(&info.acks); } - pub fn timeout(&mut self, now: Instant, timer: Timer) -> bool { + pub fn timeout(&mut self, now: Instant, timer: Timer) { match timer { Timer::Close => { self.state = State::Drained; - return self.app_closed; } Timer::Idle => { self.close_common(now); self.io.timer_stop(Timer::Close); self.events.push_back(ConnectionError::TimedOut.into()); self.state = State::Drained; - return self.app_closed; } Timer::KeepAlive => { trace!(self.log, "sending keep-alive"); @@ -568,7 +561,6 @@ impl Connection { } } } - false } fn set_key_discard_timer(&mut self, now: Instant) { @@ -2423,7 +2415,6 @@ impl Connection { self.io.close = true; } - self.app_closed = true; match self.state { State::Handshake(_) | State::Established => { self.state = State::Closed(state::Closed { reason }); diff --git a/quinn-proto/src/endpoint.rs b/quinn-proto/src/endpoint.rs index 100ad4dec..133cc9e22 100644 --- a/quinn-proto/src/endpoint.rs +++ b/quinn-proto/src/endpoint.rs @@ -42,7 +42,7 @@ pub struct Endpoint { connection_ids_initial: FnvHashMap, connection_ids: FnvHashMap, connection_remotes: FnvHashMap, - pub(crate) connections: Slab, + pub(crate) connections: Slab, config: Arc, server_config: Option>, /// Connections that might have timer updates to apply perform @@ -86,7 +86,7 @@ impl Endpoint { /// Get an application-facing event pub fn poll(&mut self) -> Option<(ConnectionHandle, Event)> { while let Some(&ch) = self.eventful_conns.iter().next() { - if let Some(e) = self.connections[ch].poll() { + if let Some(e) = self.connections[ch].conn.poll() { return Some((ch, e)); } self.eventful_conns.remove(&ch); @@ -99,7 +99,7 @@ impl Endpoint { loop { let &ch = self.dirty_timers.iter().next()?; loop { - if let Some(io) = self.connections[ch].poll_io() { + if let Some(io) = self.connections[ch].conn.poll_io() { return Some(( ch, match io { @@ -108,7 +108,7 @@ impl Endpoint { self.connection_ids.remove(&connection_id); let new_cid = self.new_cid(); self.connection_ids.insert(new_cid, ch); - self.connections[ch].issue_cid(new_cid); + self.connections[ch].conn.issue_cid(new_cid); continue; } }, @@ -128,7 +128,7 @@ impl Endpoint { } loop { let &ch = self.needs_transmit.iter().next()?; - if let Some(transmit) = self.connections[ch].poll_transmit(now) { + if let Some(transmit) = self.connections[ch].conn.poll_transmit(now) { self.dirty_timers.insert(ch); return Some(transmit); } else { @@ -204,10 +204,13 @@ impl Endpoint { .cloned() }; if let Some(ch) = known_ch { - let had_1rtt = self.connections[ch].has_1rtt(); - self.connections[ch].handle_dgram(now, remote, ecn, partial_decode, rest); + let had_1rtt = self.connections[ch].conn.has_1rtt(); + self.connections[ch] + .conn + .handle_dgram(now, remote, ecn, partial_decode, rest); if !had_1rtt - && (self.connections[ch].has_1rtt() || !self.connections[ch].is_handshaking()) + && (self.connections[ch].conn.has_1rtt() + || !self.connections[ch].conn.is_handshaking()) { self.issue_identifiers(ch); } @@ -363,13 +366,13 @@ impl Endpoint { fn add_connection( &mut self, - initial_id: ConnectionId, - remote_id: ConnectionId, + init_cid: ConnectionId, + rem_cid: ConnectionId, remote: SocketAddr, transport_config: Arc, opts: ConnectionOpts, ) -> Result { - let local_id = self.new_cid(); + let loc_cid = self.new_cid(); let params = TransportParameters::new(&transport_config); let (tls, client_config) = match opts { ConnectionOpts::Client(config) => ( @@ -380,7 +383,7 @@ impl Endpoint { ), ConnectionOpts::Server { orig_dst_cid } => { let server_params = TransportParameters { - stateless_reset_token: Some(reset_token_for(&self.config.reset_key, &local_id)), + stateless_reset_token: Some(reset_token_for(&self.config.reset_key, &loc_cid)), original_connection_id: orig_dst_cid, ..params }; @@ -398,22 +401,27 @@ impl Endpoint { let remote_validated = self.server_config.as_ref().map_or(false, |cfg| { cfg.use_stateless_retry && client_config.is_none() }); - let id = self.connections.insert(Connection::new( - self.log.new(o!("connection" => local_id)), + let conn = Connection::new( + self.log.new(o!("connection" => loc_cid)), Arc::clone(&self.config), transport_config, - initial_id, - local_id, - remote_id, + init_cid, + loc_cid, + rem_cid, remote, client_config, tls, remote_validated, - )); + ); + let id = self.connections.insert(ConnectionMeta { + conn, + init_cid, + app_closed: false, + }); let ch = ConnectionHandle(id); if self.config.local_cid_len > 0 { - self.connection_ids.insert(local_id, ch); + self.connection_ids.insert(loc_cid, ch); } self.connection_remotes.insert(remote, ch); Ok(ch) @@ -550,7 +558,7 @@ impl Endpoint { if dst_cid.len() != 0 { self.connection_ids_initial.insert(dst_cid, ch); } - match self.connections[ch].handle_initial( + match self.connections[ch].conn.handle_initial( now, remote, ecn, @@ -562,7 +570,7 @@ impl Endpoint { trace!(self.log, "connection incoming; ICID {icid}", icid = dst_cid); self.incoming_handshakes += 1; self.needs_transmit.insert(ch); - if self.connections[ch].has_1rtt() { + if self.connections[ch].conn.has_1rtt() { self.issue_identifiers(ch); } Some(ch) @@ -582,7 +590,7 @@ impl Endpoint { /// Connection is either ready to accept data or failed. fn issue_identifiers(&mut self, ch: ConnectionHandle) { - if self.config.local_cid_len != 0 && !self.connections[ch].is_closed() { + if self.config.local_cid_len != 0 && !self.connections[ch].conn.is_closed() { /// Draft 17 ยง5.1.1: endpoints SHOULD provide and maintain at least eight /// connection IDs const LOCAL_CID_COUNT: usize = 8; @@ -590,23 +598,23 @@ impl Endpoint { for _ in 1..LOCAL_CID_COUNT { let cid = self.new_cid(); self.connection_ids.insert(cid, ch); - self.connections[ch].issue_cid(cid); + self.connections[ch].conn.issue_cid(cid); } } } fn forget(&mut self, ch: ConnectionHandle) { - if self.connections[ch].side().is_server() { + if self.connections[ch].conn.side().is_server() { self.connection_ids_initial .remove(&self.connections[ch].init_cid); } if self.config.local_cid_len > 0 { - for cid in self.connections[ch].loc_cids() { + for cid in self.connections[ch].conn.loc_cids() { self.connection_ids.remove(cid); } } self.connection_remotes - .remove(&self.connections[ch].remote()); + .remove(&self.connections[ch].conn.remote()); self.dirty_timers.remove(&ch); self.eventful_conns.remove(&ch); self.needs_transmit.remove(&ch); @@ -615,7 +623,8 @@ impl Endpoint { /// Handle a timer expiring pub fn timeout(&mut self, now: Instant, ch: ConnectionHandle, timer: Timer) { - if self.connections[ch].timeout(now, timer) { + self.connections[ch].conn.timeout(now, timer); + if self.connections[ch].app_closed { self.forget(ch); return; } @@ -643,7 +652,7 @@ impl Endpoint { stream: StreamId, data: &[u8], ) -> Result { - let result = self.connections[ch].write(stream, data); + let result = self.connections[ch].conn.write(stream, data); self.needs_transmit.insert(ch); result } @@ -656,7 +665,7 @@ impl Endpoint { /// # Panics /// - when applied to a stream that does not have an active outgoing channel pub fn finish(&mut self, ch: ConnectionHandle, stream: StreamId) { - self.connections[ch].finish(stream); + self.connections[ch].conn.finish(stream); self.needs_transmit.insert(ch); } @@ -674,7 +683,7 @@ impl Endpoint { buf: &mut [u8], ) -> Result { self.needs_transmit.insert(ch); // May need to send flow control frames after reading - self.connections[ch].read(stream, buf) + self.connections[ch].conn.read(stream, buf) } /// Read data from a stream out of order @@ -695,7 +704,7 @@ impl Endpoint { stream: StreamId, ) -> Result<(Bytes, u64), ReadError> { self.needs_transmit.insert(ch); // May need to send flow control frames after reading - self.connections[ch].read_unordered(stream) + self.connections[ch].conn.read_unordered(stream) } /// Abandon transmitting data on a stream @@ -703,7 +712,7 @@ impl Endpoint { /// # Panics /// - when applied to a receive stream or an unopened send stream pub fn reset(&mut self, ch: ConnectionHandle, stream: StreamId, error_code: u16) { - self.connections[ch].reset(stream, error_code); + self.connections[ch].conn.reset(stream, error_code); self.needs_transmit.insert(ch); } @@ -712,7 +721,7 @@ impl Endpoint { /// # Panics /// - when applied to a stream that has not begun receiving data pub fn stop_sending(&mut self, ch: ConnectionHandle, stream: StreamId, error_code: u16) { - self.connections[ch].stop_sending(stream, error_code); + self.connections[ch].conn.stop_sending(stream, error_code); self.needs_transmit.insert(ch); } @@ -721,14 +730,14 @@ impl Endpoint { /// Returns `None` if the maximum number of streams currently permitted by the remote endpoint /// are already open. pub fn open(&mut self, ch: ConnectionHandle, direction: Directionality) -> Option { - self.connections[ch].open(direction) + self.connections[ch].conn.open(direction) } /// Ping the remote endpoint /// /// Useful for preventing an otherwise idle connection from timing out. pub fn ping(&mut self, ch: ConnectionHandle) { - self.connections[ch].ping(); + self.connections[ch].conn.ping(); self.needs_transmit.insert(ch); } @@ -737,11 +746,12 @@ impl Endpoint { /// This does not ensure delivery of outstanding data. It is the application's responsibility /// to call this only when all important communications have been completed. pub fn close(&mut self, now: Instant, ch: ConnectionHandle, error_code: u16, reason: Bytes) { - if self.connections[ch].is_drained() { + if self.connections[ch].conn.is_drained() { self.forget(ch); return; } - self.connections[ch].close(now, error_code, reason); + self.connections[ch].conn.close(now, error_code, reason); + self.connections[ch].app_closed = true; self.needs_transmit.insert(ch); } @@ -755,22 +765,28 @@ impl Endpoint { } pub fn accept_stream(&mut self, ch: ConnectionHandle) -> Option { - let id = self.connections[ch].accept()?; + let id = self.connections[ch].conn.accept()?; self.needs_transmit.insert(ch); Some(id) } #[doc(hidden)] pub fn force_key_update(&mut self, ch: ConnectionHandle) { - self.connections[ch].force_key_update(); + self.connections[ch].conn.force_key_update(); self.ping(ch); } pub fn connection(&self, ch: ConnectionHandle) -> &Connection { - &self.connections[ch] + &self.connections[ch].conn } } +pub(crate) struct ConnectionMeta { + pub(crate) conn: Connection, + init_cid: ConnectionId, + app_closed: bool, +} + /// Parameters governing the core QUIC state machine /// /// This should be tuned to suit the application. In particular, window sizes for streams, stream @@ -1096,15 +1112,15 @@ impl From for usize { } } -impl Index for Slab { - type Output = Connection; - fn index(&self, ch: ConnectionHandle) -> &Connection { +impl Index for Slab { + type Output = ConnectionMeta; + fn index(&self, ch: ConnectionHandle) -> &ConnectionMeta { &self[ch.0] } } -impl IndexMut for Slab { - fn index_mut(&mut self, ch: ConnectionHandle) -> &mut Connection { +impl IndexMut for Slab { + fn index_mut(&mut self, ch: ConnectionHandle) -> &mut ConnectionMeta { &mut self[ch.0] } } diff --git a/quinn-proto/src/tests.rs b/quinn-proto/src/tests.rs index c371fbfd6..ec3bafd11 100644 --- a/quinn-proto/src/tests.rs +++ b/quinn-proto/src/tests.rs @@ -863,7 +863,7 @@ fn key_update() { Ok((ref data, 0)) if data == MSG1 ); - pair.client.connections[client_ch].force_key_update(); + pair.client.connections[client_ch].conn.force_key_update(); const MSG2: &[u8] = b"hello2"; pair.client.write(client_ch, s, MSG2).unwrap(); @@ -892,7 +892,7 @@ fn key_update_reordered() { assert!(!pair.client.outbound.is_empty()); pair.client.delay_outbound(); - pair.client.connections[client_ch].force_key_update(); + pair.client.connections[client_ch].conn.force_key_update(); info!(pair.log, "updated keys"); const MSG2: &[u8] = b"two";