Queue endpoint events with an intrusive list

This commit is contained in:
Benjamin Saunders
2023-08-20 15:11:18 -07:00
parent bdcfd3bb13
commit 580331c959
9 changed files with 447 additions and 58 deletions
+79 -14
View File
@@ -24,11 +24,15 @@ use crate::{
frame::{Close, Datagram, FrameStruct},
packet::{Header, LongType, Packet, PartialDecode, SpaceId},
range_set::ArrayRangeSet,
shared::{ConnectionEvent, ConnectionEventInner, ConnectionId, EcnCodepoint, EndpointEvents},
shared::{
ConnectionEvent, ConnectionEventInner, ConnectionId, EcnCodepoint, EndpointEvents,
EndpointEventsQueue,
},
token::ResetToken,
transport_parameters::TransportParameters,
Dir, EndpointConfig, Frame, Side, StreamId, Transmit, TransportError, TransportErrorCode,
VarInt, MAX_STREAM_COUNT, MIN_INITIAL_SIZE, TIMER_GRANULARITY,
ConnectionHandle, Dir, EndpointConfig, Frame, SharedList, Side, StreamId, Transmit,
TransportError, TransportErrorCode, VarInt, MAX_STREAM_COUNT, MIN_INITIAL_SIZE,
TIMER_GRANULARITY,
};
mod ack_frequency;
@@ -86,10 +90,10 @@ use timer::{Timer, TimerTable};
/// Protocol state and logic for a single QUIC connection
///
/// 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.
/// Objects of this type receive [`ConnectionEvent`]s and emit endpoint wakeups 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.
///
/// `Connection` has roughly 4 types of methods:
///
@@ -127,7 +131,7 @@ pub struct Connection {
endpoint_config: Arc<EndpointConfig>,
server_config: Option<Arc<ServerConfig>>,
config: Arc<TransportConfig>,
endpoint_events: Arc<EndpointEvents>,
endpoint_events: EndpointEventsTracker,
rng: StdRng,
crypto: Box<dyn crypto::Session>,
/// The CID we initially chose, for use during the handshake
@@ -239,7 +243,7 @@ impl Connection {
endpoint_config: Arc<EndpointConfig>,
server_config: Option<Arc<ServerConfig>>,
config: Arc<TransportConfig>,
endpoint_events: Arc<EndpointEvents>,
endpoint_events: EndpointEventsTracker,
init_cid: ConnectionId,
loc_cid: ConnectionId,
rem_cid: ConnectionId,
@@ -404,6 +408,12 @@ impl Connection {
None
}
/// Whether the endpoint must be woken to process events
#[must_use]
pub fn poll_endpoint_events(&mut self) -> bool {
mem::replace(&mut self.endpoint_events.notify_needed, false)
}
/// Provide control over streams
#[must_use]
pub fn streams(&mut self) -> Streams<'_> {
@@ -1033,7 +1043,10 @@ impl Connection {
match timer {
Timer::Close => {
self.state = State::Drained;
self.endpoint_events.drained.store(true, Ordering::Relaxed);
self.endpoint_events
.get()
.drained
.store(true, Ordering::Relaxed);
}
Timer::Idle => {
self.kill(ConnectionError::TimedOut);
@@ -1067,6 +1080,7 @@ impl Connection {
self.local_cid_state.retire_prior_to()
);
self.endpoint_events
.get()
.need_identifiers
.fetch_add(num_new_cid, Ordering::Relaxed);
}
@@ -2186,7 +2200,10 @@ impl Connection {
}
}
if !was_drained && self.state.is_drained() {
self.endpoint_events.drained.store(true, Ordering::Relaxed);
self.endpoint_events
.get()
.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);
@@ -2362,7 +2379,7 @@ impl Connection {
}
}
if let Some(token) = params.stateless_reset_token {
*self.endpoint_events.reset_token.lock().unwrap() =
*self.endpoint_events.get().reset_token.lock().unwrap() =
Some((self.path.remote, token));
}
self.handle_peer_params(params)?;
@@ -2682,10 +2699,12 @@ impl Connection {
.on_cid_retirement(sequence, self.peer_params.issue_cids_limit())?;
if allow_more_cids {
self.endpoint_events
.get()
.need_identifiers
.fetch_add(1, Ordering::Relaxed);
}
self.endpoint_events
.get()
.retire_cids
.lock()
.unwrap()
@@ -2904,7 +2923,8 @@ impl Connection {
}
fn set_reset_token(&mut self, reset_token: ResetToken) {
*self.endpoint_events.reset_token.lock().unwrap() = Some((self.path.remote, reset_token));
*self.endpoint_events.get().reset_token.lock().unwrap() =
Some((self.path.remote, reset_token));
self.peer_params.stateless_reset_token = Some(reset_token);
}
@@ -2917,6 +2937,7 @@ 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
.get()
.need_identifiers
.fetch_add(n, Ordering::Relaxed);
}
@@ -3398,6 +3419,7 @@ 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
.get()
.need_identifiers
.fetch_add(n, Ordering::Relaxed);
}
@@ -3440,7 +3462,10 @@ impl Connection {
self.close_common();
self.error = Some(reason);
self.state = State::Drained;
self.endpoint_events.drained.store(true, Ordering::Relaxed);
self.endpoint_events
.get()
.drained
.store(true, Ordering::Relaxed);
}
}
@@ -3668,3 +3693,43 @@ impl SentFrames {
&& self.retransmits.is_empty(streams)
}
}
/// Helper struct ensuring endpoint events get queued and generate notifications
pub(super) struct EndpointEventsTracker {
events: Arc<EndpointEvents>,
queue: Arc<SharedList<EndpointEvents, EndpointEventsQueue>>,
notify_needed: bool,
}
impl EndpointEventsTracker {
pub(super) fn new(
ch: ConnectionHandle,
queue: Arc<SharedList<EndpointEvents, EndpointEventsQueue>>,
) -> Self {
Self {
events: Arc::new(EndpointEvents::new(ch)),
queue,
notify_needed: false,
}
}
fn get(&mut self) -> EndpointEventsGuard<'_> {
EndpointEventsGuard(self)
}
}
struct EndpointEventsGuard<'a>(&'a mut EndpointEventsTracker);
impl std::ops::Deref for EndpointEventsGuard<'_> {
type Target = EndpointEvents;
fn deref(&self) -> &EndpointEvents {
&self.0.events
}
}
impl Drop for EndpointEventsGuard<'_> {
fn drop(&mut self) {
self.0.notify_needed |= self.0.queue.push(self.0.events.clone());
}
}
+43 -35
View File
@@ -19,17 +19,18 @@ use crate::{
cid_generator::{ConnectionIdGenerator, RandomConnectionIdGenerator},
coding::BufMutExt,
config::{ClientConfig, EndpointConfig, ServerConfig},
connection::{Connection, ConnectionError},
connection::{self, Connection, ConnectionError},
crypto::{self, Keys, UnsupportedVersion},
frame,
packet::{Header, Packet, PacketDecodeError, PacketNumber, PartialDecode},
shared::{
ConnectionEvent, ConnectionEventInner, ConnectionId, EcnCodepoint, EndpointEvents,
IssuedCid,
EndpointEventsQueue, IssuedCid,
},
shared_list,
transport_parameters::TransportParameters,
ResetToken, RetryToken, Side, Transmit, TransportConfig, TransportError, INITIAL_MTU,
MAX_CID_SIZE, MIN_INITIAL_SIZE, RESET_TOKEN_SIZE,
ResetToken, RetryToken, SharedList, Side, Transmit, TransportConfig, TransportError,
INITIAL_MTU, MAX_CID_SIZE, MIN_INITIAL_SIZE, RESET_TOKEN_SIZE,
};
/// The main entry point to the library
@@ -45,6 +46,10 @@ pub struct Endpoint {
server_config: Option<Arc<ServerConfig>>,
/// Whether the underlying UDP socket promises not to fragment packets
allow_mtud: bool,
/// Queue of endpoint events in need of processing
event_queue: Arc<SharedList<EndpointEvents, EndpointEventsQueue>>,
/// Partially-consumed iterator from `event_queue` to be emptied before fetching a new one
event_queue_iter: shared_list::Drain<EndpointEvents, EndpointEventsQueue>,
}
impl Endpoint {
@@ -66,6 +71,8 @@ impl Endpoint {
config,
server_config,
allow_mtud,
event_queue: Arc::<SharedList<_, _>>::default(),
event_queue_iter: shared_list::Drain::default(),
}
}
@@ -74,39 +81,38 @@ 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`. Must never be called concurrently with the same `ch`
/// Must be called until `None` is returned.
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;
}
let ch = ConnectionHandle(i);
if let Some(x) = self.handle_events_for(ch) {
return Some((ch, x));
// Consume events until one requires a response, starting at `self.event_queue_iter`, then
// refreshing it at most once from `self.events_queue`. Ensures the caller can handle each
// response without us dropping any events.
fn traverse(
this: &mut Endpoint,
iter: &mut shared_list::Drain<EndpointEvents, EndpointEventsQueue>,
) -> Option<(ConnectionHandle, ConnectionEvent)> {
for events in iter {
if let Some(x) = this.handle_events_for(&events) {
return Some((events.ch, x));
}
}
None
}
None
let mut iter = mem::take(&mut self.event_queue_iter);
if let Some(x) = traverse(self, &mut iter) {
self.event_queue_iter = iter;
return Some(x);
}
iter = self.event_queue.drain();
let result = traverse(self, &mut iter);
self.event_queue_iter = iter;
result
}
fn handle_events_for(&mut self, ch: ConnectionHandle) -> Option<ConnectionEvent> {
let events = self.connections[ch].events.clone();
fn handle_events_for(&mut self, events: &EndpointEvents) -> Option<ConnectionEvent> {
let ch = events.ch;
let needed_identifers = events.need_identifiers.swap(0, Ordering::Relaxed);
let result =
(needed_identifers > 0).then(|| self.send_new_identifiers(ch, needed_identifers));
@@ -136,6 +142,11 @@ impl Endpoint {
fn handle_drained(&mut self, ch: ConnectionHandle) {
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");
}
}
@@ -604,13 +615,13 @@ impl Endpoint {
server_config: Option<Arc<ServerConfig>>,
transport_config: Arc<TransportConfig>,
) -> Connection {
let events = Arc::<EndpointEvents>::default();
let events = connection::EndpointEventsTracker::new(ch, self.event_queue.clone());
let conn = Connection::new(
self.config.clone(),
server_config,
transport_config,
events.clone(),
events,
init_cid,
loc_cid,
rem_cid,
@@ -629,7 +640,6 @@ 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");
@@ -843,8 +853,6 @@ 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
+3
View File
@@ -67,6 +67,9 @@ pub use crate::endpoint::{ConnectError, ConnectionHandle, DatagramEvent, Endpoin
mod shared;
pub use crate::shared::{ConnectionEvent, ConnectionId, EcnCodepoint};
mod shared_list;
use crate::shared_list::SharedList;
mod transport_error;
pub use crate::transport_error::{Code as TransportErrorCode, Error as TransportError};
+27 -2
View File
@@ -10,7 +10,9 @@ use std::{
use bytes::{Buf, BufMut, BytesMut};
use crate::{coding::BufExt, packet::PartialDecode, ResetToken, MAX_CID_SIZE};
use crate::{
coding::BufExt, packet::PartialDecode, shared_list, ConnectionHandle, ResetToken, MAX_CID_SIZE,
};
/// Events sent from an Endpoint to a Connection
#[derive(Debug)]
@@ -30,12 +32,35 @@ pub(crate) enum ConnectionEventInner {
NewIdentifiers(Vec<IssuedCid>),
}
#[derive(Debug, Default)]
#[derive(Debug)]
pub(crate) struct EndpointEvents {
pub(crate) ch: ConnectionHandle,
pub(crate) need_identifiers: AtomicU64,
pub(crate) reset_token: Mutex<Option<(SocketAddr, ResetToken)>>,
pub(crate) retire_cids: Mutex<Vec<u64>>,
pub(crate) drained: AtomicBool,
link: shared_list::Link<EndpointEvents>,
}
impl EndpointEvents {
pub(crate) fn new(ch: ConnectionHandle) -> Self {
Self {
ch,
need_identifiers: AtomicU64::new(0),
reset_token: Mutex::new(None),
retire_cids: Mutex::new(Vec::new()),
drained: AtomicBool::new(false),
link: shared_list::Link::default(),
}
}
}
pub(crate) enum EndpointEventsQueue {}
impl shared_list::LinkGetter<EndpointEvents> for EndpointEventsQueue {
fn get(x: &EndpointEvents) -> &shared_list::Link<EndpointEvents> {
&x.link
}
}
/// Protocol-level identifier for a connection.
+248
View File
@@ -0,0 +1,248 @@
use std::{
marker::PhantomData,
ptr,
sync::{
atomic::{AtomicBool, AtomicPtr, Ordering},
Arc,
},
};
/// A thread-safe intrusive list of `Arc<T>`s singly linked via `Getter`
///
/// An `Arc<T>` may participate in any static number of lists by having that many distinct [`Link`]
/// fields. The `Getter` helper defines a single `get` method which selects the field associated
/// with a certain list.
pub(crate) struct SharedList<T, Getter>
where
Getter: LinkGetter<T>,
{
head: AtomicPtr<T>,
_marker: PhantomData<Getter>,
}
impl<T, Getter> SharedList<T, Getter>
where
Getter: LinkGetter<T>,
{
/// Ensure an entry is in the list
///
/// Does nothing if `entry` was already in the list or one of its `Drain` iterators. Otherwise,
/// adds it to the front of the list.
///
/// Returns whether the list was previously empty, in which case a consumer might need to be
/// notified to see the new entry.
pub(crate) fn push(&self, entry: Arc<T>) -> bool {
// `Acquire` synchronizes with the `Release` in `Drain::next` to ensure we don't clobber a
// `next` pointer that the iterator is still going to read.
if Getter::get(&*entry).linked.swap(true, Ordering::Acquire) {
// Already linked
return false;
}
let entry = Arc::into_raw(entry);
// `Release` ordering ensures the write to `next` is (and any preceding writes to the `T` in
// `entry` are) visible to anyone who `Acquire`s from `self.head`
let prev = self
.head
.fetch_update(Ordering::Release, Ordering::Relaxed, |head| {
// Safety: `entry` is trivially still valid here
Getter::get(unsafe { &*entry })
.next
.store(head, Ordering::Relaxed);
Some(entry.cast_mut())
})
// Lambda always returns `Some`
.unwrap();
prev.is_null()
}
/// Consume all entries in the queue
pub(crate) fn drain(&self) -> Drain<T, Getter> {
// `Acquire` ordering ensures visibility of the `next` pointers thanks to the `Release` in
// `push`.
let head = self.head.swap(ptr::null_mut(), Ordering::Acquire);
Drain {
// Safety: above swap means we uniquely own the underlying reference
next: (!head.is_null()).then(|| unsafe { Arc::from_raw(head) }),
_marker: PhantomData,
}
}
}
impl<T, Getter> Drop for SharedList<T, Getter>
where
Getter: LinkGetter<T>,
{
fn drop(&mut self) {
self.drain();
}
}
impl<T, Getter> Default for SharedList<T, Getter>
where
Getter: LinkGetter<T>,
{
fn default() -> Self {
Self {
head: AtomicPtr::new(ptr::null_mut()),
_marker: PhantomData,
}
}
}
/// Trait of helper ZSTs that select which intrusive list for a given `T` to traverse
pub(crate) trait LinkGetter<T>: Sized + 'static {
fn get(x: &T) -> &Link<T>;
}
/// A link in a [`SharedList`]
///
/// Each `Link<T>` field in a `T` allows an `Arc<T>` to participate in a distinct list.
#[derive(Debug)]
pub(crate) struct Link<T> {
next: AtomicPtr<T>,
/// Whether the link is participating in a list
///
/// `true` when reachable through [`Drain::next`] on any existing [`Drain`] iterator, or on one
/// newly constructed via [`SharedList::drain`].
///
/// This can be `true` when `next` is null when this is the last item in a list.
linked: AtomicBool,
}
impl<T> Default for Link<T> {
fn default() -> Self {
Self {
next: AtomicPtr::new(ptr::null_mut()),
linked: AtomicBool::new(false),
}
}
}
pub(crate) struct Drain<T, Getter>
where
Getter: LinkGetter<T>,
{
next: Option<Arc<T>>,
_marker: PhantomData<Getter>,
}
impl<T, Getter> Default for Drain<T, Getter>
where
Getter: LinkGetter<T>,
{
/// Construct an empty iterator
fn default() -> Self {
Self {
next: None,
_marker: PhantomData,
}
}
}
impl<T, Getter> Iterator for Drain<T, Getter>
where
Getter: LinkGetter<T>,
{
type Item = Arc<T>;
fn next(&mut self) -> Option<Arc<T>> {
let current = self.next.take()?;
let link = Getter::get(&*current);
let next = link.next.load(Ordering::Relaxed);
// `Release` synchronizes with the `Acquire` in `SharedList::push` to ensure the above read
// gets the current value of `next` before it's clobbered by another `push`, ensuring we
// don't leak the tail of the current list.
link.linked.store(false, Ordering::Release);
// Safety: The reference represented by `next` is uniquely owned by this iterator
self.next = (!next.is_null()).then(|| unsafe { Arc::from_raw(next) });
Some(current)
}
}
impl<T, Getter> std::iter::FusedIterator for Drain<T, Getter> where Getter: LinkGetter<T> {}
impl<T, Getter> Drop for Drain<T, Getter>
where
Getter: LinkGetter<T>,
{
fn drop(&mut self) {
// Recover and drop all remaining references
for _ in self.by_ref() {}
}
}
#[cfg(test)]
mod tests {
use super::*;
struct Item {
value: u32,
link: Link<Item>,
}
impl Item {
fn new(x: u32) -> Arc<Self> {
Arc::new(Self {
value: x,
link: Link::default(),
})
}
}
impl LinkGetter<Item> for () {
fn get(x: &Item) -> &Link<Item> {
&x.link
}
}
#[test]
fn insert_and_iterate() {
let list = SharedList::<Item, ()>::default();
assert!(list.push(Item::new(1)));
assert!(!list.push(Item::new(2)));
assert!(!list.push(Item::new(3)));
let mut iter = list.drain();
assert!(list.push(Item::new(4)));
assert_eq!(iter.next().unwrap().value, 3);
assert_eq!(iter.next().unwrap().value, 2);
assert_eq!(iter.next().unwrap().value, 1);
assert!(iter.next().is_none());
let mut iter = list.drain();
assert_eq!(iter.next().unwrap().value, 4);
assert!(iter.next().is_none());
let mut iter = list.drain();
assert!(iter.next().is_none());
}
#[test]
fn no_leaks() {
let list = SharedList::<Item, ()>::default();
let a = Item::new(1);
let b = Item::new(2);
list.push(a.clone());
list.push(b.clone());
assert_eq!(Arc::strong_count(&a), 2);
assert_eq!(Arc::strong_count(&b), 2);
drop(list);
assert_eq!(Arc::strong_count(&a), 1);
assert_eq!(Arc::strong_count(&b), 1);
}
#[test]
fn reinsert() {
let list = SharedList::<Item, ()>::default();
let a = Item::new(1);
let b = Item::new(2);
list.push(a.clone());
list.push(b);
list.push(a);
let mut iter = list.drain();
assert_eq!(iter.next().unwrap().value, 2);
assert_eq!(iter.next().unwrap().value, 1);
assert!(iter.next().is_none());
}
}
+3 -1
View File
@@ -352,6 +352,7 @@ impl TestEndpoint {
}
loop {
let mut endpoint_events_pending = false;
for conn in self.connections.values_mut() {
if self.timeout.map_or(false, |x| x <= now) {
self.timeout = None;
@@ -368,9 +369,10 @@ impl TestEndpoint {
self.outbound.extend(split_transmit(x));
}
self.timeout = conn.poll_timeout();
endpoint_events_pending |= conn.poll_endpoint_events();
}
if !self.has_endpoint_events() {
if !endpoint_events_pending {
break;
}
+1
View File
@@ -35,6 +35,7 @@ maintenance = { status = "experimental" }
[dependencies]
async-io = { version = "1.6", optional = true }
async-std = { version = "1.11", optional = true }
atomic-waker = "1.1.1"
bytes = "1"
# Enables futures::io::{AsyncRead, AsyncWrite} support for streams
futures-io = { version = "0.3.19", optional = true }
+9
View File
@@ -10,6 +10,7 @@ use std::{
};
use crate::runtime::{AsyncTimer, AsyncUdpSocket, Runtime};
use atomic_waker::AtomicWaker;
use bytes::Bytes;
use pin_project_lite::pin_project;
use proto::{ConnectionError, ConnectionHandle, ConnectionStats, Dir, StreamEvent, StreamId};
@@ -43,6 +44,7 @@ impl Connecting {
conn_events: mpsc::UnboundedReceiver<ConnectionEvent>,
socket: Arc<dyn AsyncUdpSocket>,
runtime: Arc<dyn Runtime>,
endpoint: Arc<AtomicWaker>,
) -> Self {
let (on_handshake_data_send, on_handshake_data_recv) = oneshot::channel();
let (on_connected_send, on_connected_recv) = oneshot::channel();
@@ -55,6 +57,7 @@ impl Connecting {
on_connected_send,
socket,
runtime.clone(),
endpoint,
);
runtime.spawn(Box::pin(
@@ -750,6 +753,7 @@ impl ConnectionRef {
on_connected: oneshot::Sender<bool>,
socket: Arc<dyn AsyncUdpSocket>,
runtime: Arc<dyn Runtime>,
endpoint: Arc<AtomicWaker>,
) -> Self {
Self(Arc::new(ConnectionInner {
state: Mutex::new(State {
@@ -771,6 +775,7 @@ impl ConnectionRef {
ref_count: 0,
socket,
runtime,
endpoint,
}),
shared: Shared::default(),
}))
@@ -849,6 +854,7 @@ pub(crate) struct State {
ref_count: usize,
socket: Arc<dyn AsyncUdpSocket>,
runtime: Arc<dyn Runtime>,
endpoint: Arc<AtomicWaker>,
}
impl State {
@@ -881,6 +887,9 @@ impl State {
}
fn forward_endpoint_events(&mut self) {
if self.inner.poll_endpoint_events() {
self.endpoint.wake();
}
if self.inner.is_drained() {
// If the endpoint driver is gone, noop.
let _ = self
+34 -6
View File
@@ -13,6 +13,7 @@ use std::{
};
use crate::runtime::{default_runtime, AsyncUdpSocket, Runtime};
use atomic_waker::AtomicWaker;
use bytes::{Bytes, BytesMut};
use pin_project_lite::pin_project;
use proto::{
@@ -188,9 +189,13 @@ impl Endpoint {
};
let (ch, conn) = endpoint.inner.connect(config, addr, server_name)?;
let socket = endpoint.socket.clone();
Ok(endpoint
.connections
.insert(ch, conn, socket, self.runtime.clone()))
Ok(endpoint.connections.insert(
ch,
conn,
socket,
self.runtime.clone(),
self.inner.shared.endpoint_events.clone(),
))
}
/// Switch to a new UDP socket
@@ -310,6 +315,7 @@ impl Future for EndpointDriver {
#[allow(unused_mut)] // MSRV
fn poll(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Self::Output> {
self.0.shared.endpoint_events.register(cx.waker());
let mut endpoint = self.0.state.lock().unwrap();
if endpoint.driver.is_none() {
endpoint.driver = Some(cx.waker().clone());
@@ -317,7 +323,7 @@ impl Future for EndpointDriver {
let now = Instant::now();
let mut keep_going = false;
keep_going |= endpoint.drive_recv(cx, now)?;
keep_going |= endpoint.drive_recv(cx, now, &self.0.shared)?;
keep_going |= endpoint.handle_events(cx, &self.0.shared);
keep_going |= endpoint.drive_send(cx)?;
@@ -382,10 +388,16 @@ pub(crate) struct State {
pub(crate) struct Shared {
incoming: Notify,
idle: Notify,
endpoint_events: Arc<AtomicWaker>,
}
impl State {
fn drive_recv<'a>(&'a mut self, cx: &mut Context, now: Instant) -> Result<bool, io::Error> {
fn drive_recv<'a>(
&'a mut self,
cx: &mut Context,
now: Instant,
shared: &Shared,
) -> Result<bool, io::Error> {
self.recv_limiter.start_cycle();
let mut metas = [RecvMeta::default(); BATCH_SIZE];
let mut iovs = MaybeUninit::<[IoSliceMut<'a>; BATCH_SIZE]>::uninit();
@@ -420,6 +432,7 @@ impl State {
conn,
self.socket.clone(),
self.runtime.clone(),
shared.endpoint_events.clone(),
);
self.incoming.push_back(conn);
}
@@ -538,6 +551,7 @@ impl State {
}
}
let mut n = 0;
while let Some((ch, event)) = self.inner.handle_events() {
// Ignoring errors from dropped connections that haven't yet been cleaned up
let _ = self
@@ -546,6 +560,10 @@ impl State {
.get_mut(&ch)
.unwrap()
.send(ConnectionEvent::Proto(event));
n += 1;
if n > IO_LOOP_BOUND {
return true;
}
}
keep_going
@@ -598,6 +616,7 @@ impl ConnectionSet {
conn: proto::Connection,
socket: Arc<dyn AsyncUdpSocket>,
runtime: Arc<dyn Runtime>,
endpoint: Arc<AtomicWaker>,
) -> Connecting {
let (send, recv) = mpsc::unbounded_channel();
if let Some((error_code, ref reason)) = self.close {
@@ -608,7 +627,15 @@ impl ConnectionSet {
.unwrap();
}
self.senders.insert(handle, send);
Connecting::new(handle, conn, self.sender.clone(), recv, socket, runtime)
Connecting::new(
handle,
conn,
self.sender.clone(),
recv,
socket,
runtime,
endpoint,
)
}
fn is_empty(&self) -> bool {
@@ -680,6 +707,7 @@ impl EndpointRef {
shared: Shared {
incoming: Notify::new(),
idle: Notify::new(),
endpoint_events: Arc::<AtomicWaker>::default(),
},
state: Mutex::new(State {
socket,