vtcode_llm/providers/
streaming_progress.rs1use std::sync::Arc;
7use std::sync::atomic::{AtomicU8, Ordering};
8use std::time::{Duration, Instant};
9use tracing::warn;
10
11pub type StreamingProgressCallback = Box<dyn Fn(f32) + Send + Sync>;
14
15#[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 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 fn with_callback(mut self, callback: StreamingProgressCallback) -> Self {
39 self.callback = Some(Arc::new(callback));
40 self
41 }
42
43 fn with_warning_threshold(mut self, threshold: f32) -> Self {
45 self.warning_threshold = threshold.clamp(0.0, 1.0);
46 self
47 }
48
49 pub fn report_first_chunk(&self) {
51 self.report_progress(0.1);
52 }
53
54 fn report_chunk_received(&self) {
56 let elapsed = self.start_time.elapsed();
57 self.report_progress_with_elapsed(elapsed);
58 }
59
60 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)); }
69
70 fn report_error(&self) {
72 self.report_progress(1.0);
73 }
74
75 fn progress_percent(&self) -> u8 {
77 self.last_reported_progress.load(Ordering::Relaxed)
78 }
79
80 fn elapsed(&self) -> Duration {
82 self.start_time.elapsed()
83 }
84
85 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 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 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 if let Some(ref callback) = self.callback {
116 callback(progress_clamped);
117 }
118
119 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
131pub struct StreamingProgressBuilder {
133 total_timeout: Duration,
134 callback: Option<StreamingProgressCallback>,
135 warning_threshold: f32,
136}
137
138impl StreamingProgressBuilder {
139 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 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 pub fn callback(mut self, callback: StreamingProgressCallback) -> Self {
159 self.callback = Some(callback);
160 self
161 }
162
163 fn warning_threshold(mut self, threshold: f32) -> Self {
165 self.warning_threshold = threshold.clamp(0.0, 1.0);
166 self
167 }
168
169 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(); 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)); assert_eq!(tracker.progress_percent(), 99); tracker.report_error();
256 assert_eq!(tracker.progress_percent(), 100);
257 }
258}