Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions noq-proto/src/connection/assembler.rs
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,10 @@ impl Assembler {
self.data.clear();
}

pub(super) fn is_ordered(&self) -> bool {
self.state.is_ordered()
}

pub(super) fn ensure_ordering(&mut self, ordered: bool) -> Result<(), IllegalOrderedRead> {
if ordered && !self.state.is_ordered() {
return Err(IllegalOrderedRead);
Expand Down
19 changes: 18 additions & 1 deletion noq-proto/src/connection/streams/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@ use tracing::trace;
use super::spaces::Retransmits;
use crate::{
Dir, StreamId, VarInt,
connection::streams::state::{get_or_insert_recv, get_or_insert_send},
connection::streams::state::{StreamRecv, get_or_insert_recv, get_or_insert_send},
frame,
};

Expand Down Expand Up @@ -110,6 +110,23 @@ pub struct RecvStream<'a> {
}

impl RecvStream<'_> {
/// Whether this stream is still in ordered read mode.
///
/// A stream switches permanently to unordered mode when [`Self::read`] is called with
/// `ordered` set to `false`.
pub fn is_ordered(&self) -> Result<bool, ClosedStream> {
let Some(stream) = self.state.recv.get(&self.id) else {
return Err(ClosedStream { _private: () });
};
let Some(stream) = stream.as_ref().and_then(StreamRecv::as_open_recv) else {
return Ok(true);
};
if stream.stopped {
return Err(ClosedStream { _private: () });
}
Ok(stream.assembler.is_ordered())
}

/// Read from the given recv stream
///
/// `max_length` limits the maximum size of the returned `Bytes` value; passing `usize::MAX`
Expand Down
28 changes: 28 additions & 0 deletions noq-proto/src/connection/streams/state.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1194,6 +1194,7 @@ mod tests {
assert!(recv.stop(0u32.into()).is_err());
assert_eq!(recv.read(true).err(), Some(ReadableError::ClosedStream));
assert_eq!(recv.read(false).err(), Some(ReadableError::ClosedStream));
assert!(recv.is_ordered().is_err());

assert_eq!(client.local_max_data - initial_max, 32);
assert_eq!(
Expand All @@ -1214,6 +1215,33 @@ mod tests {
assert!(!client.recv.contains_key(&id));
}

#[test]
fn recv_stream_ordering_mode() {
let mut client = make(Side::Client);
let id = StreamId::new(Side::Server, Dir::Uni, 0);
let _ = client
.received(
frame::Stream {
id,
offset: 0,
fin: false,
data: Bytes::from_static(b"hello"),
},
5,
)
.unwrap();

let mut pending = Retransmits::default();
let mut recv = RecvStream {
id,
state: &mut client,
pending: &mut pending,
};
assert_eq!(recv.is_ordered(), Ok(true));
let _ = recv.read(false).unwrap().finalize();
assert_eq!(recv.is_ordered(), Ok(false));
}

#[test]
fn stopped_reset() {
let mut client = make(Side::Client);
Expand Down
27 changes: 17 additions & 10 deletions noq/src/recv_stream.rs
Original file line number Diff line number Diff line change
Expand Up @@ -200,6 +200,17 @@ impl RecvStream {
})
}

fn is_ordered(&self) -> Result<bool, ReadError> {
let mut conn = self.conn.lock_without_waking("RecvStream::is_ordered");
if self.is_0rtt {
conn.check_0rtt().map_err(|()| ReadError::ZeroRttRejected)?;
}
conn.inner
.recv_stream(self.stream)
.is_ordered()
.map_err(|_| ReadError::ClosedStream)
}

/// Reads the next segments of data.
///
/// Fills `bufs` with the segments of data beginning immediately after the last data yielded
Expand Down Expand Up @@ -252,13 +263,14 @@ impl RecvStream {
/// all data read. Uses unordered reads to be more efficient than using `AsyncRead` would
/// allow. `size_limit` should be set to limit worst-case memory use.
///
/// If unordered reads have already been made, the resulting buffer may have gaps containing
/// arbitrary data.
///
/// This operation is *not* cancel-safe.
/// This operation is *not* cancel-safe. If cancelled after it has begun reading, further read
/// operations on the stream return [`ReadError::ClosedStream`].
///
/// [`ReadToEndError::TooLong`]: crate::ReadToEndError::TooLong
pub async fn read_to_end(&mut self, size_limit: usize) -> Result<Vec<u8>, ReadToEndError> {
if !self.is_ordered()? {
return Err(ReadError::ClosedStream.into());
}
ReadToEnd {
stream: self,
size_limit,
Expand Down Expand Up @@ -391,12 +403,7 @@ impl RecvStream {
let mut recv = conn.inner.recv_stream(self.stream);
let mut chunks = recv.read(ordered).map_err(|e| match e {
ReadableError::ClosedStream => ReadError::ClosedStream,
ReadableError::IllegalOrderedRead => {
// We should never get here because the only way to do unordered reads is
// via UnorderedRecvStream, which allows only unordered reads. It is not
// possible to get a RecvStream from an UnorderedRecvStream.
unreachable!("ordered read after unordered read")
}
ReadableError::IllegalOrderedRead => ReadError::ClosedStream,
})?;
let status = read_fn(&mut chunks);
if chunks.finalize().should_transmit() {
Expand Down
49 changes: 49 additions & 0 deletions noq/src/tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1730,6 +1730,55 @@ async fn recv_stream_cancel_stop_drop() {
);
}

/// Dropping a pending `read_to_end` future makes subsequent reads fail because `read_to_end` uses
/// the unordered API internally.
#[tokio::test]
async fn recv_stream_cancel_read_to_end_then_ordered_read_is_closed() {
let _guard = subscribe();
let factory = EndpointFactory::new();
let server = factory.endpoint("server");
let server_addr = server.local_addr().unwrap();
let client = factory.endpoint("client");
let ordered_read_done = tokio::sync::SetOnce::new();

tokio::join!(
async {
let conn = server.accept().await.unwrap().await.unwrap();
let mut recv = conn.accept_uni().await.unwrap();
{
let fut = pin!(recv.read_to_end(usize::MAX));
let mut cx = Context::from_waker(Waker::noop());
assert!(fut.poll(&mut cx).is_pending());
}

{
let mut buf = [0; 1];
let fut = pin!(recv.read(&mut buf));
let mut cx = Context::from_waker(Waker::noop());
assert!(matches!(
fut.poll(&mut cx),
Poll::Ready(Err(crate::ReadError::ClosedStream))
));
}
assert_eq!(
recv.read_to_end(usize::MAX).await,
Err(crate::ReadToEndError::Read(crate::ReadError::ClosedStream))
);
ordered_read_done.set(()).unwrap();
},
async {
let conn = client
.connect(server_addr, "localhost")
.unwrap()
.await
.unwrap();
let mut send = conn.open_uni().await.unwrap();
send.write_all(b"hello").await.unwrap();
ordered_read_done.wait().await;
},
);
}

/// Regression test for an `active_connections` underflow panic in the endpoint driver.
///
/// `ConnectionSet::insert` used to only increment `active_connections` when the endpoint
Expand Down
Loading