diff --git a/noq-proto/src/connection/assembler.rs b/noq-proto/src/connection/assembler.rs index 41dd015672..ca0a8d9a06 100644 --- a/noq-proto/src/connection/assembler.rs +++ b/noq-proto/src/connection/assembler.rs @@ -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); diff --git a/noq-proto/src/connection/streams/mod.rs b/noq-proto/src/connection/streams/mod.rs index 1b4bcc4dc7..56fd178e45 100644 --- a/noq-proto/src/connection/streams/mod.rs +++ b/noq-proto/src/connection/streams/mod.rs @@ -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, }; @@ -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 { + 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` diff --git a/noq-proto/src/connection/streams/state.rs b/noq-proto/src/connection/streams/state.rs index 6eadee34f0..736a350e33 100644 --- a/noq-proto/src/connection/streams/state.rs +++ b/noq-proto/src/connection/streams/state.rs @@ -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!( @@ -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); diff --git a/noq/src/recv_stream.rs b/noq/src/recv_stream.rs index a5ba953918..26c1c878ab 100644 --- a/noq/src/recv_stream.rs +++ b/noq/src/recv_stream.rs @@ -200,6 +200,17 @@ impl RecvStream { }) } + fn is_ordered(&self) -> Result { + 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 @@ -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, ReadToEndError> { + if !self.is_ordered()? { + return Err(ReadError::ClosedStream.into()); + } ReadToEnd { stream: self, size_limit, @@ -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() { diff --git a/noq/src/tests.rs b/noq/src/tests.rs index adcd75682e..f4bc8e9afa 100755 --- a/noq/src/tests.rs +++ b/noq/src/tests.rs @@ -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