use std::cell::{Cell, RefCell};
use std::future::poll_fn;
use std::rc::Rc;
use std::task::{Context, Poll};
use crate::constants;
use crate::core::{ConnectionId, StreamId, StreamRef};
use crate::error::{ConnectionLost, ReadError, WriteError};
use crate::packet::Handshake;
use super::shared::{ConnCell, ShellLink, close_now, now, release_waker_slot, wake_settled};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Ended {
Eof,
Reset(u64),
}
fn cached_id<S: Handshake>(
cell: &RefCell<ConnCell<S>>,
r: StreamRef,
cache: &Cell<Option<StreamId>>,
) -> Option<StreamId> {
if let Some(id) = cache.get() {
return Some(id);
}
let id = cell.borrow().core.as_ref()?.stream_id(r)?;
cache.set(Some(id));
Some(id)
}
pub struct SendStream<S: Handshake> {
shell: Rc<dyn ShellLink>,
cell: Rc<RefCell<ConnCell<S>>>,
conn: ConnectionId,
r: StreamRef,
key: u64,
ack_key: u64,
local_end: LocalEnd,
id: Cell<Option<StreamId>>,
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub(crate) enum LocalEnd {
Live,
Finished,
Reset(u64),
}
impl<S: Handshake> SendStream<S> {
pub(crate) fn install(
shell: Rc<dyn ShellLink>,
cell_rc: Rc<RefCell<ConnCell<S>>>,
cell: &mut ConnCell<S>,
conn: ConnectionId,
r: StreamRef,
) -> Self {
shell.acquire();
cell.handles += 1;
let key = cell.blocked_writers.entry(r).or_default().key();
let ack_key = cell.blocked_ackers.entry(r).or_default().key();
Self {
shell,
cell: cell_rc,
conn,
r,
key,
ack_key,
local_end: LocalEnd::Live,
id: Cell::new(cell.core.as_ref().and_then(|c| c.stream_id(r))),
}
}
pub fn id(&self) -> Option<StreamId> {
cached_id(&self.cell, self.r, &self.id)
}
pub async fn write(&mut self, buf: &[u8]) -> Result<usize, WriteError> {
poll_fn(|cx| self.poll_write(cx, buf)).await
}
pub(crate) fn poll_write(
&mut self,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<Result<usize, WriteError>> {
let (outcome, dirty) = {
let mut cell = self.cell.borrow_mut();
if let Some(code) = cell.peer_resets.get(&self.r).copied() {
return Poll::Ready(Err(WriteError::Reset(code)));
}
if self.local_end != LocalEnd::Live {
return Poll::Ready(Err(WriteError::Finished));
}
if let Some(lost) = cell.closed.clone() {
return Poll::Ready(Err(WriteError::ConnectionLost(lost)));
}
if buf.is_empty() {
return Poll::Ready(Ok(0));
}
let Some(core) = cell.core.as_mut() else {
return Poll::Ready(Err(WriteError::ConnectionLost(no_core())));
};
match core.write(now(), self.r, buf) {
Ok(0) => {
cell.blocked_writers
.entry(self.r)
.or_default()
.park(self.key, cx);
cell.dirty = true;
(Poll::Pending, true)
}
Ok(n) => {
cell.dirty = true;
(Poll::Ready(Ok(n)), true)
}
Err(e) => (Poll::Ready(Err(e)), false),
}
};
if dirty {
self.shell.mark_dirty(self.conn);
}
outcome
}
pub async fn finish(&mut self) -> Result<(), WriteError> {
poll_fn(|cx| self.poll_finish(cx)).await
}
pub(crate) fn poll_finish(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), WriteError>> {
let (outcome, dirty) = {
let mut cell = self.cell.borrow_mut();
if let Some(code) = cell.peer_resets.get(&self.r).copied() {
return Poll::Ready(Err(WriteError::Reset(code)));
}
match self.local_end {
LocalEnd::Finished => return Poll::Ready(Ok(())),
LocalEnd::Reset(_) => return Poll::Ready(Err(WriteError::Finished)),
LocalEnd::Live => {}
}
if let Some(lost) = cell.closed.clone() {
return Poll::Ready(Err(WriteError::ConnectionLost(lost)));
}
let Some(core) = cell.core.as_mut() else {
return Poll::Ready(Err(WriteError::ConnectionLost(no_core())));
};
match core.finish(now(), self.r) {
Ok(()) => {
cell.dirty = true;
(Ok(()), true)
}
Err(e) => (Err(e), false),
}
};
if outcome.is_ok() {
self.local_end = LocalEnd::Finished;
}
if dirty {
self.shell.mark_dirty(self.conn);
}
Poll::Ready(outcome)
}
pub async fn acked(&mut self) -> Result<(), WriteError> {
poll_fn(|cx| self.poll_acked(cx)).await
}
pub(crate) fn poll_acked(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), WriteError>> {
let mut cell = self.cell.borrow_mut();
if let LocalEnd::Reset(code) = self.local_end {
return Poll::Ready(Err(WriteError::Reset(code)));
}
if let Some(code) = cell.peer_resets.get(&self.r).copied() {
return Poll::Ready(Err(WriteError::Reset(code)));
}
if cell.finished_senders.contains(&self.r) {
return Poll::Ready(Ok(()));
}
if let Some(lost) = cell.closed.clone() {
return Poll::Ready(Err(WriteError::ConnectionLost(lost)));
}
if cell.core.is_none() {
return Poll::Ready(Err(WriteError::ConnectionLost(no_core())));
}
cell.blocked_ackers
.entry(self.r)
.or_default()
.park(self.ack_key, cx);
Poll::Pending
}
pub fn reset(&mut self, error_code: u64) {
let dirty = {
let mut cell = self.cell.borrow_mut();
match (cell.closed.is_some(), cell.core.as_mut()) {
(false, Some(core)) => {
core.reset(now(), self.r, error_code);
cell.dirty = true;
true
}
_ => false,
}
};
self.local_end = LocalEnd::Reset(error_code);
if dirty {
self.shell.mark_dirty(self.conn);
wake_settled(&self.cell);
}
}
}
impl<S: Handshake> std::fmt::Debug for SendStream<S> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SendStream")
.field("id", &self.id())
.field("local_end", &self.local_end)
.finish_non_exhaustive()
}
}
impl<S: Handshake> Drop for SendStream<S> {
fn drop(&mut self) {
let dirty = {
let mut cell = self.cell.borrow_mut();
release_waker_slot(&mut cell.blocked_writers, self.r, self.key);
cell.peer_resets.remove(&self.r);
release_waker_slot(&mut cell.blocked_ackers, self.r, self.ack_key);
cell.finished_senders.remove(&self.r);
match (self.local_end, cell.core.as_mut()) {
(LocalEnd::Live, Some(core)) => {
core.reset(now(), self.r, constants::NO_ERROR);
cell.dirty = true;
true
}
_ => false,
}
};
if dirty {
self.shell.mark_dirty(self.conn);
wake_settled(&self.cell);
}
release_handle(&self.shell, &self.cell, self.conn);
}
}
pub struct RecvStream<S: Handshake> {
shell: Rc<dyn ShellLink>,
cell: Rc<RefCell<ConnCell<S>>>,
conn: ConnectionId,
r: StreamRef,
key: u64,
ended: Option<Ended>,
id: Cell<Option<StreamId>>,
}
impl<S: Handshake> RecvStream<S> {
pub(crate) fn install(
shell: Rc<dyn ShellLink>,
cell_rc: Rc<RefCell<ConnCell<S>>>,
cell: &mut ConnCell<S>,
conn: ConnectionId,
r: StreamRef,
) -> Self {
shell.acquire();
cell.handles += 1;
let key = cell.blocked_readers.entry(r).or_default().key();
Self {
shell,
cell: cell_rc,
conn,
r,
key,
ended: None,
id: Cell::new(cell.core.as_ref().and_then(|c| c.stream_id(r))),
}
}
pub fn id(&self) -> Option<StreamId> {
cached_id(&self.cell, self.r, &self.id)
}
pub async fn read(&mut self, buf: &mut [u8]) -> Result<Option<usize>, ReadError> {
poll_fn(|cx| self.poll_read(cx, buf)).await
}
pub(crate) fn poll_read(
&mut self,
cx: &mut Context<'_>,
buf: &mut [u8],
) -> Poll<Result<Option<usize>, ReadError>> {
let (outcome, dirty) = {
let mut cell = self.cell.borrow_mut();
match self.ended {
Some(Ended::Eof) => return Poll::Ready(Ok(None)),
Some(Ended::Reset(code)) => return Poll::Ready(Err(ReadError::Reset(code))),
None => {}
}
if buf.is_empty() {
return Poll::Ready(Ok(Some(0)));
}
let Some(core) = cell.core.as_mut() else {
return Poll::Ready(match cell.closed.clone() {
Some(lost) => Err(ReadError::ConnectionLost(lost)),
None => Err(ReadError::ConnectionLost(no_core())),
});
};
match core.read(now(), self.r, buf) {
Ok(Some(0)) => match cell.closed.clone() {
Some(lost) => (Poll::Ready(Err(ReadError::ConnectionLost(lost))), false),
None => {
cell.blocked_readers
.entry(self.r)
.or_default()
.park(self.key, cx);
cell.dirty = true;
(Poll::Pending, true)
}
},
Ok(Some(n)) => {
cell.dirty = true;
(Poll::Ready(Ok(Some(n))), true)
}
Ok(None) => {
self.ended = Some(Ended::Eof);
cell.dirty = true;
(Poll::Ready(Ok(None)), true)
}
Err(ReadError::Reset(code)) => {
self.ended = Some(Ended::Reset(code));
cell.dirty = true;
(Poll::Ready(Err(ReadError::Reset(code))), true)
}
Err(e) => (Poll::Ready(Err(e)), false),
}
};
if dirty {
self.shell.mark_dirty(self.conn);
}
outcome
}
}
impl<S: Handshake> std::fmt::Debug for RecvStream<S> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RecvStream")
.field("id", &self.id())
.field("ended", &self.ended)
.finish_non_exhaustive()
}
}
impl<S: Handshake> Drop for RecvStream<S> {
fn drop(&mut self) {
let dirty = {
let mut cell = self.cell.borrow_mut();
release_waker_slot(&mut cell.blocked_readers, self.r, self.key);
match cell.core.as_mut() {
Some(core) => {
core.abandon_recv(now(), self.r);
cell.dirty = true;
true
}
None => false,
}
};
if dirty {
self.shell.mark_dirty(self.conn);
}
release_handle(&self.shell, &self.cell, self.conn);
}
}
pub struct BiStream<S: Handshake> {
send: SendStream<S>,
recv: RecvStream<S>,
}
impl<S: Handshake> BiStream<S> {
pub(crate) fn new(send: SendStream<S>, recv: RecvStream<S>) -> Self {
Self { send, recv }
}
pub fn id(&self) -> Option<StreamId> {
self.send.id()
}
pub fn split(self) -> (SendStream<S>, RecvStream<S>) {
(self.send, self.recv)
}
pub(crate) fn send_mut(&mut self) -> &mut SendStream<S> {
&mut self.send
}
pub(crate) fn recv_mut(&mut self) -> &mut RecvStream<S> {
&mut self.recv
}
#[expect(
clippy::result_large_err,
reason = "ruling 120 fixes this signature; the `Err` is the returned pair itself"
)]
pub fn join(
send: SendStream<S>,
recv: RecvStream<S>,
) -> Result<Self, (SendStream<S>, RecvStream<S>)> {
if Rc::ptr_eq(&send.cell, &recv.cell) && send.conn == recv.conn && send.r == recv.r {
Ok(Self::new(send, recv))
} else {
Err((send, recv))
}
}
}
impl<S: Handshake> std::fmt::Debug for BiStream<S> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("BiStream")
.field("id", &self.id())
.finish_non_exhaustive()
}
}
fn no_core() -> ConnectionLost {
debug_assert!(
false,
"a connection cell held neither a core nor a close reason (§16.3)"
);
ConnectionLost::EndpointDropped
}
fn release_handle<S: Handshake>(
shell: &Rc<dyn ShellLink>,
cell: &Rc<RefCell<ConnCell<S>>>,
conn: ConnectionId,
) {
let last_for_connection = {
let mut borrow = cell.borrow_mut();
debug_assert!(borrow.handles > 0, "a stream handle was released twice");
borrow.handles = borrow.handles.saturating_sub(1);
borrow.handles == 0
};
let last_in_process = shell.release();
if last_for_connection && !last_in_process {
close_now(shell, cell, conn, constants::NO_ERROR, b"");
}
}