Early return in RecvStream::drop()

use read_exact
This commit is contained in:
Perelyn
2026-01-09 18:47:18 +01:00
committed by Benjamin Saunders
parent 4b7a039499
commit 475b55bad7
2 changed files with 136 additions and 10 deletions
+17 -5
View File
@@ -505,6 +505,19 @@ impl tokio::io::AsyncRead for RecvStream {
impl Drop for RecvStream {
fn drop(&mut self) {
if self.all_data_read {
debug_assert!(
!self
.conn
.state
.lock("RecvStream:drop")
.blocked_readers
.contains_key(&self.stream),
"Stream {} should not have a blocked reader when all data read is true",
&self.stream
);
return;
}
let mut conn = self.conn.state.lock("RecvStream::drop");
// clean up any previously registered wakers
@@ -513,11 +526,10 @@ impl Drop for RecvStream {
if conn.error.is_some() || (self.is_0rtt && conn.check_0rtt().is_err()) {
return;
}
if !self.all_data_read {
// Ignore ClosedStream errors
let _ = conn.inner.recv_stream(self.stream).stop(0u32.into());
conn.wake();
}
// Ignore ClosedStream errors
let _ = conn.inner.recv_stream(self.stream).stop(0u32.into());
conn.wake();
}
}
+119 -5
View File
@@ -7,10 +7,15 @@ use rustls::crypto::ring::default_provider;
use std::{
convert::TryInto,
future::Future,
io,
net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, UdpSocket},
str,
sync::Arc,
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
task::{Context, Poll, RawWaker, RawWakerVTable, Waker},
};
use crate::runtime::TokioRuntime;
@@ -23,6 +28,7 @@ use rustls::{
pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer},
};
use tokio::runtime::{Builder, Runtime};
use tokio::time::{sleep, timeout};
use tracing::{error_span, info};
use tracing_futures::Instrument as _;
use tracing_subscriber::EnvFilter;
@@ -157,7 +163,7 @@ fn read_after_close() {
.unwrap()
.await
.expect("connect");
tokio::time::sleep(Duration::from_millis(100)).await;
sleep(Duration::from_millis(100)).await;
let mut stream = new_conn.accept_uni().await.expect("incoming streams");
let msg = stream.read_to_end(usize::MAX).await.expect("read_to_end");
assert_eq!(msg, MSG);
@@ -900,8 +906,7 @@ async fn stream_stopped() {
let stopped3 = stopped3.await;
assert_eq!(stopped3, Ok(Some(42u32.into())));
};
let client =
tokio::time::timeout(Duration::from_millis(100), client).instrument(error_span!("client"));
let client = timeout(Duration::from_millis(100), client).instrument(error_span!("client"));
let server = async move {
let conn = server.accept().await.unwrap().await.unwrap();
let mut stream = conn.accept_uni().await.unwrap();
@@ -929,7 +934,7 @@ async fn stream_stopped_2() {
)
.unwrap();
let send_stream = conn.open_uni().await.unwrap();
let stopped = tokio::time::timeout(Duration::from_millis(100), send_stream.stopped())
let stopped = timeout(Duration::from_millis(100), send_stream.stopped())
.instrument(error_span!("stopped"));
tokio::pin!(stopped);
// poll the future once so that the waker is registered.
@@ -944,3 +949,112 @@ async fn stream_stopped_2() {
let res = stopped.await;
assert_eq!(res, Ok(Ok(None)));
}
#[tokio::test]
async fn stream_drop_removes_blocked_reader() {
let _guard = subscribe();
for drop_stream in [false, true] {
let endpoint_factory = EndpointFactory::new();
let server = endpoint_factory.endpoint();
let server_address = server.local_addr().unwrap();
let client = endpoint_factory.endpoint();
let server_task = tokio::spawn(async move {
let conn = server.accept().await.unwrap().await.unwrap();
let mut stream = conn.accept_uni().await.unwrap();
// read "hello"
let mut buf = [0u8; 5];
stream.read_exact(&mut buf).await.unwrap();
let (waker, wake_counter) = new_count_waker();
let mut cx = Context::from_waker(&waker);
// do a blocking read which will add the stream in conn.blocked_readers
{
let mut buf = [0u8; 64];
let read_fut = stream.read(&mut buf);
tokio::pin!(read_fut);
assert!(matches!(read_fut.as_mut().poll(&mut cx), Poll::Pending));
}
if !drop_stream {
assert_eq!(wake_counter.wakes(), 0);
// We have a blocked reader, closing the connection should wake it. We use this as
// a proxy to assert that the stream is in conn.blocked_readers.
conn.close(0u32.into(), b"done");
assert_eq!(wake_counter.wakes(), 1);
} else {
// dropping the stream should remove it from conn.blocked_readers, so we don't
// expect any wakeups
drop(stream);
assert_eq!(wake_counter.wakes(), 0, "no wakeups should have occurred");
conn.close(0u32.into(), b"done");
assert_eq!(wake_counter.wakes(), 0, "no wakeups should have occurred");
}
});
let conn = client
.connect(server_address, "localhost")
.unwrap()
.await
.unwrap();
let mut stream = conn.open_uni().await.unwrap();
// need to send some data to actually start the stream
stream.write_all(b"hello").await.unwrap();
server_task.await.unwrap();
}
}
#[derive(Default)]
struct WakeCounter {
wakes: AtomicUsize,
}
impl WakeCounter {
fn wakes(&self) -> usize {
self.wakes.load(Ordering::SeqCst)
}
}
fn new_count_waker() -> (Waker, Arc<WakeCounter>) {
// instance of WakeCounter
let counter = Arc::new(WakeCounter::default());
// convert
let waker = unsafe { Waker::from_raw(raw_waker(counter.clone())) };
(waker, counter)
}
fn raw_waker(counter: Arc<WakeCounter>) -> RawWaker {
// Store an Arc<WakeCounter> behind the raw pointer.
let ptr = Arc::into_raw(counter) as *const ();
RawWaker::new(ptr, &VTABLE)
}
static VTABLE: RawWakerVTable =
RawWakerVTable::new(clone_waker, wake_waker, wake_by_ref_waker, drop_waker);
unsafe fn clone_waker(data: *const ()) -> RawWaker {
let arc = Arc::<WakeCounter>::from_raw(data as *const WakeCounter);
let cloned = arc.clone();
std::mem::forget(arc);
raw_waker(cloned)
}
unsafe fn wake_waker(data: *const ()) {
let arc = Arc::<WakeCounter>::from_raw(data as *const WakeCounter);
arc.wakes.fetch_add(1, Ordering::SeqCst);
// arc drops here
}
unsafe fn wake_by_ref_waker(data: *const ()) {
let arc = Arc::<WakeCounter>::from_raw(data as *const WakeCounter);
arc.wakes.fetch_add(1, Ordering::SeqCst);
std::mem::forget(arc);
}
unsafe fn drop_waker(data: *const ()) {
drop(Arc::<WakeCounter>::from_raw(data as *const WakeCounter));
}