use crate::{ErrorHandler, RetryPolicy};
use futures::{Async, Future, Poll, Stream};
use std::time::Instant;
use tokio_timer;
pub struct StreamRetry<F, S> {
error_action: F,
stream: S,
state: RetryState,
}
pub trait StreamRetryExt: Stream {
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: Stream {}
enum RetryState {
WaitingForStream,
TimerActive(tokio_timer::Delay),
}
impl<F, S> StreamRetry<F, S> {
pub fn new(stream: S, error_action: F) -> Self
where
S: Stream,
{
Self {
error_action,
stream,
state: RetryState::WaitingForStream,
}
}
}
impl<F, S> Stream for StreamRetry<F, S>
where
S: Stream,
F: ErrorHandler<S::Error>,
{
type Item = S::Item;
type Error = F::OutError;
fn poll(&mut self) -> Poll<Option<Self::Item>, Self::Error> {
loop {
let new_state = match self.state {
RetryState::TimerActive(ref mut delay) => match delay.poll() {
Ok(Async::Ready(())) => RetryState::WaitingForStream,
Ok(Async::NotReady) => return Ok(Async::NotReady),
Err(e) => {
panic!("Timer error: {}", e)
}
},
RetryState::WaitingForStream => match self.stream.poll() {
Ok(x) => {
self.error_action.ok();
return Ok(x);
}
Err(e) => match self.error_action.handle(e) {
RetryPolicy::ForwardError(e) => return Err(e),
RetryPolicy::Repeat => RetryState::WaitingForStream,
RetryPolicy::WaitRetry(duration) => RetryState::TimerActive(
tokio_timer::Delay::new(Instant::now() + duration),
),
},
},
};
self.state = new_state;
}
}
}
#[cfg(test)]
mod test {
use super::*;
use futures::stream::iter_result;
use std::time::Duration;
use tokio;
#[test]
fn naive() {
let stream = iter_result(vec![Ok::<_, u8>(17), Ok(19)]);
let retry = StreamRetry::new(stream, |_| RetryPolicy::Repeat::<()>);
assert_eq!(Ok(vec![17, 19]), retry.collect().wait());
}
#[test]
fn repeat() {
let stream = iter_result(vec![Ok(1), Err(17), Ok(19)]);
let retry = StreamRetry::new(stream, |_| RetryPolicy::Repeat::<()>);
assert_eq!(Ok(vec![1, 19]), retry.collect().wait());
}
#[test]
fn wait() {
let stream = iter_result(vec![Err(17), Ok(19)]);
let retry = StreamRetry::new(stream, |_| {
RetryPolicy::WaitRetry::<()>(Duration::from_millis(10))
})
.collect()
.then(|x| {
assert_eq!(Ok(vec![19]), x);
Ok(())
});
tokio::run(retry);
}
#[test]
fn propagate() {
let stream = iter_result(vec![Err(17u8), Ok(19u16)]);
let mut retry = StreamRetry::new(stream, RetryPolicy::ForwardError);
assert_eq!(Err(17u8), retry.poll());
}
#[test]
fn propagate_ext() {
let mut stream = iter_result(vec![Err(17u8), Ok(19u16)]).retry(RetryPolicy::ForwardError);
assert_eq!(Err(17u8), stream.poll());
}
}