use std::{
pin::Pin,
task::{Context, Poll},
};
use std::future::Future;
use std::sync::atomic::{AtomicBool, Ordering};
use futures::Stream;
pub struct StreamWithFinalizationCallback<S, FinalizationFn, FnFut>
where
S: Stream,
FinalizationFn: FnOnce() -> FnFut,
FnFut: Future<Output = ()> + Send + 'static,
{
inner: S,
finalization_fn: Option<FinalizationFn>,
finalized: AtomicBool,
}
impl<S, FinalizationFn, FnFut> StreamWithFinalizationCallback<S, FinalizationFn, FnFut>
where
S: Stream,
FinalizationFn: FnOnce() -> FnFut,
FnFut: Future<Output = ()> + Send + 'static,
{
pub fn new(inner: S, finalization_fn: FinalizationFn) -> Self {
StreamWithFinalizationCallback {
inner,
finalization_fn: Some(finalization_fn),
finalized: AtomicBool::new(false),
}
}
}
impl<S, FinalizationFn, FnFut, T> Stream for StreamWithFinalizationCallback<S, FinalizationFn, FnFut>
where
S: Stream<Item = T> + Unpin,
FinalizationFn: FnOnce() -> FnFut,
FnFut: Future<Output = ()> + Send + 'static,
{
type Item = T;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this: &mut Self = unsafe { self.get_unchecked_mut() };
match Pin::new(&mut this.inner).poll_next(cx) {
Poll::Ready(Some(item)) => Poll::Ready(Some(item)),
Poll::Ready(None) => {
if let Some(finalization_fn) = this.finalization_fn.take() {
let finalized = this.finalized.fetch_update(Ordering::SeqCst, Ordering::SeqCst, |_| Some(true)).unwrap_or_default();
if !finalized {
tokio::runtime::Handle::current().spawn(finalization_fn());
}
}
Poll::Ready(None)
}
Poll::Pending => Poll::Pending,
}
}
}
impl<S, FinalizationFn, FnFut> Drop for StreamWithFinalizationCallback<S, FinalizationFn, FnFut>
where
S: Stream,
FinalizationFn: FnOnce() -> FnFut,
FnFut: Future<Output = ()> + Send + 'static,
{
fn drop(&mut self) {
if let Some(finalization_fn) = self.finalization_fn.take() {
let finalized = self.finalized.fetch_update(Ordering::SeqCst, Ordering::SeqCst, |_| Some(true)).unwrap_or_default();
if !finalized {
let handle = tokio::runtime::Handle::current();
let _guard = handle.enter();
handle.spawn(finalization_fn());
}
}
}
}
pub trait StreamExtFinalizationCallback: Stream + Sized {
fn on_complete_or_cancellation<FinalizationFn, FnFut>(
self,
finalization_fn: FinalizationFn,
) -> StreamWithFinalizationCallback<Self, FinalizationFn, FnFut>
where
FinalizationFn: FnOnce() -> FnFut,
FnFut: Future<Output=()> + Send + 'static,
{
StreamWithFinalizationCallback::new(self, finalization_fn)
}
}
impl<S: Stream> StreamExtFinalizationCallback for S {}