jules_core/streaming/
reconnect.rs1use crate::streaming::Stream;
2use std::future::Future;
3
4pub 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 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 let res = recon_stream.next().await;
132 assert!(res.is_some());
133 assert!(res.unwrap().is_err());
134 }
135}