From 580331c959d0e73f5d399cdf730f23c88bf389f9 Mon Sep 17 00:00:00 2001 From: Benjamin Saunders Date: Sun, 20 Aug 2023 15:11:18 -0700 Subject: [PATCH] Queue endpoint events with an intrusive list --- quinn-proto/src/connection/mod.rs | 93 +++++++++-- quinn-proto/src/endpoint.rs | 78 +++++----- quinn-proto/src/lib.rs | 3 + quinn-proto/src/shared.rs | 29 +++- quinn-proto/src/shared_list.rs | 248 ++++++++++++++++++++++++++++++ quinn-proto/src/tests/util.rs | 4 +- quinn/Cargo.toml | 1 + quinn/src/connection.rs | 9 ++ quinn/src/endpoint.rs | 40 ++++- 9 files changed, 447 insertions(+), 58 deletions(-) create mode 100644 quinn-proto/src/shared_list.rs diff --git a/quinn-proto/src/connection/mod.rs b/quinn-proto/src/connection/mod.rs index 326bd7104..8797ce0da 100644 --- a/quinn-proto/src/connection/mod.rs +++ b/quinn-proto/src/connection/mod.rs @@ -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, server_config: Option>, config: Arc, - endpoint_events: Arc, + endpoint_events: EndpointEventsTracker, rng: StdRng, crypto: Box, /// The CID we initially chose, for use during the handshake @@ -239,7 +243,7 @@ impl Connection { endpoint_config: Arc, server_config: Option>, config: Arc, - endpoint_events: Arc, + 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, + queue: Arc>, + notify_needed: bool, +} + +impl EndpointEventsTracker { + pub(super) fn new( + ch: ConnectionHandle, + queue: Arc>, + ) -> 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()); + } +} diff --git a/quinn-proto/src/endpoint.rs b/quinn-proto/src/endpoint.rs index 197bc8b54..7bac6b1d2 100644 --- a/quinn-proto/src/endpoint.rs +++ b/quinn-proto/src/endpoint.rs @@ -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>, /// Whether the underlying UDP socket promises not to fragment packets allow_mtud: bool, + /// Queue of endpoint events in need of processing + event_queue: Arc>, + /// Partially-consumed iterator from `event_queue` to be emptied before fetching a new one + event_queue_iter: shared_list::Drain, } impl Endpoint { @@ -66,6 +71,8 @@ impl Endpoint { config, server_config, allow_mtud, + event_queue: Arc::>::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, + ) -> 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 { - let events = self.connections[ch].events.clone(); + fn handle_events_for(&mut self, events: &EndpointEvents) -> Option { + 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>, transport_config: Arc, ) -> Connection { - let events = Arc::::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, } /// Internal identifier for a `Connection` currently associated with an endpoint diff --git a/quinn-proto/src/lib.rs b/quinn-proto/src/lib.rs index 78b708bb9..b78b73d8a 100644 --- a/quinn-proto/src/lib.rs +++ b/quinn-proto/src/lib.rs @@ -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}; diff --git a/quinn-proto/src/shared.rs b/quinn-proto/src/shared.rs index 0b0d8f626..7a7aef60c 100644 --- a/quinn-proto/src/shared.rs +++ b/quinn-proto/src/shared.rs @@ -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), } -#[derive(Debug, Default)] +#[derive(Debug)] pub(crate) struct EndpointEvents { + pub(crate) ch: ConnectionHandle, pub(crate) need_identifiers: AtomicU64, pub(crate) reset_token: Mutex>, pub(crate) retire_cids: Mutex>, pub(crate) drained: AtomicBool, + link: shared_list::Link, +} + +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 for EndpointEventsQueue { + fn get(x: &EndpointEvents) -> &shared_list::Link { + &x.link + } } /// Protocol-level identifier for a connection. diff --git a/quinn-proto/src/shared_list.rs b/quinn-proto/src/shared_list.rs new file mode 100644 index 000000000..32d9ce1c5 --- /dev/null +++ b/quinn-proto/src/shared_list.rs @@ -0,0 +1,248 @@ +use std::{ + marker::PhantomData, + ptr, + sync::{ + atomic::{AtomicBool, AtomicPtr, Ordering}, + Arc, + }, +}; + +/// A thread-safe intrusive list of `Arc`s singly linked via `Getter` +/// +/// An `Arc` 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 +where + Getter: LinkGetter, +{ + head: AtomicPtr, + _marker: PhantomData, +} + +impl SharedList +where + Getter: LinkGetter, +{ + /// 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) -> 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 { + // `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 Drop for SharedList +where + Getter: LinkGetter, +{ + fn drop(&mut self) { + self.drain(); + } +} + +impl Default for SharedList +where + Getter: LinkGetter, +{ + 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: Sized + 'static { + fn get(x: &T) -> &Link; +} + +/// A link in a [`SharedList`] +/// +/// Each `Link` field in a `T` allows an `Arc` to participate in a distinct list. +#[derive(Debug)] +pub(crate) struct Link { + next: AtomicPtr, + /// 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 Default for Link { + fn default() -> Self { + Self { + next: AtomicPtr::new(ptr::null_mut()), + linked: AtomicBool::new(false), + } + } +} + +pub(crate) struct Drain +where + Getter: LinkGetter, +{ + next: Option>, + _marker: PhantomData, +} + +impl Default for Drain +where + Getter: LinkGetter, +{ + /// Construct an empty iterator + fn default() -> Self { + Self { + next: None, + _marker: PhantomData, + } + } +} + +impl Iterator for Drain +where + Getter: LinkGetter, +{ + type Item = Arc; + + fn next(&mut self) -> Option> { + 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 std::iter::FusedIterator for Drain where Getter: LinkGetter {} + +impl Drop for Drain +where + Getter: LinkGetter, +{ + 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, + } + + impl Item { + fn new(x: u32) -> Arc { + Arc::new(Self { + value: x, + link: Link::default(), + }) + } + } + + impl LinkGetter for () { + fn get(x: &Item) -> &Link { + &x.link + } + } + + #[test] + fn insert_and_iterate() { + let list = SharedList::::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::::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::::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()); + } +} diff --git a/quinn-proto/src/tests/util.rs b/quinn-proto/src/tests/util.rs index c0ae3c9ff..5289c9f30 100644 --- a/quinn-proto/src/tests/util.rs +++ b/quinn-proto/src/tests/util.rs @@ -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; } diff --git a/quinn/Cargo.toml b/quinn/Cargo.toml index 784e5e292..e7e76a728 100644 --- a/quinn/Cargo.toml +++ b/quinn/Cargo.toml @@ -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 } diff --git a/quinn/src/connection.rs b/quinn/src/connection.rs index c06c26aa9..770c51da3 100644 --- a/quinn/src/connection.rs +++ b/quinn/src/connection.rs @@ -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, socket: Arc, runtime: Arc, + endpoint: Arc, ) -> 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, socket: Arc, runtime: Arc, + endpoint: Arc, ) -> 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, runtime: Arc, + endpoint: Arc, } 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 diff --git a/quinn/src/endpoint.rs b/quinn/src/endpoint.rs index 0dad3e727..baf6ceb33 100644 --- a/quinn/src/endpoint.rs +++ b/quinn/src/endpoint.rs @@ -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.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, } impl State { - fn drive_recv<'a>(&'a mut self, cx: &mut Context, now: Instant) -> Result { + fn drive_recv<'a>( + &'a mut self, + cx: &mut Context, + now: Instant, + shared: &Shared, + ) -> Result { 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, runtime: Arc, + endpoint: Arc, ) -> 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::::default(), }, state: Mutex::new(State { socket,