use super::{
byte_source_trait::ReadableByteSource, error::StreamError,
readable::ReadableByteStreamController,
};
use crate::platform::{MaybeSend, MaybeSync, SharedPtr};
use bytes::{Buf, Bytes, BytesMut};
use futures::{channel::oneshot, future::poll_fn};
use parking_lot::Mutex;
use std::{
collections::VecDeque,
sync::atomic::{AtomicBool, AtomicIsize, AtomicUsize, Ordering},
task::{Context, Poll, Waker},
};
pub struct PendingPullInto {
buf: BytesMut,
completion: oneshot::Sender<Result<(BytesMut, usize), StreamError>>,
}
impl PendingPullInto {
pub(crate) fn buf(&self) -> &[u8] {
&self.buf
}
pub(crate) fn buf_mut(&mut self) -> &mut [u8] {
&mut self.buf
}
pub(crate) fn buf_len(&self) -> usize {
self.buf.len()
}
pub(crate) fn complete(self, n: usize) {
let _ = self.completion.send(Ok((self.buf, n)));
}
pub(crate) fn complete_err(self, err: StreamError) {
let _ = self.completion.send(Err(err));
}
}
pub enum PullIntoOutcome {
Ready(BytesMut, usize),
Errored(StreamError),
Registered(oneshot::Receiver<Result<(BytesMut, usize), StreamError>>),
}
fn drain_queue_into(buffer: &mut VecDeque<Bytes>, dst: &mut [u8]) -> usize {
let mut copied = 0;
while copied < dst.len() {
let Some(front) = buffer.front_mut() else {
break;
};
let front_len = front.len();
let n = std::cmp::min(dst.len() - copied, front_len);
dst[copied..copied + n].copy_from_slice(&front[..n]);
copied += n;
if n == front_len {
buffer.pop_front();
} else {
front.advance(n);
}
}
copied
}
pub struct ByteStreamState<Source> {
buffer: Mutex<VecDeque<Bytes>>,
pub(crate) source: Mutex<Option<Source>>,
pub(crate) read_wakers: Mutex<Vec<Waker>>,
pull_waker: Mutex<Option<Waker>>,
serve_waker: Mutex<Option<Waker>>,
pub(crate) closed: AtomicBool,
pub(crate) errored: AtomicBool,
pub(crate) error: Mutex<Option<StreamError>>,
pub(crate) pull_in_progress: AtomicBool,
needs_pull: AtomicBool,
serve_pending: AtomicBool,
pub(crate) queue_total_size: AtomicUsize,
pub(crate) high_water_mark: AtomicUsize,
desired_size: AtomicIsize,
start_completed: AtomicBool,
start_wakers: Mutex<Vec<Waker>>,
pending_pull_intos: Mutex<VecDeque<PendingPullInto>>,
}
impl<Source> ByteStreamState<Source>
where
Source: ReadableByteSource + 'static,
{
pub fn new(source: Source, high_water_mark: usize) -> SharedPtr<Self> {
SharedPtr::new(Self {
buffer: Mutex::new(VecDeque::new()),
source: Mutex::new(Some(source)),
read_wakers: Mutex::new(Vec::new()),
pull_waker: Mutex::new(None),
serve_waker: Mutex::new(None),
closed: AtomicBool::new(false),
errored: AtomicBool::new(false),
error: Mutex::new(None),
pull_in_progress: AtomicBool::new(false),
needs_pull: AtomicBool::new(false),
serve_pending: AtomicBool::new(false),
queue_total_size: AtomicUsize::new(0),
high_water_mark: AtomicUsize::new(high_water_mark),
desired_size: AtomicIsize::new(high_water_mark as isize),
start_completed: AtomicBool::new(false),
start_wakers: Mutex::new(Vec::new()),
pending_pull_intos: Mutex::new(VecDeque::new()),
})
}
fn poll_start_ready(&self, cx: &mut Context<'_>) -> bool {
if self.start_completed.load(Ordering::Acquire) {
return true;
}
let mut wakers = self.start_wakers.lock();
if self.start_completed.load(Ordering::Acquire) {
return true;
}
let waker = cx.waker();
if !wakers.iter().any(|w| w.will_wake(waker)) {
wakers.push(waker.clone());
}
false
}
fn take_error(&self) -> StreamError {
self.error
.lock()
.clone()
.unwrap_or_else(|| "Stream errored".into())
}
fn register_read_waker(&self, cx: &mut Context<'_>) {
let mut wakers = self.read_wakers.lock();
let waker = cx.waker();
if !wakers.iter().any(|w| w.will_wake(waker)) {
wakers.push(waker.clone());
}
}
pub fn poll_read_into(
&self,
cx: &mut Context<'_>,
buf: &mut [u8],
) -> Poll<Result<usize, StreamError>> {
if !self.poll_start_ready(cx) {
return Poll::Pending;
}
if buf.is_empty() {
return Poll::Ready(Ok(0));
}
if self.errored.load(Ordering::Acquire) {
return Poll::Ready(Err(self.take_error()));
}
let bytes_copied = {
let mut buffer = self.buffer.lock();
let copied = drain_queue_into(&mut buffer, buf);
if copied > 0 {
let new_size = self
.queue_total_size
.load(Ordering::Relaxed)
.saturating_sub(copied);
self.queue_total_size.store(new_size, Ordering::Release);
self.update_desired_size();
}
copied
};
if bytes_copied > 0 {
self.maybe_trigger_pull();
return Poll::Ready(Ok(bytes_copied));
}
if self.closed.load(Ordering::Acquire) {
return Poll::Ready(Ok(0)); }
self.register_read_waker(cx);
self.maybe_trigger_pull();
Poll::Pending
}
pub fn poll_read_chunk(
&self,
cx: &mut Context<'_>,
) -> Poll<Result<Option<Bytes>, StreamError>> {
if !self.poll_start_ready(cx) {
return Poll::Pending;
}
if self.errored.load(Ordering::Acquire) {
return Poll::Ready(Err(self.take_error()));
}
let chunk = {
let mut buffer = self.buffer.lock();
match buffer.pop_front() {
Some(chunk) => {
let new_size = self
.queue_total_size
.load(Ordering::Relaxed)
.saturating_sub(chunk.len());
self.queue_total_size.store(new_size, Ordering::Release);
self.update_desired_size();
Some(chunk)
}
None => None,
}
};
if let Some(chunk) = chunk {
self.maybe_trigger_pull();
return Poll::Ready(Ok(Some(chunk)));
}
if self.closed.load(Ordering::Acquire) {
return Poll::Ready(Ok(None)); }
self.register_read_waker(cx);
self.maybe_trigger_pull();
Poll::Pending
}
pub fn maybe_trigger_pull(&self) {
let current_size = self.queue_total_size.load(Ordering::Acquire);
let hwm = self.high_water_mark.load(Ordering::Acquire);
if !self.pull_in_progress.load(Ordering::Acquire)
&& !self.closed.load(Ordering::Acquire)
&& !self.errored.load(Ordering::Acquire)
&& current_size < hwm
{
self.needs_pull.store(true, Ordering::Release);
if let Some(waker) = self.pull_waker.lock().take() {
waker.wake();
}
}
}
pub fn poll_pull_needed(&self, cx: &mut Context<'_>) -> Poll<()> {
if self.needs_pull.load(Ordering::Acquire) {
self.needs_pull.store(false, Ordering::Release);
Poll::Ready(())
} else {
*self.pull_waker.lock() = Some(cx.waker().clone());
Poll::Pending
}
}
fn notify_serve(&self) {
self.serve_pending.store(true, Ordering::Release);
if let Some(waker) = self.serve_waker.lock().take() {
waker.wake();
}
}
pub fn poll_serve_needed(&self, cx: &mut Context<'_>) -> Poll<()> {
if self.serve_pending.swap(false, Ordering::AcqRel) {
return Poll::Ready(());
}
*self.serve_waker.lock() = Some(cx.waker().clone());
if self.serve_pending.swap(false, Ordering::AcqRel) {
Poll::Ready(())
} else {
Poll::Pending
}
}
pub fn enqueue_bytes(&self, chunk: Bytes) {
if chunk.is_empty() {
return;
}
let len = chunk.len();
{
let mut buffer = self.buffer.lock();
buffer.push_back(chunk);
let new_size = self.queue_total_size.load(Ordering::Relaxed) + len;
self.queue_total_size.store(new_size, Ordering::Release);
}
self.update_desired_size();
self.settle_pending_pull_intos();
self.wake_readers();
self.notify_serve();
}
pub fn begin_pull_into(&self, mut buf: BytesMut) -> PullIntoOutcome {
if self.errored.load(Ordering::Acquire) {
return PullIntoOutcome::Errored(self.take_error());
}
if buf.is_empty() {
return PullIntoOutcome::Ready(buf, 0);
}
let rx = {
let mut buffer = self.buffer.lock();
let mut pending = self.pending_pull_intos.lock();
if pending.is_empty() {
if !buffer.is_empty() {
let n = drain_queue_into(&mut buffer, &mut buf);
let new_size = self
.queue_total_size
.load(Ordering::Relaxed)
.saturating_sub(n);
self.queue_total_size.store(new_size, Ordering::Release);
drop(pending);
drop(buffer);
self.update_desired_size();
self.maybe_trigger_pull();
return PullIntoOutcome::Ready(buf, n);
}
if self.closed.load(Ordering::Acquire) {
return PullIntoOutcome::Ready(buf, 0);
}
}
let (tx, rx) = oneshot::channel();
pending.push_back(PendingPullInto { buf, completion: tx });
rx
};
if self.closed.load(Ordering::Acquire) || self.errored.load(Ordering::Acquire) {
self.settle_pending_pull_intos();
}
self.force_pull();
PullIntoOutcome::Registered(rx)
}
pub fn take_pull_into(&self) -> Option<PendingPullInto> {
self.pending_pull_intos.lock().pop_front()
}
pub fn return_pull_into(&self, pending: PendingPullInto) {
self.pending_pull_intos.lock().push_front(pending);
if self.closed.load(Ordering::Acquire) || self.errored.load(Ordering::Acquire) {
self.settle_pending_pull_intos();
} else {
self.force_pull();
}
}
fn settle_pending_pull_intos(&self) {
let mut buffer = self.buffer.lock();
let mut pending = self.pending_pull_intos.lock();
if pending.is_empty() {
return;
}
if self.errored.load(Ordering::Acquire) {
let waiting: VecDeque<PendingPullInto> = std::mem::take(&mut pending);
drop(pending);
drop(buffer);
let err = self.take_error();
for p in waiting {
p.complete_err(err.clone());
}
return;
}
let mut filled: Vec<(PendingPullInto, usize)> = Vec::new();
let mut total_drained = 0;
while !buffer.is_empty() {
let Some(mut p) = pending.pop_front() else {
break;
};
let n = drain_queue_into(&mut buffer, p.buf_mut());
total_drained += n;
filled.push((p, n));
}
let eof: VecDeque<PendingPullInto> = if self.closed.load(Ordering::Acquire) {
std::mem::take(&mut pending)
} else {
VecDeque::new()
};
if total_drained > 0 {
let new_size = self
.queue_total_size
.load(Ordering::Relaxed)
.saturating_sub(total_drained);
self.queue_total_size.store(new_size, Ordering::Release);
}
drop(pending);
drop(buffer);
if total_drained > 0 {
self.update_desired_size();
}
for (p, n) in filled {
p.complete(n);
}
for p in eof {
p.complete(0); }
if total_drained > 0 {
self.maybe_trigger_pull();
}
}
fn force_pull(&self) {
if !self.pull_in_progress.load(Ordering::Acquire)
&& !self.closed.load(Ordering::Acquire)
&& !self.errored.load(Ordering::Acquire)
{
self.needs_pull.store(true, Ordering::Release);
if let Some(waker) = self.pull_waker.lock().take() {
waker.wake();
}
}
}
pub fn mark_pull_started(&self) {
self.pull_in_progress.store(true, Ordering::Release);
}
pub fn mark_pull_completed(&self, made_progress: bool) {
self.pull_in_progress.store(false, Ordering::Release);
if !self.pending_pull_intos.lock().is_empty() {
self.force_pull();
} else if made_progress {
self.maybe_trigger_pull();
}
}
pub fn close(&self) {
self.closed.store(true, Ordering::Release);
self.update_desired_size();
self.settle_pending_pull_intos();
self.wake_readers();
self.notify_serve();
if let Some(waker) = self.pull_waker.lock().take() {
waker.wake();
}
}
pub fn error(&self, err: StreamError) {
*self.error.lock() = Some(err);
self.errored.store(true, Ordering::Release);
self.update_desired_size();
self.settle_pending_pull_intos();
self.wake_readers();
self.notify_serve();
if let Some(waker) = self.pull_waker.lock().take() {
waker.wake();
}
}
fn wake_readers(&self) {
let mut wakers = self.read_wakers.lock();
for waker in wakers.drain(..) {
waker.wake();
}
}
fn update_desired_size(&self) {
if self.closed.load(Ordering::Acquire) || self.errored.load(Ordering::Acquire) {
self.desired_size.store(0, Ordering::Release);
return;
}
let hwm = self.high_water_mark.load(Ordering::Relaxed) as isize;
let current = self.queue_total_size.load(Ordering::Relaxed) as isize;
self.desired_size.store(hwm - current, Ordering::Release);
}
pub fn desired_size(&self) -> Option<isize> {
if self.errored.load(Ordering::Acquire) {
None
} else {
Some(self.desired_size.load(Ordering::Acquire))
}
}
pub fn is_buffer_empty(&self) -> bool {
self.buffer.lock().is_empty()
}
pub fn buffer_size(&self) -> usize {
self.queue_total_size.load(Ordering::Acquire)
}
pub async fn start_source(
&self,
controller: &ReadableByteStreamController,
) -> Result<(), StreamError> {
let mut source = match self.source.lock().take() {
Some(s) => s,
None => return Ok(()),
};
let mut controller = controller.clone();
let result = source.start(&mut controller).await;
match result {
Ok(()) => {
*self.source.lock() = Some(source);
self.mark_start_completed();
Ok(())
}
Err(err) => {
self.error(err.clone());
self.mark_start_completed();
Err(err)
}
}
}
pub async fn closed(&self) -> Result<(), StreamError> {
poll_fn(|cx| {
if self.is_errored() {
let error = self
.error
.lock()
.clone()
.unwrap_or_else(|| "Stream errored".into());
return Poll::Ready(Err(error));
}
if self.is_closed() {
return Poll::Ready(Ok(()));
}
let mut wakers = self.read_wakers.lock();
let waker = cx.waker();
if !wakers.iter().any(|w| w.will_wake(waker)) {
wakers.push(waker.clone());
}
Poll::Pending
})
.await
}
pub fn mark_start_completed(&self) {
if self.start_completed.swap(true, Ordering::AcqRel) {
return;
}
let mut wakers = self.start_wakers.lock();
for waker in wakers.drain(..) {
waker.wake();
}
}
}
pub trait ByteStreamStateInterface: MaybeSend + MaybeSync {
fn desired_size(&self) -> Option<isize>;
fn close(&self);
fn enqueue_bytes(&self, chunk: Bytes);
fn error(&self, error: StreamError);
fn is_buffer_empty(&self) -> bool;
fn buffer_size(&self) -> usize;
fn is_closed(&self) -> bool;
fn is_errored(&self) -> bool;
fn closed(&self) -> crate::platform::PlatformBoxFuture<'_, Result<(), StreamError>>;
fn poll_read_into(
&self,
cx: &mut Context<'_>,
buf: &mut [u8],
) -> Poll<Result<usize, StreamError>>;
fn poll_read_chunk(
&self,
cx: &mut Context<'_>,
) -> Poll<Result<Option<Bytes>, StreamError>>;
fn begin_pull_into(&self, buf: BytesMut) -> PullIntoOutcome;
fn take_pull_into(&self) -> Option<PendingPullInto>;
fn return_pull_into(&self, pending: PendingPullInto);
fn cancel_source<'a>(
&'a self,
reason: Option<String>,
) -> crate::platform::PlatformBoxFuture<'a, Result<(), StreamError>>;
}
impl<Source> ByteStreamStateInterface for ByteStreamState<Source>
where
Source: ReadableByteSource + 'static,
{
fn desired_size(&self) -> Option<isize> {
ByteStreamState::desired_size(self)
}
fn close(&self) {
ByteStreamState::close(self)
}
fn enqueue_bytes(&self, chunk: Bytes) {
ByteStreamState::enqueue_bytes(self, chunk)
}
fn error(&self, error: StreamError) {
ByteStreamState::error(self, error)
}
fn is_buffer_empty(&self) -> bool {
self.is_buffer_empty()
}
fn buffer_size(&self) -> usize {
self.buffer_size()
}
fn is_closed(&self) -> bool {
self.closed.load(Ordering::Acquire)
}
fn is_errored(&self) -> bool {
self.errored.load(Ordering::Acquire)
}
fn closed(&self) -> crate::platform::PlatformBoxFuture<'_, Result<(), StreamError>> {
Box::pin(async move { ByteStreamState::closed(self).await })
}
fn poll_read_into(
&self,
cx: &mut Context<'_>,
buf: &mut [u8],
) -> Poll<Result<usize, StreamError>> {
ByteStreamState::poll_read_into(self, cx, buf)
}
fn poll_read_chunk(
&self,
cx: &mut Context<'_>,
) -> Poll<Result<Option<Bytes>, StreamError>> {
ByteStreamState::poll_read_chunk(self, cx)
}
fn begin_pull_into(&self, buf: BytesMut) -> PullIntoOutcome {
ByteStreamState::begin_pull_into(self, buf)
}
fn take_pull_into(&self) -> Option<PendingPullInto> {
ByteStreamState::take_pull_into(self)
}
fn return_pull_into(&self, pending: PendingPullInto) {
ByteStreamState::return_pull_into(self, pending)
}
fn cancel_source<'a>(
&'a self,
reason: Option<String>,
) -> crate::platform::PlatformBoxFuture<'a, Result<(), StreamError>> {
Box::pin(async move {
if self.closed.load(Ordering::Acquire) {
return Ok(());
}
if self.errored.load(Ordering::Acquire) {
return Err(self
.error
.lock()
.clone()
.unwrap_or_else(|| "Stream errored".into()));
}
{
self.buffer.lock().clear();
self.queue_total_size.store(0, Ordering::Release);
}
self.close();
let source_opt = self.source.lock().take();
if let Some(mut s) = source_opt {
s.cancel(reason).await
} else {
Ok(())
}
})
}
}