prism_mcp_rs/utils/
async_helpers.rs1use super::{UtilError, UtilResult};
4use std::future::Future;
5use std::time::Duration;
6use tokio::time::{sleep, timeout};
7
8pub async fn with_timeout<F, T>(future: F, duration: Duration) -> UtilResult<T>
10where
11 F: Future<Output = T>,
12{
13 timeout(duration, future)
14 .await
15 .map_err(|_| UtilError::Timeout {
16 duration_ms: duration.as_millis() as u64,
17 })
18}
19
20pub async fn retry_with_backoff<F, Fut, T, E>(
22 mut operation: F,
23 max_attempts: usize,
24 initial_delay: Duration,
25 max_delay: Duration,
26) -> Result<T, E>
27where
28 F: FnMut() -> Fut,
29 Fut: Future<Output = Result<T, E>>,
30{
31 let mut delay = initial_delay;
32
33 for attempt in 0..max_attempts {
34 match operation().await {
35 Ok(result) => return Ok(result),
36 Err(e) if attempt == max_attempts - 1 => return Err(e),
37 Err(_) => {
38 sleep(delay).await;
39 delay = std::cmp::min(delay * 2, max_delay);
40 }
41 }
42 }
43
44 unreachable!()
45}
46
47pub async fn execute_with_concurrency_limit<F, Fut, T>(operations: Vec<F>, limit: usize) -> Vec<T>
49where
50 F: FnOnce() -> Fut,
51 Fut: Future<Output = T>,
52 T: Send + 'static,
53{
54 use futures::stream::{FuturesUnordered, StreamExt};
55 use std::collections::VecDeque;
56
57 let mut queue: VecDeque<_> = operations.into_iter().collect();
58 let mut active = FuturesUnordered::new();
59 let mut results = Vec::new();
60
61 for _ in 0..std::cmp::min(limit, queue.len()) {
63 if let Some(op) = queue.pop_front() {
64 active.push(op());
65 }
66 }
67
68 while let Some(result) = active.next().await {
70 results.push(result);
71
72 if let Some(op) = queue.pop_front() {
74 active.push(op());
75 }
76 }
77
78 results
79}
80
81pub async fn delay(duration: Duration) {
83 sleep(duration).await;
84}
85
86pub struct CancellableTask<T> {
88 handle: tokio::task::JoinHandle<T>,
89}
90
91impl<T> CancellableTask<T> {
92 pub fn new<F>(future: F) -> Self
93 where
94 F: Future<Output = T> + Send + 'static,
95 T: Send + 'static,
96 {
97 Self {
98 handle: tokio::spawn(future),
99 }
100 }
101
102 pub fn cancel(&self) {
103 self.handle.abort();
104 }
105
106 pub async fn wait(self) -> Result<T, tokio::task::JoinError> {
107 self.handle.await
108 }
109}
110
111#[cfg(test)]
112mod tests {
113 use super::*;
114
115 #[tokio::test]
116 async fn test_with_timeout_success() {
117 let result = with_timeout(async { "success" }, Duration::from_secs(1))
118 .await
119 .unwrap();
120
121 assert_eq!(result, "success");
122 }
123
124 #[tokio::test]
125 async fn test_with_timeout_failure() {
126 let result = with_timeout(
127 async {
128 sleep(Duration::from_secs(2)).await;
129 "too slow"
130 },
131 Duration::from_millis(100),
132 )
133 .await;
134
135 assert!(result.is_err());
136 }
137
138 #[tokio::test]
139 async fn test_retry_with_backoff() {
140 let mut attempts = 0;
141
142 let result = retry_with_backoff(
143 || {
144 attempts += 1;
145 async move {
146 if attempts < 3 {
147 Err("not ready")
148 } else {
149 Ok("success")
150 }
151 }
152 },
153 5,
154 Duration::from_millis(10),
155 Duration::from_millis(100),
156 )
157 .await;
158
159 assert_eq!(result.unwrap(), "success");
160 assert_eq!(attempts, 3);
161 }
162}