From 45deb07bf76eb43f9f3a736f5dfe7d8b32b3cff5 Mon Sep 17 00:00:00 2001 From: Benjamin Saunders Date: Sun, 20 Aug 2023 14:41:54 -0700 Subject: [PATCH] Communicate endpoint events at proto layer via shared memory Simplifies the -proto API and reasoning about memory use --- quinn-proto/src/connection/mod.rs | 64 ++++----- quinn-proto/src/endpoint.rs | 222 +++++++++++++++++++----------- quinn-proto/src/lib.rs | 2 +- quinn-proto/src/shared.rs | 48 ++----- quinn-proto/src/tests/mod.rs | 6 +- quinn-proto/src/tests/util.rs | 15 +- quinn/src/connection.rs | 11 +- quinn/src/endpoint.rs | 34 ++--- quinn/src/lib.rs | 2 +- 9 files changed, 211 insertions(+), 193 deletions(-) diff --git a/quinn-proto/src/connection/mod.rs b/quinn-proto/src/connection/mod.rs index 3def1e52d..55a13e7fe 100644 --- a/quinn-proto/src/connection/mod.rs +++ b/quinn-proto/src/connection/mod.rs @@ -4,7 +4,7 @@ use std::{ convert::TryFrom, fmt, io, mem, net::{IpAddr, SocketAddr}, - sync::Arc, + sync::{atomic::Ordering, Arc}, time::{Duration, Instant}, }; @@ -24,10 +24,7 @@ use crate::{ frame::{Close, Datagram, FrameStruct}, packet::{Header, LongType, Packet, PartialDecode, SpaceId}, range_set::ArrayRangeSet, - shared::{ - ConnectionEvent, ConnectionEventInner, ConnectionId, EcnCodepoint, EndpointEvent, - EndpointEventInner, - }, + shared::{ConnectionEvent, ConnectionEventInner, ConnectionId, EcnCodepoint, EndpointEvents}, token::ResetToken, transport_parameters::TransportParameters, Dir, EndpointConfig, Frame, Side, StreamId, Transmit, TransportError, TransportErrorCode, @@ -130,6 +127,7 @@ pub struct Connection { endpoint_config: Arc, server_config: Option>, config: Arc, + endpoint_events: Arc, rng: StdRng, crypto: Box, /// The CID we initially chose, for use during the handshake @@ -162,7 +160,6 @@ pub struct Connection { /// Total number of outgoing packets that have been deemed lost lost_packets: u64, events: VecDeque, - endpoint_events: VecDeque, /// Whether the spin bit is in use for this connection spin_enabled: bool, /// Outgoing spin bit state @@ -242,6 +239,7 @@ impl Connection { endpoint_config: Arc, server_config: Option>, config: Arc, + endpoint_events: Arc, init_cid: ConnectionId, loc_cid: ConnectionId, rem_cid: ConnectionId, @@ -272,6 +270,7 @@ impl Connection { let mut this = Self { endpoint_config, server_config, + endpoint_events, crypto, handshake_cid: loc_cid, rem_handshake_cid: rem_cid, @@ -312,7 +311,6 @@ impl Connection { retry_src_cid: None, lost_packets: 0, events: VecDeque::new(), - endpoint_events: VecDeque::new(), spin_enabled: config.allow_spin && rng.gen_ratio(7, 8), spin: false, spaces: [initial_space, PacketSpace::new(now), PacketSpace::new(now)], @@ -399,12 +397,6 @@ impl Connection { None } - /// Return endpoint-facing events - #[must_use] - pub fn poll_endpoint_events(&mut self) -> Option { - self.endpoint_events.pop_front().map(EndpointEvent) - } - /// Provide control over streams #[must_use] pub fn streams(&mut self) -> Streams<'_> { @@ -1007,13 +999,7 @@ impl Connection { self.spaces[SpaceId::Data].pending.new_cids.push(frame); }); // Update Timer::PushNewCid - if self - .timers - .get(Timer::PushNewCid) - .map_or(true, |x| x <= now) - { - self.reset_cid_retirement(); - } + self.reset_cid_retirement(); } } } @@ -1037,7 +1023,7 @@ impl Connection { match timer { Timer::Close => { self.state = State::Drained; - self.endpoint_events.push_back(EndpointEventInner::Drained); + self.endpoint_events.drained.store(true, Ordering::Relaxed); } Timer::Idle => { self.kill(ConnectionError::TimedOut); @@ -1071,7 +1057,8 @@ impl Connection { self.local_cid_state.retire_prior_to() ); self.endpoint_events - .push_back(EndpointEventInner::NeedIdentifiers(num_new_cid)); + .need_identifiers + .fetch_add(num_new_cid, Ordering::Relaxed); } } Timer::MaxAckDelay => { @@ -2194,7 +2181,7 @@ impl Connection { } } if !was_drained && self.state.is_drained() { - self.endpoint_events.push_back(EndpointEventInner::Drained); + self.endpoint_events.drained.store(true, Ordering::Relaxed); // Close timer may have been started previously, e.g. if we sent a close and got a // stateless reset in response self.timers.stop(Timer::Close); @@ -2370,8 +2357,8 @@ impl Connection { } } if let Some(token) = params.stateless_reset_token { - self.endpoint_events - .push_back(EndpointEventInner::ResetToken(self.path.remote, token)); + *self.endpoint_events.reset_token.lock().unwrap() = + Some((self.path.remote, token)); } self.handle_peer_params(params)?; self.issue_first_cids(); @@ -2688,11 +2675,16 @@ impl Connection { let allow_more_cids = self .local_cid_state .on_cid_retirement(sequence, self.peer_params.issue_cids_limit())?; + if allow_more_cids { + self.endpoint_events + .need_identifiers + .fetch_add(1, Ordering::Relaxed); + } self.endpoint_events - .push_back(EndpointEventInner::RetireConnectionId( - sequence, - allow_more_cids, - )); + .retire_cids + .lock() + .unwrap() + .push(sequence); } Frame::NewConnectionId(frame) => { trace!( @@ -2906,11 +2898,7 @@ impl Connection { } fn set_reset_token(&mut self, reset_token: ResetToken) { - self.endpoint_events - .push_back(EndpointEventInner::ResetToken( - self.path.remote, - reset_token, - )); + *self.endpoint_events.reset_token.lock().unwrap() = Some((self.path.remote, reset_token)); self.peer_params.stateless_reset_token = Some(reset_token); } @@ -2923,7 +2911,8 @@ impl Connection { // Subtract 1 to account for the CID we supplied while handshaking let n = self.peer_params.issue_cids_limit() - 1; self.endpoint_events - .push_back(EndpointEventInner::NeedIdentifiers(n)); + .need_identifiers + .fetch_add(n, Ordering::Relaxed); } fn populate_packet( @@ -3403,7 +3392,8 @@ impl Connection { pub(crate) fn rotate_local_cid(&mut self, v: u64) { let n = self.local_cid_state.assign_retire_seq(v); self.endpoint_events - .push_back(EndpointEventInner::NeedIdentifiers(n)); + .need_identifiers + .fetch_add(n, Ordering::Relaxed); } /// Check the current active remote CID sequence @@ -3444,7 +3434,7 @@ impl Connection { self.close_common(); self.error = Some(reason); self.state = State::Drained; - self.endpoint_events.push_back(EndpointEventInner::Drained); + self.endpoint_events.drained.store(true, Ordering::Relaxed); } } diff --git a/quinn-proto/src/endpoint.rs b/quinn-proto/src/endpoint.rs index 1b4c01c43..560aaec46 100644 --- a/quinn-proto/src/endpoint.rs +++ b/quinn-proto/src/endpoint.rs @@ -1,10 +1,10 @@ use std::{ collections::{hash_map, HashMap}, convert::TryFrom, - fmt, iter, + fmt, iter, mem, net::{IpAddr, SocketAddr}, ops::{Index, IndexMut}, - sync::Arc, + sync::{atomic::Ordering, Arc, Mutex, RwLock}, time::{Instant, SystemTime}, }; @@ -24,8 +24,8 @@ use crate::{ frame, packet::{Header, Packet, PacketDecodeError, PacketNumber, PartialDecode}, shared::{ - ConnectionEvent, ConnectionEventInner, ConnectionId, EcnCodepoint, EndpointEvent, - EndpointEventInner, IssuedCid, + ConnectionEvent, ConnectionEventInner, ConnectionId, EcnCodepoint, EndpointEvents, + IssuedCid, }, transport_parameters::TransportParameters, ResetToken, RetryToken, Side, Transmit, TransportConfig, TransportError, INITIAL_MTU, @@ -37,11 +37,13 @@ use crate::{ /// This object performs no I/O whatsoever. Instead, it consumes incoming packets and /// connection-generated events via `handle` and `handle_event`. pub struct Endpoint { - index: ConnectionIndex, - connections: Slab, + index: RwLock, + /// Must be locked after `index` when locks overlap + connections: Mutex>, local_cid_generator: Box, config: Arc, - server_config: Option>, + /// Must never be locked concurrently with other locks + server_config: RwLock>>, /// Whether the underlying UDP socket promises not to fragment packets allow_mtud: bool, } @@ -58,61 +60,91 @@ impl Endpoint { allow_mtud: bool, ) -> Self { Self { - index: ConnectionIndex::default(), - connections: Slab::new(), + index: RwLock::new(ConnectionIndex::default()), + connections: Mutex::new(Slab::new()), local_cid_generator: (config.connection_id_generator_factory.as_ref())(), config, - server_config, + server_config: RwLock::new(server_config), allow_mtud, } } /// Replace the server configuration, affecting new incoming connections only - pub fn set_server_config(&mut self, server_config: Option>) { - self.server_config = server_config; + pub fn set_server_config(&self, server_config: Option>) { + *self.server_config.write().unwrap() = server_config; + } + + #[cfg(test)] + pub(crate) fn has_endpoint_events(&self) -> bool { + // HACKITY HACK: Scan all connections for events. Efficient work queue coming in followup. + let conns = self.connections.lock().unwrap(); + conns.iter().any(|(_, conn)| { + let e = &*conn.events; + e.need_identifiers.load(Ordering::Relaxed) > 0 + || e.reset_token.lock().unwrap().is_some() + || !e.retire_cids.lock().unwrap().is_empty() + || e.drained.load(Ordering::Relaxed) + }) } /// Process `EndpointEvent`s emitted from related `Connection`s /// - /// In turn, processing this event may return a `ConnectionEvent` for the same `Connection`. - pub fn handle_event( - &mut self, - ch: ConnectionHandle, - event: EndpointEvent, - ) -> Option { - use EndpointEventInner::*; - match event.0 { - NeedIdentifiers(n) => { - return Some(self.send_new_identifiers(ch, n)); + /// In turn, processing this event may return a `ConnectionEvent` for the same + /// `Connection`. Must never be called concurrently with the same `ch` + pub fn handle_events(&self) -> Option<(ConnectionHandle, ConnectionEvent)> { + // HACKITY HACK: Scan all connections for events. Efficient work queue coming in followup. + let n = self.connections.lock().unwrap().capacity(); + for i in 0..n { + if !self.connections.lock().unwrap().contains(i) { + continue; } - ResetToken(remote, token) => { - if let Some(old) = self.connections[ch].reset_token.replace((remote, token)) { - self.index.connection_reset_tokens.remove(old.0, old.1); - } - if self.index.connection_reset_tokens.insert(remote, token, ch) { - warn!("duplicate reset token"); - } - } - RetireConnectionId(seq, allow_more_cids) => { - if let Some(cid) = self.connections[ch].loc_cids.remove(&seq) { - trace!("peer retired CID {}: {}", seq, cid); - self.index.retire(&cid); - if allow_more_cids { - return Some(self.send_new_identifiers(ch, 1)); - } - } - } - Drained => { - let conn = self.connections.remove(ch.0); - self.index.remove(&conn); + let ch = ConnectionHandle(i); + if let Some(x) = self.handle_events_for(ch) { + return Some((ch, x)); } } None } + fn handle_events_for(&self, ch: ConnectionHandle) -> Option { + let events = self.connections.lock().unwrap()[ch].events.clone(); + let needed_identifers = events.need_identifiers.swap(0, Ordering::Relaxed); + let result = + (needed_identifers > 0).then(|| self.send_new_identifiers(ch, needed_identifers)); + if let Some((remote, token)) = events.reset_token.lock().unwrap().take() { + let old = self.connections.lock().unwrap()[ch] + .reset_token + .replace((remote, token)); + let mut index = self.index.write().unwrap(); + if let Some(old) = old { + index.connection_reset_tokens.remove(old.0, old.1); + } + if index.connection_reset_tokens.insert(remote, token, ch) { + warn!("duplicate reset token"); + } + } + let retire_cids = mem::take(&mut *events.retire_cids.lock().unwrap()); + for seq in retire_cids { + let cid = self.connections.lock().unwrap()[ch].loc_cids.remove(&seq); + if let Some(cid) = cid { + trace!("peer retired CID {}: {}", seq, cid); + self.index.write().unwrap().retire(&cid); + } + } + if events.drained.swap(false, Ordering::Relaxed) { + self.handle_drained(ch); + } + result + } + + fn handle_drained(&self, ch: ConnectionHandle) { + let conn = self.connections.lock().unwrap().remove(ch.0); + self.index.write().unwrap().remove(&conn); + } + /// Process an incoming UDP datagram pub fn handle( - &mut self, + &self, now: Instant, remote: SocketAddr, local_ip: Option, @@ -132,7 +164,7 @@ impl Endpoint { dst_cid, version, }) => { - if self.server_config.is_none() { + if self.server_config.read().unwrap().is_none() { debug!("dropping packet with unsupported version"); return None; } @@ -173,7 +205,7 @@ impl Endpoint { // let addresses = FourTuple { remote, local_ip }; - if let Some(ch) = self.index.get(&addresses, &first_decode) { + if let Some(ch) = self.index.read().unwrap().get(&addresses, &first_decode) { return Some(DatagramEvent::ConnectionEvent( ch, ConnectionEvent(ConnectionEventInner::Datagram { @@ -191,8 +223,8 @@ impl Endpoint { // let dst_cid = first_decode.dst_cid(); - let server_config = match &self.server_config { - Some(config) => config, + let server_config = match &*self.server_config.read().unwrap() { + Some(config) => config.clone(), None => { debug!("packet for unrecognized connection {}", dst_cid); return self @@ -303,7 +335,7 @@ impl Endpoint { /// Initiate a connection pub fn connect( - &mut self, + &self, config: ClientConfig, remote: SocketAddr, server_name: &str, @@ -321,8 +353,10 @@ impl Endpoint { let remote_id = RandomConnectionIdGenerator::new(MAX_CID_SIZE).generate_cid(); trace!(initial_dcid = %remote_id); - let ch = ConnectionHandle(self.connections.vacant_key()); - let loc_cid = self.new_cid(ch); + let mut index = self.index.write().unwrap(); + let mut connections = self.connections.lock().unwrap(); + let ch = ConnectionHandle(connections.vacant_key()); + let loc_cid = new_cid(&*self.local_cid_generator, &mut index, ch); let params = TransportParameters::new( &config.transport, &self.config, @@ -349,16 +383,18 @@ impl Endpoint { None, config.transport, ); - self.connections.insert(meta); - self.index.insert_conn(addresses, loc_cid, ch); + connections.insert(meta); + index.insert_conn(addresses, loc_cid, ch); Ok((ch, conn)) } - fn send_new_identifiers(&mut self, ch: ConnectionHandle, num: u64) -> ConnectionEvent { + fn send_new_identifiers(&self, ch: ConnectionHandle, num: u64) -> ConnectionEvent { let mut ids = vec![]; + let mut index = self.index.write().unwrap(); + let mut connections = self.connections.lock().unwrap(); for _ in 0..num { - let id = self.new_cid(ch); - let meta = &mut self.connections[ch]; + let id = new_cid(&*self.local_cid_generator, &mut index, ch); + let meta = &mut connections[ch]; meta.cids_issued += 1; let sequence = meta.cids_issued; meta.loc_cids.insert(sequence, id); @@ -371,20 +407,8 @@ impl Endpoint { ConnectionEvent(ConnectionEventInner::NewIdentifiers(ids)) } - /// Generate a connection ID for `ch` - fn new_cid(&mut self, ch: ConnectionHandle) -> ConnectionId { - loop { - let cid = self.local_cid_generator.generate_cid(); - if let hash_map::Entry::Vacant(e) = self.index.connection_ids.entry(cid) { - e.insert(ch); - break cid; - } - assert!(self.local_cid_generator.cid_len() > 0); - } - } - fn handle_first_packet( - &mut self, + &self, now: Instant, addresses: FourTuple, ecn: Option, @@ -420,9 +444,10 @@ impl Endpoint { return None; } - let server_config = self.server_config.as_ref().unwrap().clone(); + let server_config = self.server_config.read().unwrap().as_ref().unwrap().clone(); - if self.connections.len() >= server_config.concurrent_connections as usize || self.is_full() + if self.connections.lock().unwrap().len() >= server_config.concurrent_connections as usize + || self.is_full() { debug!("refusing connection"); return Some(DatagramEvent::Response(self.initial_close( @@ -516,8 +541,11 @@ impl Endpoint { (None, dst_cid) }; - let ch = ConnectionHandle(self.connections.vacant_key()); - let loc_cid = self.new_cid(ch); + let mut index = self.index.write().unwrap(); + let mut connections = self.connections.lock().unwrap(); + + let ch = ConnectionHandle(connections.vacant_key()); + let loc_cid = new_cid(&*self.local_cid_generator, &mut index, ch); let mut params = TransportParameters::new( &server_config.transport, &self.config, @@ -542,11 +570,13 @@ impl Endpoint { Some(server_config), transport_config, ); - self.connections.insert(meta); - self.index.insert_conn(addresses, loc_cid, ch); + connections.insert(meta); + index.insert_conn(addresses, loc_cid, ch); if dst_cid.len() != 0 { - self.index.insert_initial(dst_cid, ch); + index.insert_initial(dst_cid, ch); } + drop((connections, index)); + match conn.handle_first_packet(now, addresses.remote, ecn, packet_number, packet, rest) { Ok(()) => { trace!(id = ch.0, icid = %dst_cid, "connection incoming"); @@ -554,7 +584,7 @@ impl Endpoint { } Err(e) => { debug!("handshake failed: {}", e); - self.handle_event(ch, EndpointEvent(EndpointEventInner::Drained)); + self.handle_drained(ch); match e { ConnectionError::TransportError(e) => Some(DatagramEvent::Response( self.initial_close(version, addresses, crypto, &src_cid, e), @@ -577,10 +607,13 @@ impl Endpoint { server_config: Option>, transport_config: Arc, ) -> (ConnectionMeta, Connection) { + let events = Arc::::default(); + let conn = Connection::new( self.config.clone(), server_config, transport_config, + events.clone(), init_cid, loc_cid, rem_cid, @@ -599,6 +632,7 @@ impl Endpoint { loc_cids: iter::once((0, loc_cid)).collect(), addresses, reset_token: None, + events, }; (meta, conn) @@ -651,8 +685,8 @@ impl Endpoint { /// [`set_server_config`](Self::set_server_config) to update /// [`concurrent_connections`](ServerConfig::concurrent_connections) to /// zero. - pub fn reject_new_connections(&mut self) { - if let Some(config) = self.server_config.as_mut() { + pub fn reject_new_connections(&self) { + if let Some(config) = self.server_config.write().unwrap().as_mut() { Arc::make_mut(config).concurrent_connections(0); } } @@ -664,18 +698,19 @@ impl Endpoint { #[cfg(test)] pub(crate) fn known_connections(&self) -> usize { - let x = self.connections.len(); - debug_assert_eq!(x, self.index.connection_ids_initial.len()); + let x = self.connections.lock().unwrap().len(); + let index = self.index.read().unwrap(); + debug_assert_eq!(x, index.connection_ids_initial.len()); // Not all connections have known reset tokens - debug_assert!(x >= self.index.connection_reset_tokens.0.len()); + debug_assert!(x >= index.connection_reset_tokens.0.len()); // Not all connections have unique remotes, and 0-length CIDs might not be in use. - debug_assert!(x >= self.index.connection_remotes.len()); + debug_assert!(x >= index.connection_remotes.len()); x } #[cfg(test)] pub(crate) fn known_cids(&self) -> usize { - self.index.connection_ids.len() + self.index.read().unwrap().connection_ids.len() } /// Whether we've used up 3/4 of the available CID space @@ -686,7 +721,7 @@ impl Endpoint { self.local_cid_generator.cid_len() <= 4 && self.local_cid_generator.cid_len() != 0 && (2usize.pow(self.local_cid_generator.cid_len() as u32 * 8) - - self.index.connection_ids.len()) + - self.index.read().unwrap().connection_ids.len()) < 2usize.pow(self.local_cid_generator.cid_len() as u32 * 8 - 2) } } @@ -702,6 +737,23 @@ impl fmt::Debug for Endpoint { } } +// Standalone method for use with split borrows of `Endpoint` +/// Generate a connection ID for `ch` if specified +fn new_cid( + generator: &dyn ConnectionIdGenerator, + index: &mut ConnectionIndex, + ch: ConnectionHandle, +) -> ConnectionId { + loop { + let cid = generator.generate_cid(); + if let hash_map::Entry::Vacant(e) = index.connection_ids.entry(cid) { + e.insert(ch); + break cid; + } + assert!(generator.cid_len() > 0); + } +} + /// Maps packets to existing connections #[derive(Default, Debug)] struct ConnectionIndex { @@ -808,6 +860,8 @@ pub(crate) struct ConnectionMeta { /// Reset token provided by the peer for the CID we're currently sending to, and the address /// being sent to reset_token: Option<(SocketAddr, ResetToken)>, + // TODO: queue of dirty connections to drive reading these (ArcSwap based?) + events: Arc, } /// Internal identifier for a `Connection` currently associated with an endpoint diff --git a/quinn-proto/src/lib.rs b/quinn-proto/src/lib.rs index e3eb51d47..78b708bb9 100644 --- a/quinn-proto/src/lib.rs +++ b/quinn-proto/src/lib.rs @@ -65,7 +65,7 @@ mod endpoint; pub use crate::endpoint::{ConnectError, ConnectionHandle, DatagramEvent, Endpoint}; mod shared; -pub use crate::shared::{ConnectionEvent, ConnectionId, EcnCodepoint, EndpointEvent}; +pub use crate::shared::{ConnectionEvent, ConnectionId, EcnCodepoint}; mod transport_error; pub use crate::transport_error::{Code as TransportErrorCode, Error as TransportError}; diff --git a/quinn-proto/src/shared.rs b/quinn-proto/src/shared.rs index fd3634813..0b0d8f626 100644 --- a/quinn-proto/src/shared.rs +++ b/quinn-proto/src/shared.rs @@ -1,4 +1,12 @@ -use std::{fmt, net::SocketAddr, time::Instant}; +use std::{ + fmt, + net::SocketAddr, + sync::{ + atomic::{AtomicBool, AtomicU64}, + Mutex, + }, + time::Instant, +}; use bytes::{Buf, BufMut, BytesMut}; @@ -22,38 +30,12 @@ pub(crate) enum ConnectionEventInner { NewIdentifiers(Vec), } -/// Events sent from a Connection to an Endpoint -#[derive(Debug)] -pub struct EndpointEvent(pub(crate) EndpointEventInner); - -impl EndpointEvent { - /// Construct an event that indicating that a `Connection` will no longer emit events - /// - /// Useful for notifying an `Endpoint` that a `Connection` has been destroyed outside of the - /// usual state machine flow, e.g. when being dropped by the user. - pub fn drained() -> Self { - Self(EndpointEventInner::Drained) - } - - /// Determine whether this is the last event a `Connection` will emit - /// - /// Useful for determining when connection-related event loop state can be freed. - pub fn is_drained(&self) -> bool { - self.0 == EndpointEventInner::Drained - } -} - -#[derive(Clone, Debug, Eq, PartialEq)] -pub(crate) enum EndpointEventInner { - /// The connection has been drained - Drained, - /// The reset token and/or address eligible for generating resets has been updated - ResetToken(SocketAddr, ResetToken), - /// The connection needs connection identifiers - NeedIdentifiers(u64), - /// Stop routing connection ID for this sequence number to the connection - /// When `bool == true`, a new connection ID will be issued to peer - RetireConnectionId(u64, bool), +#[derive(Debug, Default)] +pub(crate) struct EndpointEvents { + pub(crate) need_identifiers: AtomicU64, + pub(crate) reset_token: Mutex>, + pub(crate) retire_cids: Mutex>, + pub(crate) drained: AtomicBool, } /// Protocol-level identifier for a connection. diff --git a/quinn-proto/src/tests/mod.rs b/quinn-proto/src/tests/mod.rs index 2e64e5782..276fbbcec 100644 --- a/quinn-proto/src/tests/mod.rs +++ b/quinn-proto/src/tests/mod.rs @@ -27,7 +27,7 @@ use util::*; fn version_negotiate_server() { let _guard = subscribe(); let client_addr = "[::2]:7890".parse().unwrap(); - let mut server = Endpoint::new(Default::default(), Some(Arc::new(server_config())), true); + let server = Endpoint::new(Default::default(), Some(Arc::new(server_config())), true); let now = Instant::now(); let event = server.handle( now, @@ -54,7 +54,7 @@ fn version_negotiate_client() { // packet let cid_generator_factory: fn() -> Box = || Box::new(RandomConnectionIdGenerator::new(0)); - let mut client = Endpoint::new( + let client = Endpoint::new( Arc::new(EndpointConfig { connection_id_generator_factory: Arc::new(cid_generator_factory), ..Default::default() @@ -1919,7 +1919,7 @@ fn big_cert_and_key() -> (rustls::Certificate, rustls::PrivateKey) { fn malformed_token_len() { let _guard = subscribe(); let client_addr = "[::2]:7890".parse().unwrap(); - let mut server = Endpoint::new(Default::default(), Some(Arc::new(server_config())), true); + let server = Endpoint::new(Default::default(), Some(Arc::new(server_config())), true); server.handle( Instant::now(), client_addr, diff --git a/quinn-proto/src/tests/util.rs b/quinn-proto/src/tests/util.rs index e73323c55..a3ba3bdde 100644 --- a/quinn-proto/src/tests/util.rs +++ b/quinn-proto/src/tests/util.rs @@ -344,7 +344,6 @@ impl TestEndpoint { } loop { - let mut endpoint_events: Vec<(ConnectionHandle, EndpointEvent)> = vec![]; for (ch, conn) in self.connections.iter_mut() { if self.timeout.map_or(false, |x| x <= now) { self.timeout = None; @@ -357,25 +356,19 @@ impl TestEndpoint { } } - while let Some(event) = conn.poll_endpoint_events() { - endpoint_events.push((*ch, event)); - } - while let Some(x) = conn.poll_transmit(now, MAX_DATAGRAMS) { self.outbound.extend(split_transmit(x)); } self.timeout = conn.poll_timeout(); } - if endpoint_events.is_empty() { + if !self.has_endpoint_events() { break; } - for (ch, event) in endpoint_events { - if let Some(event) = self.handle_event(ch, event) { - if let Some(conn) = self.connections.get_mut(&ch) { - conn.handle_event(event, now); - } + while let Some((ch, event)) = self.handle_events() { + if let Some(conn) = self.connections.get_mut(&ch) { + conn.handle_event(event, now); } } } diff --git a/quinn/src/connection.rs b/quinn/src/connection.rs index 21aaca90e..68056b035 100644 --- a/quinn/src/connection.rs +++ b/quinn/src/connection.rs @@ -879,11 +879,11 @@ impl State { } fn forward_endpoint_events(&mut self) { - while let Some(event) = self.inner.poll_endpoint_events() { + if self.inner.is_drained() { // If the endpoint driver is gone, noop. let _ = self .endpoint_events - .send((self.handle, EndpointEvent::Proto(event))); + .send((self.handle, EndpointEvent::Drained)); } } @@ -1098,10 +1098,9 @@ impl Drop for State { fn drop(&mut self) { if !self.inner.is_drained() { // Ensure the endpoint can tidy up - let _ = self.endpoint_events.send(( - self.handle, - EndpointEvent::Proto(proto::EndpointEvent::drained()), - )); + let _ = self + .endpoint_events + .send((self.handle, EndpointEvent::Drained)); } } } diff --git a/quinn/src/endpoint.rs b/quinn/src/endpoint.rs index f6252604f..0dad3e727 100644 --- a/quinn/src/endpoint.rs +++ b/quinn/src/endpoint.rs @@ -513,24 +513,14 @@ impl State { fn handle_events(&mut self, cx: &mut Context, shared: &Shared) -> bool { use EndpointEvent::*; + let mut keep_going = true; for _ in 0..IO_LOOP_BOUND { match self.events.poll_recv(cx) { Poll::Ready(Some((ch, event))) => match event { - Proto(e) => { - if e.is_drained() { - self.connections.senders.remove(&ch); - if self.connections.is_empty() { - shared.idle.notify_waiters(); - } - } - if let Some(event) = self.inner.handle_event(ch, e) { - // Ignoring errors from dropped connections that haven't yet been cleaned up - let _ = self - .connections - .senders - .get_mut(&ch) - .unwrap() - .send(ConnectionEvent::Proto(event)); + Drained => { + self.connections.senders.remove(&ch); + if self.connections.is_empty() { + shared.idle.notify_waiters(); } } Transmit(t) => { @@ -543,12 +533,22 @@ impl State { }, Poll::Ready(None) => unreachable!("EndpointInner owns one sender"), Poll::Pending => { - return false; + keep_going = false; } } } - true + while let Some((ch, event)) = self.inner.handle_events() { + // Ignoring errors from dropped connections that haven't yet been cleaned up + let _ = self + .connections + .senders + .get_mut(&ch) + .unwrap() + .send(ConnectionEvent::Proto(event)); + } + + keep_going } } diff --git a/quinn/src/lib.rs b/quinn/src/lib.rs index ad1da0114..0f528c610 100644 --- a/quinn/src/lib.rs +++ b/quinn/src/lib.rs @@ -97,7 +97,7 @@ enum ConnectionEvent { #[derive(Debug)] enum EndpointEvent { - Proto(proto::EndpointEvent), + Drained, Transmit(proto::Transmit), }