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 4f5b27b195
commit bdcfd3bb13
8 changed files with 143 additions and 152 deletions
+31 -41
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,
@@ -89,10 +86,10 @@ use timer::{Timer, TimerTable};
/// Protocol state and logic for a single QUIC connection
///
/// Objects of this type receive [`ConnectionEvent`]s and emit [`EndpointEvent`]s and application
/// [`Event`]s to make progress. To handle timeouts, a `Connection` returns timer updates and
/// expects timeouts through various methods. A number of simple getter methods are exposed
/// to allow callers to inspect some of the connection state.
/// Objects of this type receive [`ConnectionEvent`]s and emit application [`Event`]s to make
/// progress. To handle timeouts, a `Connection` returns timer updates and expects timeouts through
/// various methods. A number of simple getter methods are exposed to allow callers to inspect some
/// of the connection state.
///
/// `Connection` has roughly 4 types of methods:
///
@@ -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,
@@ -313,7 +312,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)],
@@ -406,12 +404,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<'_> {
@@ -1017,13 +1009,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();
}
}
}
@@ -1047,7 +1033,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);
@@ -1081,7 +1067,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 => {
@@ -2199,7 +2186,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);
@@ -2375,8 +2362,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();
@@ -2693,11 +2680,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!(
@@ -2912,11 +2904,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);
}
@@ -2929,7 +2917,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(
@@ -3409,7 +3398,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
@@ -3450,7 +3440,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);
}
}
+68 -41
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},
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,
@@ -74,50 +74,71 @@ impl Endpoint {
self.server_config = 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.
self.connections.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(&mut self) -> Option<(ConnectionHandle, ConnectionEvent)> {
// HACKITY HACK: Scan all connections for events. Efficient work queue coming in followup.
let n = self.connections.capacity();
for i in 0..n {
if !self.connections.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 => {
if let Some(conn) = self.connections.try_remove(ch.0) {
self.index.remove(&conn);
} else {
// This indicates a bug in downstream code, which could cause spurious
// connection loss instead of this error if the CID was (re)allocated prior to
// the illegal call.
error!(id = ch.0, "unknown connection drained");
}
let ch = ConnectionHandle(i);
if let Some(x) = self.handle_events_for(ch) {
return Some((ch, x));
}
}
None
}
fn handle_events_for(&mut self, ch: ConnectionHandle) -> Option<ConnectionEvent> {
let events = self.connections[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[ch].reset_token.replace((remote, token));
if let Some(old) = old {
self.index.connection_reset_tokens.remove(old.0, old.1);
}
if self.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[ch].loc_cids.remove(&seq);
if let Some(cid) = cid {
trace!("peer retired CID {}: {}", seq, cid);
self.index.retire(&cid);
}
}
if events.drained.swap(false, Ordering::Relaxed) {
self.handle_drained(ch);
}
result
}
fn handle_drained(&mut self, ch: ConnectionHandle) {
if let Some(conn) = self.connections.try_remove(ch.0) {
self.index.remove(&conn);
}
}
/// Process an incoming UDP datagram
pub fn handle(
&mut self,
@@ -559,7 +580,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),
@@ -583,10 +604,13 @@ impl Endpoint {
server_config: Option<Arc<ServerConfig>>,
transport_config: Arc<TransportConfig>,
) -> 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,
@@ -605,6 +629,7 @@ impl Endpoint {
loc_cids: iter::once((0, loc_cid)).collect(),
addresses,
reset_token: None,
events,
});
debug_assert_eq!(id, ch.0, "connection handle allocation out of sync");
@@ -818,6 +843,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: Intrusive queue of dirty connections to drive reading these
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.
+5 -12
View File
@@ -352,8 +352,7 @@ impl TestEndpoint {
}
loop {
let mut endpoint_events: Vec<(ConnectionHandle, EndpointEvent)> = vec![];
for (ch, conn) in self.connections.iter_mut() {
for conn in self.connections.values_mut() {
if self.timeout.map_or(false, |x| x <= now) {
self.timeout = None;
conn.handle_timeout(now);
@@ -365,25 +364,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
@@ -881,11 +881,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));
}
}
@@ -1100,10 +1100,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),
}