1use hashbrown::HashMap;
7use std::sync::{Arc, Mutex, PoisonError};
8
9use std::time::{Duration, Instant};
10use tokio::sync::RwLock;
11use tokio::time;
12
13#[derive(Debug, Clone, Hash, Eq, PartialEq, serde::Serialize, serde::Deserialize)]
15pub enum OperationType {
16 ApiCall,
18 FileOperation,
20 CodeAnalysis,
22 ToolExecution,
24 NetworkRequest,
26 Processing,
28 Custom(String),
30}
31
32#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
34pub struct TimeoutConfig {
35 pub timeout_duration: Duration,
37 pub max_retries: u32,
39 pub initial_retry_delay: Duration,
41 pub max_retry_delay: Duration,
43 pub backoff_multiplier: f64,
45 pub use_jitter: bool,
47 pub retry_on_timeout: bool,
49 pub retry_on_errors: Vec<String>,
51}
52
53const DEFAULT_RETRY_ERRORS: [&str; 4] = ["timeout", "connection", "network", "server_error"];
55
56impl Default for TimeoutConfig {
57 fn default() -> Self {
58 Self {
59 timeout_duration: Duration::from_secs(30),
60 max_retries: 3,
61 initial_retry_delay: Duration::from_millis(100),
62 max_retry_delay: Duration::from_secs(30),
63 backoff_multiplier: 2.0,
64 use_jitter: true,
65 retry_on_timeout: true,
66 retry_on_errors: DEFAULT_RETRY_ERRORS.iter().map(|s| (*s).into()).collect(),
67 }
68 }
69}
70
71impl TimeoutConfig {
72 pub fn api_call() -> Self {
74 Self {
75 timeout_duration: Duration::from_secs(60),
76 max_retries: 5,
77 initial_retry_delay: Duration::from_millis(200),
78 max_retry_delay: Duration::from_secs(10),
79 backoff_multiplier: 1.5,
80 ..Default::default()
81 }
82 }
83
84 pub fn file_operation() -> Self {
86 Self {
87 timeout_duration: Duration::from_secs(10),
88 max_retries: 2,
89 initial_retry_delay: Duration::from_millis(50),
90 max_retry_delay: Duration::from_secs(2),
91 backoff_multiplier: 2.0,
92 retry_on_timeout: false, ..Default::default()
94 }
95 }
96
97 pub fn analysis() -> Self {
99 Self {
100 timeout_duration: Duration::from_secs(120),
101 max_retries: 1,
102 initial_retry_delay: Duration::from_secs(5),
103 max_retry_delay: Duration::from_secs(10),
104 backoff_multiplier: 2.0,
105 ..Default::default()
106 }
107 }
108}
109
110#[derive(Debug, Clone)]
112pub struct TimeoutEvent {
113 pub operation_id: String,
115 pub operation_type: OperationType,
117 pub start_time: Instant,
119 pub timeout_duration: Duration,
121 pub retry_count: u32,
123 pub error_message: Option<String>,
125}
126
127#[derive(Debug, Clone, Default)]
129pub struct TimeoutStats {
130 pub total_operations: usize,
132 pub timed_out_operations: usize,
134 pub successful_retries: usize,
136 pub failed_retries: usize,
138 pub average_timeout_duration: Duration,
140 pub total_retry_attempts: usize,
142}
143
144type PendingEnds = Mutex<Vec<(String, Instant)>>;
146
147pub struct TimeoutDetector {
153 configs: Arc<RwLock<HashMap<OperationType, TimeoutConfig>>>,
154 active_operations: Arc<RwLock<HashMap<String, TimeoutEvent>>>,
155 stats: Arc<RwLock<TimeoutStats>>,
156 pending_ends: Arc<PendingEnds>,
158}
159
160impl Default for TimeoutDetector {
161 fn default() -> Self {
162 Self::new()
163 }
164}
165
166impl TimeoutDetector {
167 pub fn new() -> Self {
171 let mut configs_map = HashMap::new();
173 configs_map.insert(OperationType::ApiCall, TimeoutConfig::api_call());
174 configs_map.insert(OperationType::FileOperation, TimeoutConfig::file_operation());
175 configs_map.insert(OperationType::CodeAnalysis, TimeoutConfig::analysis());
176 configs_map.insert(OperationType::ToolExecution, TimeoutConfig::default());
177 configs_map.insert(OperationType::NetworkRequest, TimeoutConfig::api_call());
178 configs_map.insert(OperationType::Processing, TimeoutConfig::analysis());
179
180 let configs = Arc::new(RwLock::new(configs_map));
181 let active_operations: Arc<RwLock<HashMap<String, TimeoutEvent>>> = Arc::new(RwLock::new(HashMap::new()));
182 let stats = Arc::new(RwLock::new(TimeoutStats::default()));
183 let pending_ends = Arc::new(Mutex::new(Vec::new()));
184
185 Self { configs, active_operations, stats, pending_ends }
186 }
187
188 async fn drain_pending_ends(&self) {
190 let ended = std::mem::take(&mut *self.pending_ends.lock().unwrap_or_else(PoisonError::into_inner));
191 for (operation_id, ended_at) in ended {
192 self.end_operation_inner(&operation_id, ended_at).await;
193 }
194 }
195
196 async fn end_operation_inner(&self, operation_id: &str, ended_at: Instant) {
197 let mut active_ops = self.active_operations.write().await;
198 if let Some(event) = active_ops.remove(operation_id) {
199 let duration = ended_at.saturating_duration_since(event.start_time);
200 let mut stats = self.stats.write().await;
201 if stats.total_operations > 0 {
202 let total_duration = stats.average_timeout_duration * (stats.total_operations - 1) as u32;
203 stats.average_timeout_duration = (total_duration + duration) / stats.total_operations as u32;
204 }
205 }
206 }
207
208 pub async fn set_config(&self, operation_type: OperationType, config: TimeoutConfig) {
210 let mut configs = self.configs.write().await;
211 configs.insert(operation_type, config);
212 }
213
214 pub async fn get_config(&self, operation_type: &OperationType) -> TimeoutConfig {
216 let configs = self.configs.read().await;
217 configs.get(operation_type).cloned().unwrap_or_default()
218 }
219
220 pub async fn start_operation(&self, operation_id: String, operation_type: OperationType) -> TimeoutHandle {
222 self.drain_pending_ends().await;
223 let config = self.get_config(&operation_type).await;
224
225 let event = TimeoutEvent {
226 operation_id: operation_id.clone(),
227 operation_type,
228 start_time: Instant::now(),
229 timeout_duration: config.timeout_duration,
230 retry_count: 0,
231 error_message: None,
232 };
233
234 let mut active_ops = self.active_operations.write().await;
235 active_ops.insert(operation_id.clone(), event);
236
237 let mut stats = self.stats.write().await;
238 stats.total_operations += 1;
239
240 TimeoutHandle {
241 operation_id,
242 pending_ends: Some(Arc::clone(&self.pending_ends)),
243 }
244 }
245
246 pub async fn check_timeout(&self, operation_id: &str) -> Option<TimeoutEvent> {
248 self.drain_pending_ends().await;
249 let active_ops = self.active_operations.read().await;
250 active_ops
251 .get(operation_id)
252 .filter(|event| event.start_time.elapsed() >= event.timeout_duration)
253 .cloned()
254 }
255
256 pub async fn record_timeout(&self, operation_id: &str, error_message: Option<String>) {
258 let mut active_ops = self.active_operations.write().await;
259 if let Some(event) = active_ops.get_mut(operation_id) {
260 event.error_message = error_message;
261 }
262
263 let mut stats = self.stats.write().await;
264 stats.timed_out_operations += 1;
265 }
266
267 pub async fn record_successful_retry(&self, _operation_id: &str) {
269 let mut stats = self.stats.write().await;
270 stats.successful_retries += 1;
271 stats.total_retry_attempts += 1;
272 }
273
274 pub async fn record_failed_retry(&self, _operation_id: &str) {
276 let mut stats = self.stats.write().await;
277 stats.failed_retries += 1;
278 stats.total_retry_attempts += 1;
279 }
280
281 pub async fn end_operation(&self, operation_id: &str) {
283 self.drain_pending_ends().await;
284 self.end_operation_inner(operation_id, Instant::now()).await;
285 }
286
287 pub async fn get_stats(&self) -> TimeoutStats {
289 self.drain_pending_ends().await;
290 self.stats.read().await.clone()
291 }
292
293 pub async fn calculate_retry_delay(&self, operation_type: &OperationType, attempt: u32) -> Duration {
295 let config = self.get_config(operation_type).await;
296
297 let base_delay = config.initial_retry_delay.as_millis() as f64;
298 let multiplier = config.backoff_multiplier.powi(attempt as i32);
299 #[allow(
300 clippy::cast_sign_loss,
301 reason = "Intentional compatibility, platform, or test-only suppression."
302 )]
303 let delay_ms = (base_delay * multiplier) as u64;
304
305 let mut delay = Duration::from_millis(delay_ms.min(config.max_retry_delay.as_millis() as u64));
306
307 if config.use_jitter {
309 use std::time::SystemTime;
310 let seed = SystemTime::now()
311 .duration_since(std::time::UNIX_EPOCH)
312 .unwrap_or_default()
313 .as_nanos() as u64;
314 let jitter_factor = (seed % 100) as f64 / 100.0; #[allow(
316 clippy::cast_sign_loss,
317 reason = "Intentional compatibility, platform, or test-only suppression."
318 )]
319 let jitter_ms = (delay.as_millis() as f64 * 0.1 * jitter_factor) as u64; delay += Duration::from_millis(jitter_ms);
321 }
322
323 delay
324 }
325
326 pub async fn should_retry(&self, operation_type: &OperationType, error: &anyhow::Error, attempt: u32) -> bool {
329 if vtcode_commons::detect_misconfiguration_in_anyhow(error).is_some() {
330 return false;
331 }
332
333 let config = self.get_config(operation_type).await;
334
335 if attempt >= config.max_retries {
336 return false;
337 }
338
339 let error_str = error.to_string();
340
341 let contains_ci = |pattern: &str| {
343 error_str
344 .as_bytes()
345 .windows(pattern.len())
346 .any(|window| window.eq_ignore_ascii_case(pattern.as_bytes()))
347 };
348
349 for retry_error in &config.retry_on_errors {
351 if contains_ci(retry_error) {
352 return true;
353 }
354 }
355
356 if config.retry_on_timeout && (contains_ci("timeout") || contains_ci("timed out")) {
358 return true;
359 }
360
361 false
362 }
363
364 pub async fn execute_with_timeout_retry<F, Fut, T>(
366 &self,
367 operation_id: String,
368 operation_type: OperationType,
369 mut operation: F,
370 ) -> Result<T, anyhow::Error>
371 where
372 F: FnMut() -> Fut,
373 Fut: Future<Output = Result<T, anyhow::Error>>,
374 {
375 let config = self.get_config(&operation_type).await;
376 let mut attempt = 0;
377 let _last_error: Option<anyhow::Error> = None;
378
379 loop {
380 let handle = self
381 .start_operation(format!("{operation_id}_{attempt}"), operation_type.clone())
382 .await;
383
384 let result = match time::timeout(config.timeout_duration, operation()).await {
385 Ok(result) => result,
386 Err(_) => {
387 self.record_timeout(&handle.operation_id, Some("Operation timed out".to_owned()))
388 .await;
389 Err(anyhow::anyhow!("Operation '{}' timed out after {:?}", operation_id, config.timeout_duration))
390 }
391 };
392
393 handle.end().await;
394
395 match result {
396 Ok(value) => {
397 if attempt > 0 {
398 self.record_successful_retry(&format!("{operation_id}_{attempt}")).await;
399 }
400 return Ok(value);
401 }
402 Err(error) => {
403 let should_retry_op = self.should_retry(&operation_type, &error, attempt).await;
404
405 if !should_retry_op {
406 if attempt > 0 {
407 self.record_failed_retry(&format!("{operation_id}_{attempt}")).await;
408 }
409 return Err(error);
410 }
411
412 attempt += 1;
413 self.record_failed_retry(&format!("{operation_id}_{attempt}")).await;
414
415 let delay = self.calculate_retry_delay(&operation_type, attempt).await;
416 tracing::warn!(
417 operation_id,
418 attempt,
419 max_retries = config.max_retries,
420 delay = ?delay,
421 "Operation failed and will be retried"
422 );
423 time::sleep(delay).await;
424 }
425 }
426 }
427 }
428}
429
430impl Clone for TimeoutDetector {
431 fn clone(&self) -> Self {
432 Self {
433 configs: Arc::clone(&self.configs),
434 active_operations: Arc::clone(&self.active_operations),
435 stats: Arc::clone(&self.stats),
436 pending_ends: Arc::clone(&self.pending_ends),
437 }
438 }
439}
440
441pub struct TimeoutHandle {
446 operation_id: String,
447 pending_ends: Option<Arc<PendingEnds>>,
450}
451
452impl TimeoutHandle {
453 pub async fn end(mut self) {
457 self.enqueue_end();
458 }
459
460 fn enqueue_end(&mut self) {
461 if let Some(queue) = self.pending_ends.take() {
462 queue
463 .lock()
464 .unwrap_or_else(PoisonError::into_inner)
465 .push((std::mem::take(&mut self.operation_id), Instant::now()));
466 }
467 }
468
469 pub fn operation_id(&self) -> &str {
471 &self.operation_id
472 }
473}
474
475impl Drop for TimeoutHandle {
476 fn drop(&mut self) {
477 self.enqueue_end();
478 }
479}
480
481use once_cell::sync::Lazy;
483pub static TIMEOUT_DETECTOR: Lazy<TimeoutDetector> = Lazy::new(TimeoutDetector::new);
484
485#[cfg(test)]
486mod tests {
487 use super::*;
488 use std::sync::atomic::{AtomicUsize, Ordering};
489 use std::time::Duration;
490 use tokio::time::sleep;
491
492 #[tokio::test]
493 async fn test_timeout_detection() {
494 let detector = TimeoutDetector::new();
495
496 let config = TimeoutConfig {
498 timeout_duration: Duration::from_millis(10),
499 max_retries: 0,
500 ..Default::default()
501 };
502
503 detector.set_config(OperationType::ApiCall, config).await;
504
505 let result = detector
506 .execute_with_timeout_retry("test_operation".to_owned(), OperationType::ApiCall, || async {
507 sleep(Duration::from_millis(20)).await;
508 Ok("success")
509 })
510 .await;
511
512 assert!(result.is_err());
513 assert!(result.unwrap_err().to_string().contains("timed out"));
514 }
515
516 #[tokio::test]
517 async fn test_successful_retry() {
518 let detector = TimeoutDetector::new();
519
520 let config = TimeoutConfig {
521 timeout_duration: Duration::from_millis(50),
522 max_retries: 2,
523 initial_retry_delay: Duration::from_millis(5),
524 retry_on_timeout: true,
525 ..Default::default()
526 };
527
528 detector.set_config(OperationType::ApiCall, config).await;
529
530 let call_count = Arc::new(AtomicUsize::new(0));
531 let call_count_clone = call_count.clone();
532 let result = detector
533 .execute_with_timeout_retry("test_retry".to_owned(), OperationType::ApiCall, move || {
534 let call_count = call_count_clone.clone();
535 async move {
536 let count = call_count.fetch_add(1, Ordering::SeqCst) + 1;
537 if count == 1 {
538 sleep(Duration::from_millis(60)).await;
540 Ok("should not reach here")
541 } else {
542 sleep(Duration::from_millis(10)).await;
544 Ok("success")
545 }
546 }
547 })
548 .await;
549
550 assert!(result.is_ok());
551 assert_eq!(result.unwrap(), "success");
552 assert_eq!(call_count.load(Ordering::SeqCst), 2);
553
554 let stats = detector.get_stats().await;
555 assert_eq!(stats.successful_retries, 1);
556 assert_eq!(stats.total_retry_attempts, 2);
557 }
558
559 #[test]
560 fn detector_constructs_and_cleans_up_across_runtimes_without_ambient_runtime() {
561 let detector = TimeoutDetector::new();
563
564 let first = tokio::runtime::Builder::new_current_thread()
565 .enable_all()
566 .build()
567 .expect("runtime");
568 let handle = first.block_on(detector.start_operation("op-a".to_owned(), OperationType::ApiCall));
569 drop(first);
570 drop(handle);
572
573 let second = tokio::runtime::Builder::new_current_thread()
574 .enable_all()
575 .build()
576 .expect("runtime");
577 second.block_on(async {
578 let kept = detector.start_operation("op-b".to_owned(), OperationType::ApiCall).await;
579 let stats = detector.get_stats().await;
580 assert_eq!(stats.total_operations, 2);
581 let active = detector.active_operations.read().await;
582 assert!(!active.contains_key("op-a"));
583 assert!(active.contains_key("op-b"));
584 drop(active);
585 kept.end().await;
586 detector.get_stats().await;
587 assert!(detector.active_operations.read().await.is_empty());
588 });
589 }
590
591 #[tokio::test]
592 async fn queued_end_records_duration_at_end_not_at_drain() {
593 let detector = TimeoutDetector::new();
594 let handle = detector.start_operation("op".to_owned(), OperationType::ApiCall).await;
595 handle.end().await;
596 sleep(Duration::from_millis(200)).await;
597
598 let stats = detector.get_stats().await;
599
600 assert!(stats.average_timeout_duration < Duration::from_millis(100));
601 }
602
603 #[tokio::test]
604 async fn test_calculate_retry_delay() {
605 let detector = TimeoutDetector::new();
606
607 let delay = detector.calculate_retry_delay(&OperationType::ApiCall, 0).await;
608 assert!(delay >= Duration::from_millis(200)); let delay2 = detector.calculate_retry_delay(&OperationType::ApiCall, 1).await;
611 assert!(delay2 > delay); }
613
614 #[tokio::test]
615 async fn test_misconfiguration_does_not_retry_even_when_pattern_matches() {
616 let detector = TimeoutDetector::new();
617 let config = TimeoutConfig {
618 retry_on_errors: vec!["network".to_string()],
619 ..Default::default()
620 };
621 detector.set_config(OperationType::ApiCall, config).await;
622
623 let error = anyhow::anyhow!("network error: invalid endpoint in base_url");
624 assert!(!detector.should_retry(&OperationType::ApiCall, &error, 0).await);
625 }
626}