mirror of
https://github.com/n0-computer/noq.git
synced 2026-10-05 12:41:30 +00:00
Queue endpoint events with an intrusive list
This commit is contained in:
@@ -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
@@ -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
|
||||
|
||||
@@ -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};
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
|
||||
@@ -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 }
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user