use super::super::{CountQueuingStrategy, QueuingStrategy, Unlocked};
use super::{
error::StreamError,
readable::{DefaultStream, ReadableSource, ReadableStream, ReadableStreamDefaultController},
writable::{WritableSink, WritableStream, WritableStreamDefaultController},
};
use crate::platform::{MaybeSend, SharedPtr};
use futures::{
channel::{
mpsc::{UnboundedReceiver, UnboundedSender, unbounded},
oneshot,
},
future::{self, Future, poll_fn},
stream::StreamExt,
FutureExt,
};
use std::sync::atomic::{AtomicBool, Ordering as AtomicOrdering};
use super::shared::StreamResult;
#[derive(Debug)]
enum TransformCommand<I> {
Write {
chunk: I,
completion: oneshot::Sender<StreamResult<()>>,
},
Close {
completion: oneshot::Sender<StreamResult<()>>,
},
Abort {
reason: Option<String>,
completion: oneshot::Sender<StreamResult<()>>,
},
}
pub struct TransformStream<I: MaybeSend + 'static, O: MaybeSend + 'static> {
readable: ReadableStream<O, TransformReadableSource<O>, DefaultStream, Unlocked>,
writable: WritableStream<I, TransformWritableSink<I>, Unlocked>,
}
impl<I: MaybeSend + 'static, O: MaybeSend + 'static> TransformStream<I, O> {
fn new_inner<T>(
transformer: T,
writable_strategy: crate::platform::BoxedStrategy<I>,
readable_strategy: crate::platform::BoxedStrategy<O>,
) -> (
Self,
impl Future<Output = ()>, // readable task
impl Future<Output = ()>, // writable task
impl Future<Output = ()>, // transform task
)
where
T: Transformer<I, O> + 'static,
{
let (transform_tx, transform_rx) = unbounded::<TransformCommand<I>>();
let (cancel_tx, cancel_rx): (CancelTx, CancelRx) = unbounded();
let hwm = readable_strategy.high_water_mark();
let space_signal = SharedPtr::new(SpaceSignal::new(hwm));
let writable_sink = TransformWritableSink::new(transform_tx);
let (writable, writable_fut) = WritableStream::new_inner(writable_sink, writable_strategy);
let readable_source = TransformReadableSource::new(
writable.controller.clone(),
cancel_tx,
space_signal.clone(),
);
let (readable, readable_fut) =
ReadableStream::new_inner(readable_source, readable_strategy);
let controller = TransformStreamDefaultController::new(
readable.controller.clone(),
writable.controller.clone(),
space_signal,
);
let transform_fut = transform_task(transformer, transform_rx, cancel_rx, controller);
(
TransformStream { readable, writable },
readable_fut,
writable_fut,
transform_fut,
)
}
pub fn readable(
self,
) -> ReadableStream<O, TransformReadableSource<O>, DefaultStream, Unlocked> {
self.readable
}
pub fn writable(self) -> WritableStream<I, TransformWritableSink<I>, Unlocked> {
self.writable
}
pub fn split(
self,
) -> (
ReadableStream<O, TransformReadableSource<O>, DefaultStream, Unlocked>,
WritableStream<I, TransformWritableSink<I>, Unlocked>,
) {
(self.readable, self.writable)
}
}
pub struct TransformStreamDefaultController<O: MaybeSend + 'static> {
readable_controller: SharedPtr<ReadableStreamDefaultController<O>>,
writable_controller: SharedPtr<WritableStreamDefaultController>,
space_signal: SharedPtr<SpaceSignal>,
errored_with: SharedPtr<std::sync::Mutex<Option<StreamError>>>,
}
impl<O: MaybeSend + 'static> Clone for TransformStreamDefaultController<O> {
fn clone(&self) -> Self {
Self {
readable_controller: self.readable_controller.clone(),
writable_controller: self.writable_controller.clone(),
space_signal: self.space_signal.clone(),
errored_with: self.errored_with.clone(),
}
}
}
impl<O: MaybeSend + 'static> TransformStreamDefaultController<O> {
fn new(
readable_controller: SharedPtr<ReadableStreamDefaultController<O>>,
writable_controller: SharedPtr<WritableStreamDefaultController>,
space_signal: SharedPtr<SpaceSignal>,
) -> Self {
Self {
readable_controller,
writable_controller,
space_signal,
errored_with: SharedPtr::new(std::sync::Mutex::new(None)),
}
}
pub fn enqueue(&self, chunk: O) -> StreamResult<()> {
let result = self.readable_controller.enqueue(chunk);
self.space_signal.record_enqueue();
result
}
pub fn error(&self, error: StreamError) -> StreamResult<()> {
self.record_error(&error);
self.readable_controller.error(error.clone())?;
self.writable_controller.error(error);
self.space_signal.wake_pull();
Ok(())
}
fn record_error(&self, error: &StreamError) {
let mut guard = self.errored_with.lock().unwrap();
if guard.is_none() {
*guard = Some(error.clone());
}
}
pub(super) fn error_raised(&self) -> Option<StreamError> {
self.errored_with.lock().unwrap().clone()
}
pub(super) fn error_writable(&self, error: StreamError) {
self.writable_controller.error(error);
}
pub fn terminate(&self) -> StreamResult<()> {
self.readable_controller.close()?;
self.writable_controller.error("Terminated".into());
self.space_signal.wake_pull();
Ok(())
}
pub fn desired_size(&self) -> Option<isize> {
if self.readable_controller.is_closed_or_errored() {
return None;
}
let pending = self.space_signal.pending.load(std::sync::atomic::Ordering::Acquire);
Some(self.space_signal.hwm as isize - pending as isize)
}
pub(super) async fn wait_for_readable_space(&self) {
if self.readable_controller.is_closed_or_errored() {
return;
}
futures::select! {
_ = poll_fn(|cx| self.space_signal.poll_space(cx)).fuse() => {}
_ = self.writable_controller.abort_future().fuse() => {}
}
}
}
pub trait Transformer<I: MaybeSend + 'static, O: MaybeSend + 'static>: MaybeSend + 'static {
fn start(
&mut self,
#[allow(unused)] controller: &mut TransformStreamDefaultController<O>,
) -> impl Future<Output = StreamResult<()>> + MaybeSend {
future::ready(Ok(()))
}
fn transform(
&mut self,
chunk: I,
controller: &mut TransformStreamDefaultController<O>,
) -> impl Future<Output = StreamResult<()>> + MaybeSend;
fn flush(
&mut self,
#[allow(unused)] controller: &mut TransformStreamDefaultController<O>,
) -> impl Future<Output = StreamResult<()>> + MaybeSend {
future::ready(Ok(()))
}
fn cancel(
&mut self,
#[allow(unused)] reason: Option<String>,
) -> impl Future<Output = StreamResult<()>> + MaybeSend {
future::ready(Ok(()))
}
}
type CancelTx = UnboundedSender<(Option<String>, oneshot::Sender<StreamResult<()>>)>;
type CancelRx = UnboundedReceiver<(Option<String>, oneshot::Sender<StreamResult<()>>)>;
struct SpaceSignal {
has_space: AtomicBool,
waker: futures::task::AtomicWaker,
pull_waker: futures::task::AtomicWaker,
pending: std::sync::atomic::AtomicUsize,
hwm: usize,
}
impl SpaceSignal {
fn new(hwm: usize) -> Self {
Self {
has_space: AtomicBool::new(hwm > 0),
waker: futures::task::AtomicWaker::new(),
pull_waker: futures::task::AtomicWaker::new(),
pending: std::sync::atomic::AtomicUsize::new(0),
hwm,
}
}
fn record_enqueue(&self) {
let prev = self.pending.fetch_add(1, AtomicOrdering::AcqRel);
if self.hwm == 0 || prev + 1 >= self.hwm {
self.has_space.store(false, AtomicOrdering::Release);
self.pull_waker.wake();
}
}
fn wake_pull(&self) {
self.pull_waker.wake();
}
fn poll_backpressure(&self, cx: &mut std::task::Context<'_>, closed: bool) -> std::task::Poll<()> {
if closed || !self.has_space.load(AtomicOrdering::Acquire) {
return std::task::Poll::Ready(());
}
self.pull_waker.register(cx.waker());
if closed || !self.has_space.load(AtomicOrdering::Acquire) {
std::task::Poll::Ready(())
} else {
std::task::Poll::Pending
}
}
fn sync_from_desired_size(&self, ds: isize) {
let actual_pending = (self.hwm as isize - ds).max(0) as usize;
self.pending.store(actual_pending, AtomicOrdering::Release);
let open = actual_pending < self.hwm || self.hwm == 0;
self.has_space.store(open, AtomicOrdering::Release);
if open {
self.waker.wake();
}
}
fn cancel_fired(&self) {
self.pending.store(0, AtomicOrdering::Release);
self.has_space.store(true, AtomicOrdering::Release);
self.waker.wake();
}
fn poll_space(&self, cx: &mut std::task::Context<'_>) -> std::task::Poll<()> {
if self.has_space.load(AtomicOrdering::Acquire) {
return std::task::Poll::Ready(());
}
self.waker.register(cx.waker());
if self.has_space.load(AtomicOrdering::Acquire) {
std::task::Poll::Ready(())
} else {
std::task::Poll::Pending
}
}
}
pub struct TransformReadableSource<O: MaybeSend + 'static> {
writable_controller: crate::platform::SharedPtr<WritableStreamDefaultController>,
cancel_tx: CancelTx,
space_signal: SharedPtr<SpaceSignal>,
_phantom: std::marker::PhantomData<O>,
}
impl<O: MaybeSend + 'static> TransformReadableSource<O> {
fn new(
writable_controller: crate::platform::SharedPtr<WritableStreamDefaultController>,
cancel_tx: CancelTx,
space_signal: SharedPtr<SpaceSignal>,
) -> Self {
Self {
writable_controller,
cancel_tx,
space_signal,
_phantom: std::marker::PhantomData,
}
}
}
impl<O: MaybeSend + 'static> ReadableSource<O> for TransformReadableSource<O> {
async fn pull(
&mut self,
controller: &mut ReadableStreamDefaultController<O>,
) -> StreamResult<()> {
if let Some(ds) = controller.desired_size() {
self.space_signal.sync_from_desired_size(ds);
}
let space_signal = self.space_signal.clone();
poll_fn(move |cx| space_signal.poll_backpressure(cx, controller.is_closed_or_errored()))
.await;
Ok(())
}
async fn cancel(&mut self, reason: Option<String>) -> StreamResult<()> {
self.space_signal.cancel_fired();
let (result_tx, result_rx) = oneshot::channel();
match self.cancel_tx.unbounded_send((reason.clone(), result_tx)) {
Ok(()) => result_rx.await.unwrap_or(Ok(())),
Err(_) => {
let error = match &reason {
Some(r) => StreamError::from(r.as_str()),
None => StreamError::Canceled,
};
self.writable_controller.error(error);
Ok(())
}
}
}
}
pub struct TransformWritableSink<I: MaybeSend + 'static> {
transform_tx: UnboundedSender<TransformCommand<I>>,
}
impl<I: MaybeSend + 'static> TransformWritableSink<I> {
fn new(transform_tx: UnboundedSender<TransformCommand<I>>) -> Self {
Self { transform_tx }
}
}
impl<I: MaybeSend + 'static> WritableSink<I> for TransformWritableSink<I> {
async fn write(
&mut self,
chunk: I,
_controller: &mut WritableStreamDefaultController,
) -> StreamResult<()> {
let (tx, rx) = oneshot::channel();
self.transform_tx
.unbounded_send(TransformCommand::Write {
chunk,
completion: tx,
})
.map_err(|_| StreamError::TaskDropped)?;
rx.await.unwrap_or_else(|_| Err(StreamError::TaskDropped))
}
async fn close(self) -> StreamResult<()> {
let (tx, rx) = oneshot::channel();
self.transform_tx
.unbounded_send(TransformCommand::Close { completion: tx })
.map_err(|_| StreamError::TaskDropped)?;
rx.await.unwrap_or_else(|_| Err(StreamError::TaskDropped))
}
async fn abort(&mut self, reason: Option<String>) -> StreamResult<()> {
let (tx, rx) = oneshot::channel();
self.transform_tx
.unbounded_send(TransformCommand::Abort {
reason,
completion: tx,
})
.map_err(|_| StreamError::TaskDropped)?;
rx.await.unwrap_or_else(|_| Err(StreamError::TaskDropped))
}
}
async fn transform_task<I: MaybeSend + 'static, O: MaybeSend + 'static, T>(
mut transformer: T,
mut transform_rx: UnboundedReceiver<TransformCommand<I>>,
mut cancel_rx: CancelRx,
mut controller: TransformStreamDefaultController<O>,
) where
T: Transformer<I, O>,
{
if let Err(error) = transformer.start(&mut controller).await {
let _ = controller.error(error);
return;
}
loop {
futures::select! {
cmd = transform_rx.next().fuse() => {
match cmd {
Some(TransformCommand::Write { chunk, completion }) => {
controller.wait_for_readable_space().await;
let result = transformer.transform(chunk, &mut controller).await;
match result {
Ok(()) => { let _ = completion.send(Ok(())); }
Err(error) => {
let _ = controller.error(error.clone());
let _ = completion.send(Err(error));
break;
}
}
}
Some(TransformCommand::Close { completion: close_completion }) => {
if let Some(Some((reason, cancel_completion))) =
cancel_rx.next().now_or_never()
{
let result = transformer.cancel(reason).await;
let effective = match result {
Ok(()) => controller.error_raised().map_or(Ok(()), Err),
err => err,
};
let reject_err = effective
.clone()
.err()
.unwrap_or(StreamError::Canceled);
let _ = cancel_completion.send(effective);
let _ = controller.error(reject_err.clone());
let _ = close_completion.send(Err(reject_err));
break;
}
let flush_result = transformer.flush(&mut controller).await;
if let Err(error) = flush_result {
let _ = controller.error(error.clone());
let _ = close_completion.send(Err(error));
} else {
let _ = controller.terminate();
let _ = close_completion.send(Ok(()));
}
break;
}
Some(TransformCommand::Abort { reason, completion }) => {
let cancel_result = transformer.cancel(reason.clone()).await;
let effective = match cancel_result {
Ok(()) => controller.error_raised().map_or(Ok(()), Err),
err => err,
};
match &effective {
Ok(()) => {
let _ = controller.error(StreamError::Aborted(reason));
}
Err(e) => {
let _ = controller.error(e.clone());
}
}
let _ = completion.send(effective);
break;
}
None => break,
}
}
cancel_msg = cancel_rx.next().fuse() => {
match cancel_msg {
Some((reason, completion)) => {
let result = transformer.cancel(reason.clone()).await;
let effective = match result {
Ok(()) => controller.error_raised().map_or(Ok(()), Err),
err => err,
};
let cancel_err = effective.clone().err();
let _ = completion.send(effective);
match &cancel_err {
Some(e) => {
let _ = controller.error(e.clone());
}
None => {
let reason_err = reason
.map_or(StreamError::Canceled, |r| StreamError::from(r.as_str()));
controller.error_writable(reason_err);
}
}
let reject_err = cancel_err.unwrap_or(StreamError::Canceled);
while let Some(Some(cmd)) = transform_rx.next().now_or_never() {
match cmd {
TransformCommand::Write { completion, .. } => {
let _ = completion.send(Err(reject_err.clone()));
}
TransformCommand::Close { completion } => {
let _ = completion.send(Err(reject_err.clone()));
}
TransformCommand::Abort { completion, .. } => {
let _ = completion.send(Err(reject_err.clone()));
}
}
}
break;
}
None => break,
}
}
}
}
}
pub struct IdentityTransformer<T> {
_phantom: std::marker::PhantomData<T>,
}
impl<T> IdentityTransformer<T> {
pub fn new() -> Self {
Self {
_phantom: std::marker::PhantomData,
}
}
}
impl<T: MaybeSend + 'static> Transformer<T, T> for IdentityTransformer<T> {
fn transform(
&mut self,
chunk: T,
controller: &mut TransformStreamDefaultController<T>,
) -> impl Future<Output = StreamResult<()>> {
let result = controller.enqueue(chunk);
future::ready(result)
}
}
pub struct TransformStreamBuilder<I, O, T> {
transformer: T,
writable_strategy: crate::platform::BoxedStrategyStatic<I>,
readable_strategy: crate::platform::BoxedStrategyStatic<O>,
}
impl<I: MaybeSend + 'static, O: MaybeSend + 'static, T: Transformer<I, O> + 'static>
TransformStreamBuilder<I, O, T>
{
fn new(transformer: T) -> Self {
Self {
transformer,
writable_strategy: Box::new(CountQueuingStrategy::new(1)),
readable_strategy: Box::new(CountQueuingStrategy::new(0)),
}
}
pub fn writable_strategy<S: QueuingStrategy<I> + MaybeSend + 'static>(mut self, s: S) -> Self {
self.writable_strategy = Box::new(s);
self
}
pub fn readable_strategy<S: QueuingStrategy<O> + MaybeSend + 'static>(mut self, s: S) -> Self {
self.readable_strategy = Box::new(s);
self
}
pub fn prepare(
self,
) -> (
TransformStream<I, O>,
impl Future<Output = ()>,
impl Future<Output = ()>,
impl Future<Output = ()>,
) {
TransformStream::new_inner(
self.transformer,
self.writable_strategy,
self.readable_strategy,
)
}
pub fn spawn<F, R>(self, spawn_fn: F) -> TransformStream<I, O>
where
F: FnOnce(crate::platform::PlatformFuture<'static, ()>) -> R,
{
let (stream, rfut, wfut, tfut) = self.prepare();
let fut = async move {
futures::join!(rfut, wfut, tfut);
};
spawn_fn(Box::pin(fut));
stream
}
pub fn spawn_ref<F, R>(self, spawn_fn: &'static F) -> TransformStream<I, O>
where
F: Fn(crate::platform::PlatformFuture<'static, ()>) -> R,
{
let (stream, rfut, wfut, tfut) = self.prepare();
let fut = async move {
futures::join!(rfut, wfut, tfut);
};
spawn_fn(Box::pin(fut));
stream
}
pub fn spawn_parts<F1, F2, F3, R1, R2, R3>(
self,
rf: F1,
wf: F2,
tf: F3,
) -> TransformStream<I, O>
where
F1: FnOnce(crate::platform::PlatformFuture<'static, ()>) -> R1,
F2: FnOnce(crate::platform::PlatformFuture<'static, ()>) -> R2,
F3: FnOnce(crate::platform::PlatformFuture<'static, ()>) -> R3,
{
let (stream, rfut, wfut, tfut) = self.prepare();
rf(Box::pin(rfut));
wf(Box::pin(wfut));
tf(Box::pin(tfut));
stream
}
pub fn spawn_parts_ref<R1, R2, R3>(
self,
rf: &'static dyn Fn(crate::platform::PlatformFuture<'static, ()>) -> R1,
wf: &'static dyn Fn(crate::platform::PlatformFuture<'static, ()>) -> R2,
tf: &'static dyn Fn(crate::platform::PlatformFuture<'static, ()>) -> R3,
) -> TransformStream<I, O> {
let (stream, rfut, wfut, tfut) = self.prepare();
rf(Box::pin(rfut));
wf(Box::pin(wfut));
tf(Box::pin(tfut));
stream
}
}
impl<I: MaybeSend + 'static, O: MaybeSend + 'static> TransformStream<I, O> {
pub fn builder<T: Transformer<I, O> + 'static>(
transformer: T,
) -> TransformStreamBuilder<I, O, T> {
TransformStreamBuilder::new(transformer)
}
}
#[cfg(test)]
mod tests {
use super::*;
use futures::future;
use std::time::Duration;
use tokio::time::timeout;
pub 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());
future::ready(result)
}
}
pub 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);
future::ready(result)
}
}
pub struct OddFilterTransformer;
impl Transformer<i32, i32> for OddFilterTransformer {
fn transform(
&mut self,
chunk: i32,
controller: &mut TransformStreamDefaultController<i32>,
) -> impl Future<Output = StreamResult<()>> {
let result = if chunk % 2 != 0 {
controller.enqueue(chunk)
} else {
Ok(()) };
future::ready(result)
}
}
pub struct ErrorOnThreeTransformer;
impl Transformer<i32, i32> for ErrorOnThreeTransformer {
fn transform(
&mut self,
chunk: i32,
controller: &mut TransformStreamDefaultController<i32>,
) -> impl Future<Output = StreamResult<()>> {
if chunk == 3 {
future::ready(Err("Cannot process 3".into()))
} else {
let result = controller.enqueue(chunk);
future::ready(result)
}
}
}
#[tokio_localset_test::localset_test]
async fn test_basic_transform() {
let transformer = UppercaseTransformer;
let transform_stream = TransformStream::builder(transformer)
.readable_strategy(CountQueuingStrategy::new(10))
.spawn(tokio::task::spawn_local);
let (readable, writable) = transform_stream.split();
let (_stream, writer) = writable.get_writer().unwrap();
let (_, reader) = readable.get_reader().unwrap();
writer.write("hello".to_string()).await.unwrap();
writer.write("world".to_string()).await.unwrap();
writer.close().await.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 test_numeric_transform() {
let transformer = DoubleTransformer;
let transform_stream = TransformStream::builder(transformer)
.readable_strategy(CountQueuingStrategy::new(10))
.spawn(tokio::task::spawn_local);
let (readable, writable) = transform_stream.split();
let (_, writer) = writable.get_writer().unwrap();
let (_, reader) = readable.get_reader().unwrap();
writer.write(5).await.unwrap();
writer.write(10).await.unwrap();
writer.write(-3).await.unwrap();
writer.close().await.unwrap();
assert_eq!(reader.read().await.unwrap(), Some(10));
assert_eq!(reader.read().await.unwrap(), Some(20));
assert_eq!(reader.read().await.unwrap(), Some(-6));
assert_eq!(reader.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn test_filtering_transform() {
let transformer = OddFilterTransformer;
let transform_stream = TransformStream::builder(transformer)
.readable_strategy(CountQueuingStrategy::new(10))
.spawn(tokio::task::spawn_local);
let (readable, writable) = transform_stream.split();
let (_, writer) = writable.get_writer().unwrap();
let (_, reader) = readable.get_reader().unwrap();
writer.write(1).await.unwrap(); writer.write(2).await.unwrap(); writer.write(3).await.unwrap(); writer.write(4).await.unwrap(); writer.write(5).await.unwrap(); writer.close().await.unwrap();
assert_eq!(reader.read().await.unwrap(), Some(1));
assert_eq!(reader.read().await.unwrap(), Some(3));
assert_eq!(reader.read().await.unwrap(), Some(5));
assert_eq!(reader.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn test_transform_error_handling() {
let transformer = ErrorOnThreeTransformer;
let transform_stream = TransformStream::builder(transformer)
.readable_strategy(CountQueuingStrategy::new(10))
.spawn(tokio::task::spawn_local);
let (readable, writable) = transform_stream.split();
let (_, writer) = writable.get_writer().unwrap();
let (_, reader) = readable.get_reader().unwrap();
writer.write(1).await.unwrap();
writer.write(2).await.unwrap();
assert_eq!(reader.read().await.unwrap(), Some(1));
assert_eq!(reader.read().await.unwrap(), Some(2));
let write_result = writer.write(3).await;
assert!(write_result.is_err());
let read_result = reader.read().await;
assert!(read_result.is_err());
}
#[tokio_localset_test::localset_test]
async fn test_empty_stream() {
let transformer = UppercaseTransformer;
let transform_stream =
TransformStream::builder(transformer).spawn(tokio::task::spawn_local);
let (readable, writable) = transform_stream.split();
let (_, writer) = writable.get_writer().unwrap();
let (_, reader) = readable.get_reader().unwrap();
writer.close().await.unwrap();
assert_eq!(reader.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn test_multiple_writes_before_read() {
let transformer = DoubleTransformer;
let transform_stream = TransformStream::builder(transformer)
.readable_strategy(CountQueuingStrategy::new(10))
.spawn(tokio::task::spawn_local);
let (readable, writable) = transform_stream.split();
let (_, writer) = writable.get_writer().unwrap();
let (_, reader) = readable.get_reader().unwrap();
writer.write(1).await.unwrap();
writer.write(2).await.unwrap();
writer.write(3).await.unwrap();
writer.close().await.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(), None);
}
#[tokio_localset_test::localset_test]
async fn test_abort_stream() {
let transformer = UppercaseTransformer;
let transform_stream = TransformStream::builder(transformer)
.readable_strategy(CountQueuingStrategy::new(1)) .spawn(tokio::task::spawn_local);
let (readable, writable) = transform_stream.split();
let (_, writer) = writable.get_writer().unwrap();
let (_, reader) = readable.get_reader().unwrap();
writer.write("hello".to_string()).await.unwrap();
assert_eq!(reader.read().await.unwrap(), Some("HELLO".to_string()));
let abort_result = writer.abort(Some("Test abort".to_string())).await;
assert!(abort_result.is_ok(), "abort() must resolve per spec");
let read_result = reader.read().await;
assert!(read_result.is_err());
}
#[tokio_localset_test::localset_test]
async fn test_identity_transform_default() {
let transform_stream = TransformStream::builder(IdentityTransformer::new())
.readable_strategy(CountQueuingStrategy::new(10))
.spawn(tokio::task::spawn_local);
let (readable, writable) = transform_stream.split();
let (_stream, writer) = writable.get_writer().unwrap();
let (_, reader) = readable.get_reader().unwrap();
let numbers = vec![1, 2, 3, 4, 5];
for &num in numbers.iter() {
writer.write(num).await.unwrap();
}
writer.close().await.unwrap();
for &num in numbers.iter() {
let result = timeout(Duration::from_secs(1), reader.read())
.await
.unwrap()
.unwrap();
assert_eq!(result, Some(num));
}
let result_none = timeout(Duration::from_secs(1), reader.read())
.await
.unwrap()
.unwrap();
assert_eq!(result_none, None);
}
}
#[cfg(test)]
mod builder_tests {
use super::tests::*;
use super::*;
#[tokio_localset_test::localset_test]
async fn test_builder_spawn() {
let transform_stream = TransformStream::builder(DoubleTransformer)
.readable_strategy(CountQueuingStrategy::new(10))
.spawn(tokio::task::spawn_local);
let (readable, writable) = transform_stream.split();
let (_, writer) = writable.get_writer().unwrap();
let (_, reader) = readable.get_reader().unwrap();
writer.write(1).await.unwrap();
writer.write(2).await.unwrap();
writer.close().await.unwrap();
assert_eq!(reader.read().await.unwrap(), Some(2));
assert_eq!(reader.read().await.unwrap(), Some(4));
assert_eq!(reader.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn test_builder_spawn_parts() {
let transform_stream = TransformStream::builder(DoubleTransformer)
.readable_strategy(CountQueuingStrategy::new(10))
.spawn_parts(
tokio::task::spawn_local, tokio::task::spawn_local, tokio::task::spawn_local, );
let (readable, writable) = transform_stream.split();
let (_, writer) = writable.get_writer().unwrap();
let (_, reader) = readable.get_reader().unwrap();
writer.write(3).await.unwrap();
writer.write(4).await.unwrap();
writer.close().await.unwrap();
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 test_builder_prepare() {
let (stream, rfut, wfut, tfut) = TransformStream::builder(DoubleTransformer)
.readable_strategy(CountQueuingStrategy::new(1)) .prepare();
tokio::task::spawn_local(rfut);
tokio::task::spawn_local(wfut);
tokio::task::spawn_local(tfut);
let (readable, writable) = stream.split();
let (_, writer) = writable.get_writer().unwrap();
let (_, reader) = readable.get_reader().unwrap();
writer.write(2).await.unwrap();
writer.close().await.unwrap();
assert_eq!(reader.read().await.unwrap(), Some(4));
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 test_builder_spawn_ref() {
let stream = TransformStream::builder(DoubleTransformer)
.readable_strategy(CountQueuingStrategy::new(1)) .spawn_ref(&spawn_local_fn);
let (readable, writable) = stream.split();
let (_, writer) = writable.get_writer().unwrap();
let (_, reader) = readable.get_reader().unwrap();
writer.write(3).await.unwrap();
writer.close().await.unwrap();
assert_eq!(reader.read().await.unwrap(), Some(6));
assert_eq!(reader.read().await.unwrap(), None);
}
#[tokio_localset_test::localset_test]
async fn test_builder_spawn_parts_ref() {
let stream = TransformStream::builder(DoubleTransformer)
.readable_strategy(CountQueuingStrategy::new(1)) .spawn_parts_ref(&spawn_local_fn, &spawn_local_fn, &spawn_local_fn);
let (readable, writable) = stream.split();
let (_, writer) = writable.get_writer().unwrap();
let (_, reader) = readable.get_reader().unwrap();
writer.write(4).await.unwrap();
writer.close().await.unwrap();
assert_eq!(reader.read().await.unwrap(), Some(8));
assert_eq!(reader.read().await.unwrap(), None);
}
}