Skip to main content

prism_mcp_rs/utils/
async_helpers.rs

1//! Async utility functions and helpers
2
3use super::{UtilError, UtilResult};
4use std::future::Future;
5use std::time::Duration;
6use tokio::time::{sleep, timeout};
7
8/// Execute a future with a timeout
9pub 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
20/// Retry a future with exponential backoff
21pub 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
47/// Execute multiple futures concurrently with a limit
48pub 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    // Start initial batch
62    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    // Process results and start new operations
69    while let Some(result) = active.next().await {
70        results.push(result);
71
72        // Start next operation if available
73        if let Some(op) = queue.pop_front() {
74            active.push(op());
75        }
76    }
77
78    results
79}
80
81/// Create a future that completes after a delay
82pub async fn delay(duration: Duration) {
83    sleep(duration).await;
84}
85
86/// Create a cancellable future
87pub 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}