Skip to main content

vtcode_core/tools/registry/
progress_facade.rs

1//! Progress callback accessors for ToolRegistry.
2
3use super::{ToolProgressCallback, ToolRegistry};
4
5impl ToolRegistry {
6    /// Replace the callback for streaming tool output and progress, returning the previous callback.
7    pub fn replace_progress_callback(&self, callback: Option<ToolProgressCallback>) -> Option<ToolProgressCallback> {
8        let Ok(mut slot) = self.progress_callback.write() else {
9            return None;
10        };
11        std::mem::replace(&mut *slot, callback)
12    }
13
14    /// Set the callback for streaming tool output and progress
15    pub fn set_progress_callback(&self, callback: ToolProgressCallback) {
16        let _ = self.replace_progress_callback(Some(callback));
17    }
18
19    /// Clear the progress callback
20    pub fn clear_progress_callback(&self) {
21        let _ = self.replace_progress_callback(None);
22    }
23
24    /// Get the current progress callback if set
25    pub fn progress_callback(&self) -> Option<ToolProgressCallback> {
26        self.progress_callback.read().unwrap_or_else(|e| e.into_inner()).clone()
27    }
28}
29
30#[cfg(test)]
31mod tests {
32    use super::*;
33    use std::sync::Arc;
34    use std::sync::atomic::{AtomicUsize, Ordering};
35    use tempfile::TempDir;
36
37    #[tokio::test]
38    async fn replace_progress_callback_restores_previous() {
39        let temp_dir = TempDir::new().expect("create temp dir");
40        let registry = ToolRegistry::new(temp_dir.path().to_path_buf()).await;
41
42        let first_hits = Arc::new(AtomicUsize::new(0));
43        let first_hits_clone = Arc::clone(&first_hits);
44        registry.set_progress_callback(Arc::new(move |_, _| {
45            let _ = first_hits_clone.fetch_add(1, Ordering::SeqCst);
46        }));
47
48        let second_hits = Arc::new(AtomicUsize::new(0));
49        let second_hits_clone = Arc::clone(&second_hits);
50        let previous = registry.replace_progress_callback(Some(Arc::new(move |_, _| {
51            let _ = second_hits_clone.fetch_add(1, Ordering::SeqCst);
52        })));
53
54        if let Some(current) = registry.progress_callback() {
55            current("run_pty_cmd", "chunk");
56        }
57        assert_eq!(second_hits.load(Ordering::SeqCst), 1);
58
59        let _ = registry.replace_progress_callback(previous);
60        if let Some(current) = registry.progress_callback() {
61            current("run_pty_cmd", "chunk");
62        }
63        assert_eq!(first_hits.load(Ordering::SeqCst), 1);
64    }
65
66    #[tokio::test]
67    async fn clear_progress_callback_removes_registered_callback() {
68        let temp_dir = TempDir::new().expect("create temp dir");
69        let registry = ToolRegistry::new(temp_dir.path().to_path_buf()).await;
70
71        registry.set_progress_callback(Arc::new(|_, _| {}));
72        assert!(registry.progress_callback().is_some());
73
74        registry.clear_progress_callback();
75        assert!(registry.progress_callback().is_none());
76    }
77}