use std::{
fmt::Debug,
task::{Context, Poll, ready},
};
use crate::{Error, StreamError, coding::*, ietf};
pub struct Writer<S: crate::transport::poll::SendStream, V> {
stream: Option<S>,
buffer: bytes::BytesMut,
version: V,
}
impl<S: crate::transport::poll::SendStream, V: StreamCodes> Writer<S, V> {
pub fn new(stream: S, version: V) -> Self {
Self {
stream: Some(stream),
buffer: Default::default(),
version,
}
}
pub fn buffer<T: Encode<V> + Debug>(&mut self, msg: &T) -> Result<(), Error>
where
V: Clone,
{
let start = self.buffer.len();
if let Err(err) = msg.encode(&mut self.buffer, self.version.clone()) {
self.buffer.truncate(start);
return Err(err.into());
}
Ok(())
}
pub fn buffer_raw(&mut self, bytes: &[u8]) {
self.buffer.extend_from_slice(bytes);
}
pub fn poll_flush(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
while !self.buffer.is_empty() {
ready!(self.stream.as_mut().unwrap().poll_write_buf(cx, &mut self.buffer))
.map_err(|err| self.version.transport_error(err))?;
}
Poll::Ready(Ok(()))
}
pub async fn encode<T: Encode<V> + Debug>(&mut self, msg: &T) -> Result<(), Error>
where
V: Clone,
{
self.buffer(msg)?;
std::future::poll_fn(|cx| self.poll_flush(cx)).await
}
pub fn poll_write<Buf: bytes::Buf>(&mut self, cx: &mut Context<'_>, buf: &mut Buf) -> Poll<Result<usize, Error>> {
ready!(self.poll_flush(cx))?;
self.stream
.as_mut()
.unwrap()
.poll_write_buf(cx, buf)
.map_err(|err| self.version.transport_error(err))
}
pub fn poll_write_all<Buf: bytes::Buf>(&mut self, cx: &mut Context<'_>, buf: &mut Buf) -> Poll<Result<(), Error>> {
while buf.has_remaining() {
ready!(self.poll_write(cx, buf))?;
}
Poll::Ready(Ok(()))
}
pub async fn write_all<Buf: bytes::Buf + Send>(&mut self, buf: &mut Buf) -> Result<(), Error> {
std::future::poll_fn(|cx| self.poll_write_all(cx, buf)).await
}
pub fn finish(&mut self) -> Result<(), Error> {
debug_assert!(self.buffer.is_empty(), "finish with unflushed bytes");
self.stream
.as_mut()
.unwrap()
.finish()
.map_err(|err| self.version.transport_error(err))
}
pub fn abort(mut self, err: &Error) {
if let Some(mut stream) = self.stream.take() {
stream.reset(self.version.encode_stream_code(&StreamError::from(err)));
}
}
pub async fn close(mut self) -> Result<(), Error> {
let Some(mut stream) = self.stream.take() else {
return Ok(());
};
stream.finish().map_err(|err| self.version.transport_error(err))?;
let version = &self.version;
std::future::poll_fn(|cx| stream.poll_closed(cx).map_err(|err| version.transport_error(err))).await
}
pub fn poll_closed(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
let version = &self.version;
self.stream
.as_mut()
.unwrap()
.poll_closed(cx)
.map_err(|err| version.transport_error(err))
}
pub fn poll_close(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
let res = ready!(self.poll_closed(cx));
self.stream = None;
Poll::Ready(res)
}
pub async fn closed(&mut self) -> Result<(), Error> {
std::future::poll_fn(|cx| self.poll_closed(cx)).await
}
pub fn set_priority(&mut self, send_order: u8) {
self.stream.as_mut().unwrap().set_priority(send_order);
}
pub fn with_version<O>(mut self, version: O) -> Writer<S, O> {
Writer {
stream: self.stream.take(),
buffer: std::mem::take(&mut self.buffer),
version,
}
}
}
impl<S: crate::transport::poll::SendStream> Writer<S, ietf::Version> {
pub async fn encode_message<T: ietf::Message>(&mut self, msg: &T) -> Result<(), Error> {
self.buffer(&T::ID)?;
self.encode(msg).await
}
}
impl<S: crate::transport::poll::SendStream, V> Drop for Writer<S, V> {
fn drop(&mut self) {
if let Some(mut stream) = self.stream.take() {
stream.reset(StreamError::Cancel.to_code());
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::lite::test_transport::{Log, SinkSend};
use std::task::Waker;
#[derive(Debug, Clone, thiserror::Error)]
#[error("stream stopped with {0}")]
struct Stopped(u32);
impl web_transport_trait::Error for Stopped {
fn session_error(&self) -> Option<(u32, String)> {
None
}
fn stream_error(&self) -> Option<u32> {
Some(self.0)
}
}
impl web_transport_trait::poll::SendStream for Stopped {
type Error = Self;
fn poll_write(&mut self, _: &mut Context<'_>, _: &[u8]) -> Poll<Result<usize, Self::Error>> {
Poll::Ready(Err(self.clone()))
}
fn set_priority(&mut self, _: u8) {}
fn finish(&mut self) -> Result<(), Self::Error> {
Err(self.clone())
}
fn reset(&mut self, _: u32) {}
fn poll_closed(&mut self, _: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Err(self.clone()))
}
}
#[tokio::test]
async fn raw_payload_stop_uses_the_negotiated_registry() {
for (version, expected) in [
(
crate::Version::Ietf(crate::ietf::Version::Draft17),
Error::Stream(crate::StreamError::Unknown(0x4)),
),
(
crate::Version::Ietf(crate::ietf::Version::Draft20),
Error::Stream(crate::StreamError::GoingAway),
),
(
crate::Version::Lite(crate::lite::Version::Lite05),
Error::Stream(crate::StreamError::GoingAway),
),
] {
let mut writer = Writer::new(Stopped(0x4), version);
let err = writer.write_all(&mut b"payload".as_slice()).await.unwrap_err();
assert!(
matches!(
(&err, &expected),
(
Error::Stream(crate::StreamError::Unknown(4)),
Error::Stream(crate::StreamError::Unknown(4))
) | (
Error::Stream(crate::StreamError::GoingAway),
Error::Stream(crate::StreamError::GoingAway)
)
),
"{version} decoded the STOP_SENDING with the wrong registry: {err:?}"
);
}
}
#[test]
fn set_priority_forwards_send_order() {
let log = Log::default();
let mut writer = Writer::new(SinkSend::new(log.clone()), crate::lite::Version::Lite05);
for send_order in 0u8..=255 {
writer.set_priority(send_order);
}
assert_eq!(log.priorities(), (0u8..=255).collect::<Vec<_>>());
}
#[derive(Debug)]
struct Poison;
impl Encode<crate::lite::Version> for Poison {
fn encode<W: bytes::BufMut>(&self, w: &mut W, _: crate::lite::Version) -> Result<(), EncodeError> {
w.put_slice(b"junk");
Err(EncodeError::BoundsExceeded)
}
}
#[test]
fn a_failed_encode_leaves_no_partial_bytes() {
let mut writer = Writer::new(SinkSend::new(Log::default()), crate::lite::Version::Lite05);
writer.buffer(&5u8).unwrap();
writer.buffer(&Poison).unwrap_err();
writer.buffer(&7u8).unwrap();
let log = writer.stream.as_ref().unwrap().log.clone();
let mut cx = std::task::Context::from_waker(Waker::noop());
assert!(writer.poll_flush(&mut cx).is_ready());
assert_eq!(*log.writes.lock().unwrap(), vec![5, 7], "the partial encode leaked");
}
#[test]
fn flush_resumes_after_pending() {
let gate = kio::Producer::new(false);
let mut writer = Writer::new(
SinkSend::gated(Log::default(), gate.consume()),
crate::lite::Version::Lite05,
);
let log = writer.stream.as_ref().unwrap().log.clone();
writer.buffer(&5u8).unwrap();
let mut cx = std::task::Context::from_waker(Waker::noop());
assert!(writer.poll_flush(&mut cx).is_pending());
assert!(log.writes.lock().unwrap().is_empty());
writer.buffer(&7u8).unwrap();
let Ok(mut open) = gate.write() else {
panic!("gate closed")
};
*open = true;
drop(open);
assert!(writer.poll_flush(&mut cx).is_ready());
assert_eq!(
*log.writes.lock().unwrap(),
vec![5, 7],
"the flush lost or reordered bytes"
);
}
}