From 0a07eaba20890a89ea5bf332cc2a8a2e31ba05ef Mon Sep 17 00:00:00 2001 From: Dirkjan Ochtman Date: Mon, 25 Jan 2021 22:07:44 +0100 Subject: [PATCH] quinn-proto: let Assembler take responsibility for reads from stopped streams --- quinn-proto/src/connection/assembler.rs | 61 ++++++++++++++-------- quinn-proto/src/connection/streams/recv.rs | 25 ++++----- 2 files changed, 47 insertions(+), 39 deletions(-) diff --git a/quinn-proto/src/connection/assembler.rs b/quinn-proto/src/connection/assembler.rs index fcf70c3e6..fb384ea74 100644 --- a/quinn-proto/src/connection/assembler.rs +++ b/quinn-proto/src/connection/assembler.rs @@ -29,7 +29,11 @@ impl Assembler { Self::default() } - pub(crate) fn read_unordered(&mut self) -> Option<(u64, Bytes)> { + pub(crate) fn read_unordered(&mut self) -> Result, AssembleError> { + if self.is_stopped() { + return Err(AssembleError::UnknownStream); + } + if let State::Ordered = self.state { // Enter unordered mode let mut recvd = RangeSet::new(); @@ -39,15 +43,19 @@ impl Assembler { } self.state = State::Unordered { recvd }; } - let (n, data) = self.pop()?; - self.bytes_read += data.len() as u64; - Some((n, data)) + + Ok(self.pop().map(|(offset, data)| { + self.bytes_read += data.len() as u64; + (offset, data) + })) } // Get the the next ordered chunk - pub(crate) fn read(&mut self, max_length: usize) -> Result, IllegalOrderedRead> { - if let State::Unordered { .. } = self.state { - return Err(IllegalOrderedRead); + pub(crate) fn read(&mut self, max_length: usize) -> Result, AssembleError> { + if self.is_stopped() { + return Err(AssembleError::UnknownStream); + } else if let State::Unordered { .. } = self.state { + return Err(AssembleError::IllegalOrderedRead); } loop { @@ -237,7 +245,10 @@ impl Default for State { /// Error indicating that an ordered read was performed on a stream after an unordered read #[derive(Debug, Copy, Clone)] -pub struct IllegalOrderedRead; +pub enum AssembleError { + IllegalOrderedRead, + UnknownStream, +} #[cfg(test)] mod test { @@ -433,34 +444,34 @@ mod test { fn unordered_happy_path() { let mut x = Assembler::new(); x.insert(0, Bytes::from_static(b"abc")); - assert_eq!(x.read_unordered(), Some((0, Bytes::from_static(b"abc")))); - assert_eq!(x.read_unordered(), None); + assert_eq!(next_unordered(&mut x), (0, Bytes::from_static(b"abc"))); + assert_eq!(x.read_unordered().unwrap(), None); x.insert(3, Bytes::from_static(b"def")); - assert_eq!(x.read_unordered(), Some((3, Bytes::from_static(b"def")))); - assert_eq!(x.read_unordered(), None); + assert_eq!(next_unordered(&mut x), (3, Bytes::from_static(b"def"))); + assert_eq!(x.read_unordered().unwrap(), None); } #[test] fn unordered_dedup() { let mut x = Assembler::new(); x.insert(3, Bytes::from_static(b"def")); - assert_eq!(x.read_unordered(), Some((3, Bytes::from_static(b"def")))); - assert_eq!(x.read_unordered(), None); + assert_eq!(next_unordered(&mut x), (3, Bytes::from_static(b"def"))); + assert_eq!(x.read_unordered().unwrap(), None); x.insert(0, Bytes::from_static(b"a")); x.insert(0, Bytes::from_static(b"abcdefghi")); x.insert(0, Bytes::from_static(b"abcd")); - assert_eq!(x.read_unordered(), Some((0, Bytes::from_static(b"a")))); - assert_eq!(x.read_unordered(), Some((1, Bytes::from_static(b"bc")))); - assert_eq!(x.read_unordered(), Some((6, Bytes::from_static(b"ghi")))); - assert_eq!(x.read_unordered(), None); + assert_eq!(next_unordered(&mut x), (0, Bytes::from_static(b"a"))); + assert_eq!(next_unordered(&mut x), (1, Bytes::from_static(b"bc"))); + assert_eq!(next_unordered(&mut x), (6, Bytes::from_static(b"ghi"))); + assert_eq!(x.read_unordered().unwrap(), None); x.insert(8, Bytes::from_static(b"ijkl")); - assert_eq!(x.read_unordered(), Some((9, Bytes::from_static(b"jkl")))); - assert_eq!(x.read_unordered(), None); + assert_eq!(next_unordered(&mut x), (9, Bytes::from_static(b"jkl"))); + assert_eq!(x.read_unordered().unwrap(), None); x.insert(12, Bytes::from_static(b"mno")); - assert_eq!(x.read_unordered(), Some((12, Bytes::from_static(b"mno")))); - assert_eq!(x.read_unordered(), None); + assert_eq!(next_unordered(&mut x), (12, Bytes::from_static(b"mno"))); + assert_eq!(x.read_unordered().unwrap(), None); x.insert(2, Bytes::from_static(b"cde")); - assert_eq!(x.read_unordered(), None); + assert_eq!(x.read_unordered().unwrap(), None); } #[test] @@ -496,6 +507,10 @@ mod test { assert_eq!(x.read(usize::MAX).unwrap(), None); } + fn next_unordered(x: &mut Assembler) -> (u64, Bytes) { + x.read_unordered().unwrap().unwrap() + } + fn next(x: &mut Assembler, size: usize) -> Option { x.read(size).unwrap() } diff --git a/quinn-proto/src/connection/streams/recv.rs b/quinn-proto/src/connection/streams/recv.rs index bddce94fa..db1b83c35 100644 --- a/quinn-proto/src/connection/streams/recv.rs +++ b/quinn-proto/src/connection/streams/recv.rs @@ -2,7 +2,7 @@ use bytes::Bytes; use thiserror::Error; use tracing::debug; -use crate::connection::assembler::{Assembler, IllegalOrderedRead}; +use crate::connection::assembler::{AssembleError, Assembler}; use crate::{ frame::{self, ShouldTransmit}, TransportError, VarInt, @@ -62,11 +62,8 @@ impl Recv { } pub(super) fn read_unordered(&mut self) -> StreamReadResult<(Bytes, u64)> { - if self.assembler.is_stopped() { - return Err(ReadError::UnknownStream); - } // Return data we already have buffered, regardless of state - if let Some((offset, bytes)) = self.assembler.read_unordered() { + if let Some((offset, bytes)) = self.assembler.read_unordered()? { Ok(Some((bytes, offset))) } else { self.read_blocked().map(|()| None) @@ -74,10 +71,6 @@ impl Recv { } pub(super) fn read(&mut self, max_length: usize) -> StreamReadResult { - if self.assembler.is_stopped() { - return Err(ReadError::UnknownStream); - } - match self.assembler.read(max_length)? { Some(bytes) => Ok(Some(bytes)), None => self.read_blocked().map(|()| None), @@ -88,10 +81,6 @@ impl Recv { &mut self, chunks: &mut [Bytes], ) -> Result, ReadError> { - if self.assembler.is_stopped() { - return Err(ReadError::UnknownStream); - } - let mut out = ReadChunks { bufs: 0, read: 0 }; if chunks.is_empty() { return Ok(Some(out)); @@ -313,9 +302,13 @@ pub enum ReadError { IllegalOrderedRead, } -impl From for ReadError { - fn from(_: IllegalOrderedRead) -> Self { - ReadError::IllegalOrderedRead +impl From for ReadError { + fn from(e: AssembleError) -> Self { + use AssembleError::*; + match e { + IllegalOrderedRead => ReadError::IllegalOrderedRead, + UnknownStream => ReadError::UnknownStream, + } } }