Communicate endpoint events at proto layer via shared memory

Simplifies the -proto API and reasoning about memory use
This commit is contained in:
Benjamin Saunders
2023-08-20 14:41:54 -07:00
parent 34dfec147e
commit 45deb07bf7
9 changed files with 211 additions and 193 deletions
+27 -37
View File
@@ -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<EndpointConfig>,
server_config: Option<Arc<ServerConfig>>,
config: Arc<TransportConfig>,
endpoint_events: Arc<EndpointEvents>,
rng: StdRng,
crypto: Box<dyn crypto::Session>,
/// 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<Event>,
endpoint_events: VecDeque<EndpointEventInner>,
/// 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<EndpointConfig>,
server_config: Option<Arc<ServerConfig>>,
config: Arc<TransportConfig>,
endpoint_events: Arc<EndpointEvents>,
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<EndpointEvent> {
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);
}
}
+138 -84
View File
@@ -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<ConnectionMeta>,
index: RwLock<ConnectionIndex>,
/// Must be locked after `index` when locks overlap
connections: Mutex<Slab<ConnectionMeta>>,
local_cid_generator: Box<dyn ConnectionIdGenerator>,
config: Arc<EndpointConfig>,
server_config: Option<Arc<ServerConfig>>,
/// Must never be locked concurrently with other locks
server_config: RwLock<Option<Arc<ServerConfig>>>,
/// 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<Arc<ServerConfig>>) {
self.server_config = server_config;
pub fn set_server_config(&self, server_config: Option<Arc<ServerConfig>>) {
*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<ConnectionEvent> {
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<ConnectionEvent> {
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<IpAddr>,
@@ -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<EcnCodepoint>,
@@ -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<Arc<ServerConfig>>,
transport_config: Arc<TransportConfig>,
) -> (ConnectionMeta, Connection) {
let events = Arc::<EndpointEvents>::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<EndpointEvents>,
}
/// Internal identifier for a `Connection` currently associated with an endpoint
+1 -1
View File
@@ -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};
+15 -33
View File
@@ -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<IssuedCid>),
}
/// 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<Option<(SocketAddr, ResetToken)>>,
pub(crate) retire_cids: Mutex<Vec<u64>>,
pub(crate) drained: AtomicBool,
}
/// Protocol-level identifier for a connection.
+3 -3
View File
@@ -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<dyn ConnectionIdGenerator> =
|| 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,
+4 -11
View File
@@ -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);
}
}
}
+5 -6
View File
@@ -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));
}
}
}
+17 -17
View File
@@ -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
}
}
+1 -1
View File
@@ -97,7 +97,7 @@ enum ConnectionEvent {
#[derive(Debug)]
enum EndpointEvent {
Proto(proto::EndpointEvent),
Drained,
Transmit(proto::Transmit),
}