Skip to main content

vtcode_llm/providers/
streaming_progress.rs

1//! Provider-agnostic streaming timeout progress tracking
2//!
3//! This module provides a unified interface for tracking streaming timeout progress
4//! across all LLM providers (OpenAI, Anthropic, Gemini, Ollama, etc.)
5
6use std::sync::Arc;
7use std::sync::atomic::{AtomicU8, Ordering};
8use std::time::{Duration, Instant};
9use tracing::warn;
10
11/// Callback for streaming timeout progress updates
12/// Progress value is 0.0-1.0 representing elapsed / total_timeout
13pub type StreamingProgressCallback = Box<dyn Fn(f32) + Send + Sync>;
14
15/// Unified streaming progress tracker for all LLM providers
16#[derive(Clone)]
17pub struct StreamingProgressTracker {
18    callback: Option<Arc<StreamingProgressCallback>>,
19    warning_threshold: f32,
20    total_timeout: Duration,
21    start_time: Arc<Instant>,
22    last_reported_progress: Arc<AtomicU8>,
23}
24
25impl StreamingProgressTracker {
26    /// Create a new streaming progress tracker
27    fn new(total_timeout: Duration) -> Self {
28        Self {
29            callback: None,
30            warning_threshold: 0.8,
31            total_timeout,
32            start_time: Arc::new(Instant::now()),
33            last_reported_progress: Arc::new(AtomicU8::new(0)),
34        }
35    }
36
37    /// Set a progress callback
38    fn with_callback(mut self, callback: StreamingProgressCallback) -> Self {
39        self.callback = Some(Arc::new(callback));
40        self
41    }
42
43    /// Set the warning threshold (0.0-1.0)
44    fn with_warning_threshold(mut self, threshold: f32) -> Self {
45        self.warning_threshold = threshold.clamp(0.0, 1.0);
46        self
47    }
48
49    /// Report that the first chunk has been received
50    pub fn report_first_chunk(&self) {
51        self.report_progress(0.1);
52    }
53
54    /// Report progress with elapsed time
55    fn report_chunk_received(&self) {
56        let elapsed = self.start_time.elapsed();
57        self.report_progress_with_elapsed(elapsed);
58    }
59
60    /// Report progress at a specific elapsed duration
61    fn report_progress_with_elapsed(&self, elapsed: Duration) {
62        if self.total_timeout.as_secs() == 0 {
63            return;
64        }
65
66        let progress = elapsed.as_secs_f32() / self.total_timeout.as_secs_f32();
67        self.report_progress(progress.min(0.99)); // Cap at 99%
68    }
69
70    /// Report error or timeout (100% progress)
71    fn report_error(&self) {
72        self.report_progress(1.0);
73    }
74
75    /// Get current progress as percentage (0-100)
76    fn progress_percent(&self) -> u8 {
77        self.last_reported_progress.load(Ordering::Relaxed)
78    }
79
80    /// Get elapsed time since start
81    fn elapsed(&self) -> Duration {
82        self.start_time.elapsed()
83    }
84
85    /// Check if warning threshold has been exceeded
86    fn is_approaching_timeout(&self) -> bool {
87        let elapsed = self.start_time.elapsed();
88        if self.total_timeout.as_secs() == 0 {
89            return false;
90        }
91
92        let elapsed_progress = elapsed.as_secs_f32() / self.total_timeout.as_secs_f32();
93        let reported_progress = f32::from(self.last_reported_progress.load(Ordering::Relaxed)) / 100.0;
94        elapsed_progress.max(reported_progress) >= self.warning_threshold
95    }
96
97    // Private: Report progress with clamping and threshold checking
98    fn report_progress(&self, progress: f32) {
99        let progress_clamped = progress.clamp(0.0, 1.0);
100        #[allow(
101            clippy::cast_sign_loss,
102            reason = "Intentional compatibility, platform, or test-only suppression."
103        )]
104        let percent = (progress_clamped * 100.0) as u8;
105
106        // Only update if progress changed by at least 1%
107        let last_percent = self.last_reported_progress.load(Ordering::Relaxed);
108        if percent <= last_percent {
109            return;
110        }
111
112        self.last_reported_progress.store(percent, Ordering::Relaxed);
113
114        // Call the callback if set
115        if let Some(ref callback) = self.callback {
116            callback(progress_clamped);
117        }
118
119        // Warn if approaching threshold
120        if progress_clamped >= self.warning_threshold && progress_clamped < 1.0 {
121            warn!(
122                "Streaming operation at {:.0}% of timeout limit ({:?}/{:?} elapsed). Approaching timeout.",
123                progress_clamped * 100.0,
124                self.elapsed(),
125                self.total_timeout
126            );
127        }
128    }
129}
130
131/// Builder for creating streaming progress trackers with fluent API
132pub struct StreamingProgressBuilder {
133    total_timeout: Duration,
134    callback: Option<StreamingProgressCallback>,
135    warning_threshold: f32,
136}
137
138impl StreamingProgressBuilder {
139    /// Create a new builder with total timeout in seconds
140    fn new(timeout_secs: u64) -> Self {
141        Self {
142            total_timeout: Duration::from_secs(timeout_secs),
143            callback: None,
144            warning_threshold: 0.8,
145        }
146    }
147
148    /// Create a new builder with a specific duration
149    pub fn with_duration(duration: Duration) -> Self {
150        Self {
151            total_timeout: duration,
152            callback: None,
153            warning_threshold: 0.8,
154        }
155    }
156
157    /// Set the progress callback
158    pub fn callback(mut self, callback: StreamingProgressCallback) -> Self {
159        self.callback = Some(callback);
160        self
161    }
162
163    /// Set the warning threshold (0.0-1.0)
164    fn warning_threshold(mut self, threshold: f32) -> Self {
165        self.warning_threshold = threshold.clamp(0.0, 1.0);
166        self
167    }
168
169    /// Build the tracker
170    fn build(self) -> StreamingProgressTracker {
171        let mut tracker = StreamingProgressTracker::new(self.total_timeout);
172        if let Some(callback) = self.callback {
173            tracker.callback = Some(Arc::new(callback));
174        }
175        tracker.warning_threshold = self.warning_threshold;
176        tracker
177    }
178}
179
180#[cfg(test)]
181mod tests {
182    use super::*;
183    use std::sync::Mutex;
184
185    #[test]
186    fn test_progress_tracker_creation() {
187        let tracker = StreamingProgressTracker::new(Duration::from_secs(600));
188        assert_eq!(tracker.progress_percent(), 0);
189        assert!(!tracker.is_approaching_timeout());
190    }
191
192    #[test]
193    fn test_progress_reporting() {
194        let tracker = StreamingProgressTracker::new(Duration::from_secs(100));
195
196        tracker.report_progress_with_elapsed(Duration::from_secs(30));
197        assert_eq!(tracker.progress_percent(), 30);
198
199        tracker.report_progress_with_elapsed(Duration::from_secs(80));
200        assert_eq!(tracker.progress_percent(), 80);
201    }
202
203    #[test]
204    fn test_warning_threshold() {
205        let tracker = StreamingProgressTracker::new(Duration::from_secs(100)).with_warning_threshold(0.8);
206
207        tracker.report_progress_with_elapsed(Duration::from_secs(50));
208        assert!(!tracker.is_approaching_timeout());
209
210        tracker.report_progress_with_elapsed(Duration::from_secs(85));
211        assert!(tracker.is_approaching_timeout());
212    }
213
214    #[test]
215    fn test_callback_execution() {
216        let progress_log = Arc::new(Mutex::new(Vec::new()));
217        let progress_clone = progress_log.clone();
218
219        let tracker =
220            StreamingProgressTracker::new(Duration::from_secs(100)).with_callback(Box::new(move |progress: f32| {
221                progress_clone.lock().unwrap().push(progress);
222            }));
223
224        tracker.report_progress_with_elapsed(Duration::from_secs(30));
225        tracker.report_progress_with_elapsed(Duration::from_secs(60));
226        tracker.report_progress_with_elapsed(Duration::from_secs(90));
227
228        let log = progress_log.lock().unwrap();
229        assert!(!log.is_empty());
230        assert!(log.iter().all(|&p| (0.0..=1.0).contains(&p)));
231    }
232
233    #[test]
234    fn test_builder_pattern() {
235        let tracker = StreamingProgressBuilder::new(300).warning_threshold(0.75).build();
236
237        assert_eq!(tracker.total_timeout.as_secs(), 300);
238        assert!((tracker.warning_threshold - 0.75).abs() < f32::EPSILON);
239    }
240
241    #[test]
242    fn test_zero_timeout_safety() {
243        let tracker = StreamingProgressTracker::new(Duration::from_secs(0));
244        tracker.report_chunk_received(); // Should not panic
245        assert!(!tracker.is_approaching_timeout());
246    }
247
248    #[test]
249    fn test_progress_clamping() {
250        let tracker = StreamingProgressTracker::new(Duration::from_secs(100));
251
252        tracker.report_progress_with_elapsed(Duration::from_secs(150)); // Beyond timeout
253        assert_eq!(tracker.progress_percent(), 99); // Clamped at 99%
254
255        tracker.report_error();
256        assert_eq!(tracker.progress_percent(), 100);
257    }
258}