use crate::{ErrorHandler, RetryPolicy};
use futures::{ready, Stream, TryStream};
use pin_project_lite::pin_project;
use std::{
future::Future,
pin::Pin,
task::{Context, Poll},
};
use tokio::time;
pin_project! {
pub struct StreamRetry<F, S> {
error_action: F,
#[pin]
stream: S,
attempt: usize,
#[pin]
state: RetryState,
}
}
pub trait StreamRetryExt: TryStream {
fn retry<F>(self, error_action: F) -> StreamRetry<F, Self>
where
Self: Sized,
{
StreamRetry::new(self, error_action)
}
}
impl<S: ?Sized> StreamRetryExt for S where S: TryStream {}
pin_project! {
#[project = RetryStateProj]
enum RetryState {
WaitingForStream,
TimerActive { #[pin] delay: time::Sleep },
}
}
impl<F, S> StreamRetry<F, S> {
pub fn new(stream: S, error_action: F) -> Self
where
S: TryStream,
{
Self::with_counter(stream, error_action, 1)
}
pub fn with_counter(stream: S, error_action: F, attempt_counter: usize) -> Self {
Self {
error_action,
stream,
attempt: attempt_counter,
state: RetryState::WaitingForStream,
}
}
}
impl<F, S> Stream for StreamRetry<F, S>
where
S: TryStream,
F: ErrorHandler<S::Error>,
{
type Item = Result<(S::Ok, usize), (F::OutError, usize)>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Option<Self::Item>> {
loop {
let this = self.as_mut().project();
let attempt = *this.attempt;
let new_state = match this.state.project() {
RetryStateProj::TimerActive { delay } => {
ready!(delay.poll(cx));
RetryState::WaitingForStream
}
RetryStateProj::WaitingForStream => match ready!(this.stream.try_poll_next(cx)) {
Some(Ok(x)) => {
*this.attempt = 1;
this.error_action.ok(attempt);
return Poll::Ready(Some(Ok((x, attempt))));
}
None => {
return Poll::Ready(None);
}
Some(Err(e)) => {
*this.attempt += 1;
match this.error_action.handle(attempt, e) {
RetryPolicy::ForwardError(e) => {
return Poll::Ready(Some(Err((e, attempt))))
}
RetryPolicy::Repeat => RetryState::WaitingForStream,
RetryPolicy::WaitRetry(duration) => RetryState::TimerActive {
delay: time::sleep(duration),
},
}
}
},
};
self.as_mut().project().state.set(new_state);
}
}
}
#[cfg(test)]
mod test {
use super::*;
use futures::{pin_mut, prelude::*};
use std::time::Duration;
#[tokio::test]
async fn naive() {
let stream = stream::iter(vec![Ok::<_, u8>(17u8), Ok(19u8)]);
let retry = StreamRetry::new(stream, |_| RetryPolicy::Repeat::<()>);
assert_eq!(
Ok(vec![(17, 1), (19, 1)]),
retry.try_collect::<Vec<_>>().await,
);
}
#[tokio::test]
async fn repeat() {
let stream = stream::iter(vec![Ok(1), Err(17), Ok(19)]);
let retry = StreamRetry::new(stream, |_| RetryPolicy::Repeat::<()>);
assert_eq!(
Ok(vec![(1, 1), (19, 2)]),
retry.try_collect::<Vec<_>>().await,
);
}
#[tokio::test]
async fn wait() {
let stream = stream::iter(vec![Err(17), Ok(19)]);
let retry = StreamRetry::new(stream, |_| {
RetryPolicy::WaitRetry::<()>(Duration::from_millis(10))
})
.try_collect()
.into_future();
assert_eq!(Ok(vec!((19, 2))), retry.await);
}
#[tokio::test]
async fn propagate() {
let stream = stream::iter(vec![Err(17u8), Ok(19u16)]);
let retry = StreamRetry::new(stream, RetryPolicy::ForwardError);
pin_mut!(retry);
assert_eq!(Some(Err((17u8, 1))), retry.next().await,);
}
}