use futures::Stream;
use futures::future::FusedFuture;
use futures::stream::FusedStream;
use parking_lot::Mutex;
use pin_project_lite::pin_project;
use std::ops::DerefMut;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
pub fn async_stream<T, F: Future<Output = ()>>(
generator: impl FnOnce(Emitter<T>) -> F,
) -> impl FusedStream<Item = T> {
let (emitter, receiver) = tx_rx();
AsyncStream::new(receiver, generator(emitter))
}
pub fn async_try_stream<T, E, F: Future<Output = Result<(), E>>>(
generator: impl FnOnce(TryEmitter<T, E>) -> F,
) -> impl FusedStream<Item = Result<T, E>> {
let (try_emitter, mut emitter, receiver) = try_tx_rx::<T, E>();
AsyncStream::new(receiver, async move {
if let Err(e) = generator(try_emitter).await {
emitter.set(Err(e));
}
})
}
fn tx_rx<T>() -> (Emitter<T>, Receiver<T>) {
let slot = Arc::new(Mutex::new(None));
(
Emitter {
slot: Arc::clone(&slot),
},
Receiver { slot },
)
}
#[expect(
clippy::type_complexity,
reason = "three-element tuple is clearer than an alias here"
)]
fn try_tx_rx<T, E>() -> (
TryEmitter<T, E>,
Emitter<Result<T, E>>,
Receiver<Result<T, E>>,
) {
let slot = Arc::new(Mutex::new(None));
(
TryEmitter {
slot: Arc::clone(&slot),
},
Emitter {
slot: Arc::clone(&slot),
},
Receiver { slot },
)
}
type SlotRef<T> = Arc<Mutex<Option<T>>>;
pub struct Emitter<T> {
slot: SlotRef<T>,
}
pub struct TryEmitter<T, E> {
slot: SlotRef<Result<T, E>>,
}
struct Receiver<T> {
slot: SlotRef<T>,
}
impl<T> Emitter<T> {
pub fn emit(&mut self, value: T) -> impl FusedFuture<Output = ()> {
self.set(value);
Emit { done: false }
}
fn set(&mut self, value: T) {
let mut guard = self.slot.lock();
match guard.deref_mut() {
Some(_) => panic!("Misuse: await was not called after calling emit"),
slot => *slot = Some(value),
}
}
}
impl<T, E> TryEmitter<T, E> {
pub fn emit(&mut self, value: T) -> impl FusedFuture<Output = ()> {
let mut guard = self.slot.lock();
match guard.deref_mut() {
Some(_) => panic!("Misuse: await was not called after calling emit"),
slot => *slot = Some(Ok::<T, E>(value)),
}
Emit { done: false }
}
}
struct Emit {
done: bool,
}
impl FusedFuture for Emit {
fn is_terminated(&self) -> bool {
self.done
}
}
impl Future for Emit {
type Output = ();
fn poll(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<()> {
if !self.done {
self.done = true;
Poll::Pending
} else {
Poll::Ready(())
}
}
}
pin_project! {
struct AsyncStream<T, U> {
rx: Receiver<T>,
done: bool,
#[pin]
generator: U,
}
}
impl<T, U> AsyncStream<T, U> {
fn new(rx: Receiver<T>, generator: U) -> AsyncStream<T, U> {
AsyncStream {
rx,
done: false,
generator,
}
}
}
impl<T, U> FusedStream for AsyncStream<T, U>
where
U: Future<Output = ()>,
{
fn is_terminated(&self) -> bool {
self.done
}
}
impl<T, U> Stream for AsyncStream<T, U>
where
U: Future<Output = ()>,
{
type Item = T;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = self.project();
if *this.done {
return Poll::Ready(None);
}
debug_assert!(this.rx.slot.lock().is_none());
let res = this.generator.poll(cx);
*this.done = res.is_ready();
match this.rx.slot.lock().take() {
Some(v) => Poll::Ready(Some(v)),
None if *this.done => Poll::Ready(None),
None => Poll::Pending,
}
}
fn size_hint(&self) -> (usize, Option<usize>) {
if self.done { (0, Some(0)) } else { (0, None) }
}
}
#[cfg(test)]
mod test {
use crate::async_stream::Emitter;
use crate::{async_stream, async_try_stream};
use futures::stream::FusedStream;
use futures::{Stream, StreamExt, pin_mut};
use std::assert_matches;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::task::{Context, Poll};
use tokio::sync::mpsc;
#[tokio::test]
async fn noop_stream() {
let s = async_stream(|_: Emitter<()>| async {});
pin_mut!(s);
assert_eq!(s.next().await, None);
}
#[tokio::test]
async fn empty_stream() {
let mut ran = false;
{
let r = &mut ran;
let s = async_stream(|_: Emitter<()>| async {
*r = true;
println!("hello world!");
});
pin_mut!(s);
assert_eq!(s.next().await, None);
}
assert!(ran);
}
#[tokio::test]
async fn emit_single_value() {
let s = async_stream(|mut emitter| async move {
emitter.emit("hello").await;
});
let values: Vec<_> = s.collect().await;
assert_eq!(1, values.len());
assert_eq!("hello", values[0]);
}
#[tokio::test]
async fn fused() {
let s = async_stream(|mut emitter| async move {
emitter.emit("hello").await;
});
pin_mut!(s);
assert!(!s.is_terminated());
assert_eq!(s.next().await, Some("hello"));
assert_eq!(s.next().await, None);
assert!(s.is_terminated());
assert_eq!(s.next().await, None);
}
#[tokio::test]
async fn emit_multi_value() {
let s = async_stream(|mut emitter| async move {
emitter.emit("hello").await;
emitter.emit("world").await;
emitter.emit("dizzy").await;
});
let values: Vec<_> = s.collect().await;
assert_eq!(3, values.len());
assert_eq!("hello", values[0]);
assert_eq!("world", values[1]);
assert_eq!("dizzy", values[2]);
}
#[tokio::test]
#[should_panic = "await was not called after calling emit"]
async fn emit_without_await() {
let s = async_stream(|mut emitter| async move {
#[expect(clippy::let_underscore_future)]
{
let _ = emitter.emit("hello");
let _ = emitter.emit("world");
}
});
let _: Vec<_> = s.collect().await;
}
#[tokio::test]
async fn unit_emit_in_select() {
use tokio::select;
#[expect(clippy::unused_async)]
async fn do_stuff_async() {}
let s = async_stream(|mut emitter| async move {
select! {
_ = do_stuff_async() => emitter.emit(()).await,
else => emitter.emit(()).await,
}
});
let values: Vec<_> = s.collect().await;
assert_eq!(values.len(), 1);
}
#[tokio::test]
async fn emit_with_select() {
use tokio::select;
#[expect(clippy::unused_async)]
async fn do_stuff_async() {}
#[expect(clippy::unused_async)]
async fn more_async_work() {}
let s = async_stream(|mut emitter| async move {
select! {
_ = do_stuff_async() => emitter.emit("hey").await,
_ = more_async_work() => emitter.emit("hey").await,
else => emitter.emit("hey").await,
}
});
let values: Vec<_> = s.collect().await;
assert_eq!(values, vec!["hey"]);
}
#[tokio::test]
async fn return_stream() {
fn build_stream() -> impl Stream<Item = u32> {
async_stream(|mut emitter| async move {
emitter.emit(1).await;
emitter.emit(2).await;
emitter.emit(3).await;
})
}
let s = build_stream();
let values: Vec<_> = s.collect().await;
assert_eq!(3, values.len());
assert_eq!(1, values[0]);
assert_eq!(2, values[1]);
assert_eq!(3, values[2]);
}
#[tokio::test]
async fn consume_channel() {
let (tx, mut rx) = mpsc::channel(10);
let s = async_stream(|mut emitter| async move {
while let Some(v) = rx.recv().await {
emitter.emit(v).await;
}
});
pin_mut!(s);
for i in 0..3 {
assert_matches!(tx.send(i).await, Ok(_));
assert_eq!(Some(i), s.next().await);
}
drop(tx);
assert_eq!(None, s.next().await);
}
#[tokio::test]
async fn borrow_self() {
struct Data(String);
impl Data {
fn stream(&self) -> impl Stream<Item = &str> + '_ {
async_stream(move |mut emitter| async move {
emitter.emit(&self.0[..]).await;
})
}
}
let data = Data("hello".to_string());
let s = data.stream();
pin_mut!(s);
assert_eq!(Some("hello"), s.next().await);
}
#[tokio::test]
async fn stream_in_stream() {
let s = async_stream(|mut emitter| async move {
let s = async_stream(|mut inner_emitter| async move {
for i in 0..3 {
inner_emitter.emit(i).await;
}
});
pin_mut!(s);
while let Some(v) = s.next().await {
emitter.emit(v).await;
}
});
let values: Vec<_> = s.collect().await;
assert_eq!(3, values.len());
}
#[tokio::test]
async fn stream_in_stream_misuse() {
let s = async_stream(|mut emitter| async move {
let s = async_stream(|_inner_emitter: Emitter<i32>| async move {
for _i in 0..3 {
emitter.emit("foo").await;
}
});
pin_mut!(s);
while let Some(v) = s.next().await {
println!("{v}");
}
});
let values: Vec<_> = s.collect().await;
assert_eq!(3, values.len());
}
#[tokio::test]
async fn emit_non_unpin_value() {
let s: Vec<_> = async_stream(|mut emitter| async move {
for i in 0..3 {
emitter.emit(async move { i }).await;
}
})
.buffered(1)
.collect()
.await;
assert_eq!(s, vec![0, 1, 2]);
}
#[tokio::test]
async fn should_not_call_handler_function_if_not_polled() {
let _ = async_stream(|_: Emitter<()>| async move {
panic!("should not be called");
});
}
#[tokio::test]
async fn should_not_continue_until_next_poll() {
let s = async_stream(|mut emitter| async move {
emitter.emit("hey").await;
panic!("make sure poll based and not push based");
});
pin_mut!(s);
let _ = s.next().await;
}
#[test]
fn inner_try_stream() {
use tokio::select;
#[expect(clippy::unused_async)]
async fn do_stuff_async() {}
let _ = async_stream(|mut emitter| async move {
select! {
_ = do_stuff_async() => {
let another_s = async_try_stream(|mut inner_emitter| async move {
inner_emitter.emit(()).await;
Ok(())
});
let _: Result<(), ()> = Box::pin(another_s).next().await.unwrap();
},
else => {},
}
emitter.emit(()).await;
});
}
#[tokio::test]
async fn single_err() {
let s = async_try_stream(|mut emitter| async move {
if true {
Err("hello")?;
} else {
emitter.emit("world").await;
}
unreachable!();
});
let values: Vec<_> = s.collect().await;
assert_eq!(1, values.len());
assert_eq!(Err("hello"), values[0]);
}
#[tokio::test]
async fn emit_then_err() {
let s = async_try_stream(|mut emitter| async move {
emitter.emit("hello").await;
Err("world")?;
unreachable!();
});
let values: Vec<_> = s.collect().await;
assert_eq!(2, values.len());
assert_eq!(Ok("hello"), values[0]);
assert_eq!(Err("world"), values[1]);
}
#[tokio::test]
async fn convert_err() {
struct ErrorA(u8);
#[derive(PartialEq, Debug)]
struct ErrorB(u8);
impl From<ErrorA> for ErrorB {
fn from(a: ErrorA) -> ErrorB {
ErrorB(a.0)
}
}
fn test() -> impl Stream<Item = Result<&'static str, ErrorB>> {
async_try_stream(|mut emitter| async move {
if true {
Err(ErrorA(1))?;
} else {
Err(ErrorB(2))?;
}
emitter.emit("unreachable").await;
Ok(())
})
}
let values: Vec<_> = test().collect().await;
assert_eq!(1, values.len());
assert_eq!(Err(ErrorB(1)), values[0]);
}
#[tokio::test]
async fn multi_try() {
fn test() -> impl Stream<Item = Result<i32, String>> {
async_try_stream(|mut emitter| async move {
let a = Ok::<_, String>(Ok::<_, String>(123))??;
for _ in 1..10 {
emitter.emit(a).await;
}
Ok(())
})
}
let values: Vec<_> = test().collect().await;
assert_eq!(9, values.len());
assert_eq!(
std::iter::repeat_n(123, 9).map(Ok).collect::<Vec<_>>(),
values
);
}
struct DropGuard(Arc<AtomicUsize>);
impl Drop for DropGuard {
fn drop(&mut self) {
self.0.fetch_add(1, Ordering::SeqCst);
}
}
#[tokio::test]
async fn generator_freed_on_done() {
let drops = Arc::new(AtomicUsize::new(0));
let guard = DropGuard(Arc::clone(&drops));
let s = async_stream(|mut emitter| async move {
let _guard = guard;
emitter.emit(1).await;
});
pin_mut!(s);
assert_eq!(s.next().await, Some(1));
assert_eq!(s.next().await, None);
assert_eq!(drops.load(Ordering::SeqCst), 1);
assert_eq!(s.next().await, None);
}
#[tokio::test]
async fn generator_freed_on_emitted_error() {
let drops = Arc::new(AtomicUsize::new(0));
let guard = DropGuard(Arc::clone(&drops));
let s = async_try_stream(|mut emitter| async move {
let _guard = guard;
emitter.emit(1).await;
Err("boom")
});
pin_mut!(s);
assert_eq!(s.next().await, Some(Ok(1)));
assert_eq!(s.next().await, Some(Err("boom")));
assert!(s.is_terminated());
assert_eq!(drops.load(Ordering::SeqCst), 1);
assert_eq!(s.next().await, None);
}
use pin_project_lite::pin_project;
pin_project! {
struct MyStream<T: Stream> {
#[pin]
input: T,
}
}
impl<T: Stream> Stream for MyStream<T> {
type Item = T::Item;
fn poll_next(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Self::Item>> {
let this = self.project();
this.input.poll_next(cx)
}
}
#[tokio::test]
async fn emit_does_not_hold_on_value() {
let waker = futures::task::noop_waker_ref();
let mut cx = Context::from_waker(waker);
let run = Arc::<AtomicUsize>::new(AtomicUsize::new(0));
let moved = Arc::clone(&run);
let s = async_stream(|mut emitter| async move {
for _ in 0..2 {
let before = moved.fetch_add(1, Ordering::SeqCst);
emitter.emit(before).await;
}
});
let mut my_stream = Box::pin(MyStream { input: s });
#[derive(Debug, PartialEq)]
struct Item {
before: usize,
result: Poll<Option<usize>>,
after: usize,
}
let mut results = vec![];
assert_eq!(run.load(Ordering::SeqCst), 0);
while run.load(Ordering::SeqCst) < 2 {
let before = run.load(Ordering::SeqCst);
let result = my_stream.poll_next_unpin(&mut cx);
let after = run.load(Ordering::SeqCst);
results.push(Item {
before,
result,
after,
});
}
assert_eq!(
results,
vec![
Item {
before: 0,
result: Poll::Ready(Some(0)),
after: 1,
},
Item {
before: 1,
result: Poll::Ready(Some(1)),
after: 2,
}
]
);
}
}