From 2e95861ce8e69a0d0af9000000053142b476ea24 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Philipp=20Kr=C3=BCger?= Date: Tue, 3 Jun 2025 16:11:44 +0200 Subject: [PATCH] Store a `Box` instead of `Arc`ing it, make `poll_recv` take `&mut self` --- quinn/src/endpoint.rs | 18 ++++++++---------- quinn/src/runtime.rs | 2 +- quinn/src/runtime/async_io.rs | 17 +++++++---------- quinn/src/runtime/tokio.rs | 15 ++++++--------- 4 files changed, 22 insertions(+), 30 deletions(-) diff --git a/quinn/src/endpoint.rs b/quinn/src/endpoint.rs index 879ba3645..43541805e 100644 --- a/quinn/src/endpoint.rs +++ b/quinn/src/endpoint.rs @@ -224,7 +224,7 @@ impl Endpoint { .inner .connect(self.runtime.now(), config, addr, server_name)?; - let sender = endpoint.socket.clone().create_sender(); + let sender = endpoint.socket.create_sender(); endpoint.stats.outgoing_handshakes += 1; Ok(endpoint .recv_state @@ -255,9 +255,7 @@ impl Endpoint { // Update connection socket references for sender in inner.recv_state.connections.senders.values() { // Ignoring errors from dropped connections - let _ = sender.send(ConnectionEvent::Rebind( - inner.socket.clone().create_sender(), - )); + let _ = sender.send(ConnectionEvent::Rebind(inner.socket.create_sender())); } if let Some(driver) = inner.driver.take() { // Ensure the driver can register for wake-ups from the new socket @@ -426,7 +424,7 @@ impl EndpointInner { { Ok((handle, conn)) => { state.stats.accepted_handshakes += 1; - let sender = state.socket.clone().create_sender(); + let sender = state.socket.create_sender(); let runtime = state.runtime.clone(); Ok(state .recv_state @@ -467,7 +465,7 @@ impl EndpointInner { #[derive(Debug)] pub(crate) struct State { - socket: Arc, + socket: Box, sender: Pin>, /// During an active migration, abandoned_socket receives traffic /// until the first packet arrives on the new socket. @@ -499,7 +497,7 @@ impl State { let poll_res = self.recv_state.poll_socket( cx, &mut self.inner, - &**socket, + &mut **socket, &mut self.sender, &*self.runtime, now, @@ -511,7 +509,7 @@ impl State { let poll_res = self.recv_state.poll_socket( cx, &mut self.inner, - &*self.socket, + &mut *self.socket, &mut self.sender, &*self.runtime, now, @@ -719,7 +717,7 @@ impl EndpointRef { ) -> Self { let (sender, events) = mpsc::unbounded_channel(); let recv_state = RecvState::new(sender, socket.max_receive_segments(), &inner); - let sender = socket.clone().create_sender(); + let sender = socket.create_sender(); Self(Arc::new(EndpointInner { shared: Shared { incoming: Notify::new(), @@ -809,7 +807,7 @@ impl RecvState { &mut self, cx: &mut Context, endpoint: &mut proto::Endpoint, - socket: &dyn AsyncUdpSocket, + socket: &mut dyn AsyncUdpSocket, sender: &mut Pin>, runtime: &dyn Runtime, now: Instant, diff --git a/quinn/src/runtime.rs b/quinn/src/runtime.rs index b23b7baf2..a6bb185ef 100644 --- a/quinn/src/runtime.rs +++ b/quinn/src/runtime.rs @@ -49,7 +49,7 @@ pub trait AsyncUdpSocket: Send + Sync + Debug + 'static { /// [`Waker`]. /// /// [`Waker`]: std::task::Waker - fn create_sender(self: Arc) -> Pin>; + fn create_sender(&self) -> Pin>; /// Receive UDP datagrams, or register to be woken if receiving may succeed in the future fn poll_recv( diff --git a/quinn/src/runtime/async_io.rs b/quinn/src/runtime/async_io.rs index 9d10043ca..eb5380665 100644 --- a/quinn/src/runtime/async_io.rs +++ b/quinn/src/runtime/async_io.rs @@ -55,7 +55,7 @@ impl AsyncTimer for Timer { } } -#[cfg(feature = "runtime-smol")] +#[cfg(any(feature = "runtime-smol"))] #[derive(Debug, Clone)] struct UdpSocket { io: Arc>, @@ -73,7 +73,7 @@ impl UdpSocket { } #[cfg(feature = "runtime-smol")] -impl UdpSenderHelperSocket for Arc { +impl UdpSenderHelperSocket for UdpSocket { fn max_transmit_segments(&self) -> usize { self.inner.max_gso_segments() } @@ -85,14 +85,11 @@ impl UdpSenderHelperSocket for Arc { #[cfg(feature = "runtime-smol")] impl AsyncUdpSocket for UdpSocket { - fn create_sender(self: Arc) -> Pin> { - Box::pin(UdpSenderHelper::new( - Arc::clone(&self), - |socket: &Arc| { - let socket = socket.clone(); - async move { socket.io.writable().await } - }, - )) + fn create_sender(&self) -> Pin> { + Box::pin(UdpSenderHelper::new(self.clone(), |socket: &Self| { + let socket = socket.clone(); + async move { socket.io.writable().await } + })) } fn poll_recv( diff --git a/quinn/src/runtime/tokio.rs b/quinn/src/runtime/tokio.rs index 980990934..3fa7a27b1 100644 --- a/quinn/src/runtime/tokio.rs +++ b/quinn/src/runtime/tokio.rs @@ -55,7 +55,7 @@ struct UdpSocket { inner: Arc, } -impl UdpSenderHelperSocket for Arc { +impl UdpSenderHelperSocket for UdpSocket { fn max_transmit_segments(&self) -> usize { self.inner.max_gso_segments() } @@ -68,14 +68,11 @@ impl UdpSenderHelperSocket for Arc { } impl AsyncUdpSocket for UdpSocket { - fn create_sender(self: Arc) -> Pin> { - Box::pin(UdpSenderHelper::new( - Arc::clone(&self), - |socket: &Arc| { - let socket = socket.clone(); - async move { socket.io.writable().await } - }, - )) + fn create_sender(&self) -> Pin> { + Box::pin(UdpSenderHelper::new(self.clone(), |socket: &Self| { + let socket = socket.clone(); + async move { socket.io.writable().await } + })) } fn poll_recv(