Skip to main content

jules_core/streaming/
reconnect.rs

1use crate::streaming::Stream;
2use std::future::Future;
3
4/// A stream that automatically attempts to reconnect upon encountering an error.
5pub struct ReconnectableStream<S, F> {
6    stream: S,
7    reconnect_fn: F,
8    max_retries: usize,
9    current_retries: usize,
10}
11
12impl<S, F, Fut, T, E> ReconnectableStream<S, F>
13where
14    S: Stream<Item = Result<T, E>> + Send,
15    F: FnMut() -> Fut + Send,
16    Fut: Future<Output = Result<S, E>> + Send,
17    T: Send,
18    E: Send,
19{
20    /// Creates a new `ReconnectableStream`.
21    pub fn new(stream: S, reconnect_fn: F, max_retries: usize) -> Self {
22        Self {
23            stream,
24            reconnect_fn,
25            max_retries,
26            current_retries: 0,
27        }
28    }
29}
30
31impl<S, F, Fut, T, E> Stream for ReconnectableStream<S, F>
32where
33    S: Stream<Item = Result<T, E>> + Send,
34    F: FnMut() -> Fut + Send,
35    Fut: Future<Output = Result<S, E>> + Send,
36    T: Send,
37    E: Send,
38{
39    type Item = Result<T, E>;
40
41    async fn next(&mut self) -> Option<Self::Item> {
42        loop {
43            match self.stream.next().await {
44                Some(Ok(item)) => {
45                    self.current_retries = 0;
46                    return Some(Ok(item));
47                }
48                Some(Err(e)) => {
49                    if self.current_retries >= self.max_retries {
50                        return Some(Err(e));
51                    }
52                    self.current_retries += 1;
53                    match (self.reconnect_fn)().await {
54                        Ok(new_stream) => {
55                            self.stream = new_stream;
56                        }
57                        Err(reconnect_err) => {
58                            return Some(Err(reconnect_err));
59                        }
60                    }
61                }
62                None => return None,
63            }
64        }
65    }
66}
67
68#[cfg(test)]
69mod tests {
70    use super::*;
71    use crate::errors::StreamingError;
72
73    struct MockErrorStream {
74        items: Vec<Result<String, StreamingError>>,
75    }
76
77    impl Stream for MockErrorStream {
78        type Item = Result<String, StreamingError>;
79
80        async fn next(&mut self) -> Option<Self::Item> {
81            if self.items.is_empty() {
82                None
83            } else {
84                Some(self.items.remove(0))
85            }
86        }
87    }
88
89    #[tokio::test]
90    async fn test_reconnect_success() {
91        let stream = MockErrorStream {
92            items: vec![
93                Ok("chunk 1".to_string()),
94                Err(StreamingError::new("connection lost")),
95            ],
96        };
97
98        let mut reconnects = 0;
99        let reconnect_fn = || {
100            reconnects += 1;
101            async move {
102                Ok(MockErrorStream {
103                    items: vec![Ok("chunk 2".to_string()), Ok("chunk 3".to_string())],
104                })
105            }
106        };
107
108        let mut recon_stream = ReconnectableStream::new(stream, reconnect_fn, 3);
109
110        assert_eq!(recon_stream.next().await.unwrap().unwrap(), "chunk 1");
111        assert_eq!(recon_stream.next().await.unwrap().unwrap(), "chunk 2");
112        assert_eq!(recon_stream.next().await.unwrap().unwrap(), "chunk 3");
113        assert!(recon_stream.next().await.is_none());
114    }
115
116    #[tokio::test]
117    async fn test_reconnect_failure_max_retries() {
118        let stream = MockErrorStream {
119            items: vec![Err(StreamingError::new("connection lost"))],
120        };
121
122        let reconnect_fn = || async move {
123            Ok(MockErrorStream {
124                items: vec![Err(StreamingError::new("still broken"))],
125            })
126        };
127
128        let mut recon_stream = ReconnectableStream::new(stream, reconnect_fn, 2);
129
130        // Fails after retries exhaust
131        let res = recon_stream.next().await;
132        assert!(res.is_some());
133        assert!(res.unwrap().is_err());
134    }
135}