use std::{
cmp,
fmt::Debug,
io,
task::{Context, Poll, ready},
};
use bytes::{Buf, Bytes, BytesMut};
use crate::{Error, StreamError, coding::*};
pub struct Reader<S: crate::transport::poll::RecvStream, V> {
stream: S,
buffer: BytesMut,
version: V,
}
const WAKE_BUDGET: usize = 64 * 1024;
impl<S: crate::transport::poll::RecvStream, V: StreamCodes> Reader<S, V> {
pub fn new(stream: S, version: V) -> Self {
Self {
stream,
buffer: Default::default(),
version,
}
}
pub fn poll_decode<T: Decode<V> + Debug>(&mut self, cx: &mut Context<'_>) -> Poll<Result<T, Error>>
where
V: Clone,
{
loop {
let mut cursor = io::Cursor::new(&self.buffer);
match T::decode(&mut cursor, self.version.clone()) {
Ok(msg) => {
self.buffer.advance(cursor.position() as usize);
return Poll::Ready(Ok(msg));
}
Err(DecodeError::Short) if !ready!(self.poll_read_more(cx))? => {
return Poll::Ready(Err(DecodeError::Short.into()));
}
Err(DecodeError::Short) => {}
Err(e) => return Poll::Ready(Err(e.into())),
}
}
}
pub async fn decode<T: Decode<V> + Debug>(&mut self) -> Result<T, Error>
where
V: Clone,
{
std::future::poll_fn(|cx| self.poll_decode(cx)).await
}
pub fn poll_decode_maybe<T: Decode<V> + Debug>(&mut self, cx: &mut Context<'_>) -> Poll<Result<Option<T>, Error>>
where
V: Clone,
{
if !ready!(self.poll_has_more(cx))? {
return Poll::Ready(Ok(None));
}
self.poll_decode(cx).map_ok(Some)
}
pub async fn decode_maybe<T: Decode<V> + Debug>(&mut self) -> Result<Option<T>, Error>
where
V: Clone,
{
std::future::poll_fn(|cx| self.poll_decode_maybe(cx)).await
}
pub fn poll_decode_peek<T: Decode<V> + Debug>(&mut self, cx: &mut Context<'_>) -> Poll<Result<T, Error>>
where
V: Clone,
{
loop {
let mut cursor = io::Cursor::new(&self.buffer);
match T::decode(&mut cursor, self.version.clone()) {
Ok(msg) => return Poll::Ready(Ok(msg)),
Err(DecodeError::Short) if !ready!(self.poll_read_more(cx))? => {
return Poll::Ready(Err(DecodeError::Short.into()));
}
Err(DecodeError::Short) => {}
Err(e) => return Poll::Ready(Err(e.into())),
}
}
}
pub async fn decode_peek<T: Decode<V> + Debug>(&mut self) -> Result<T, Error>
where
V: Clone,
{
std::future::poll_fn(|cx| self.poll_decode_peek(cx)).await
}
pub fn poll_decode_peek_maybe<T: Decode<V> + Debug>(
&mut self,
cx: &mut Context<'_>,
) -> Poll<Result<Option<T>, Error>>
where
V: Clone,
{
if !ready!(self.poll_has_more(cx))? {
return Poll::Ready(Ok(None));
}
self.poll_decode_peek(cx).map_ok(Some)
}
pub async fn decode_peek_maybe<T: Decode<V> + Debug>(&mut self) -> Result<Option<T>, Error>
where
V: Clone,
{
std::future::poll_fn(|cx| self.poll_decode_peek_maybe(cx)).await
}
pub fn poll_read_chunk(&mut self, cx: &mut Context<'_>, max: usize) -> Poll<Result<Option<Bytes>, Error>> {
if !self.buffer.is_empty() {
let n = cmp::min(self.buffer.len(), max);
return Poll::Ready(Ok(Some(self.buffer.split_to(n).freeze())));
}
self.stream
.poll_read_chunk(cx, max)
.map_err(|err| self.version.transport_error(err))
}
pub(crate) fn poll_read_frame(
&mut self,
cx: &mut Context<'_>,
frame: &mut crate::frame::ProducerOwned,
) -> Poll<Result<(), Error>> {
let mut owed = 0;
let result = loop {
if frame.remaining() == 0 {
break Ok(());
}
match self.poll_read_chunk(cx, frame.remaining()) {
Poll::Pending => {
if owed > 0 {
frame.notify();
}
return Poll::Pending;
}
Poll::Ready(Ok(Some(chunk))) if !chunk.is_empty() => {
owed += chunk.len();
if let Err(err) = frame.write(chunk) {
break Err(err);
}
if owed >= WAKE_BUDGET {
frame.notify();
owed = 0;
}
}
Poll::Ready(Ok(_)) => break Err(Error::WrongSize),
Poll::Ready(Err(err)) => break Err(err),
}
};
Poll::Ready(result)
}
pub fn poll_read_exact(&mut self, cx: &mut Context<'_>, size: usize) -> Poll<Result<Bytes, Error>> {
while self.buffer.len() < size {
if !ready!(self.poll_read_more(cx))? {
return Poll::Ready(Err(DecodeError::Short.into()));
}
}
Poll::Ready(Ok(self.buffer.split_to(size).freeze()))
}
pub async fn read_exact(&mut self, size: usize) -> Result<Bytes, Error> {
std::future::poll_fn(|cx| self.poll_read_exact(cx, size)).await
}
pub fn poll_closed(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
if ready!(self.poll_has_more(cx))? {
return Poll::Ready(Err(DecodeError::Short.into()));
}
Poll::Ready(Ok(()))
}
fn poll_has_more(&mut self, cx: &mut Context<'_>) -> Poll<Result<bool, Error>> {
if !self.buffer.is_empty() {
return Poll::Ready(Ok(true));
}
self.poll_read_more(cx)
}
fn poll_read_more(&mut self, cx: &mut Context<'_>) -> Poll<Result<bool, Error>> {
match ready!(self.stream.poll_read_buf(cx, &mut self.buffer)) {
Ok(Some(_)) => Poll::Ready(Ok(true)),
Ok(None) => Poll::Ready(Ok(false)),
Err(e) => Poll::Ready(Err(self.version.transport_error(e))),
}
}
pub fn abort(&mut self, err: impl Into<StreamError>) {
self.stream.stop(self.version.encode_stream_code(&err.into()));
}
pub fn with_version<V2>(self, version: V2) -> Reader<S, V2> {
Reader {
stream: self.stream,
buffer: self.buffer,
version,
}
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use futures::FutureExt;
use super::*;
use crate::StreamError;
#[derive(Default)]
struct StopLog {
stops: Vec<u32>,
}
impl web_transport_trait::poll::RecvStream for StopLog {
type Error = crate::lite::test_transport::SinkError;
fn poll_read(
&mut self,
_cx: &mut std::task::Context<'_>,
_dst: &mut [u8],
) -> std::task::Poll<Result<Option<usize>, Self::Error>> {
std::task::Poll::Ready(Ok(None))
}
fn stop(&mut self, code: u32) {
self.stops.push(code);
}
fn poll_closed(&mut self, _cx: &mut std::task::Context<'_>) -> std::task::Poll<Result<(), Self::Error>> {
std::task::Poll::Ready(Ok(()))
}
}
#[test]
fn abort_stops_with_a_stream_code() {
const VERSION: crate::lite::Version = crate::lite::Version::Lite05;
for (err, expected) in [
(Error::Cancel, StreamError::Cancel.to_code()),
(Error::Lagged, StreamError::TooFarBehind.to_code()),
(
Error::Unauthorized,
StreamError::Session(crate::SessionError::Unauthorized).to_code(),
),
] {
let mut reader = Reader::new(StopLog::default(), VERSION);
reader.abort(&err);
assert_eq!(reader.stream.stops, vec![expected], "{err:?} used the wrong registry");
}
assert_ne!(StreamError::Cancel.to_code(), crate::SessionError::Cancel.to_code());
}
#[test]
fn abort_stops_with_the_negotiated_protocols_code() {
let mut lite = Reader::new(StopLog::default(), crate::lite::Version::Lite05);
lite.abort(&Error::Old);
let mut ietf = Reader::new(StopLog::default(), crate::ietf::Version::Draft20);
ietf.abort(&Error::Old);
assert_eq!(lite.stream.stops, vec![StreamError::Old.to_code()]);
assert_eq!(ietf.stream.stops, vec![crate::ietf::error::INTERNAL_ERROR]);
assert_ne!(lite.stream.stops, ietf.stream.stops);
}
#[derive(Default)]
struct CountWaker(std::sync::atomic::AtomicUsize);
impl std::task::Wake for CountWaker {
fn wake(self: Arc<Self>) {
self.0.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
}
}
impl CountWaker {
fn count(&self) -> usize {
self.0.load(std::sync::atomic::Ordering::SeqCst)
}
}
fn parked_frame(
size: usize,
) -> (
crate::frame::ProducerOwned,
crate::frame::Consumer,
kio::Waiter,
Arc<CountWaker>,
) {
let mut group = crate::group::Info { sequence: 0 }.produce();
let mut consumer = group.consume();
let frame = group
.create_frame_owned(crate::frame::Info {
size: size as u64,
timestamp: crate::Timestamp::ZERO,
})
.unwrap();
let payload = consumer.next_frame().now_or_never().unwrap().unwrap().unwrap();
let wakes = Arc::new(CountWaker::default());
let waiter = kio::Waiter::new(wakes.clone().into());
(frame, payload, waiter, wakes)
}
#[derive(Default)]
struct Chunks(std::collections::VecDeque<&'static [u8]>);
impl web_transport_trait::poll::RecvStream for Chunks {
type Error = crate::lite::test_transport::SinkError;
fn poll_read(&mut self, _cx: &mut Context<'_>, dst: &mut [u8]) -> Poll<Result<Option<usize>, Self::Error>> {
let Some(chunk) = self.0.pop_front() else {
return Poll::Pending;
};
let n = chunk.len().min(dst.len());
dst[..n].copy_from_slice(&chunk[..n]);
if n < chunk.len() {
self.0.push_front(&chunk[n..]);
}
Poll::Ready(Ok(Some(n)))
}
fn stop(&mut self, _code: u32) {}
fn poll_closed(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Pending
}
}
#[test]
fn frame_payload_wakes_once_per_poll_turn() {
let (mut frame, mut payload, waiter, wakes) = parked_frame(9);
assert!(payload.poll_read_chunk(&waiter).is_pending());
let mut reader = Reader::new(Chunks([b"foo".as_slice(), b"bar"].into()), crate::lite::Version::Lite05);
let mut cx = Context::from_waker(std::task::Waker::noop());
assert!(reader.poll_read_frame(&mut cx, &mut frame).is_pending());
assert_eq!(wakes.count(), 1, "two chunks in one turn owe one wake");
let Poll::Ready(Ok(Some(chunk))) = payload.poll_read_chunk(&waiter) else {
panic!("the boundary wake did not publish both chunks");
};
assert_eq!(chunk, Bytes::from_static(b"foobar"));
assert!(payload.poll_read_chunk(&waiter).is_pending());
assert!(reader.poll_read_frame(&mut cx, &mut frame).is_pending());
assert_eq!(wakes.count(), 1, "an empty poll turn woke a consumer");
reader.stream.0.push_back(b"baz");
assert!(matches!(
reader.poll_read_frame(&mut cx, &mut frame),
Poll::Ready(Ok(()))
));
assert_eq!(wakes.count(), 1, "a full frame woke before it was committed");
frame.finish().unwrap();
assert_eq!(wakes.count(), 2);
}
struct Burst {
chunks: Chunks,
payload: crate::frame::Consumer,
waiter: kio::Waiter,
}
impl web_transport_trait::poll::RecvStream for Burst {
type Error = crate::lite::test_transport::SinkError;
fn poll_read(&mut self, cx: &mut Context<'_>, dst: &mut [u8]) -> Poll<Result<Option<usize>, Self::Error>> {
self.waiter = kio::Waiter::new(self.waiter.waker().clone());
while self.payload.poll_read_chunk(&self.waiter).is_ready() {}
self.chunks.poll_read(cx, dst)
}
fn stop(&mut self, _code: u32) {}
fn poll_closed(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Pending
}
}
#[test]
fn a_burst_past_the_budget_wakes_before_the_boundary() {
const CHUNK: &[u8] = &[0u8; 8 * 1024];
let chunks = 2 * WAKE_BUDGET / CHUNK.len();
let (mut frame, mut payload, waiter, wakes) = parked_frame(chunks * CHUNK.len() + 1);
assert!(payload.poll_read_chunk(&waiter).is_pending());
let mut reader = Reader::new(
Burst {
chunks: Chunks(std::iter::repeat_n(CHUNK, chunks).collect()),
payload,
waiter,
},
crate::lite::Version::Lite05,
);
let mut cx = Context::from_waker(std::task::Waker::noop());
assert!(reader.poll_read_frame(&mut cx, &mut frame).is_pending());
assert_eq!(wakes.count(), 2, "the burst was withheld past its budget");
}
}