use super::super::{CountQueuingStrategy, Locked, QueuingStrategy, Unlocked};
use super::byte_state::{PendingPullInto, PullIntoOutcome};
use super::shared::AbortSignal;
use super::shared::StreamResult;
use super::shared::WakerSet;
pub use super::{
byte_source_trait::ReadableByteSource,
byte_state::{ByteStreamState, ByteStreamStateInterface},
error::StreamError,
transform::{TransformReadableSource, TransformStream},
writable::{WritableSink, WritableStream},
};
use crate::platform::{MaybeSend, MaybeSync, SharedPtr};
use bytes::{Bytes, BytesMut};
use futures::{
FutureExt,
channel::{
mpsc::{UnboundedReceiver, UnboundedSender, unbounded},
oneshot,
},
future::{Either, poll_fn, select},
io::AsyncRead,
stream::{Stream, StreamExt},
};
use parking_lot::{Mutex, RwLock};
use std::{
collections::VecDeque,
future::Future,
io::Result as IoResult,
marker::PhantomData,
pin::Pin,
sync::atomic::{AtomicBool, AtomicIsize, AtomicUsize, Ordering},
task::{Context, Poll, Waker},
};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum StreamState {
Readable,
Closed,
Errored,
}
pub struct DefaultStream;
pub struct ByteStream;
pub trait StreamTypeMarker: MaybeSend + 'static {
type Controller<T: MaybeSend + 'static>: MaybeSend + MaybeSync + Clone + 'static;
}
impl StreamTypeMarker for DefaultStream {
type Controller<T: MaybeSend + 'static> = ReadableStreamDefaultController<T>;
}
impl StreamTypeMarker for ByteStream {
type Controller<T: MaybeSend + 'static> = ReadableByteStreamController;
}
pub trait ReadableSource<T: MaybeSend + 'static>: MaybeSend + 'static {
fn start(
&mut self,
#[allow(unused)] controller: &mut ReadableStreamDefaultController<T>,
) -> impl Future<Output = StreamResult<()>> + MaybeSend {
async { Ok(()) }
}
fn pull(
&mut self,
#[allow(unused)] controller: &mut ReadableStreamDefaultController<T>,
) -> impl Future<Output = StreamResult<()>> + MaybeSend {
async { Ok(()) }
}
fn cancel(
&mut self,
#[allow(unused)] reason: Option<String>,
) -> impl Future<Output = StreamResult<()>> + MaybeSend {
async { Ok(()) }
}
}
enum StreamCommand<T> {
Read {
completion: oneshot::Sender<StreamResult<Option<T>>>,
},
Cancel {
reason: Option<String>,
completion: oneshot::Sender<StreamResult<()>>,
},
RegisterReadyWaker {
waker: Waker,
},
RegisterClosedWaker {
waker: Waker,
},
}
enum ControllerMsg<T> {
Enqueue { chunk: T },
Close,
Error(StreamError),
}
pub struct ReadableStreamDefaultController<T: MaybeSend + 'static> {
tx: UnboundedSender<ControllerMsg<T>>,
queue_total_size: SharedPtr<AtomicUsize>,
high_water_mark: SharedPtr<AtomicUsize>,
desired_size: SharedPtr<AtomicIsize>,
closed: SharedPtr<AtomicBool>,
errored: SharedPtr<AtomicBool>,
close_requested: SharedPtr<AtomicBool>,
error_requested: SharedPtr<AtomicBool>,
}
impl<T: MaybeSend + 'static> Clone for ReadableStreamDefaultController<T> {
fn clone(&self) -> Self {
Self {
tx: self.tx.clone(),
queue_total_size: self.queue_total_size.clone(),
high_water_mark: self.high_water_mark.clone(),
desired_size: self.desired_size.clone(),
closed: self.closed.clone(),
errored: self.errored.clone(),
close_requested: self.close_requested.clone(),
error_requested: self.error_requested.clone(),
}
}
}
impl<T: MaybeSend + 'static> ReadableStreamDefaultController<T> {
fn new(
tx: UnboundedSender<ControllerMsg<T>>,
queue_total_size: SharedPtr<AtomicUsize>,
high_water_mark: SharedPtr<AtomicUsize>,
desired_size: SharedPtr<AtomicIsize>,
closed: SharedPtr<AtomicBool>,
errored: SharedPtr<AtomicBool>,
close_requested: SharedPtr<AtomicBool>,
error_requested: SharedPtr<AtomicBool>,
) -> Self {
Self {
tx,
queue_total_size,
high_water_mark,
desired_size,
closed,
errored,
close_requested,
error_requested,
}
}
pub fn is_closed_or_errored(&self) -> bool {
self.is_stream_unusable()
}
pub fn desired_size(&self) -> Option<isize> {
if self.errored.load(Ordering::Acquire) || self.error_requested.load(Ordering::Acquire) {
return None;
}
if self.closed.load(Ordering::Acquire)
|| (self.close_requested.load(Ordering::Acquire)
&& self.queue_total_size.load(Ordering::Acquire) == 0)
{
return Some(0);
}
Some(self.desired_size.load(Ordering::Acquire))
}
fn is_stream_unusable(&self) -> bool {
self.close_requested.load(Ordering::Acquire)
|| self.closed.load(Ordering::Acquire)
|| self.error_requested.load(Ordering::Acquire)
|| self.errored.load(Ordering::Acquire)
}
pub fn close(&self) -> StreamResult<()> {
if self.is_stream_unusable() {
return Err(StreamError::from(
"Cannot close a ReadableStream that is already closed or errored",
));
}
self.close_requested.store(true, Ordering::Release);
self.tx
.unbounded_send(ControllerMsg::Close)
.map_err(|_| StreamError::from("Failed to close stream"))?;
Ok(())
}
pub fn enqueue(&self, chunk: T) -> StreamResult<()> {
if self.close_requested.load(Ordering::Acquire) || self.closed.load(Ordering::Acquire) {
return Err(StreamError::from(
"Cannot enqueue a chunk into a closed ReadableStream",
));
}
if self.error_requested.load(Ordering::Acquire) || self.errored.load(Ordering::Acquire) {
return Err(StreamError::from(
"Cannot enqueue a chunk into an errored ReadableStream",
));
}
self.tx
.unbounded_send(ControllerMsg::Enqueue { chunk })
.map_err(|_| StreamError::from("Failed to enqueue chunk"))?;
Ok(())
}
pub fn error(&self, error: StreamError) -> StreamResult<()> {
if self.closed.load(Ordering::Acquire)
|| self.errored.load(Ordering::Acquire)
|| self.error_requested.load(Ordering::Acquire)
{
return Ok(());
}
self.error_requested.store(true, Ordering::Release);
self.tx
.unbounded_send(ControllerMsg::Error(error))
.map_err(|_| StreamError::from("Failed to error stream"))?;
Ok(())
}
}
pub struct ReadableByteStreamController {
byte_state: SharedPtr<dyn ByteStreamStateInterface>,
}
impl ReadableByteStreamController {
pub fn new<Source>(byte_state: SharedPtr<ByteStreamState<Source>>) -> Self
where
Source: ReadableByteSource,
{
Self {
byte_state: byte_state as SharedPtr<dyn ByteStreamStateInterface>,
}
}
pub fn desired_size(&self) -> Option<isize> {
self.byte_state.desired_size()
}
pub fn close(&self) -> StreamResult<()> {
if self.byte_state.is_closed() {
return Err("Stream is already closed".into());
}
if self.byte_state.is_errored() {
return Err("Stream is errored".into());
}
self.byte_state.close();
Ok(())
}
pub fn enqueue(&self, chunk: impl Into<Bytes>) -> StreamResult<()> {
if self.byte_state.is_closed() {
return Err("Stream is closed".into());
}
if self.byte_state.is_errored() {
return Err("Stream is errored".into());
}
self.byte_state.enqueue_bytes(chunk.into());
Ok(())
}
pub fn error(&self, error: StreamError) -> StreamResult<()> {
self.byte_state.error(error);
Ok(())
}
pub fn byob_request(&mut self) -> Option<ByobRequest> {
self.byte_state.take_pull_into().map(|pending| ByobRequest {
pending: Some(pending),
byte_state: self.byte_state.clone(),
})
}
}
impl Clone for ReadableByteStreamController {
fn clone(&self) -> Self {
Self {
byte_state: self.byte_state.clone(),
}
}
}
pub struct ByobRequest {
pending: Option<PendingPullInto>,
byte_state: SharedPtr<dyn ByteStreamStateInterface>,
}
impl ByobRequest {
pub fn respond(mut self, n: usize) -> StreamResult<()> {
let pending = self.pending.take().expect("respond on an active request");
if n > pending.buf_len() {
let err = StreamError::from("respond() byte count exceeds the request buffer");
pending.complete_err(err.clone());
return Err(err);
}
pending.complete(n);
Ok(())
}
}
impl std::ops::Deref for ByobRequest {
type Target = [u8];
fn deref(&self) -> &[u8] {
self.pending.as_ref().expect("active request").buf()
}
}
impl std::ops::DerefMut for ByobRequest {
fn deref_mut(&mut self) -> &mut [u8] {
self.pending.as_mut().expect("active request").buf_mut()
}
}
impl Drop for ByobRequest {
fn drop(&mut self) {
if let Some(pending) = self.pending.take() {
self.byte_state.return_pull_into(pending);
}
}
}
struct ReadableStreamInner<T, Source> {
state: StreamState,
queue: VecDeque<T>,
queue_total_size: usize,
strategy: crate::platform::BoxedStrategyStatic<T>,
source: Option<Source>,
cancel_requested: bool,
cancel_reason: Option<String>,
cancel_completions: Vec<oneshot::Sender<StreamResult<()>>>,
pending_reads: VecDeque<oneshot::Sender<StreamResult<Option<T>>>>,
ready_wakers: WakerSet,
closed_wakers: WakerSet,
stored_error: Option<StreamError>,
pulling: bool,
}
impl<T: MaybeSend + 'static, Source> ReadableStreamInner<T, Source> {
fn new(source: Source, strategy: crate::platform::BoxedStrategyStatic<T>) -> Self {
Self {
state: StreamState::Readable,
queue: VecDeque::new(),
queue_total_size: 0,
strategy,
source: Some(source),
cancel_requested: false,
cancel_reason: None,
cancel_completions: Vec::new(),
pending_reads: VecDeque::new(),
ready_wakers: WakerSet::new(),
closed_wakers: WakerSet::new(),
stored_error: None,
pulling: false,
}
}
fn get_stored_error(&self) -> StreamError {
self.stored_error
.clone()
.unwrap_or_else(|| StreamError::from("Stream is errored"))
}
}
pub struct ReadableStream<T: MaybeSend + 'static, Source, StreamType, LockState = Unlocked>
where
StreamType: StreamTypeMarker,
{
command_tx: UnboundedSender<StreamCommand<T>>,
queue_total_size: SharedPtr<AtomicUsize>,
high_water_mark: SharedPtr<AtomicUsize>,
closed: SharedPtr<AtomicBool>,
errored: SharedPtr<AtomicBool>,
locked: SharedPtr<AtomicBool>,
stored_error: SharedPtr<RwLock<Option<StreamError>>>,
desired_size: SharedPtr<AtomicIsize>,
pub(crate) controller: SharedPtr<StreamType::Controller<T>>,
pub(crate) byte_state: Option<SharedPtr<dyn ByteStreamStateInterface>>,
poll_read_rx: Option<oneshot::Receiver<StreamResult<Option<T>>>>,
_phantom: PhantomData<fn() -> (T, Source, StreamType, LockState)>,
}
impl<T: MaybeSend + 'static, Source> ReadableStream<T, Source, DefaultStream, Unlocked> {
pub fn locked(&self) -> bool {
self.locked.load(Ordering::Acquire)
}
}
impl<T: MaybeSend + 'static, Source, S> ReadableStream<T, Source, S, Unlocked>
where
S: StreamTypeMarker,
{
pub async fn cancel(&self, reason: Option<String>) -> StreamResult<()> {
if self.locked.load(Ordering::Acquire) {
return Err(StreamError::from("Cannot cancel a locked ReadableStream"));
}
let (tx, rx) = oneshot::channel();
self.command_tx
.unbounded_send(StreamCommand::Cancel {
reason,
completion: tx,
})
.map_err(|_| StreamError::TaskDropped)?;
rx.await.unwrap_or_else(|_| Err(StreamError::TaskDropped))
}
async fn pipe_to_inner<Sink>(
self,
destination: &WritableStream<T, Sink>,
options: Option<StreamPipeOptions>,
) -> StreamResult<()>
where
Sink: WritableSink<T> + 'static,
{
let options = options.unwrap_or_default();
let (_, writer) = destination.get_writer()?;
let (_stream, reader) = self.get_reader()?;
let pipe_loop = async {
use futures::FutureExt;
loop {
if let Err(write_err) = writer.ready().await {
if !options.prevent_cancel {
if let Err(cancel_err) = reader.cancel(Some(write_err.to_string())).await {
return Err(cancel_err);
}
}
return Err(write_err);
}
let read_result = {
let read_fut = reader.read().fuse();
let mut closed_fut = writer.closed().fuse();
futures::pin_mut!(read_fut);
futures::select! {
r = read_fut => r,
c = closed_fut => {
match c {
Err(write_err) => {
if !options.prevent_cancel {
if let Err(cancel_err) =
reader.cancel(Some(write_err.to_string())).await
{
return Err(cancel_err);
}
}
return Err(write_err);
}
Ok(()) => {
let reason =
"cannot pipe to a closed writable stream".to_string();
let mut result_err = StreamError::from(reason.clone());
if !options.prevent_cancel {
if let Err(cancel_err) =
reader.cancel(Some(reason)).await
{
result_err = cancel_err;
}
}
return Err(result_err);
}
}
}
}
};
match read_result {
Ok(Some(chunk)) => {
let _ = writer.write(chunk);
}
Ok(None) => {
if !options.prevent_close {
writer.close().await?;
}
return Ok(());
}
Err(read_err) => {
if !options.prevent_abort {
if let Err(abort_err) = writer.abort(Some(read_err.to_string())).await {
return Err(abort_err);
}
}
return Err(read_err);
}
}
}
};
if let Some(signal) = options.signal {
let abort_fut = signal.aborted_future();
futures::pin_mut!(pipe_loop, abort_fut);
match select(abort_fut, pipe_loop).await {
Either::Right((result, _)) => result,
Either::Left(((), _pipe_loop)) => {
let reason = signal.reason();
let mut action_error = None;
if !options.prevent_cancel {
if let Err(e) = reader.cancel(reason.clone()).await {
action_error = Some(e);
}
}
if !options.prevent_abort {
if let Err(e) = writer.abort(reason.clone()).await {
action_error = Some(e);
}
}
Err(action_error.unwrap_or(StreamError::Aborted(reason)))
}
}
} else {
pipe_loop.await
}
}
pub async fn pipe_to<Sink>(
self,
destination: &WritableStream<T, Sink>,
options: Option<StreamPipeOptions>,
) -> StreamResult<()>
where
Sink: WritableSink<T> + 'static,
{
self.pipe_to_inner(destination, options).await
}
}
#[derive(Debug, Clone)]
pub struct TeeConfig {
pub backpressure_mode: BackpressureMode,
pub branch1_hwm: Option<usize>,
pub branch2_hwm: Option<usize>,
pub max_buffer_per_branch: Option<usize>,
}
impl Default for TeeConfig {
fn default() -> Self {
Self {
backpressure_mode: BackpressureMode::Unbounded,
branch1_hwm: None,
branch2_hwm: None,
max_buffer_per_branch: Some(1000),
}
}
}
#[derive(Debug, Clone)]
pub enum BackpressureMode {
SpecCompliant,
Unbounded,
SlowestConsumer,
Aggregate,
}
#[derive(Clone)]
struct AsyncSignal {
waker: SharedPtr<Mutex<Option<Waker>>>,
signaled: SharedPtr<AtomicBool>,
}
impl AsyncSignal {
pub fn new() -> Self {
Self {
waker: SharedPtr::new(Mutex::new(None)),
signaled: SharedPtr::new(AtomicBool::new(false)),
}
}
pub async fn wait(&self) {
poll_fn(|cx| {
if self.signaled.swap(false, Ordering::AcqRel) {
return Poll::Ready(());
}
*self.waker.lock() = Some(cx.waker().clone());
if self.signaled.swap(false, Ordering::AcqRel) {
Poll::Ready(())
} else {
Poll::Pending
}
})
.await
}
pub fn signal(&self) {
self.signaled.store(true, Ordering::Release);
if let Some(w) = self.waker.lock().take() {
w.wake();
}
}
}
struct TeeCancelResult {
result: Mutex<Option<StreamResult<()>>>,
waker: futures::task::AtomicWaker,
}
impl TeeCancelResult {
fn new() -> Self {
Self {
result: Mutex::new(None),
waker: futures::task::AtomicWaker::new(),
}
}
fn complete(&self, result: StreamResult<()>) {
let mut lock = self.result.lock();
if lock.is_some() {
return; }
*lock = Some(result);
drop(lock);
self.waker.wake();
}
fn is_done(&self) -> bool {
self.result.lock().is_some()
}
}
trait TeeBranch<T>: Clone {
fn enqueue(&self, chunk: T) -> StreamResult<()>;
fn desired_size(&self) -> Option<isize>;
fn close(&self) -> StreamResult<()>;
fn error(&self, error: StreamError) -> StreamResult<()>;
}
impl<T: MaybeSend + 'static> TeeBranch<T> for ReadableStreamDefaultController<T> {
fn enqueue(&self, chunk: T) -> StreamResult<()> {
ReadableStreamDefaultController::enqueue(self, chunk)
}
fn desired_size(&self) -> Option<isize> {
ReadableStreamDefaultController::desired_size(self)
}
fn close(&self) -> StreamResult<()> {
ReadableStreamDefaultController::close(self)
}
fn error(&self, error: StreamError) -> StreamResult<()> {
ReadableStreamDefaultController::error(self, error)
}
}
impl TeeBranch<Bytes> for ReadableByteStreamController {
fn enqueue(&self, chunk: Bytes) -> StreamResult<()> {
ReadableByteStreamController::enqueue(self, chunk)
}
fn desired_size(&self) -> Option<isize> {
ReadableByteStreamController::desired_size(self)
}
fn close(&self) -> StreamResult<()> {
ReadableByteStreamController::close(self)
}
fn error(&self, error: StreamError) -> StreamResult<()> {
ReadableByteStreamController::error(self, error)
}
}
pub struct TeeSource<T: MaybeSend + 'static> {
branch_canceled: SharedPtr<AtomicBool>,
pull_signal: AsyncSignal,
first_cancel_done: SharedPtr<AtomicBool>,
cancel_result: SharedPtr<TeeCancelResult>,
_phantom: PhantomData<fn() -> T>,
}
impl<T: MaybeSend + 'static> TeeSource<T> {
async fn cancel_coordinated(&mut self) -> StreamResult<()> {
self.branch_canceled.store(true, Ordering::Release);
self.pull_signal.signal();
let is_second = self.first_cancel_done.swap(true, Ordering::AcqRel);
if !is_second {
return Ok(());
}
let cancel_result = SharedPtr::clone(&self.cancel_result);
poll_fn(move |cx| {
if cancel_result.is_done() {
return std::task::Poll::Ready(());
}
cancel_result.waker.register(cx.waker());
if cancel_result.is_done() {
std::task::Poll::Ready(())
} else {
std::task::Poll::Pending
}
})
.await;
self.cancel_result.result.lock().clone().unwrap_or(Ok(()))
}
}
impl<T: MaybeSend + 'static> ReadableSource<T> for TeeSource<T> {
async fn pull(
&mut self,
_controller: &mut ReadableStreamDefaultController<T>,
) -> StreamResult<()> {
self.pull_signal.signal();
Ok(())
}
async fn cancel(&mut self, _reason: Option<String>) -> StreamResult<()> {
self.cancel_coordinated().await
}
}
impl ReadableByteSource for TeeSource<Bytes> {
async fn pull(
&mut self,
_controller: &mut ReadableByteStreamController,
) -> StreamResult<()> {
self.pull_signal.signal();
Ok(())
}
async fn cancel(&mut self, _reason: Option<String>) -> StreamResult<()> {
self.cancel_coordinated().await
}
}
struct TeeCoordinator<T, Source, StreamType, LockState, B>
where
T: MaybeSend + Clone + 'static,
StreamType: StreamTypeMarker,
B: TeeBranch<T>,
{
reader: ReadableStreamDefaultReader<T, Source, StreamType, LockState>,
branch1: B,
branch2: B,
branch1_canceled: SharedPtr<AtomicBool>,
branch2_canceled: SharedPtr<AtomicBool>,
cancel_result: SharedPtr<TeeCancelResult>,
backpressure_mode: BackpressureMode,
pull_signal: AsyncSignal,
}
impl<T: MaybeSend + 'static, Source, StreamType, LockState, B>
TeeCoordinator<T, Source, StreamType, LockState, B>
where
T: Clone,
StreamType: StreamTypeMarker,
B: TeeBranch<T>,
{
fn branch1_active(&self) -> bool {
!self.branch1_canceled.load(Ordering::Acquire)
}
fn branch2_active(&self) -> bool {
!self.branch2_canceled.load(Ordering::Acquire)
}
fn both_branches_dead(&self) -> bool {
!self.branch1_active() && !self.branch2_active()
}
fn should_pull(&self) -> bool {
let b1_active = self.branch1_active();
let b2_active = self.branch2_active();
if !b1_active && !b2_active {
return false;
}
let ds1 = self.branch1.desired_size().unwrap_or(0);
let ds2 = self.branch2.desired_size().unwrap_or(0);
match self.backpressure_mode {
BackpressureMode::Unbounded => b1_active || b2_active,
BackpressureMode::SpecCompliant => (b1_active && ds1 > 0) || (b2_active && ds2 > 0),
BackpressureMode::SlowestConsumer => (!b1_active || ds1 > 0) && (!b2_active || ds2 > 0),
BackpressureMode::Aggregate => {
let mut total = 0isize;
if b1_active {
total += ds1;
}
if b2_active {
total += ds2;
}
total > 0
}
}
}
fn distribute_chunk(&self, chunk: T) {
if self.branch1_active() && self.branch1.enqueue(chunk.clone()).is_err() {
self.branch1_canceled.store(true, Ordering::Release);
}
if self.branch2_active() && self.branch2.enqueue(chunk).is_err() {
self.branch2_canceled.store(true, Ordering::Release);
}
}
fn close_branches(&self) {
let _ = self.branch1.close();
let _ = self.branch2.close();
}
fn error_branches(&self, err: StreamError) {
let _ = self.branch1.error(err.clone());
let _ = self.branch2.error(err);
}
async fn finish_both_canceled(&self) {
let result = self
.reader
.cancel(Some("Both tee branches terminated".to_string()))
.await;
self.cancel_result.complete(result);
}
async fn run(self) {
self.run_inner().await;
self.cancel_result.complete(Ok(()));
}
async fn run_inner(&self) {
loop {
if self.both_branches_dead() {
self.finish_both_canceled().await;
return;
}
self.pull_signal.wait().await;
if self.both_branches_dead() {
self.finish_both_canceled().await;
return;
}
if !self.should_pull() {
continue;
}
match self.reader.read().await {
Ok(Some(chunk)) => self.distribute_chunk(chunk),
Ok(None) => {
self.close_branches();
return;
}
Err(err) => {
self.error_branches(err);
return;
}
}
}
}
}
struct TeeSetup {
branch1_canceled: SharedPtr<AtomicBool>,
branch2_canceled: SharedPtr<AtomicBool>,
pull_signal: AsyncSignal,
first_cancel_done: SharedPtr<AtomicBool>,
cancel_result: SharedPtr<TeeCancelResult>,
}
impl TeeSetup {
fn new() -> Self {
Self {
branch1_canceled: SharedPtr::new(AtomicBool::new(false)),
branch2_canceled: SharedPtr::new(AtomicBool::new(false)),
pull_signal: AsyncSignal::new(),
first_cancel_done: SharedPtr::new(AtomicBool::new(false)),
cancel_result: SharedPtr::new(TeeCancelResult::new()),
}
}
fn sources<T: MaybeSend + 'static>(&self) -> (TeeSource<T>, TeeSource<T>) {
let source1 = TeeSource {
branch_canceled: self.branch1_canceled.clone(),
pull_signal: self.pull_signal.clone(),
first_cancel_done: self.first_cancel_done.clone(),
cancel_result: self.cancel_result.clone(),
_phantom: PhantomData,
};
let source2 = TeeSource {
branch_canceled: self.branch2_canceled.clone(),
pull_signal: self.pull_signal.clone(),
first_cancel_done: self.first_cancel_done.clone(),
cancel_result: self.cancel_result.clone(),
_phantom: PhantomData,
};
(source1, source2)
}
fn coordinator<T, Source, StreamType, LockState, B>(
self,
reader: ReadableStreamDefaultReader<T, Source, StreamType, LockState>,
branch1: B,
branch2: B,
mode: BackpressureMode,
) -> TeeCoordinator<T, Source, StreamType, LockState, B>
where
T: MaybeSend + Clone + 'static,
StreamType: StreamTypeMarker,
B: TeeBranch<T>,
{
TeeCoordinator {
reader,
branch1,
branch2,
branch1_canceled: self.branch1_canceled,
branch2_canceled: self.branch2_canceled,
cancel_result: self.cancel_result,
backpressure_mode: mode,
pull_signal: self.pull_signal,
}
}
}
pub struct TeeBuilder<T, Source, S>
where
T: MaybeSend + Clone + 'static,
Source: MaybeSend + 'static,
S: StreamTypeMarker,
{
mode: BackpressureMode,
stream: ReadableStream<T, Source, S, Unlocked>,
branch1_strategy: crate::platform::BoxedStrategyStatic<T>,
branch2_strategy: crate::platform::BoxedStrategyStatic<T>,
}
impl<T: MaybeSend + 'static, Source, S> TeeBuilder<T, Source, S>
where
T: Clone,
Source: MaybeSend + 'static,
S: StreamTypeMarker,
{
fn new(stream: ReadableStream<T, Source, S, Unlocked>) -> Self {
Self {
stream,
mode: BackpressureMode::SpecCompliant,
branch1_strategy: Box::new(CountQueuingStrategy::new(1)),
branch2_strategy: Box::new(CountQueuingStrategy::new(1)),
}
}
pub fn backpressure_mode(mut self, mode: BackpressureMode) -> Self {
self.mode = mode;
self
}
pub fn branch1_strategy<Strategy: QueuingStrategy<T> + MaybeSend + 'static>(
mut self,
strategy: Strategy,
) -> Self {
self.branch1_strategy = Box::new(strategy);
self
}
pub fn branch2_strategy<Strategy: QueuingStrategy<T> + MaybeSend + 'static>(
mut self,
strategy: Strategy,
) -> Self {
self.branch2_strategy = Box::new(strategy);
self
}
pub fn strategy<Strategy: QueuingStrategy<T> + MaybeSend + 'static + Clone>(
mut self,
strategy: Strategy,
) -> Self {
self.branch1_strategy = Box::new(strategy.clone());
self.branch2_strategy = Box::new(strategy);
self
}
}
impl<T: MaybeSend + 'static, Source> TeeBuilder<T, Source, DefaultStream>
where
T: Clone,
Source: MaybeSend + 'static,
{
pub fn prepare(
self,
) -> Result<
(
ReadableStream<T, TeeSource<T>, DefaultStream, Unlocked>,
ReadableStream<T, TeeSource<T>, DefaultStream, Unlocked>,
impl Future<Output = ()>, // coordinator future
impl Future<Output = ()>, // branch1 future
impl Future<Output = ()>, // branch2 future
),
StreamError,
> {
self.stream
.tee_inner(self.mode, self.branch1_strategy, self.branch2_strategy)
}
pub fn spawn<F, R>(
self,
spawn_fn: F,
) -> Result<
(
ReadableStream<T, TeeSource<T>, DefaultStream, Unlocked>,
ReadableStream<T, TeeSource<T>, DefaultStream, Unlocked>,
),
StreamError,
>
where
F: FnOnce(crate::platform::PlatformFuture<'static, ()>) -> R,
{
let (stream1, stream2, coord_fut, rfut1, rfut2) = self.prepare()?;
let fut = async move {
futures::join!(coord_fut, rfut1, rfut2);
};
spawn_fn(Box::pin(fut));
Ok((stream1, stream2))
}
pub fn spawn_ref<F, R>(
self,
spawn_fn: &'static F,
) -> Result<
(
ReadableStream<T, TeeSource<T>, DefaultStream, Unlocked>,
ReadableStream<T, TeeSource<T>, DefaultStream, Unlocked>,
),
StreamError,
>
where
F: Fn(crate::platform::PlatformFuture<'static, ()>) -> R,
{
let (stream1, stream2, coord_fut, rfut1, rfut2) = self.prepare()?;
let fut = async move {
futures::join!(coord_fut, rfut1, rfut2);
};
spawn_fn(Box::pin(fut));
Ok((stream1, stream2))
}
pub fn spawn_parts<R1, R2, R3, F1, F2, F3>(
self,
coordinator_spawn: F1,
branch1_spawn: F2,
branch2_spawn: F3,
) -> Result<
(
ReadableStream<T, TeeSource<T>, DefaultStream, Unlocked>,
ReadableStream<T, TeeSource<T>, DefaultStream, Unlocked>,
),
StreamError,
>
where
F1: FnOnce(crate::platform::PlatformFuture<'static, ()>) -> R1,
F2: FnOnce(crate::platform::PlatformFuture<'static, ()>) -> R2,
F3: FnOnce(crate::platform::PlatformFuture<'static, ()>) -> R3,
{
let (stream1, stream2, coord_fut, rfut1, rfut2) = self.prepare()?;
coordinator_spawn(Box::pin(coord_fut));
branch1_spawn(Box::pin(rfut1));
branch2_spawn(Box::pin(rfut2));
Ok((stream1, stream2))
}
pub fn spawn_parts_ref<R1, R2, R3, F1, F2, F3>(
self,
coordinator_spawn: &'static F1,
branch1_spawn: &'static F2,
branch2_spawn: &'static F3,
) -> Result<
(
ReadableStream<T, TeeSource<T>, DefaultStream, Unlocked>,
ReadableStream<T, TeeSource<T>, DefaultStream, Unlocked>,
),
StreamError,
>
where
F1: Fn(crate::platform::PlatformFuture<'static, ()>) -> R1,
F2: Fn(crate::platform::PlatformFuture<'static, ()>) -> R2,
F3: Fn(crate::platform::PlatformFuture<'static, ()>) -> R3,
{
let (stream1, stream2, coord_fut, rfut1, rfut2) = self.prepare()?;
coordinator_spawn(Box::pin(coord_fut));
branch1_spawn(Box::pin(rfut1));
branch2_spawn(Box::pin(rfut2));
Ok((stream1, stream2))
}
}
impl<Source> TeeBuilder<Bytes, Source, ByteStream>
where
Source: ReadableByteSource,
{
pub fn prepare(
self,
) -> Result<
(
ReadableStream<Bytes, TeeSource<Bytes>, ByteStream, Unlocked>,
ReadableStream<Bytes, TeeSource<Bytes>, ByteStream, Unlocked>,
impl Future<Output = ()>, // coordinator future
impl Future<Output = ()>, // branch1 future
impl Future<Output = ()>, // branch2 future
),
StreamError,
> {
self.stream
.tee_inner_bytes(self.mode, self.branch1_strategy, self.branch2_strategy)
}
pub fn spawn<F, R>(
self,
spawn_fn: F,
) -> Result<
(
ReadableStream<Bytes, TeeSource<Bytes>, ByteStream, Unlocked>,
ReadableStream<Bytes, TeeSource<Bytes>, ByteStream, Unlocked>,
),
StreamError,
>
where
F: FnOnce(crate::platform::PlatformFuture<'static, ()>) -> R,
{
let (stream1, stream2, coord_fut, rfut1, rfut2) = self.prepare()?;
let fut = async move {
futures::join!(coord_fut, rfut1, rfut2);
};
spawn_fn(Box::pin(fut));
Ok((stream1, stream2))
}
pub fn spawn_ref<F, R>(
self,
spawn_fn: &'static F,
) -> Result<
(
ReadableStream<Bytes, TeeSource<Bytes>, ByteStream, Unlocked>,
ReadableStream<Bytes, TeeSource<Bytes>, ByteStream, Unlocked>,
),
StreamError,
>
where
F: Fn(crate::platform::PlatformFuture<'static, ()>) -> R,
{
let (stream1, stream2, coord_fut, rfut1, rfut2) = self.prepare()?;
let fut = async move {
futures::join!(coord_fut, rfut1, rfut2);
};
spawn_fn(Box::pin(fut));
Ok((stream1, stream2))
}
pub fn spawn_parts<R1, R2, R3, F1, F2, F3>(
self,
coordinator_spawn: F1,
branch1_spawn: F2,
branch2_spawn: F3,
) -> Result<
(
ReadableStream<Bytes, TeeSource<Bytes>, ByteStream, Unlocked>,
ReadableStream<Bytes, TeeSource<Bytes>, ByteStream, Unlocked>,
),
StreamError,
>
where
F1: FnOnce(crate::platform::PlatformFuture<'static, ()>) -> R1,
F2: FnOnce(crate::platform::PlatformFuture<'static, ()>) -> R2,
F3: FnOnce(crate::platform::PlatformFuture<'static, ()>) -> R3,
{
let (stream1, stream2, coord_fut, rfut1, rfut2) = self.prepare()?;
coordinator_spawn(Box::pin(coord_fut));
branch1_spawn(Box::pin(rfut1));
branch2_spawn(Box::pin(rfut2));
Ok((stream1, stream2))
}
pub fn spawn_parts_ref<R1, R2, R3, F1, F2, F3>(
self,
coordinator_spawn: &'static F1,
branch1_spawn: &'static F2,
branch2_spawn: &'static F3,
) -> Result<
(
ReadableStream<Bytes, TeeSource<Bytes>, ByteStream, Unlocked>,
ReadableStream<Bytes, TeeSource<Bytes>, ByteStream, Unlocked>,
),
StreamError,
>
where
F1: Fn(crate::platform::PlatformFuture<'static, ()>) -> R1,
F2: Fn(crate::platform::PlatformFuture<'static, ()>) -> R2,
F3: Fn(crate::platform::PlatformFuture<'static, ()>) -> R3,
{
let (stream1, stream2, coord_fut, rfut1, rfut2) = self.prepare()?;
coordinator_spawn(Box::pin(coord_fut));
branch1_spawn(Box::pin(rfut1));
branch2_spawn(Box::pin(rfut2));
Ok((stream1, stream2))
}
}
impl<T: MaybeSend + 'static, Source, S> ReadableStream<T, Source, S, Unlocked>
where
T: Clone,
Source: MaybeSend + 'static,
S: StreamTypeMarker,
{
fn tee_inner(
self,
mode: BackpressureMode,
branch1_strategy: crate::platform::BoxedStrategyStatic<T>,
branch2_strategy: crate::platform::BoxedStrategyStatic<T>,
) -> Result<
(
ReadableStream<T, TeeSource<T>, DefaultStream, Unlocked>,
ReadableStream<T, TeeSource<T>, DefaultStream, Unlocked>,
impl Future<Output = ()>, // coordinator future
impl Future<Output = ()>, // branch1 future
impl Future<Output = ()>, // branch2 future
),
StreamError,
> {
let (_, reader) = self.get_reader()?;
let setup = TeeSetup::new();
let (source1, source2) = setup.sources();
let (stream1, rfut1) = ReadableStream::new_inner(source1, branch1_strategy);
let (stream2, rfut2) = ReadableStream::new_inner(source2, branch2_strategy);
let coordinator = setup.coordinator(
reader,
(*stream1.controller).clone(),
(*stream2.controller).clone(),
mode,
);
let coordinator_fut = coordinator.run();
Ok((stream1, stream2, coordinator_fut, rfut1, rfut2))
}
pub fn tee(self) -> TeeBuilder<T, Source, S> {
TeeBuilder::new(self)
}
}
impl<Source> ReadableStream<Bytes, Source, ByteStream, Unlocked>
where
Source: ReadableByteSource,
{
fn tee_inner_bytes(
self,
mode: BackpressureMode,
branch1_strategy: crate::platform::BoxedStrategyStatic<Bytes>,
branch2_strategy: crate::platform::BoxedStrategyStatic<Bytes>,
) -> Result<
(
ReadableStream<Bytes, TeeSource<Bytes>, ByteStream, Unlocked>,
ReadableStream<Bytes, TeeSource<Bytes>, ByteStream, Unlocked>,
impl Future<Output = ()>, // coordinator future
impl Future<Output = ()>, // branch1 future
impl Future<Output = ()>, // branch2 future
),
StreamError,
> {
let (_, reader) = self.get_reader()?;
let setup = TeeSetup::new();
let (source1, source2) = setup.sources::<Bytes>();
let (stream1, rfut1) = ReadableStream::new_bytes_inner(source1, branch1_strategy);
let (stream2, rfut2) = ReadableStream::new_bytes_inner(source2, branch2_strategy);
let coordinator = setup.coordinator(
reader,
(*stream1.controller).clone(),
(*stream2.controller).clone(),
mode,
);
let coordinator_fut = coordinator.run();
Ok((stream1, stream2, coordinator_fut, rfut1, rfut2))
}
}
pub struct PipeBuilder<T, O, Source, S>
where
T: MaybeSend + 'static,
O: MaybeSend + 'static,
S: StreamTypeMarker,
{
source_stream: ReadableStream<T, Source, S, Unlocked>,
transform: TransformStream<T, O>,
options: Option<StreamPipeOptions>,
}
impl<T: MaybeSend + 'static, O: MaybeSend + 'static, Source, S> PipeBuilder<T, O, Source, S>
where
Source: MaybeSend + 'static,
S: StreamTypeMarker,
{
pub fn new(
source_stream: ReadableStream<T, Source, S, Unlocked>,
transform: TransformStream<T, O>,
options: Option<StreamPipeOptions>,
) -> Self {
Self {
source_stream,
transform,
options,
}
}
pub fn prepare(
self,
) -> (
ReadableStream<O, TransformReadableSource<O>, DefaultStream, Unlocked>,
impl Future<Output = StreamResult<()>>,
) {
let (readable, writable) = self.transform.split();
let pipe_future = async move { self.source_stream.pipe_to(&writable, self.options).await };
(readable, pipe_future)
}
pub fn spawn<SpawnFn, R>(
self,
spawn_fn: SpawnFn,
) -> ReadableStream<O, TransformReadableSource<O>, DefaultStream, Unlocked>
where
SpawnFn: FnOnce(crate::platform::PlatformFuture<'static, ()>) -> R,
{
let (readable, pipe_future) = self.prepare();
let fut = Box::pin(async move {
let _ = pipe_future.await;
});
spawn_fn(fut);
readable
}
pub fn spawn_ref<SpawnFn, R>(
self,
spawn_fn: &'static SpawnFn,
) -> ReadableStream<O, TransformReadableSource<O>, DefaultStream, Unlocked>
where
SpawnFn: Fn(crate::platform::PlatformFuture<'static, ()>) -> R,
{
let (readable, pipe_future) = self.prepare();
let fut = Box::pin(async move {
let _ = pipe_future.await;
});
spawn_fn(fut);
readable
}
}
impl<T: MaybeSend + 'static, Source, S> ReadableStream<T, Source, S, Unlocked>
where
Source: MaybeSend + 'static,
S: StreamTypeMarker,
{
pub fn pipe_through<O>(
self,
transform: TransformStream<T, O>,
options: Option<StreamPipeOptions>,
) -> PipeBuilder<T, O, Source, S>
where
O: MaybeSend + 'static,
{
PipeBuilder::new(self, transform, options)
}
}
#[derive(Default)]
pub struct StreamPipeOptions {
pub prevent_close: bool,
pub prevent_abort: bool,
pub prevent_cancel: bool,
pub signal: Option<AbortSignal>,
}
impl<Source> ReadableStream<Bytes, Source, ByteStream, Unlocked>
where
Source: ReadableByteSource,
{
pub(crate) fn new_bytes_inner(
source: Source,
strategy: Box<dyn QueuingStrategy<Bytes> + 'static>,
) -> (Self, impl Future<Output = ()>) {
let (command_tx, command_rx) = unbounded();
let queue_total_size = SharedPtr::new(AtomicUsize::new(0));
let closed = SharedPtr::new(AtomicBool::new(false));
let errored = SharedPtr::new(AtomicBool::new(false));
let locked = SharedPtr::new(AtomicBool::new(false));
let stored_error = SharedPtr::new(RwLock::new(None));
let high_water_mark = SharedPtr::new(AtomicUsize::new(strategy.high_water_mark()));
let desired_size = SharedPtr::new(AtomicIsize::new(strategy.high_water_mark() as isize));
let byte_state = ByteStreamState::new(source, strategy.high_water_mark());
let controller = ReadableByteStreamController::new(byte_state.clone());
let task_state = byte_state.clone();
let task_fut = readable_byte_stream_task(task_state, command_rx, controller.clone());
let stream = Self {
command_tx,
queue_total_size,
high_water_mark,
desired_size,
closed,
errored,
locked,
stored_error,
controller: SharedPtr::new(controller),
byte_state: Some(byte_state),
poll_read_rx: None,
_phantom: PhantomData,
};
(stream, task_fut)
}
}
impl<T: MaybeSend + 'static, Source: ReadableSource<T>>
ReadableStream<T, Source, DefaultStream, Unlocked>
{
pub(crate) fn new_inner(
source: Source,
strategy: crate::platform::BoxedStrategy<T>,
) -> (Self, impl Future<Output = ()>) {
let (command_tx, command_rx) = unbounded();
let (ctrl_tx, ctrl_rx) = unbounded();
let queue_total_size = SharedPtr::new(AtomicUsize::new(0));
let closed = SharedPtr::new(AtomicBool::new(false));
let errored = SharedPtr::new(AtomicBool::new(false));
let locked = SharedPtr::new(AtomicBool::new(false));
let stored_error = SharedPtr::new(RwLock::new(None));
let high_water_mark = SharedPtr::new(AtomicUsize::new(strategy.high_water_mark()));
let desired_size = SharedPtr::new(AtomicIsize::new(strategy.high_water_mark() as isize));
let inner = ReadableStreamInner::new(source, strategy);
let close_requested = SharedPtr::new(AtomicBool::new(false));
let error_requested = SharedPtr::new(AtomicBool::new(false));
let controller = ReadableStreamDefaultController::new(
ctrl_tx.clone(),
SharedPtr::clone(&queue_total_size),
SharedPtr::clone(&high_water_mark),
SharedPtr::clone(&desired_size),
SharedPtr::clone(&closed),
SharedPtr::clone(&errored),
SharedPtr::clone(&close_requested),
SharedPtr::clone(&error_requested),
);
let task_fut = readable_stream_task(
command_rx,
ctrl_rx,
inner,
SharedPtr::clone(&queue_total_size),
SharedPtr::clone(&high_water_mark),
SharedPtr::clone(&desired_size),
SharedPtr::clone(&closed),
SharedPtr::clone(&errored),
SharedPtr::clone(&stored_error),
ctrl_tx,
controller.clone(),
);
let stream = Self {
command_tx,
queue_total_size,
high_water_mark,
desired_size,
closed,
errored,
locked,
stored_error,
controller: controller.into(),
byte_state: None,
poll_read_rx: None,
_phantom: PhantomData,
};
(stream, Box::pin(task_fut))
}
}
impl<T: MaybeSend + 'static, Source, StreamType> ReadableStream<T, Source, StreamType, Unlocked>
where
StreamType: StreamTypeMarker,
{
pub fn get_reader(
&self,
) -> Result<
(
ReadableStream<T, Source, StreamType, Locked>,
ReadableStreamDefaultReader<T, Source, StreamType, Locked>,
),
StreamError,
> {
if self
.locked
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.is_err()
{
return Err("Stream already locked".into());
}
let locked_stream = ReadableStream {
command_tx: self.command_tx.clone(),
queue_total_size: SharedPtr::clone(&self.queue_total_size),
high_water_mark: SharedPtr::clone(&self.high_water_mark),
desired_size: SharedPtr::clone(&self.desired_size),
closed: SharedPtr::clone(&self.closed),
errored: SharedPtr::clone(&self.errored),
locked: SharedPtr::clone(&self.locked),
stored_error: SharedPtr::clone(&self.stored_error),
controller: self.controller.clone(),
byte_state: self.byte_state.clone(),
poll_read_rx: None,
_phantom: PhantomData,
};
let reader = ReadableStreamDefaultReader::new(ReadableStream {
command_tx: self.command_tx.clone(),
queue_total_size: self.queue_total_size.clone(),
high_water_mark: self.high_water_mark.clone(),
desired_size: self.desired_size.clone(),
closed: self.closed.clone(),
errored: self.errored.clone(),
locked: self.locked.clone(),
stored_error: self.stored_error.clone(),
controller: self.controller.clone(),
byte_state: self.byte_state.clone(),
poll_read_rx: None,
_phantom: PhantomData,
});
Ok((locked_stream, reader))
}
}
impl<Source> ReadableStream<Bytes, Source, ByteStream, Unlocked>
where
Source: ReadableByteSource,
{
pub fn get_byob_reader(
&self,
) -> Result<
(
ReadableStream<Bytes, Source, ByteStream, Locked>,
ReadableStreamBYOBReader<Source, Locked>,
),
StreamError,
> {
if self
.locked
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.is_err()
{
return Err(StreamError::from("Stream already locked"));
}
let locked_stream = ReadableStream {
command_tx: self.command_tx.clone(),
queue_total_size: SharedPtr::clone(&self.queue_total_size),
high_water_mark: SharedPtr::clone(&self.high_water_mark),
desired_size: SharedPtr::clone(&self.desired_size),
closed: SharedPtr::clone(&self.closed),
errored: SharedPtr::clone(&self.errored),
locked: SharedPtr::clone(&self.locked),
stored_error: SharedPtr::clone(&self.stored_error),
controller: self.controller.clone(),
byte_state: self.byte_state.clone(),
poll_read_rx: None,
_phantom: PhantomData,
};
let reader = ReadableStreamBYOBReader::new(ReadableStream {
command_tx: self.command_tx.clone(),
queue_total_size: self.queue_total_size.clone(),
high_water_mark: self.high_water_mark.clone(),
desired_size: self.desired_size.clone(),
closed: self.closed.clone(),
errored: self.errored.clone(),
locked: self.locked.clone(),
stored_error: self.stored_error.clone(),
controller: self.controller.clone(),
byte_state: self.byte_state.clone(),
poll_read_rx: None,
_phantom: PhantomData,
});
Ok((locked_stream, reader))
}
}
impl<T: MaybeSend + 'static, Source, StreamType, LockState> Stream
for ReadableStream<T, Source, StreamType, LockState>
where
StreamType: StreamTypeMarker,
{
type Item = StreamResult<T>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = self.get_mut();
if let Some(rx) = this.poll_read_rx.as_mut() {
match Pin::new(rx).poll(cx) {
Poll::Ready(Ok(Ok(Some(chunk)))) => {
this.poll_read_rx = None;
return Poll::Ready(Some(Ok(chunk)));
}
Poll::Ready(Ok(Ok(None))) => {
this.poll_read_rx = None;
return Poll::Ready(None); }
Poll::Ready(Ok(Err(e))) => {
this.poll_read_rx = None;
return Poll::Ready(Some(Err(e)));
}
Poll::Ready(Err(_)) => {
this.poll_read_rx = None;
return Poll::Ready(Some(Err(StreamError::TaskDropped)));
}
Poll::Pending => return Poll::Pending,
}
}
let (tx, rx) = oneshot::channel();
if this
.command_tx
.unbounded_send(StreamCommand::Read { completion: tx })
.is_err()
{
return Poll::Ready(Some(Err(StreamError::TaskDropped)));
}
this.poll_read_rx = Some(rx);
match Pin::new(this.poll_read_rx.as_mut().unwrap()).poll(cx) {
Poll::Ready(Ok(Ok(Some(chunk)))) => {
this.poll_read_rx = None;
Poll::Ready(Some(Ok(chunk)))
}
Poll::Ready(Ok(Ok(None))) => {
this.poll_read_rx = None;
Poll::Ready(None)
}
Poll::Ready(Ok(Err(e))) => {
this.poll_read_rx = None;
Poll::Ready(Some(Err(e)))
}
Poll::Ready(Err(_)) => {
this.poll_read_rx = None;
Poll::Ready(Some(Err(StreamError::TaskDropped)))
}
Poll::Pending => Poll::Pending,
}
}
}
impl<T: MaybeSend + 'static, Source, LockState> AsyncRead
for ReadableStream<T, Source, ByteStream, LockState>
where
Source: ReadableByteSource,
{
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut [u8],
) -> Poll<IoResult<usize>> {
match self.byte_state.as_ref().unwrap().poll_read_into(cx, buf) {
Poll::Ready(Ok(bytes)) => Poll::Ready(Ok(bytes)),
Poll::Ready(Err(stream_err)) => Poll::Ready(Err(std::io::Error::new(
std::io::ErrorKind::Other,
stream_err.to_string(),
))),
Poll::Pending => Poll::Pending,
}
}
}
pub struct IteratorSource<I: MaybeSend + 'static> {
iter: I,
}
impl<I: MaybeSend + 'static, T: MaybeSend + 'static> ReadableSource<T> for IteratorSource<I>
where
I: Iterator<Item = T> + MaybeSend + 'static,
{
async fn pull(
&mut self,
controller: &mut ReadableStreamDefaultController<T>,
) -> StreamResult<()> {
if let Some(item) = self.iter.next() {
controller.enqueue(item)?;
} else {
controller.close()?;
}
Ok(())
}
}
pub struct AsyncStreamSource<S: MaybeSend + 'static> {
stream: S,
}
impl<S: MaybeSend + 'static, T: MaybeSend + 'static> ReadableSource<T> for AsyncStreamSource<S>
where
S: Stream<Item = T> + Unpin + MaybeSend + 'static,
{
async fn pull(
&mut self,
controller: &mut ReadableStreamDefaultController<T>,
) -> StreamResult<()> {
if let Some(item) = self.stream.next().await {
controller.enqueue(item)?;
} else {
controller.close()?;
}
Ok(())
}
}
pub struct ReadableStreamDefaultReader<T: MaybeSend + 'static, Source, StreamType, LockState>(
ReadableStream<T, Source, StreamType, LockState>,
)
where
StreamType: StreamTypeMarker;
impl<T: MaybeSend + 'static, Source, StreamType, LockState>
ReadableStreamDefaultReader<T, Source, StreamType, LockState>
where
StreamType: StreamTypeMarker,
{
pub fn new(stream: ReadableStream<T, Source, StreamType, LockState>) -> Self {
ReadableStreamDefaultReader(stream)
}
fn is_byte_stream(&self) -> bool {
self.0.byte_state.is_some()
}
pub async fn closed(&self) -> StreamResult<()> {
if self.is_byte_stream() {
return self.0.byte_state.as_ref().unwrap().closed().await;
}
poll_fn(|cx| {
if self.0.errored.load(Ordering::Acquire) {
let error = self
.0
.stored_error
.read()
.clone()
.unwrap_or_else(|| "Stream is errored".into());
return Poll::Ready(Err(error));
}
if self.0.closed.load(Ordering::Acquire) {
return Poll::Ready(Ok(()));
}
let waker = cx.waker().clone();
let _ = self
.0
.command_tx
.unbounded_send(StreamCommand::RegisterClosedWaker { waker });
Poll::Pending
})
.await
}
pub async fn cancel(&self, reason: Option<String>) -> StreamResult<()> {
let (tx, rx) = oneshot::channel();
self.0
.command_tx
.unbounded_send(StreamCommand::Cancel {
reason,
completion: tx,
})
.map_err(|_| StreamError::TaskDropped)?;
rx.await.unwrap_or_else(|_| Err(StreamError::TaskDropped))
}
pub fn ready(&self) -> impl Future<Output = StreamResult<()>> + '_ {
poll_fn(|cx| {
if self.0.errored.load(Ordering::Acquire) {
let error = self
.0
.stored_error
.read()
.clone()
.unwrap_or_else(|| "Stream is errored".into());
return Poll::Ready(Err(error));
}
if self.0.closed.load(Ordering::Acquire)
|| self.0.queue_total_size.load(Ordering::Acquire) > 0
{
return Poll::Ready(Ok(()));
}
let waker = cx.waker().clone();
let _ = self
.0
.command_tx
.unbounded_send(StreamCommand::RegisterReadyWaker { waker });
if self.0.errored.load(Ordering::Acquire) {
let error = self
.0
.stored_error
.read()
.clone()
.unwrap_or_else(|| "Stream is errored".into());
return Poll::Ready(Err(error));
}
if self.0.closed.load(Ordering::Acquire)
|| self.0.queue_total_size.load(Ordering::Acquire) > 0
{
return Poll::Ready(Ok(()));
}
Poll::Pending
})
}
pub async fn read(&self) -> StreamResult<Option<T>> {
let (tx, rx) = oneshot::channel();
self.0
.command_tx
.unbounded_send(StreamCommand::Read { completion: tx })
.map_err(|_| StreamError::TaskDropped)?;
rx.await.unwrap_or_else(|_| Err(StreamError::TaskDropped))
}
pub fn release_lock(self) -> ReadableStream<T, Source, StreamType, Unlocked> {
self.0.locked.store(false, Ordering::Release);
ReadableStream {
command_tx: self.0.command_tx.clone(),
queue_total_size: self.0.queue_total_size.clone(),
high_water_mark: self.0.high_water_mark.clone(),
desired_size: self.0.desired_size.clone(),
closed: self.0.closed.clone(),
errored: self.0.errored.clone(),
locked: self.0.locked.clone(),
stored_error: self.0.stored_error.clone(),
controller: self.0.controller.clone(),
byte_state: self.0.byte_state.clone(),
poll_read_rx: None,
_phantom: PhantomData,
}
}
}
impl<T: MaybeSend + 'static, Source, StreamType, LockState> Drop
for ReadableStreamDefaultReader<T, Source, StreamType, LockState>
where
StreamType: StreamTypeMarker,
{
fn drop(&mut self) {
self.0.locked.store(false, Ordering::Release);
}
}
pub struct ReadableStreamBYOBReader<Source, LockState>(
ReadableStream<Bytes, Source, ByteStream, LockState>,
);
impl<Source, LockState> ReadableStreamBYOBReader<Source, LockState>
where
Source: ReadableByteSource,
{
pub fn new(stream: ReadableStream<Bytes, Source, ByteStream, LockState>) -> Self {
ReadableStreamBYOBReader(stream)
}
pub async fn closed(&self) -> Result<(), StreamError> {
self.0.controller.byte_state.closed().await
}
pub async fn cancel(&self, reason: Option<String>) -> Result<(), StreamError> {
let (tx, rx) = oneshot::channel();
self.0
.command_tx
.unbounded_send(StreamCommand::Cancel {
reason,
completion: tx,
})
.map_err(|_| StreamError::TaskDropped)?;
rx.await.unwrap_or_else(|_| Err(StreamError::TaskDropped))
}
pub async fn read(&self, buf: &mut [u8]) -> StreamResult<usize> {
poll_fn(|cx| self.0.byte_state.as_ref().unwrap().poll_read_into(cx, buf)).await
}
pub async fn read_owned(&self, buf: BytesMut) -> StreamResult<(BytesMut, usize)> {
match self.0.byte_state.as_ref().unwrap().begin_pull_into(buf) {
PullIntoOutcome::Ready(buf, n) => Ok((buf, n)),
PullIntoOutcome::Errored(err) => Err(err),
PullIntoOutcome::Registered(rx) => rx.await.unwrap_or_else(|_canceled| {
Err(StreamError::from("stream dropped before the read completed"))
}),
}
}
pub fn release_lock(self) -> ReadableStream<Bytes, Source, ByteStream, Unlocked> {
self.0.locked.store(false, Ordering::Release);
ReadableStream {
command_tx: self.0.command_tx.clone(),
queue_total_size: self.0.queue_total_size.clone(),
high_water_mark: self.0.high_water_mark.clone(),
desired_size: self.0.desired_size.clone(),
closed: self.0.closed.clone(),
errored: self.0.errored.clone(),
locked: self.0.locked.clone(),
stored_error: self.0.stored_error.clone(),
controller: self.0.controller.clone(),
byte_state: self.0.byte_state.clone(),
poll_read_rx: None,
_phantom: PhantomData,
}
}
}
impl<Source, LockState> Drop for ReadableStreamBYOBReader<Source, LockState> {
fn drop(&mut self) {
self.0.locked.store(false, Ordering::Release);
}
}
fn update_desired_size(
queue_total_size: &SharedPtr<AtomicUsize>,
high_water_mark: &SharedPtr<AtomicUsize>,
desired_size: &SharedPtr<AtomicIsize>,
closed: &SharedPtr<AtomicBool>,
errored: &SharedPtr<AtomicBool>,
) {
if closed.load(Ordering::Acquire) || errored.load(Ordering::Acquire) {
desired_size.store(0, Ordering::Release);
return;
}
let hwm = high_water_mark.load(Ordering::Acquire) as isize;
let current = queue_total_size.load(Ordering::Acquire) as isize;
let new_desired_size = hwm - current;
desired_size.store(new_desired_size, Ordering::Release);
}
async fn pull_with_cancel<T, Source>(
source: &mut Source,
controller: &mut ReadableStreamDefaultController<T>,
cancel: oneshot::Receiver<()>,
) -> StreamResult<()>
where
T: MaybeSend + 'static,
Source: ReadableSource<T>,
{
use futures::FutureExt;
let pull_fut = source.pull(controller).fuse();
let mut cancel_fut = cancel.fuse();
futures::pin_mut!(pull_fut);
futures::select! {
r = pull_fut => r,
_ = cancel_fut => Ok(()),
}
}
async fn readable_stream_task<T: 'static, Source>(
mut command_rx: UnboundedReceiver<StreamCommand<T>>,
mut ctrl_rx: UnboundedReceiver<ControllerMsg<T>>,
mut inner: ReadableStreamInner<T, Source>,
queue_total_size: SharedPtr<AtomicUsize>,
high_water_mark: SharedPtr<AtomicUsize>,
desired_size: SharedPtr<AtomicIsize>,
closed: SharedPtr<AtomicBool>,
errored: SharedPtr<AtomicBool>,
stored_error: SharedPtr<RwLock<Option<StreamError>>>,
_ctrl_tx: UnboundedSender<ControllerMsg<T>>,
mut controller: ReadableStreamDefaultController<T>,
) where
T: MaybeSend,
Source: ReadableSource<T>,
{
if let Some(mut source) = inner.source.take() {
match source.start(&mut controller).await {
Ok(()) => {
inner.source = Some(source);
}
Err(err) => {
inner.state = StreamState::Errored;
errored.store(true, Ordering::Release);
inner.stored_error = Some(err.clone());
{
let mut guard = stored_error.write();
*guard = Some(err.clone());
}
desired_size.store(0, Ordering::Release);
inner.closed_wakers.wake_all();
inner.ready_wakers.wake_all();
}
}
}
let mut pull_future: Option<
crate::platform::PlatformBoxFutureStatic<(Source, StreamResult<()>)>,
> = None;
let mut pull_cancel_tx: Option<oneshot::Sender<()>> = None;
let mut cancel_future: Option<crate::platform::PlatformBoxFutureStatic<StreamResult<()>>> =
None;
let mut needs_pull = true;
let mut close_pending = false;
poll_fn(|cx| {
while let Poll::Ready(Some(msg)) = ctrl_rx.poll_next_unpin(cx) {
match msg {
ControllerMsg::Enqueue { chunk } => {
if close_pending || inner.state != StreamState::Readable {
continue;
}
if let Some(tx) = inner.pending_reads.pop_front() {
let _ = tx.send(Ok(Some(chunk)));
} else {
let chunk_size = inner.strategy.size(&chunk);
inner.queue.push_back(chunk);
inner.queue_total_size += chunk_size;
queue_total_size.store(inner.queue_total_size, Ordering::Release);
update_desired_size(
&queue_total_size,
&high_water_mark,
&desired_size,
&closed,
&errored,
);
inner.ready_wakers.wake_all();
}
needs_pull = true;
}
ControllerMsg::Close => {
if inner.state == StreamState::Readable {
desired_size.store(0, Ordering::Release);
if inner.queue.is_empty() {
inner.state = StreamState::Closed;
closed.store(true, Ordering::Release);
while let Some(tx) = inner.pending_reads.pop_front() {
let _ = tx.send(Ok(None));
}
inner.closed_wakers.wake_all();
inner.ready_wakers.wake_all();
} else {
close_pending = true;
}
}
}
ControllerMsg::Error(err) => {
if inner.state == StreamState::Readable {
close_pending = false;
inner.state = StreamState::Errored;
errored.store(true, Ordering::Release);
inner.stored_error = Some(err.clone());
desired_size.store(0, Ordering::Release);
{
let mut guard = stored_error.write();
*guard = Some(err.clone());
}
inner.queue.clear();
inner.queue_total_size = 0;
queue_total_size.store(0, Ordering::Release);
while let Some(tx) = inner.pending_reads.pop_front() {
let _ = tx.send(Err(err.clone()));
}
inner.closed_wakers.wake_all();
inner.ready_wakers.wake_all();
}
}
}
}
while let Poll::Ready(Some(cmd)) = command_rx.poll_next_unpin(cx) {
match cmd {
StreamCommand::Read { completion } => {
if inner.state == StreamState::Errored {
let _ = completion.send(Err(inner.get_stored_error()));
continue;
}
if inner.state == StreamState::Closed && inner.queue.is_empty() {
let _ = completion.send(Ok(None));
continue;
}
if let Some(chunk) = inner.queue.pop_front() {
let chunk_size = inner.strategy.size(&chunk);
inner.queue_total_size -= chunk_size;
queue_total_size.store(inner.queue_total_size, Ordering::Release);
update_desired_size(
&queue_total_size,
&high_water_mark,
&desired_size,
&closed,
&errored,
);
let _ = completion.send(Ok(Some(chunk)));
if close_pending && inner.queue.is_empty() {
close_pending = false;
inner.state = StreamState::Closed;
closed.store(true, Ordering::Release);
desired_size.store(0, Ordering::Release);
inner.closed_wakers.wake_all();
inner.ready_wakers.wake_all();
}
} else {
inner.pending_reads.push_back(completion);
}
needs_pull = true;
}
StreamCommand::Cancel { reason, completion } => {
if inner.state == StreamState::Closed {
let _ = completion.send(Ok(()));
continue;
}
if inner.state == StreamState::Errored {
let _ = completion.send(Err(inner.get_stored_error()));
continue;
}
if inner.cancel_requested {
inner.cancel_completions.push(completion);
} else {
close_pending = false;
inner.cancel_requested = true;
inner.cancel_reason = reason.clone();
inner.cancel_completions.push(completion);
inner.state = StreamState::Closed;
closed.store(true, Ordering::Release);
inner.queue.clear();
inner.queue_total_size = 0;
queue_total_size.store(0, Ordering::Release);
while let Some(tx) = inner.pending_reads.pop_front() {
let _ = tx.send(Err(StreamError::Canceled));
}
inner.closed_wakers.wake_all();
inner.ready_wakers.wake_all();
if let Some(mut source) = inner.source.take() {
let reason_clone = reason;
cancel_future =
Some(Box::pin(async move { source.cancel(reason_clone).await }));
} else if let Some(tx) = pull_cancel_tx.take() {
drop(tx);
}
}
}
StreamCommand::RegisterReadyWaker { waker } => {
inner.ready_wakers.register(&waker);
if !inner.queue.is_empty()
|| inner.state == StreamState::Closed
|| inner.state == StreamState::Errored
{
inner.ready_wakers.wake_all();
}
}
StreamCommand::RegisterClosedWaker { waker } => {
inner.closed_wakers.register(&waker);
if inner.state == StreamState::Closed || inner.state == StreamState::Errored {
inner.closed_wakers.wake_all();
}
}
}
}
if let Some(ref mut fut) = cancel_future {
match fut.as_mut().poll(cx) {
Poll::Ready(result) => {
inner.cancel_requested = false;
for tx in inner.cancel_completions.drain(..) {
let _ = tx.send(result.clone());
}
cancel_future = None;
cx.waker().wake_by_ref();
}
Poll::Pending => {}
}
}
let current_ds = desired_size.load(std::sync::atomic::Ordering::Acquire);
if needs_pull
&& inner.state == StreamState::Readable
&& !close_pending
&& !inner.pulling
&& !inner.cancel_requested
{
needs_pull = false;
if (current_ds > 0 || !inner.pending_reads.is_empty()) && inner.source.is_some() {
let source = inner.source.take().unwrap();
inner.pulling = true;
let mut ctrl = controller.clone();
let (cancel_tx, cancel_rx) = oneshot::channel::<()>();
pull_cancel_tx = Some(cancel_tx);
pull_future = Some(Box::pin(async move {
let mut source = source;
let result = pull_with_cancel(&mut source, &mut ctrl, cancel_rx).await;
(source, result)
}));
}
}
if let Some(ref mut fut) = pull_future {
match fut.as_mut().poll(cx) {
Poll::Ready((source, result)) => {
inner.pulling = false;
pull_cancel_tx = None; match result {
Ok(()) => {
inner.source = Some(source);
if inner.cancel_requested && cancel_future.is_none() {
if let Some(mut src) = inner.source.take() {
let reason = inner.cancel_reason.clone();
cancel_future =
Some(Box::pin(async move { src.cancel(reason).await }));
}
}
}
Err(err) => {
inner.state = StreamState::Errored;
errored.store(true, Ordering::Release);
inner.stored_error = Some(err.clone());
{
let mut guard = stored_error.write();
*guard = Some(err.clone());
}
while let Some(tx) = inner.pending_reads.pop_front() {
let _ = tx.send(Err(err.clone()));
}
if inner.cancel_requested {
inner.cancel_requested = false;
for tx in inner.cancel_completions.drain(..) {
let _ = tx.send(Ok(()));
}
}
inner.closed_wakers.wake_all();
inner.ready_wakers.wake_all();
}
}
pull_future = None;
cx.waker().wake_by_ref();
}
Poll::Pending => {}
}
}
Poll::Pending::<()>
})
.await;
}
async fn readable_byte_stream_task<Source>(
byte_state: SharedPtr<ByteStreamState<Source>>,
mut command_rx: UnboundedReceiver<StreamCommand<Bytes>>,
mut controller: ReadableByteStreamController,
) where
Source: ReadableByteSource,
{
let _ = byte_state.start_source(&controller).await;
let mut pending_reads: VecDeque<oneshot::Sender<StreamResult<Option<Bytes>>>> =
VecDeque::new();
loop {
let pull_fut = poll_fn(|cx| byte_state.poll_pull_needed(cx)).fuse();
let serve_fut = poll_fn(|cx| byte_state.poll_serve_needed(cx)).fuse();
let cmd_fut = command_rx.next().fuse();
futures::pin_mut!(pull_fut, serve_fut, cmd_fut);
futures::select! {
_ = pull_fut => {
if byte_state.closed.load(std::sync::atomic::Ordering::Acquire)
|| byte_state.errored.load(std::sync::atomic::Ordering::Acquire)
{
continue;
}
let mut source = match byte_state.source.lock().take() {
Some(s) => s,
None => {
continue;
}
};
byte_state.mark_pull_started();
let size_before = byte_state.buffer_size();
let made_progress = match source.pull(&mut controller).await {
Ok(()) => {
let made_progress = byte_state.buffer_size() > size_before;
if made_progress {
while let Some(completion) = pending_reads.pop_front() {
match poll_fn(|cx| byte_state.poll_read_chunk(cx)).now_or_never() {
Some(Ok(None)) => {
let _ = completion.send(Ok(None));
}
Some(Ok(Some(chunk))) => {
let _ = completion.send(Ok(Some(chunk)));
break; }
Some(Err(err)) => {
let _ = completion.send(Err(err));
}
None => {
pending_reads.push_front(completion);
break;
}
}
}
}
if byte_state.is_closed() {
while let Some(completion) = pending_reads.pop_front() {
let _ = completion.send(Ok(None));
}
}
if byte_state.errored.load(std::sync::atomic::Ordering::Acquire) {
let error = byte_state
.error
.lock()
.clone()
.unwrap_or_else(|| "Stream errored".into());
while let Some(completion) = pending_reads.pop_front() {
let _ = completion.send(Err(error.clone()));
}
}
made_progress
}
Err(err) => {
byte_state.error(err.clone());
while let Some(completion) = pending_reads.pop_front() {
let _ = completion.send(Err(err.clone()));
}
false
}
};
byte_state.mark_pull_completed(made_progress);
if !byte_state.closed.load(std::sync::atomic::Ordering::Acquire)
&& !byte_state.errored.load(std::sync::atomic::Ordering::Acquire)
{
*byte_state.source.lock() = Some(source);
}
}
_ = serve_fut => {
while let Some(completion) = pending_reads.pop_front() {
match poll_fn(|cx| byte_state.poll_read_chunk(cx)).now_or_never() {
Some(Ok(Some(chunk))) => {
let _ = completion.send(Ok(Some(chunk)));
}
Some(Ok(None)) => {
let _ = completion.send(Ok(None));
}
Some(Err(err)) => {
let _ = completion.send(Err(err));
}
None => {
pending_reads.push_front(completion);
break;
}
}
}
}
cmd = cmd_fut => {
match cmd {
Some(StreamCommand::Read { completion }) => {
if byte_state.errored.load(Ordering::Acquire) {
let error = byte_state.error.lock()
.clone()
.unwrap_or_else(|| "Stream errored".into());
let _ = completion.send(Err(error));
continue;
}
if byte_state.closed.load(Ordering::Acquire) && byte_state.is_buffer_empty() {
let _ = completion.send(Ok(None));
continue;
}
match poll_fn(|cx| byte_state.poll_read_chunk(cx)).now_or_never() {
Some(Ok(None)) => {
let _ = completion.send(Ok(None));
}
Some(Ok(Some(chunk))) => {
let _ = completion.send(Ok(Some(chunk)));
}
Some(Err(err)) => {
let _ = completion.send(Err(err));
}
None => {
if byte_state.closed.load(Ordering::Acquire) {
let _ = completion.send(Ok(None));
} else {
pending_reads.push_back(completion);
if !byte_state.errored.load(Ordering::Acquire) {
byte_state.maybe_trigger_pull();
}
}
}
}
}
Some(StreamCommand::Cancel { reason, completion }) => {
let cancel_result = byte_state.cancel_source(reason).await;
while let Some(tx) = pending_reads.pop_front() {
let _ = tx.send(Ok(None));
}
let _ = completion.send(cancel_result);
}
Some(_) => {}
None => {
if pending_reads.is_empty() {
break;
}
if byte_state.closed.load(Ordering::Acquire) {
while let Some(completion) = pending_reads.pop_front() {
let _ = completion.send(Ok(None));
}
break;
}
}
}
}
}
}
}
pub struct ReadableStreamBuilder<T, Source, StreamType = DefaultStream>
where
T: MaybeSend + 'static,
StreamType: StreamTypeMarker,
{
source: Source,
strategy: crate::platform::BoxedStrategyStatic<T>,
_phantom: PhantomData<fn() -> (T, StreamType)>,
}
impl<T: MaybeSend + 'static, Source> ReadableStreamBuilder<T, Source, DefaultStream>
where
Source: ReadableSource<T>,
{
fn new(source: Source) -> Self {
Self {
source,
strategy: Box::new(CountQueuingStrategy::new(1)),
_phantom: PhantomData,
}
}
pub fn strategy<S: QueuingStrategy<T> + MaybeSend + 'static>(mut self, s: S) -> Self {
self.strategy = Box::new(s);
self
}
pub fn prepare(
self,
) -> (
ReadableStream<T, Source, DefaultStream, Unlocked>,
impl Future<Output = ()>,
) {
ReadableStream::new_inner(self.source, self.strategy)
}
pub fn spawn<F, R>(self, spawn_fn: F) -> ReadableStream<T, Source, DefaultStream, Unlocked>
where
F: FnOnce(crate::platform::PlatformFuture<'static, ()>) -> R,
{
let (stream, fut) = self.prepare();
spawn_fn(Box::pin(fut));
stream
}
pub fn spawn_ref<F, R>(
self,
spawn_fn: &'static F,
) -> ReadableStream<T, Source, DefaultStream, Unlocked>
where
F: Fn(crate::platform::PlatformFuture<'static, ()>) -> R,
{
let (stream, fut) = self.prepare();
spawn_fn(Box::pin(fut));
stream
}
}
impl<Source> ReadableStreamBuilder<Bytes, Source, ByteStream>
where
Source: ReadableByteSource,
{
fn new_bytes(source: Source) -> Self {
Self {
source,
strategy: Box::new(CountQueuingStrategy::new(1)),
_phantom: PhantomData,
}
}
pub fn strategy<S: QueuingStrategy<Bytes> + MaybeSend + 'static>(mut self, s: S) -> Self {
self.strategy = Box::new(s);
self
}
pub fn prepare(
self,
) -> (
ReadableStream<Bytes, Source, ByteStream, Unlocked>,
impl Future<Output = ()>,
) {
ReadableStream::new_bytes_inner(self.source, self.strategy)
}
pub fn spawn<F, R>(self, spawn_fn: F) -> ReadableStream<Bytes, Source, ByteStream, Unlocked>
where
F: FnOnce(crate::platform::PlatformFuture<'static, ()>) -> R,
{
let (stream, fut) = self.prepare();
spawn_fn(Box::pin(fut));
stream
}
pub fn spawn_ref<F, R>(
self,
spawn_fn: &'static F,
) -> ReadableStream<Bytes, Source, ByteStream, Unlocked>
where
F: Fn(crate::platform::PlatformFuture<'static, ()>) -> R,
{
let (stream, fut) = self.prepare();
spawn_fn(Box::pin(fut));
stream
}
}
impl<T: MaybeSend + 'static, Source> ReadableStream<T, Source, DefaultStream, Unlocked>
where
Source: ReadableSource<T>,
{
pub fn builder(source: Source) -> ReadableStreamBuilder<T, Source, DefaultStream> {
ReadableStreamBuilder::new(source)
}
}
impl<T: MaybeSend + 'static>
ReadableStream<T, IteratorSource<std::vec::IntoIter<T>>, DefaultStream, Unlocked>
{
pub fn from_vec(
vec: Vec<T>,
) -> ReadableStreamBuilder<T, IteratorSource<std::vec::IntoIter<T>>, DefaultStream> {
ReadableStreamBuilder::from_vec(vec)
}
}
impl<T: MaybeSend + 'static, I> ReadableStream<T, IteratorSource<I>, DefaultStream, Unlocked>
where
I: Iterator<Item = T> + MaybeSend + 'static,
{
pub fn from_iterator(iter: I) -> ReadableStreamBuilder<T, IteratorSource<I>, DefaultStream> {
ReadableStreamBuilder::from_iterator(iter)
}
}
impl<T: MaybeSend + 'static, S> ReadableStream<T, AsyncStreamSource<S>, DefaultStream, Unlocked>
where
S: Stream<Item = T> + Unpin + MaybeSend + 'static,
{
pub fn from_stream(stream: S) -> ReadableStreamBuilder<T, AsyncStreamSource<S>, DefaultStream> {
ReadableStreamBuilder::from_stream(stream)
}
}
impl<Source> ReadableStream<Bytes, Source, ByteStream, Unlocked>
where
Source: ReadableByteSource,
{
pub fn builder_bytes(source: Source) -> ReadableStreamBuilder<Bytes, Source, ByteStream> {
ReadableStreamBuilder::new_bytes(source)
}
}
impl<T: MaybeSend + 'static>
ReadableStreamBuilder<T, IteratorSource<std::vec::IntoIter<T>>, DefaultStream>
{
pub fn from_vec(vec: Vec<T>) -> Self {
Self::new(IteratorSource {
iter: vec.into_iter(),
})
}
}
impl<T: MaybeSend + 'static, I> ReadableStreamBuilder<T, IteratorSource<I>, DefaultStream>
where
I: Iterator<Item = T> + MaybeSend + 'static,
{
pub fn from_iterator(iter: I) -> Self {
Self::new(IteratorSource { iter })
}
}
impl<T: MaybeSend + 'static, S> ReadableStreamBuilder<T, AsyncStreamSource<S>, DefaultStream>
where
S: Stream<Item = T> + Unpin + MaybeSend + 'static,
{
pub fn from_stream(stream: S) -> Self {
Self::new(AsyncStreamSource { stream })
}
}
#[cfg(test)]
mod tests {
use super::*;
use futures::stream;
use std::time::Duration;
use tokio::time::timeout;
#[tokio_localset_test::localset_test]
async fn reads_items_sequentially_from_iterator() {
let data = vec![1, 2, 3, 4, 5];
let stream = ReadableStream::from_iterator(data.clone().into_iter()).spawn(|fut| {
tokio::task::spawn_local(fut);
});
let (_locked, reader) = stream.get_reader().unwrap();
for expected in data {
assert_eq!(reader.read().await.unwrap(), Some(expected));
}
assert_eq!(reader.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn transitions_to_closed_state_after_exhaustion() {
let data = vec![1, 2, 3];
let stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
let (_locked, reader) = stream.get_reader().unwrap();
assert!(!reader.0.closed.load(std::sync::atomic::Ordering::Acquire));
assert!(!reader.0.errored.load(std::sync::atomic::Ordering::Acquire));
while reader.read().await.unwrap().is_some() {}
reader.closed().await.unwrap();
assert!(reader.0.closed.load(std::sync::atomic::Ordering::Acquire));
}
#[tokio_localset_test::localset_test]
async fn handles_empty_stream_immediately() {
let empty: Vec<i32> = vec![];
let stream = ReadableStream::from_iterator(empty.into_iter()).spawn(|fut| {
tokio::task::spawn_local(fut);
});
let (_locked, reader) = stream.get_reader().unwrap();
assert_eq!(reader.read().await.unwrap(), None);
reader.closed().await.unwrap();
}
#[tokio_localset_test::localset_test]
async fn enforces_stream_locking_correctly() {
let data = vec![1, 2, 3];
let stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
assert!(!stream.locked());
let (_locked_stream, reader) = stream.get_reader().unwrap();
let unlocked_stream = reader.release_lock();
assert!(!unlocked_stream.locked());
}
#[tokio_localset_test::localset_test]
async fn auto_unlocks_on_reader_drop() {
let data = vec![1, 2, 3];
let stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
let locked_ref = SharedPtr::clone(&stream.locked);
{
let (_locked_stream, _reader) = stream.get_reader().unwrap();
}
assert!(!locked_ref.load(std::sync::atomic::Ordering::Acquire));
}
#[tokio_localset_test::localset_test]
async fn cancels_stream_and_stops_reading() {
let data = vec![1, 2, 3, 4, 5];
let stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
let (_locked, reader) = stream.get_reader().unwrap();
assert_eq!(reader.read().await.unwrap(), Some(1));
reader
.cancel(Some("test cancellation".to_string()))
.await
.unwrap();
assert_eq!(reader.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn cancels_without_reason() {
let data = vec![1, 2, 3];
let stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
let (_locked, reader) = stream.get_reader().unwrap();
reader.cancel(None).await.unwrap();
assert_eq!(reader.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn propagates_source_errors() {
struct ErroringSource {
call_count: std::cell::RefCell<usize>,
}
impl ReadableSource<i32> for ErroringSource {
async fn pull(
&mut self,
controller: &mut ReadableStreamDefaultController<i32>,
) -> StreamResult<()> {
let count = *self.call_count.borrow();
*self.call_count.borrow_mut() += 1;
if count == 0 {
controller.enqueue(42)?;
Ok(())
} else {
controller.error("Source error".into())?;
Ok(())
}
}
}
let source = ErroringSource {
call_count: std::cell::RefCell::new(0),
};
let stream = ReadableStream::builder(source).spawn(|fut| {
tokio::task::spawn_local(fut);
});
let (_locked, reader) = stream.get_reader().unwrap();
assert_eq!(reader.read().await.unwrap(), Some(42));
let read_result = reader.read().await;
assert!(read_result.is_err(), "Second read should propagate error");
assert!(reader.0.errored.load(std::sync::atomic::Ordering::Acquire));
}
#[tokio_localset_test::localset_test]
async fn integrates_with_async_streams() {
let items = vec![10, 20, 30];
let async_stream = stream::iter(items.clone());
let readable_stream =
ReadableStream::from_stream(async_stream).spawn(tokio::task::spawn_local);
let (_locked, reader) = readable_stream.get_reader().unwrap();
for expected in items {
assert_eq!(reader.read().await.unwrap(), Some(expected));
}
assert_eq!(reader.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn futures_stream_trait_yields_all_chunks() {
use futures::stream::StreamExt;
let items = vec![1, 2, 3, 4, 5];
let mut stream = ReadableStream::from_vec(items.clone()).spawn(tokio::task::spawn_local);
let mut collected = Vec::new();
while let Some(result) = stream.next().await {
collected.push(result.unwrap());
}
assert_eq!(collected, items);
}
#[tokio_localset_test::localset_test]
async fn futures_stream_trait_propagates_errors() {
use futures::stream::StreamExt;
struct FailSource {
count: usize,
}
impl ReadableSource<String> for FailSource {
async fn pull(
&mut self,
controller: &mut ReadableStreamDefaultController<String>,
) -> StreamResult<()> {
self.count += 1;
if self.count <= 2 {
controller.enqueue(format!("item-{}", self.count))?;
Ok(())
} else {
Err(StreamError::from("source failed"))
}
}
}
let mut stream =
ReadableStream::builder(FailSource { count: 0 }).spawn(tokio::task::spawn_local);
assert_eq!(stream.next().await.unwrap().unwrap(), "item-1");
assert_eq!(stream.next().await.unwrap().unwrap(), "item-2");
let err = stream.next().await.unwrap().unwrap_err();
assert!(err.to_string().contains("source failed"));
}
#[tokio_localset_test::localset_test]
async fn handles_byte_stream_operations() {
struct ChunkedByteSource {
chunks: Vec<Vec<u8>>,
index: std::cell::RefCell<usize>,
}
impl ReadableByteSource for ChunkedByteSource {
async fn pull(
&mut self,
controller: &mut ReadableByteStreamController,
) -> StreamResult<()> {
let idx = *self.index.borrow();
if idx >= self.chunks.len() {
controller.close()?;
return Ok(());
}
let chunk = self.chunks[idx].clone();
*self.index.borrow_mut() = idx + 1;
controller.enqueue(chunk)?;
Ok(())
}
}
let source = ChunkedByteSource {
chunks: vec![b"hello".to_vec(), b" world".to_vec()],
index: std::cell::RefCell::new(0),
};
let stream = ReadableStream::builder_bytes(source).spawn(tokio::task::spawn_local);
let (_locked, reader) = stream.get_reader().unwrap();
let mut all_data = Vec::new();
while let Some(chunk) = reader.read().await.unwrap() {
all_data.extend(chunk);
}
assert_eq!(all_data, b"hello world");
}
#[tokio_localset_test::localset_test]
async fn supports_byob_reader_operations() {
struct SingleChunkByteSource {
data: Vec<u8>,
consumed: std::cell::RefCell<bool>,
}
impl ReadableByteSource for SingleChunkByteSource {
async fn pull(
&mut self,
controller: &mut ReadableByteStreamController,
) -> StreamResult<()> {
if *self.consumed.borrow() {
controller.close()?;
return Ok(());
}
let chunk = self.data.clone();
*self.consumed.borrow_mut() = true;
controller.enqueue(chunk)?;
Ok(())
}
}
let source = SingleChunkByteSource {
data: b"byob test data".to_vec(),
consumed: std::cell::RefCell::new(false),
};
let stream = ReadableStream::builder_bytes(source).spawn(tokio::task::spawn_local);
let (_locked, byob_reader) = stream.get_byob_reader().unwrap();
let mut buffer = [0u8; 20];
let bytes_read = byob_reader.read(&mut buffer).await.unwrap();
assert!(bytes_read > 0, "BYOB reader should return bytes read");
assert_eq!(byob_reader.read(&mut buffer).await.unwrap(), 0);
}
#[tokio_localset_test::localset_test]
async fn read_owned_direct_fill_single_threaded() {
struct DirectSource {
data: Vec<u8>,
sent: std::cell::RefCell<bool>,
}
impl ReadableByteSource for DirectSource {
async fn pull(
&mut self,
controller: &mut ReadableByteStreamController,
) -> StreamResult<()> {
if *self.sent.borrow() {
controller.close()?;
return Ok(());
}
if let Some(mut req) = controller.byob_request() {
let n = self.data.len().min(req.len());
req[..n].copy_from_slice(&self.data[..n]);
*self.sent.borrow_mut() = true;
req.respond(n)?;
} else {
controller.enqueue(self.data.clone())?;
*self.sent.borrow_mut() = true;
}
Ok(())
}
}
let source = DirectSource {
data: b"owned bytes".to_vec(),
sent: std::cell::RefCell::new(false),
};
let stream = ReadableStream::builder_bytes(source).spawn(tokio::task::spawn_local);
let (_locked, byob) = stream.get_byob_reader().unwrap();
let (buf, n) = byob.read_owned(BytesMut::zeroed(32)).await.unwrap();
assert_eq!(&buf[..n], b"owned bytes");
let (_buf, n) = byob.read_owned(BytesMut::zeroed(32)).await.unwrap();
assert_eq!(n, 0);
}
#[tokio_localset_test::localset_test]
async fn controller_closes_stream_properly() {
struct ControlledSource {
items: Vec<i32>,
index: std::cell::RefCell<usize>,
}
impl ReadableSource<i32> for ControlledSource {
async fn pull(
&mut self,
controller: &mut ReadableStreamDefaultController<i32>,
) -> StreamResult<()> {
let idx = *self.index.borrow();
if idx >= self.items.len() {
controller.close()?;
return Ok(());
}
controller.enqueue(self.items[idx])?;
*self.index.borrow_mut() = idx + 1;
Ok(())
}
}
let source = ControlledSource {
items: vec![1, 2],
index: std::cell::RefCell::new(0),
};
let stream = ReadableStream::builder(source).spawn(|fut| {
tokio::task::spawn_local(fut);
});
let (_locked, reader) = stream.get_reader().unwrap();
assert_eq!(reader.read().await.unwrap(), Some(1));
assert_eq!(reader.read().await.unwrap(), Some(2));
assert_eq!(reader.read().await.unwrap(), None);
reader.closed().await.unwrap();
}
#[tokio_localset_test::localset_test]
async fn controller_errors_stream_correctly() {
struct ErrorAfterItemsSource {
sent_items: std::cell::RefCell<bool>,
}
impl ReadableSource<String> for ErrorAfterItemsSource {
async fn pull(
&mut self,
controller: &mut ReadableStreamDefaultController<String>,
) -> StreamResult<()> {
if !*self.sent_items.borrow() {
controller.enqueue("valid item".to_string())?;
*self.sent_items.borrow_mut() = true;
Ok(())
} else {
controller.error("Controller error".into())?;
Ok(())
}
}
}
let source = ErrorAfterItemsSource {
sent_items: std::cell::RefCell::new(false),
};
let stream = ReadableStream::builder(source).spawn(|fut| {
tokio::task::spawn_local(fut);
});
let (_locked, reader) = stream.get_reader().unwrap();
assert_eq!(reader.read().await.unwrap(), Some("valid item".to_string()));
let read_result = reader.read().await;
assert!(read_result.is_err(), "Controller error should propagate");
}
#[tokio_localset_test::localset_test]
async fn handles_concurrent_read_attempts() {
let data: Vec<i32> = (0..10).collect();
let stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
let (_locked, reader) = stream.get_reader().unwrap();
let read1 = reader.read();
let read2 = reader.read();
let result1 = timeout(Duration::from_millis(100), read1).await.unwrap();
let result2 = timeout(Duration::from_millis(100), read2).await.unwrap();
assert!(result1.is_ok());
assert!(result2.is_ok());
}
#[tokio_localset_test::localset_test]
async fn notifies_when_closed() {
let data = vec![1];
let stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
let (_locked, reader) = stream.get_reader().unwrap();
let close_future = reader.closed();
assert_eq!(reader.read().await.unwrap(), Some(1));
assert_eq!(reader.read().await.unwrap(), None);
timeout(Duration::from_millis(100), close_future)
.await
.expect("Stream should close within timeout")
.expect("Close should succeed");
}
#[tokio_localset_test::localset_test]
async fn processes_large_streams_efficiently() {
let large_data: Vec<i32> = (0..1000).collect();
let expected_sum: i32 = large_data.iter().sum();
let stream = ReadableStream::from_iterator(large_data.into_iter()).spawn(|fut| {
tokio::task::spawn_local(fut);
});
let (_locked, reader) = stream.get_reader().unwrap();
let mut actual_sum = 0;
while let Some(item) = reader.read().await.unwrap() {
actual_sum += item;
}
assert_eq!(actual_sum, expected_sum);
}
#[tokio_localset_test::localset_test]
async fn supports_push_based_sources() {
use parking_lot::Mutex;
struct PushStartSource {
data: Vec<i32>,
enqueued: SharedPtr<Mutex<bool>>,
}
impl ReadableSource<i32> for PushStartSource {
fn start(
&mut self,
controller: &mut ReadableStreamDefaultController<i32>,
) -> impl std::future::Future<Output = StreamResult<()>> {
let enqueued = self.enqueued.clone();
let data = self.data.clone();
async move {
let mut enq_lock = enqueued.lock();
if !*enq_lock {
for item in data {
controller.enqueue(item)?;
}
controller.close()?;
*enq_lock = true;
}
Ok(())
}
}
fn pull(
&mut self,
_controller: &mut ReadableStreamDefaultController<i32>,
) -> impl std::future::Future<Output = StreamResult<()>> {
async { Ok(()) }
}
}
let source = PushStartSource {
data: vec![10, 20, 30],
enqueued: SharedPtr::new(Mutex::new(false)),
};
let stream = ReadableStream::builder(source).spawn(|fut| {
tokio::task::spawn_local(fut);
});
let (_locked_stream, reader) = stream.get_reader().unwrap();
assert_eq!(reader.read().await.unwrap(), Some(10));
assert_eq!(reader.read().await.unwrap(), Some(20));
assert_eq!(reader.read().await.unwrap(), Some(30));
assert_eq!(reader.read().await.unwrap(), None);
reader.closed().await.unwrap();
}
#[tokio_localset_test::localset_test]
async fn errors_when_start_method_fails() {
struct FailingStartSource;
impl ReadableSource<i32> for FailingStartSource {
fn start(
&mut self,
_controller: &mut ReadableStreamDefaultController<i32>,
) -> impl std::future::Future<Output = StreamResult<()>> {
async { Err("Start initialization failed".into()) }
}
fn pull(
&mut self,
_controller: &mut ReadableStreamDefaultController<i32>,
) -> impl std::future::Future<Output = StreamResult<()>> {
async { Ok(()) }
}
}
let source = FailingStartSource;
let stream = ReadableStream::builder(source).spawn(|fut| {
tokio::task::spawn_local(fut);
});
let (_locked_stream, reader) = stream.get_reader().unwrap();
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
let read_result = reader.read().await;
assert!(read_result.is_err());
if let Err(err) = read_result {
assert_eq!(err.to_string(), "Start initialization failed");
}
let read_result2 = reader.read().await;
assert!(read_result2.is_err());
let closed_result = reader.closed().await;
assert!(closed_result.is_err());
if let Err(err) = closed_result {
assert_eq!(err.to_string(), "Start initialization failed");
}
}
#[tokio_localset_test::localset_test]
async fn start_blocks_read_operations() {
use std::sync::atomic::{AtomicBool, Ordering};
use tokio::sync::Barrier;
struct SlowStartByteSource {
data: Vec<u8>,
start_barrier: SharedPtr<Barrier>,
start_completed: SharedPtr<AtomicBool>,
}
impl ReadableByteSource for SlowStartByteSource {
async fn start(
&mut self,
_controller: &mut ReadableByteStreamController,
) -> StreamResult<()> {
self.start_barrier.wait().await;
tokio::time::sleep(Duration::from_millis(100)).await;
self.start_completed.store(true, Ordering::Release);
Ok(())
}
async fn pull(
&mut self,
controller: &mut ReadableByteStreamController,
) -> StreamResult<()> {
assert!(
self.start_completed.load(Ordering::Acquire),
"pull called before start completed"
);
if !self.data.is_empty() {
let chunk = self.data.clone();
self.data.clear();
controller.enqueue(chunk)?;
} else {
controller.close()?;
}
Ok(())
}
}
let barrier = SharedPtr::new(Barrier::new(2));
let start_completed = SharedPtr::new(AtomicBool::new(false));
let source = SlowStartByteSource {
data: b"delayed data".to_vec(),
start_barrier: SharedPtr::clone(&barrier),
start_completed: SharedPtr::clone(&start_completed),
};
let stream = ReadableStream::builder_bytes(source).spawn(|fut| {
tokio::task::spawn_local(fut);
});
let (_locked, reader) = stream.get_reader().unwrap();
let read_future = reader.read();
assert!(!start_completed.load(Ordering::Acquire));
barrier.wait().await;
let result = timeout(Duration::from_millis(500), read_future)
.await
.expect("Read should complete after start finishes")
.unwrap();
assert!(
result.is_some(),
"Should receive data after start completes"
);
assert!(start_completed.load(Ordering::Acquire));
}
}
#[cfg(test)]
mod pipe_to_tests {
use super::super::writable::WritableStreamDefaultController;
use super::*;
use parking_lot::Mutex;
#[derive(Clone)]
struct CountingSink {
written: SharedPtr<Mutex<Vec<Vec<u8>>>>,
}
impl CountingSink {
fn new() -> Self {
Self {
written: SharedPtr::new(Mutex::new(Vec::new())),
}
}
fn get_written(&self) -> Vec<Vec<u8>> {
self.written.lock().clone()
}
}
impl WritableSink<Vec<u8>> for CountingSink {
fn write(
&mut self,
chunk: Vec<u8>,
_controller: &mut WritableStreamDefaultController,
) -> impl std::future::Future<Output = StreamResult<()>> {
let written = self.written.clone();
async move {
written.lock().push(chunk);
Ok(())
}
}
}
#[tokio_localset_test::localset_test]
async fn pipes_data_from_readable_to_writable() {
let data = vec![vec![1u8, 2, 3], vec![4u8, 5], vec![6u8]];
let readable =
ReadableStream::from_iterator(data.clone().into_iter()).spawn(tokio::task::spawn_local);
let sink = CountingSink::new();
let writable = WritableStream::builder(sink.clone())
.strategy(CountQueuingStrategy::new(10))
.spawn(|fut| {
tokio::task::spawn_local(fut);
});
readable
.pipe_to(&writable, None)
.await
.expect("pipe_to failed");
let written = sink.get_written();
assert_eq!(written, data);
}
#[tokio_localset_test::localset_test]
async fn handles_destination_write_errors() {
#[derive(Clone)]
struct FailingWriteSink {
written: SharedPtr<Mutex<Vec<Vec<u8>>>>,
fail_after: usize,
write_count: SharedPtr<Mutex<usize>>,
abort_called: SharedPtr<Mutex<bool>>,
}
impl FailingWriteSink {
fn new(fail_after: usize) -> Self {
Self {
written: SharedPtr::new(Mutex::new(Vec::new())),
fail_after,
write_count: SharedPtr::new(Mutex::new(0)),
abort_called: SharedPtr::new(Mutex::new(false)),
}
}
fn get_written(&self) -> Vec<Vec<u8>> {
self.written.lock().clone()
}
}
impl WritableSink<Vec<u8>> for FailingWriteSink {
fn write(
&mut self,
chunk: Vec<u8>,
_controller: &mut WritableStreamDefaultController,
) -> impl std::future::Future<Output = StreamResult<()>> {
let written = SharedPtr::clone(&self.written);
let write_count = SharedPtr::clone(&self.write_count);
let fail_after = self.fail_after;
async move {
let mut count = write_count.lock();
*count += 1;
if *count > fail_after {
return Err("Write failed".into());
}
written.lock().push(chunk);
Ok(())
}
}
fn abort(
&mut self,
_reason: Option<String>,
) -> impl std::future::Future<Output = StreamResult<()>> {
let abort_called = self.abort_called.clone();
async move {
*abort_called.lock() = true;
Ok(())
}
}
}
struct TrackingSource {
data: Vec<Vec<u8>>,
index: usize,
cancelled: SharedPtr<Mutex<bool>>,
}
impl TrackingSource {
fn new(data: Vec<Vec<u8>>) -> Self {
Self {
data,
index: 0,
cancelled: SharedPtr::new(Mutex::new(false)),
}
}
}
impl ReadableSource<Vec<u8>> for TrackingSource {
async fn pull(
&mut self,
controller: &mut ReadableStreamDefaultController<Vec<u8>>,
) -> StreamResult<()> {
if self.index >= self.data.len() {
controller.close()?;
return Ok(());
}
controller.enqueue(self.data[self.index].clone())?;
self.index += 1;
Ok(())
}
async fn cancel(&mut self, _reason: Option<String>) -> StreamResult<()> {
*self.cancelled.lock() = true;
Ok(())
}
}
let data = vec![
vec![1u8, 2, 3],
vec![4u8, 5, 6],
vec![7u8, 8, 9],
vec![10u8, 11],
];
let source = TrackingSource::new(data.clone());
let cancelled_flag = SharedPtr::clone(&source.cancelled);
let readable = ReadableStream::builder(source).spawn(|fut| {
tokio::task::spawn_local(fut);
});
let sink = FailingWriteSink::new(2);
let writable = WritableStream::builder(sink.clone())
.strategy(CountQueuingStrategy::new(10))
.spawn(tokio::task::spawn_local);
let pipe_result = readable.pipe_to(&writable, None).await;
assert!(
pipe_result.is_err(),
"pipe_to should fail when write errors"
);
let written = sink.get_written();
assert!(
written.len() <= 2,
"At most 2 chunks should have been written successfully, got {}",
written.len()
);
}
#[tokio_localset_test::localset_test]
async fn handles_source_read_errors() {
struct ErroringSource {
data: Vec<Vec<u8>>,
index: usize,
error_after: usize,
}
impl ErroringSource {
fn new(data: Vec<Vec<u8>>, error_after: usize) -> Self {
Self {
data,
index: 0,
error_after,
}
}
}
impl ReadableSource<Vec<u8>> for ErroringSource {
async fn pull(
&mut self,
controller: &mut ReadableStreamDefaultController<Vec<u8>>,
) -> StreamResult<()> {
if self.index >= self.error_after {
return Err("Source error".into());
}
if self.index >= self.data.len() {
controller.close()?;
return Ok(());
}
controller.enqueue(self.data[self.index].clone())?;
self.index += 1;
Ok(())
}
}
#[derive(Clone)]
struct TrackingSink {
written: SharedPtr<Mutex<Vec<Vec<u8>>>>,
abort_called: SharedPtr<Mutex<bool>>,
abort_reason: SharedPtr<Mutex<Option<String>>>,
}
impl TrackingSink {
fn new() -> Self {
Self {
written: SharedPtr::new(Mutex::new(Vec::new())),
abort_called: SharedPtr::new(Mutex::new(false)),
abort_reason: SharedPtr::new(Mutex::new(None)),
}
}
fn get_written(&self) -> Vec<Vec<u8>> {
self.written.lock().clone()
}
fn was_abort_called(&self) -> bool {
*self.abort_called.lock()
}
fn get_abort_reason(&self) -> Option<String> {
self.abort_reason.lock().clone()
}
}
impl WritableSink<Vec<u8>> for TrackingSink {
fn write(
&mut self,
chunk: Vec<u8>,
_controller: &mut WritableStreamDefaultController,
) -> impl std::future::Future<Output = StreamResult<()>> {
let written = SharedPtr::clone(&self.written);
async move {
written.lock().push(chunk);
Ok(())
}
}
fn abort(
&mut self,
reason: Option<String>,
) -> impl std::future::Future<Output = StreamResult<()>> {
let abort_called = self.abort_called.clone();
let abort_reason = self.abort_reason.clone();
async move {
*abort_called.lock() = true;
*abort_reason.lock() = reason;
Ok(())
}
}
}
let data = vec![vec![1u8, 2, 3], vec![4u8, 5, 6]];
let source = ErroringSource::new(data.clone(), 2);
let readable = ReadableStream::builder(source).spawn(|fut| {
tokio::task::spawn_local(fut);
});
let sink = TrackingSink::new();
let writable = WritableStream::builder(sink.clone())
.strategy(CountQueuingStrategy::new(10))
.spawn(tokio::task::spawn_local);
let pipe_result = readable.pipe_to(&writable, None).await;
assert!(
pipe_result.is_err(),
"pipe_to should fail when source errors"
);
assert!(
sink.was_abort_called(),
"Destination should be aborted when source errors"
);
let written = sink.get_written();
assert_eq!(written, data, "All chunks before error should be written");
let abort_reason = sink.get_abort_reason();
assert!(
abort_reason.is_some() && abort_reason.unwrap().contains("Source error"),
"Abort reason should contain source error information"
);
}
#[tokio_localset_test::localset_test]
async fn respects_prevent_options() {
#[derive(Clone)]
struct TestSource {
data: Vec<Vec<u8>>,
index: SharedPtr<Mutex<usize>>,
cancelled: SharedPtr<Mutex<bool>>,
should_error: bool,
error_after: usize,
}
impl TestSource {
fn new(data: Vec<Vec<u8>>) -> Self {
Self {
data,
index: SharedPtr::new(Mutex::new(0)),
cancelled: SharedPtr::new(Mutex::new(false)),
should_error: false,
error_after: 0,
}
}
fn with_error_after(mut self, count: usize) -> Self {
self.should_error = true;
self.error_after = count;
self
}
fn was_cancelled(&self) -> bool {
*self.cancelled.lock()
}
}
impl ReadableSource<Vec<u8>> for TestSource {
async fn pull(
&mut self,
controller: &mut ReadableStreamDefaultController<Vec<u8>>,
) -> StreamResult<()> {
let mut idx = self.index.lock();
if self.should_error && *idx >= self.error_after {
return Err("Source error".into());
}
if *idx >= self.data.len() {
controller.close()?;
return Ok(());
}
controller.enqueue(self.data[*idx].clone())?;
*idx += 1;
Ok(())
}
async fn cancel(&mut self, _reason: Option<String>) -> StreamResult<()> {
*self.cancelled.lock() = true;
Ok(())
}
}
#[derive(Clone)]
struct TestSink {
written: SharedPtr<Mutex<Vec<Vec<u8>>>>,
aborted: SharedPtr<Mutex<bool>>,
closed: SharedPtr<Mutex<bool>>,
should_error: bool,
}
impl TestSink {
fn new() -> Self {
Self {
written: SharedPtr::new(Mutex::new(Vec::new())),
aborted: SharedPtr::new(Mutex::new(false)),
closed: SharedPtr::new(Mutex::new(false)),
should_error: false,
}
}
fn with_error(mut self) -> Self {
self.should_error = true;
self
}
fn was_aborted(&self) -> bool {
*self.aborted.lock()
}
fn was_closed(&self) -> bool {
*self.closed.lock()
}
fn written_data(&self) -> Vec<Vec<u8>> {
self.written.lock().clone()
}
}
impl WritableSink<Vec<u8>> for TestSink {
fn write(
&mut self,
chunk: Vec<u8>,
_controller: &mut WritableStreamDefaultController,
) -> impl std::future::Future<Output = StreamResult<()>> {
let written = self.written.clone();
let should_error = self.should_error;
async move {
if should_error {
return Err("Sink error".into());
}
written.lock().push(chunk);
Ok(())
}
}
fn abort(
&mut self,
_reason: Option<String>,
) -> impl std::future::Future<Output = StreamResult<()>> {
let aborted = self.aborted.clone();
async move {
*aborted.lock() = true;
Ok(())
}
}
fn close(self) -> impl std::future::Future<Output = StreamResult<()>> {
let closed = self.closed;
async move {
*closed.lock() = true;
Ok(())
}
}
}
let data = vec![vec![1u8, 2, 3], vec![4u8, 5, 6]];
let source = TestSource::new(data.clone());
let readable = ReadableStream::builder(source.clone()).spawn(|fut| {
tokio::task::spawn_local(fut);
});
let sink = TestSink::new();
let writable = WritableStream::builder(sink.clone())
.strategy(CountQueuingStrategy::new(10))
.spawn(tokio::task::spawn_local);
let result = readable.pipe_to(&writable, None).await;
assert!(result.is_ok(), "Normal pipe should succeed");
assert!(
!source.was_cancelled(),
"Source should not be cancelled on success"
);
assert!(sink.was_closed(), "Sink should be closed on success");
assert!(!sink.was_aborted(), "Sink should not be aborted on success");
assert_eq!(sink.written_data(), data, "All data should be written");
let data = vec![vec![7u8, 8, 9]];
let source = TestSource::new(data.clone());
let readable = ReadableStream::builder(source).spawn(|fut| {
tokio::task::spawn_local(fut);
});
let sink = TestSink::new();
let writable = WritableStream::builder(sink.clone())
.strategy(CountQueuingStrategy::new(10))
.spawn(tokio::task::spawn_local);
let options = StreamPipeOptions {
prevent_close: true,
..Default::default()
};
let result = readable.pipe_to(&writable, Some(options)).await;
assert!(result.is_ok(), "Pipe with prevent_close should succeed");
assert!(
!sink.was_closed(),
"Sink should NOT be closed when prevent_close=true"
);
}
}
#[cfg(test)]
mod tee_tests {
use super::*;
use std::time::Duration;
#[tokio_localset_test::localset_test]
async fn splits_stream_into_two_identical_branches() {
let data = vec![1, 2, 3, 4];
let source_stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
let (stream1, stream2) = source_stream.tee().spawn(tokio::task::spawn_local).unwrap();
let (_, reader1) = stream1.get_reader().unwrap();
let (_, reader2) = stream2.get_reader().unwrap();
assert_eq!(reader1.read().await.unwrap(), Some(1));
assert_eq!(reader2.read().await.unwrap(), Some(1));
assert_eq!(reader1.read().await.unwrap(), Some(2));
assert_eq!(reader2.read().await.unwrap(), Some(2));
assert_eq!(reader1.read().await.unwrap(), Some(3));
assert_eq!(reader2.read().await.unwrap(), Some(3));
assert_eq!(reader1.read().await.unwrap(), Some(4));
assert_eq!(reader2.read().await.unwrap(), Some(4));
assert_eq!(reader1.read().await.unwrap(), None);
assert_eq!(reader2.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn handles_different_consumption_speeds() {
let data = vec![1, 2, 3];
let source_stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
let (stream1, stream2) = source_stream.tee().spawn(tokio::task::spawn_local).unwrap();
let (_, reader1) = stream1.get_reader().unwrap();
let (_, reader2) = stream2.get_reader().unwrap();
assert_eq!(reader1.read().await.unwrap(), Some(1));
assert_eq!(reader1.read().await.unwrap(), Some(2));
assert_eq!(reader1.read().await.unwrap(), Some(3));
assert_eq!(reader1.read().await.unwrap(), None);
assert_eq!(reader2.read().await.unwrap(), Some(1));
assert_eq!(reader2.read().await.unwrap(), Some(2));
assert_eq!(reader2.read().await.unwrap(), Some(3));
assert_eq!(reader2.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn continues_when_one_branch_cancels() {
let data = vec![1, 2, 3, 4];
let source_stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
let (stream1, stream2) = source_stream.tee().spawn(tokio::task::spawn_local).unwrap();
let (_, reader1) = stream1.get_reader().unwrap();
let (_, reader2) = stream2.get_reader().unwrap();
assert_eq!(reader1.read().await.unwrap(), Some(1));
assert_eq!(reader2.read().await.unwrap(), Some(1));
reader1
.cancel(Some("Branch 1 canceled".to_string()))
.await
.unwrap();
assert_eq!(reader2.read().await.unwrap(), Some(2));
assert_eq!(reader2.read().await.unwrap(), Some(3));
assert_eq!(reader2.read().await.unwrap(), Some(4));
assert_eq!(reader2.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn stops_when_both_branches_cancel() {
let data = vec![1, 2, 3, 4];
let source_stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
let (stream1, stream2) = source_stream.tee().spawn(tokio::task::spawn_local).unwrap();
let (_, reader1) = stream1.get_reader().unwrap();
let (_, reader2) = stream2.get_reader().unwrap();
assert_eq!(reader1.read().await.unwrap(), Some(1));
assert_eq!(reader2.read().await.unwrap(), Some(1));
reader1
.cancel(Some("Branch 1 canceled".to_string()))
.await
.unwrap();
reader2
.cancel(Some("Branch 2 canceled".to_string()))
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(10)).await;
}
#[tokio_localset_test::localset_test]
async fn handles_empty_source_stream() {
let data: Vec<i32> = vec![];
let source_stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
let (stream1, stream2) = source_stream.tee().spawn(tokio::task::spawn_local).unwrap();
let (_, reader1) = stream1.get_reader().unwrap();
let (_, reader2) = stream2.get_reader().unwrap();
assert_eq!(reader1.read().await.unwrap(), None);
assert_eq!(reader2.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn propagates_source_errors_to_both_branches() {
struct ErrorSource {
count: i32,
}
impl ReadableSource<i32> for ErrorSource {
async fn pull(
&mut self,
controller: &mut ReadableStreamDefaultController<i32>,
) -> StreamResult<()> {
if self.count < 2 {
controller.enqueue(self.count)?;
self.count += 1;
Ok(())
} else {
Err("Test error".into())
}
}
}
let error_stream =
ReadableStream::builder(ErrorSource { count: 0 }).spawn(tokio::task::spawn_local);
let (stream1, stream2) = error_stream.tee().spawn(tokio::task::spawn_local).unwrap();
let (_, reader1) = stream1.get_reader().unwrap();
let (_, reader2) = stream2.get_reader().unwrap();
assert_eq!(reader1.read().await.unwrap(), Some(0));
assert_eq!(reader2.read().await.unwrap(), Some(0));
assert_eq!(reader1.read().await.unwrap(), Some(1));
let err2 = match reader2.read().await {
Ok(Some(v)) => {
assert_eq!(v, 1);
reader2.read().await.unwrap_err()
}
Err(e) => e,
Ok(None) => panic!("expected the source error on branch2, got EOF"),
};
let err1 = reader1.read().await.unwrap_err();
assert_eq!(err1.to_string(), "Test error");
assert_eq!(err2.to_string(), "Test error");
}
#[tokio_localset_test::localset_test]
async fn supports_nested_tee_operations() {
let data = vec![1, 2];
let source_stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
let (stream1, stream2) = source_stream.tee().spawn(tokio::task::spawn_local).unwrap();
let (stream1a, stream1b) = stream1.tee().spawn(tokio::task::spawn_local).unwrap();
let (_, reader1a) = stream1a.get_reader().unwrap();
let (_, reader1b) = stream1b.get_reader().unwrap();
let (_, reader2) = stream2.get_reader().unwrap();
assert_eq!(reader1a.read().await.unwrap(), Some(1));
assert_eq!(reader1b.read().await.unwrap(), Some(1));
assert_eq!(reader2.read().await.unwrap(), Some(1));
assert_eq!(reader1a.read().await.unwrap(), Some(2));
assert_eq!(reader1b.read().await.unwrap(), Some(2));
assert_eq!(reader2.read().await.unwrap(), Some(2));
assert_eq!(reader1a.read().await.unwrap(), None);
assert_eq!(reader1b.read().await.unwrap(), None);
assert_eq!(reader2.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn works_with_different_data_types() {
let data = vec!["hello".to_string(), "world".to_string()];
let source_stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
let (stream1, stream2) = source_stream.tee().spawn(tokio::task::spawn_local).unwrap();
let (_, reader1) = stream1.get_reader().unwrap();
let (_, reader2) = stream2.get_reader().unwrap();
assert_eq!(reader1.read().await.unwrap(), Some("hello".to_string()));
assert_eq!(reader2.read().await.unwrap(), Some("hello".to_string()));
assert_eq!(reader1.read().await.unwrap(), Some("world".to_string()));
assert_eq!(reader2.read().await.unwrap(), Some("world".to_string()));
assert_eq!(reader1.read().await.unwrap(), None);
assert_eq!(reader2.read().await.unwrap(), None);
}
}
#[cfg(test)]
mod pipe_through_tests {
use super::super::transform::*;
use super::*;
use std::time::Duration;
use tokio::time::timeout;
struct UppercaseTransformer;
impl Transformer<String, String> for UppercaseTransformer {
fn transform(
&mut self,
chunk: String,
controller: &mut TransformStreamDefaultController<String>,
) -> impl Future<Output = StreamResult<()>> {
let result = controller.enqueue(chunk.to_uppercase());
futures::future::ready(result)
}
}
struct DoubleTransformer;
impl Transformer<i32, i32> for DoubleTransformer {
fn transform(
&mut self,
chunk: i32,
controller: &mut TransformStreamDefaultController<i32>,
) -> impl Future<Output = StreamResult<()>> {
let result = controller.enqueue(chunk * 2);
futures::future::ready(result)
}
}
#[tokio_localset_test::localset_test]
async fn transforms_data_through_pipe() {
let data = vec!["hello".to_string(), "world".to_string()];
let source_stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
let transform =
TransformStream::builder(UppercaseTransformer).spawn(tokio::task::spawn_local);
let result_stream = source_stream
.pipe_through(transform, None)
.spawn(tokio::task::spawn_local);
let (_locked, reader) = result_stream.get_reader().unwrap();
let result1 = timeout(Duration::from_secs(1), reader.read())
.await
.unwrap()
.unwrap();
assert_eq!(result1, Some("HELLO".to_string()));
let result2 = timeout(Duration::from_secs(1), reader.read())
.await
.unwrap()
.unwrap();
assert_eq!(result2, Some("WORLD".to_string()));
let result3 = timeout(Duration::from_secs(1), reader.read())
.await
.unwrap()
.unwrap();
assert_eq!(result3, None);
}
#[tokio_localset_test::localset_test]
async fn transforms_numeric_data() {
let data = vec![1, 2, 3, 4];
let source_stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
let transform = TransformStream::builder(DoubleTransformer).spawn(tokio::task::spawn_local);
let result_stream = source_stream.pipe_through(transform, None).spawn(|fut| {
tokio::task::spawn_local(fut);
});
let (_locked, reader) = result_stream.get_reader().unwrap();
assert_eq!(reader.read().await.unwrap(), Some(2));
assert_eq!(reader.read().await.unwrap(), Some(4));
assert_eq!(reader.read().await.unwrap(), Some(6));
assert_eq!(reader.read().await.unwrap(), Some(8));
assert_eq!(reader.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn handles_empty_input_stream() {
let data: Vec<String> = vec![];
let source_stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
let transform =
TransformStream::builder(UppercaseTransformer).spawn(tokio::task::spawn_local);
let result_stream = source_stream.pipe_through(transform, None).spawn(|fut| {
tokio::task::spawn_local(fut);
});
let (_locked, reader) = result_stream.get_reader().unwrap();
assert_eq!(reader.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn chains_multiple_transformations() {
let data = vec![1, 2, 3];
let source_stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
let transform1 =
TransformStream::builder(DoubleTransformer).spawn(tokio::task::spawn_local);
let intermediate_stream = source_stream.pipe_through(transform1, None).spawn(|fut| {
tokio::task::spawn_local(fut);
});
let transform2 =
TransformStream::builder(DoubleTransformer).spawn(tokio::task::spawn_local);
let result_stream = intermediate_stream
.pipe_through(transform2, None)
.spawn(|fut| {
tokio::task::spawn_local(fut);
});
let (_locked, reader) = result_stream.get_reader().unwrap();
assert_eq!(reader.read().await.unwrap(), Some(4)); assert_eq!(reader.read().await.unwrap(), Some(8)); assert_eq!(reader.read().await.unwrap(), Some(12)); assert_eq!(reader.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn respects_pipe_options() {
let data = vec!["test".to_string()];
let source_stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
let transform =
TransformStream::builder(UppercaseTransformer).spawn(tokio::task::spawn_local);
let options = StreamPipeOptions {
prevent_close: false,
prevent_abort: false,
prevent_cancel: false,
signal: None,
};
let result_stream = source_stream
.pipe_through(transform, Some(options))
.spawn(|fut| {
tokio::task::spawn_local(fut);
});
let (_locked, reader) = result_stream.get_reader().unwrap();
assert_eq!(reader.read().await.unwrap(), Some("TEST".to_string()));
assert_eq!(reader.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn handles_transformation_errors() {
struct ErrorTransformer;
impl Transformer<i32, i32> for ErrorTransformer {
fn transform(
&mut self,
chunk: i32,
controller: &mut TransformStreamDefaultController<i32>,
) -> impl Future<Output = StreamResult<()>> {
if chunk == 3 {
futures::future::ready(Err("Error on 3".into()))
} else {
let result = controller.enqueue(chunk);
futures::future::ready(result)
}
}
}
let data = vec![1, 2, 3, 4];
let source_stream =
ReadableStream::from_iterator(data.into_iter()).spawn(tokio::task::spawn_local);
let transform = TransformStream::builder(ErrorTransformer).spawn(tokio::task::spawn_local);
let result_stream = source_stream.pipe_through(transform, None).spawn(|fut| {
tokio::task::spawn_local(fut);
});
let (_locked, reader) = result_stream.get_reader().unwrap();
match reader.read().await {
Ok(Some(v)) => assert_eq!(v, 1),
_ => panic!("Expected first value"),
}
match reader.read().await {
Ok(Some(v)) => assert_eq!(v, 2),
Err(e) => assert_eq!(e.to_string(), "Error on 3"),
Ok(None) => panic!("Expected value or error, got end of stream"),
}
let read_result = reader.read().await;
assert!(read_result.is_err());
}
}
#[cfg(test)]
mod builder_tests {
use super::*;
pub struct TestSource {
pub data: Vec<String>,
pub index: usize,
}
impl TestSource {
pub fn new(data: Vec<String>) -> Self {
Self { data, index: 0 }
}
}
impl ReadableSource<String> for TestSource {
async fn pull(
&mut self,
controller: &mut ReadableStreamDefaultController<String>,
) -> Result<(), StreamError> {
if self.index < self.data.len() {
let item = self.data[self.index].clone();
self.index += 1;
controller.enqueue(item)?;
} else {
controller.close()?;
}
Ok(())
}
}
#[tokio_localset_test::localset_test]
async fn builder_spawn_creates_working_stream() {
let source = TestSource::new(vec!["hello".to_string(), "world".to_string()]);
let stream = ReadableStream::builder(source).spawn(tokio::task::spawn_local);
let (_, reader) = stream.get_reader().unwrap();
assert_eq!(reader.read().await.unwrap(), Some("hello".to_string()));
assert_eq!(reader.read().await.unwrap(), Some("world".to_string()));
assert_eq!(reader.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn builder_prepare_allows_manual_task_spawn() {
let source = TestSource::new(vec!["test".to_string()]);
let (stream, fut) = ReadableStream::builder(source).prepare();
tokio::task::spawn_local(fut);
let (_, reader) = stream.get_reader().unwrap();
assert_eq!(reader.read().await.unwrap(), Some("test".to_string()));
assert_eq!(reader.read().await.unwrap(), None);
}
fn spawn_local_fn(fut: crate::platform::PlatformFuture<'static, ()>) {
tokio::task::spawn_local(fut);
}
#[tokio_localset_test::localset_test]
async fn builder_spawn_ref_works_with_function_pointer() {
let source = TestSource::new(vec!["reference".to_string()]);
let stream = ReadableStream::builder(source).spawn_ref(&spawn_local_fn);
let (_, reader) = stream.get_reader().unwrap();
assert_eq!(reader.read().await.unwrap(), Some("reference".to_string()));
assert_eq!(reader.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn builder_accepts_custom_strategy() {
let source = TestSource::new(vec!["custom".to_string()]);
let custom_strategy = CountQueuingStrategy::new(5);
let stream = ReadableStream::builder(source)
.strategy(custom_strategy)
.spawn(tokio::task::spawn_local);
let (_, reader) = stream.get_reader().unwrap();
assert_eq!(reader.read().await.unwrap(), Some("custom".to_string()));
assert_eq!(reader.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn builds_from_vec_data() {
let data = vec!["a".to_string(), "b".to_string(), "c".to_string()];
let stream = ReadableStreamBuilder::from_vec(data).spawn(tokio::task::spawn_local);
let (_, reader) = stream.get_reader().unwrap();
assert_eq!(reader.read().await.unwrap(), Some("a".to_string()));
assert_eq!(reader.read().await.unwrap(), Some("b".to_string()));
assert_eq!(reader.read().await.unwrap(), Some("c".to_string()));
assert_eq!(reader.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn builds_from_iterator_data() {
let numbers = vec![1, 2, 3];
let stream = ReadableStreamBuilder::from_iterator(numbers.into_iter())
.spawn(tokio::task::spawn_local);
let (_, reader) = stream.get_reader().unwrap();
assert_eq!(reader.read().await.unwrap(), Some(1));
assert_eq!(reader.read().await.unwrap(), Some(2));
assert_eq!(reader.read().await.unwrap(), Some(3));
assert_eq!(reader.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn builds_from_async_stream() {
let async_stream = futures::stream::iter(vec!["x", "y", "z"]);
let stream =
ReadableStreamBuilder::from_stream(async_stream).spawn(tokio::task::spawn_local);
let (_, reader) = stream.get_reader().unwrap();
assert_eq!(reader.read().await.unwrap(), Some("x"));
assert_eq!(reader.read().await.unwrap(), Some("y"));
assert_eq!(reader.read().await.unwrap(), Some("z"));
assert_eq!(reader.read().await.unwrap(), None);
}
}