#![doc = include_str!("../README.md")]
use core::fmt::{self, Debug, Display};
use futures_channel::mpsc::{self, Receiver, SendError};
use futures_util::{FutureExt, SinkExt, Stream, StreamExt, TryStream, future};
use tokio::task::JoinError;
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(Receiver<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> + 'static,
{
let (mut 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(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(ErrorRepr);
#[derive(Debug)]
pub(crate) enum ErrorRepr {
Join(JoinError),
Send(SendError),
}
enum InternalSendError<E> {
Item(E),
Send(SendError),
}
impl Display for Error {
#[inline]
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
Display::fmt("unable to transpose item", f)
}
}
impl core::error::Error for Error {
#[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 From<JoinError> for Error {
#[inline]
fn from(value: JoinError) -> Self {
Self(ErrorRepr::Join(value))
}
}
impl From<SendError> for Error {
#[inline]
fn from(value: SendError) -> Self {
Self(ErrorRepr::Send(value))
}
}
impl<E> From<E> for InternalSendError<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),
}
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"
);
}
}