use std::cell::UnsafeCell;
use std::marker::{PhantomData, PhantomPinned};
use std::pin::Pin;
use std::ptr::NonNull;
use std::task::{Context, Poll};
use futures::Stream;
use futures::stream::FusedStream;
pub trait StreamFn<'a, Y: 'a>: FnOnce(Yielder<'a, Y>) -> <Self as StreamFn<Y>>::Fut {
type Error;
type Fut: Future<Output = Result<(), Self::Error>>;
}
impl<'a, T, Fut, Y: 'a, E> StreamFn<'a, Y> for T
where
T: FnOnce(Yielder<'a, Y>) -> Fut,
Fut: Future<Output = Result<(), E>>,
{
type Error = E;
type Fut = Fut;
}
pub struct AsyncStream<'a, F: StreamFn<'a, Y>, Y> {
call: Option<F>,
future: Option<F::Fut>,
place: UnsafeCell<Option<Y>>,
_marker: PhantomPinned,
}
impl<'a, F, Y> Stream for AsyncStream<'a, F, Y>
where
F: StreamFn<'a, Y>,
{
type Item = Result<Y, F::Error>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = unsafe { self.get_unchecked_mut() };
if let Some(x) = unsafe { &mut (*this.place.get()) }.take() {
return Poll::Ready(Some(Ok(x)));
}
if let Some(f) = this.call.take() {
std::hint::cold_path();
let yielder = Yielder {
ptr: unsafe { NonNull::new_unchecked(this.place.get()) },
_marker: PhantomData,
};
let future = f(yielder);
this.future = Some(future);
}
if let Some(fut) = this.future.as_mut() {
let future = unsafe { Pin::new_unchecked(fut) };
match future.poll(cx) {
Poll::Ready(x) => {
this.future.take();
match x {
Ok(_) => Poll::Ready(None),
Err(e) => Poll::Ready(Some(Err(e))),
}
}
Poll::Pending => {
if let Some(x) = unsafe { &mut (*this.place.get()) }.take() {
Poll::Ready(Some(Ok(x)))
} else {
Poll::Pending
}
}
}
} else {
Poll::Ready(None)
}
}
}
impl<'a, F, Y> FusedStream for AsyncStream<'a, F, Y>
where
F: StreamFn<'a, Y>,
{
fn is_terminated(&self) -> bool {
self.future.is_none()
}
}
pub struct Yielder<'a, T> {
ptr: NonNull<Option<T>>,
_marker: PhantomData<&'a mut T>,
}
unsafe impl<Y: Send> Send for Yielder<'_, Y> {}
impl<'a, T> Yielder<'a, T> {
#[inline]
#[must_use = "The value won't be emitted unless the returned future completes."]
pub fn emit<'b>(&'b mut self, value: T) -> YielderFuture<'b, 'a, T> {
unsafe { self.ptr.replace(Some(value)) };
YielderFuture {
ptr: self.ptr,
_marker: PhantomData,
}
}
}
pub struct YielderFuture<'a, 'b, T> {
ptr: NonNull<Option<T>>,
_marker: PhantomData<&'a mut Yielder<'b, T>>,
}
unsafe impl<Y: Send> Send for YielderFuture<'_, '_, Y> {}
impl<T> Future for YielderFuture<'_, '_, T> {
type Output = ();
#[inline]
fn poll(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.get_mut();
if unsafe { this.ptr.as_ref().is_some() } {
Poll::Pending
} else {
Poll::Ready(())
}
}
}
#[must_use = "A stream does nothing unless polled"]
pub fn try_async_stream<'s, Y, F>(f: F) -> AsyncStream<'s, F, Y>
where
F: for<'a> StreamFn<'a, Y>,
{
AsyncStream {
call: Some(f),
future: None,
place: UnsafeCell::new(None),
_marker: PhantomPinned,
}
}
#[cfg(test)]
mod test {
use std::cell::Cell;
use std::future;
use std::time::Duration;
use futures::future::poll_fn;
use futures::{StreamExt, TryStreamExt};
use super::*;
#[test]
fn sequence() {
tokio::runtime::Builder::new_current_thread().build().unwrap().block_on(async {
let stream = try_async_stream(async |mut yielder: Yielder<usize>| -> Result<(), ()> {
for i in 0..10 {
yielder.emit(i).await
}
Ok(())
});
let res = stream.try_collect::<Vec<_>>().await.unwrap();
assert_eq!(res, vec![0, 1, 2, 3, 4, 5, 6, 7, 8, 9])
})
}
#[test]
fn do_other_stuff_between() {
tokio::runtime::Builder::new_current_thread().enable_time().build().unwrap().block_on(
async {
let stream =
try_async_stream(async |mut yielder: Yielder<usize>| -> Result<(), ()> {
yielder.emit(0).await;
tokio::time::sleep(Duration::from_millis(10)).await;
yielder.emit(1).await;
Ok(())
});
let res = stream.try_collect::<Vec<_>>().await.unwrap();
assert_eq!(res, vec![0, 1])
},
)
}
#[test]
fn wait_on_channel() {
tokio::runtime::Builder::new_current_thread().build().unwrap().block_on(async {
let (send, recv) = futures::channel::oneshot::channel::<()>();
let stream =
try_async_stream(async move |mut yielder: Yielder<usize>| -> Result<(), ()> {
yielder.emit(0).await;
recv.await.unwrap();
yielder.emit(1).await;
Ok(())
});
let mut stream = Box::pin(stream);
assert_eq!(stream.next().await.unwrap().unwrap(), 0);
poll_fn(|cx| {
match stream.poll_next_unpin(cx) {
Poll::Ready(_) => panic!("Did not wait correctly"),
Poll::Pending => {}
}
Poll::Ready(())
})
.await;
send.send(()).unwrap();
assert_eq!(stream.next().await.unwrap().unwrap(), 1);
assert_eq!(stream.next().await, None);
})
}
#[test]
fn return_error() {
tokio::runtime::Builder::new_current_thread().build().unwrap().block_on(async {
let stream =
try_async_stream(async move |mut yielder: Yielder<usize>| -> Result<(), usize> {
yielder.emit(0).await;
Err(1)
});
let mut stream = Box::pin(stream);
assert_eq!(stream.next().await.unwrap().unwrap(), 0);
assert_eq!(stream.next().await.unwrap().unwrap_err(), 1);
assert!(stream.is_terminated());
poll_fn(|cx| {
match stream.poll_next_unpin(cx) {
Poll::Ready(None) => {}
_ => panic!("wrong value"),
}
Poll::Ready(())
})
.await;
})
}
#[test]
fn immediate_done() {
tokio::runtime::Builder::new_current_thread().build().unwrap().block_on(async {
let stream =
try_async_stream(async move |_: Yielder<usize>| -> Result<(), ()> { Ok(()) });
let mut stream = Box::pin(stream);
assert_eq!(stream.next().await, None);
})
}
#[test]
fn drop_mid_stream() {
thread_local! {
static DROPPED: Cell<usize> = const{ Cell::new(0) };
}
tokio::runtime::Builder::new_current_thread().build().unwrap().block_on(async {
DROPPED.with(|x| x.set(0));
struct Dropped;
impl Drop for Dropped {
fn drop(&mut self) {
DROPPED.with(|x| x.update(|x| x + 1));
}
}
let stream =
try_async_stream(async move |mut yielder: Yielder<Dropped>| -> Result<(), ()> {
yielder.emit(Dropped).await;
struct DropYield<'a>(Yielder<'a, Dropped>);
impl Drop for DropYield<'_> {
fn drop(&mut self) {
let f = self.0.emit(Dropped);
#[allow(clippy::drop_non_drop)]
std::mem::drop(f);
}
}
let _drop = DropYield(yielder);
let _ = future::pending::<()>().await;
Ok(())
});
let mut stream = Box::pin(stream);
let _: Dropped = stream.next().await.unwrap().unwrap();
assert_eq!(DROPPED.with(|x| x.get()), 1);
poll_fn(|cx| {
let Poll::Pending = stream.poll_next_unpin(cx) else {
panic!("Incorrect poll result")
};
Poll::Ready(())
})
.await;
std::mem::drop(stream);
assert_eq!(DROPPED.with(|x| x.get()), 2);
})
}
#[test]
fn fused_stream() {
tokio::runtime::Builder::new_current_thread().build().unwrap().block_on(async {
let stream =
try_async_stream(async move |mut yielder: Yielder<usize>| -> Result<(), ()> {
yielder.emit(0).await;
yielder.emit(1).await;
Ok(())
});
let mut stream = Box::pin(stream);
poll_fn(|cx| {
let Poll::Ready(Some(Ok(0))) = stream.poll_next_unpin(cx) else {
panic!("Wrong value")
};
let Poll::Ready(Some(Ok(1))) = stream.poll_next_unpin(cx) else {
panic!("Wrong value")
};
let Poll::Ready(None) = stream.poll_next_unpin(cx) else {
panic!("Wrong value")
};
assert!(stream.is_terminated());
let Poll::Ready(None) = stream.poll_next_unpin(cx) else {
panic!("Wrong value")
};
Poll::Ready(())
})
.await;
})
}
}