#![doc = include_str!("../README.md")]
use core::fmt::{self, Debug, Display};
use futures_util::{FutureExt, Stream, StreamExt, TryStream, future};
use tokio::{
sync::mpsc::{self, error::SendError},
task::JoinError,
};
use tokio_stream::wrappers::ReceiverStream;
pub trait TryStreamTranspose: TryStream {
fn transpose_with<Fut, F, T, E>(
self,
buf_size: usize,
mut f: F,
) -> impl Future<Output = Result<T, E>>
where
F: FnMut(ReceiverStream<Self::Ok>) -> Fut,
Fut: Future<Output = Result<T, E>>,
Self: Send
+ Sized
+ Stream<Item = Result<Self::Ok, Self::Error>>
+ 'static,
Self::Ok: Send,
Self::Error: Send,
E: Send + From<Self::Error> + From<Error<Self::Ok>> + 'static,
{
let (sender, recver) = mpsc::channel(buf_size);
let send_handle = tokio::spawn(async move {
tokio::pin! {
let stream = self;
}
while let Some(line) = stream.next().await {
sender.send(line?).await.map_err(InternalSendError::Send)?;
}
Ok(())
});
let recv_handle = f(ReceiverStream::new(recver));
future::join(send_handle, recv_handle).map(|result| match result {
(Err(err), _) => Err(Error::from(err).into()),
(Ok(Err(InternalSendError::Item(err))), _) => Err(err.into()),
(_, Err(err)) => Err(err),
(Ok(Err(InternalSendError::Send(err))), _) => {
Err(Error::from(err).into())
}
(_, ok) => ok,
})
}
}
#[derive(Debug)]
pub struct Error<Item>(ErrorRepr<Item>);
#[derive(Debug)]
pub(crate) enum ErrorRepr<Item> {
Join(JoinError),
Send(SendError<Item>),
}
enum InternalSendError<I, E> {
Item(E),
Send(SendError<I>),
}
impl<Item> Display for Error<Item> {
#[inline]
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
Display::fmt("unable to transpose item", f)
}
}
impl<Item> core::error::Error for Error<Item>
where
Item: Debug + 'static,
{
#[inline]
fn source(&self) -> Option<&(dyn core::error::Error + 'static)> {
match &self.0 {
ErrorRepr::Join(err) => Some(err),
ErrorRepr::Send(err) => Some(err),
}
}
}
impl<S: TryStream> TryStreamTranspose for S {}
impl<Item> From<JoinError> for Error<Item> {
#[inline]
fn from(value: JoinError) -> Self {
Self(ErrorRepr::Join(value))
}
}
impl<Item> From<SendError<Item>> for Error<Item> {
#[inline]
fn from(value: SendError<Item>) -> Self {
Self(ErrorRepr::Send(value))
}
}
impl<I, E> From<E> for InternalSendError<I, E> {
#[inline]
fn from(err: E) -> Self {
Self::Item(err)
}
}
#[cfg(test)]
mod tests {
use super::*;
use futures_util::stream;
#[tokio::test]
async fn error_propagation() {
#[derive(Debug, thiserror::Error)]
#[error(transparent)]
enum TestError {
#[error("{0}")]
Message(&'static str),
Internal(#[from] Error<i32>),
}
let nums: [Result<i32, TestError>; 4] = [Ok(4), Ok(5), Ok(6), Ok(7)];
core::assert_matches!(
stream::iter(nums)
.transpose_with(1024, |_| {
future::err(TestError::Message("stream err"))
})
.await,
Err::<i32, TestError>(TestError::Message("stream err")),
"error propagation from inside the closure"
);
}
}