use std::task::{Context, Poll};
use bytes::{Buf, Bytes, BytesMut};
use moq_noq_proto::{StreamId, VarInt};
use super::super::Error;
use super::{End, Shared};
pub struct SendStream {
shared: Shared,
id: StreamId,
park: kio::Park,
fin: bool,
reset: bool,
}
impl SendStream {
pub(crate) fn new(shared: Shared, id: StreamId) -> Self {
shared.track(id);
Self {
shared,
id,
park: kio::Park::default(),
fin: false,
reset: false,
}
}
pub(crate) fn id(&self) -> u64 {
self.id.into()
}
pub(crate) fn ended(&self) -> bool {
self.fin || self.reset
}
pub(crate) fn try_write(&mut self, buf: &[u8]) -> usize {
if self.fin || self.reset {
return 0;
}
let n = self
.shared
.conn
.borrow_mut()
.send_stream(self.id)
.write(buf)
.unwrap_or(0);
if n > 0 {
self.shared.kick();
}
n
}
pub(crate) fn reset_code(&mut self, code: u64) {
if self.reset {
return;
}
let _ = self
.shared
.conn
.borrow_mut()
.send_stream(self.id)
.reset(VarInt::from_u64(code).unwrap_or(VarInt::MAX));
self.reset = true;
self.shared.kick();
}
}
impl web_transport_trait::poll::SendStream for SendStream {
type Error = Error;
fn poll_write(&mut self, cx: &mut Context<'_>, buf: &[u8]) -> Poll<Result<usize, Self::Error>> {
let waiter = self.park.hold(cx);
if self.fin || self.reset {
return Poll::Ready(Err(Error::Quic("stream already finished".to_string())));
}
if let Some(err) = self.shared.closed() {
return Poll::Ready(Err(err));
}
let result = self.shared.conn.borrow_mut().send_stream(self.id).write(buf);
match result {
Ok(n) => {
self.shared.kick();
Poll::Ready(Ok(n))
}
Err(moq_noq_proto::WriteError::Blocked) => {
self.shared.park_writable(self.id, waiter);
Poll::Pending
}
Err(moq_noq_proto::WriteError::Stopped(code)) => Poll::Ready(Err(Error::Stop(code.into_inner()))),
Err(moq_noq_proto::WriteError::ClosedStream) => {
Poll::Ready(Err(Error::Quic("stream already finished".to_string())))
}
}
}
fn set_priority(&mut self, order: u8) {
let _ = self
.shared
.conn
.borrow_mut()
.send_stream(self.id)
.set_priority(i32::from(order));
}
fn finish(&mut self) -> Result<(), Self::Error> {
if self.fin || self.reset {
return Ok(());
}
match self.shared.conn.borrow_mut().send_stream(self.id).finish() {
Ok(()) => {}
Err(moq_noq_proto::FinishError::Stopped(code)) => {
self.reset = true;
return Err(Error::Stop(code.into_inner()));
}
Err(moq_noq_proto::FinishError::ClosedStream) => {}
}
self.fin = true;
self.shared.kick();
Ok(())
}
fn reset(&mut self, code: u32) {
self.reset_code(u64::from(code));
}
fn poll_closed(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
let waiter = self.park.hold(cx);
if self.reset {
return Poll::Ready(Ok(()));
}
match self.shared.ended(self.id) {
Some(End::Stopped(code)) => Poll::Ready(Err(Error::Stop(code))),
Some(End::Delivered) => Poll::Ready(Ok(())),
None => match self.shared.closed() {
Some(err) => Poll::Ready(Err(err)),
None => {
self.shared.park_finishing(self.id, waiter);
Poll::Pending
}
},
}
}
}
impl Drop for SendStream {
fn drop(&mut self) {
self.shared.forget_send(self.id);
if !self.fin && !self.reset {
let _ = self
.shared
.conn
.borrow_mut()
.send_stream(self.id)
.reset(VarInt::from_u32(0));
self.shared.kick();
}
}
}
impl std::fmt::Debug for SendStream {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SendStream").field("id", &self.id).finish()
}
}
const READ_AHEAD: usize = 64 * 1024;
const READ_CHUNK: usize = 8 * 1024;
enum Read {
Chunk(Bytes),
Finished,
Blocked,
Reset(u64),
}
pub struct RecvStream {
shared: Shared,
id: StreamId,
park: kio::Park,
finished: bool,
stopped: bool,
backlog: BytesMut,
}
impl RecvStream {
pub(crate) fn new(shared: Shared, id: StreamId) -> Self {
Self {
shared,
id,
park: kio::Park::default(),
finished: false,
stopped: false,
backlog: BytesMut::new(),
}
}
pub(crate) fn ended(&self) -> bool {
self.finished || self.stopped
}
pub(crate) fn stop_code(&mut self, code: u64) {
self.backlog.clear();
if self.stopped || self.finished {
return;
}
let _ = self
.shared
.conn
.borrow_mut()
.recv_stream(self.id)
.stop(VarInt::from_u64(code).unwrap_or(VarInt::MAX));
self.stopped = true;
self.shared.kick();
}
fn read(shared: &Shared, id: StreamId, max: usize) -> Read {
let mut conn = shared.conn.borrow_mut();
let mut recv = conn.recv_stream(id);
let mut chunks = match recv.read(true) {
Ok(chunks) => chunks,
Err(_) => return Read::Finished,
};
let read = match chunks.next(max) {
Ok(Some(chunk)) => Read::Chunk(chunk.bytes),
Ok(None) => Read::Finished,
Err(moq_noq_proto::ReadError::Blocked) => Read::Blocked,
Err(moq_noq_proto::ReadError::Reset(code)) => Read::Reset(code.into_inner()),
};
let transmit = chunks.finalize().should_transmit();
drop(conn);
if transmit {
shared.kick();
}
read
}
fn drain(&mut self, dst: &mut [u8]) -> usize {
let n = dst.len().min(self.backlog.len());
dst[..n].copy_from_slice(&self.backlog[..n]);
self.backlog.advance(n);
if n > 0 {
self.shared.wake_readable(self.id);
}
n
}
}
impl web_transport_trait::poll::RecvStream for RecvStream {
type Error = Error;
fn poll_read(&mut self, cx: &mut Context<'_>, dst: &mut [u8]) -> Poll<Result<Option<usize>, Self::Error>> {
let waiter = self.park.hold(cx);
if dst.is_empty() {
return Poll::Ready(Ok(Some(0)));
}
if !self.backlog.is_empty() {
return Poll::Ready(Ok(Some(self.drain(dst))));
}
if self.finished {
return Poll::Ready(Ok(None));
}
match Self::read(&self.shared, self.id, dst.len()) {
Read::Chunk(bytes) => {
let n = bytes.len().min(dst.len());
dst[..n].copy_from_slice(&bytes[..n]);
Poll::Ready(Ok(Some(n)))
}
Read::Finished => {
self.finished = true;
Poll::Ready(Ok(None))
}
Read::Blocked => {
if let Some(err) = self.shared.closed() {
return Poll::Ready(Err(err));
}
self.shared.park_readable(self.id, waiter);
Poll::Pending
}
Read::Reset(code) => Poll::Ready(Err(Error::Reset(code))),
}
}
fn stop(&mut self, code: u32) {
self.stop_code(u64::from(code));
}
fn poll_closed(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
let waiter = self.park.hold(cx);
if self.finished || self.stopped {
return Poll::Ready(Ok(()));
}
loop {
if self.backlog.len() >= READ_AHEAD {
self.shared.park_readable(self.id, waiter);
return Poll::Pending;
}
match Self::read(&self.shared, self.id, READ_CHUNK) {
Read::Chunk(bytes) => self.backlog.extend_from_slice(&bytes),
Read::Finished => {
self.finished = true;
return Poll::Ready(Ok(()));
}
Read::Blocked => {
if let Some(err) = self.shared.closed() {
return Poll::Ready(Err(err));
}
self.shared.park_readable(self.id, waiter);
return Poll::Pending;
}
Read::Reset(code) => return Poll::Ready(Err(Error::Reset(code))),
}
}
}
}
impl Drop for RecvStream {
fn drop(&mut self) {
self.shared.forget_recv(self.id);
if !self.finished && !self.stopped {
let _ = self
.shared
.conn
.borrow_mut()
.recv_stream(self.id)
.stop(VarInt::from_u32(0));
self.shared.kick();
}
}
}
impl std::fmt::Debug for RecvStream {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RecvStream").field("id", &self.id).finish()
}
}