Skip to main content

vtcode_core/tools/registry/
approval_recorder.rs

1/// Approval Decision Recording and Learning
2///
3/// Records user approval decisions for high-risk tools and enables pattern learning
4/// to reduce approval friction over time.
5use super::justification::{ApprovalPattern, JustificationManager};
6use anyhow::Result;
7use std::path::PathBuf;
8use std::sync::Arc;
9use tokio::sync::RwLock;
10use vtcode_commons::VtCodePaths;
11
12/// Records tool approval decisions for learning
13#[derive(Clone)]
14pub struct ApprovalRecorder {
15    manager: Arc<RwLock<JustificationManager>>,
16}
17
18impl ApprovalRecorder {
19    /// Create a new approval recorder
20    pub fn new(cache_dir: PathBuf) -> Self {
21        let manager = JustificationManager::new(cache_dir);
22        Self { manager: Arc::new(RwLock::new(manager)) }
23    }
24
25    /// Create a recorder that can recover approval patterns from older cache directories.
26    pub fn new_with_legacy_cache_dirs(
27        cache_dir: PathBuf,
28        legacy_cache_dirs: impl IntoIterator<Item = PathBuf>,
29    ) -> Self {
30        let legacy_pattern_files = legacy_cache_dirs
31            .into_iter()
32            .map(|directory| directory.join("approval_patterns.json"));
33        let manager = JustificationManager::new_with_legacy_pattern_files(cache_dir, legacy_pattern_files);
34        Self { manager: Arc::new(RwLock::new(manager)) }
35    }
36
37    /// Create a recorder without reading approval-pattern files.
38    ///
39    /// First-paint path uses this; hydration must call [`Self::reload`] before
40    /// the first model turn so auto-approval history matches prior behavior.
41    pub fn new_deferred(cache_dir: PathBuf, legacy_cache_dirs: impl IntoIterator<Item = PathBuf>) -> Self {
42        let legacy_pattern_files = legacy_cache_dirs
43            .into_iter()
44            .map(|directory| directory.join("approval_patterns.json"));
45        let manager = JustificationManager::new_without_load(cache_dir, legacy_pattern_files);
46        Self { manager: Arc::new(RwLock::new(manager)) }
47    }
48
49    /// Load (or re-load) approval patterns from disk, merging with in-memory state.
50    pub async fn reload(&self) {
51        let manager = self.manager.read().await;
52        if let Err(err) = manager.refresh_patterns() {
53            tracing::debug!(error = %err, "Failed to load approval patterns during hydration");
54        }
55    }
56}
57
58impl Default for ApprovalRecorder {
59    fn default() -> Self {
60        match VtCodePaths::resolve() {
61            Ok(paths) => match paths.ensure_cache_child_dir("approval") {
62                Ok(cache_dir) => {
63                    let legacy_cache_dirs = [
64                        paths.cache_dir().to_path_buf(),
65                        paths.config_dir().join("cache"),
66                        paths.legacy_dir().join("cache"),
67                    ];
68                    Self::new_with_legacy_cache_dirs(cache_dir, legacy_cache_dirs)
69                }
70                Err(_) => Self::fallback(),
71            },
72            Err(_) => Self::fallback(),
73        }
74    }
75}
76
77impl ApprovalRecorder {
78    fn fallback() -> Self {
79        let cache_dir = std::env::temp_dir()
80            .join(format!("vtcode-{}", std::process::id()))
81            .join("approval");
82        Self::new(cache_dir)
83    }
84}
85
86impl ApprovalRecorder {
87    /// Record a user's approval decision for a learned approval key
88    pub async fn record_approval(
89        &self,
90        approval_key: &str,
91        display_name: Option<&str>,
92        approved: bool,
93        reason: Option<String>,
94    ) -> Result<()> {
95        let manager = self.manager.write().await;
96        manager.record_decision(approval_key, display_name, approved, reason);
97        Ok(())
98    }
99
100    /// Get the approval pattern for a learned approval key
101    pub async fn get_pattern(&self, approval_key: &str) -> Option<ApprovalPattern> {
102        let manager = self.manager.read().await;
103        manager.get_pattern(approval_key)
104    }
105
106    /// Check if a key has high approval rate from history
107    pub async fn has_high_approval_rate(&self, approval_key: &str) -> bool {
108        let manager = self.manager.read().await;
109        if let Some(pattern) = manager.get_pattern(approval_key) {
110            pattern.has_high_approval_rate()
111        } else {
112            false
113        }
114    }
115
116    /// Get learning summary for a learned approval key
117    pub async fn get_learning_summary(&self, approval_key: &str) -> Option<String> {
118        let manager = self.manager.read().await;
119        manager.get_learning_summary(approval_key)
120    }
121
122    /// Get approval count for a learned approval key
123    pub async fn get_approval_count(&self, approval_key: &str) -> u32 {
124        let manager = self.manager.read().await;
125        if let Some(pattern) = manager.get_pattern(approval_key) {
126            pattern.approval_count()
127        } else {
128            0
129        }
130    }
131
132    /// Should auto-approve based on approval pattern
133    /// Rules:
134    /// - At least 3 approvals
135    /// - Approval rate > 80%
136    ///
137    /// Refreshes the in-memory pattern map from disk first so we observe
138    /// approvals recorded by concurrent sessions (e.g. another running vtcode
139    /// instance sharing the same user cache approval-pattern file).
140    pub async fn should_auto_approve(&self, approval_key: &str) -> bool {
141        let manager = self.manager.write().await;
142        if let Err(err) = manager.refresh_patterns() {
143            tracing::debug!(
144                approval_key = %approval_key,
145                error = %err,
146                "Failed to refresh approval patterns before auto-approve check"
147            );
148        }
149        if let Some(pattern) = manager.get_pattern(approval_key) {
150            pattern.has_high_approval_rate()
151        } else {
152            false
153        }
154    }
155
156    /// Suggest auto-approval message if user has approved this target many times
157    pub async fn get_auto_approval_suggestion(
158        &self,
159        approval_key: &str,
160        fallback_display_name: &str,
161    ) -> Option<String> {
162        let manager = self.manager.read().await;
163        if let Some(pattern) = manager.get_pattern(approval_key) {
164            let rate = pattern.approval_rate();
165            if pattern.approval_count() >= 5 {
166                let display_name = pattern.display_name(fallback_display_name);
167                return Some(format!(
168                    "You've approved {} {} times ({:.0}% approval rate)",
169                    display_name,
170                    pattern.approval_count(),
171                    rate * 100.0
172                ));
173            }
174        }
175        None
176    }
177}
178
179#[cfg(test)]
180mod tests {
181    use super::*;
182    use std::collections::HashMap;
183    use tempfile::TempDir;
184
185    fn temp_cache_dir() -> TempDir {
186        TempDir::new().expect("temp approval cache")
187    }
188
189    #[tokio::test]
190    async fn test_approval_recording() {
191        let temp_dir = temp_cache_dir();
192        let recorder = ApprovalRecorder::new(temp_dir.path().to_path_buf());
193
194        // Record some approvals
195        recorder
196            .record_approval("read_file", Some("Read File"), true, None)
197            .await
198            .unwrap();
199        recorder
200            .record_approval("read_file", Some("Read File"), true, None)
201            .await
202            .unwrap();
203        recorder
204            .record_approval("read_file", Some("Read File"), false, None)
205            .await
206            .unwrap();
207
208        // Check pattern
209        let pattern = recorder.get_pattern("read_file").await;
210        assert!(pattern.is_some());
211        assert_eq!(pattern.unwrap().approval_count(), 2);
212    }
213
214    #[tokio::test]
215    async fn test_auto_approval_suggestion() {
216        let temp_dir = temp_cache_dir();
217        let recorder = ApprovalRecorder::new(temp_dir.path().to_path_buf());
218
219        // Not enough approvals initially
220        assert!(recorder.get_auto_approval_suggestion("read_file", "Read File").await.is_none());
221
222        // Add 5 approvals
223        for _ in 0..5 {
224            let _ = recorder.record_approval("read_file", Some("Read File"), true, None).await;
225        }
226
227        // Now we should get a suggestion
228        let suggestion = recorder.get_auto_approval_suggestion("read_file", "Read File").await;
229        assert!(suggestion.is_some());
230        assert!(suggestion.unwrap().contains("100%"));
231    }
232
233    #[tokio::test]
234    async fn test_should_auto_approve() {
235        let temp_dir = temp_cache_dir();
236        let recorder = ApprovalRecorder::new(temp_dir.path().to_path_buf());
237
238        // Not approved initially
239        assert!(!recorder.should_auto_approve("run_command").await);
240
241        // Add 3 approvals (minimum threshold)
242        for _ in 0..3 {
243            let _ = recorder.record_approval("run_command", Some("Run Command"), true, None).await;
244        }
245
246        // Now should auto-approve
247        assert!(recorder.should_auto_approve("run_command").await);
248    }
249
250    #[tokio::test]
251    async fn test_auto_approval_suggestion_uses_display_name() {
252        let temp_dir = temp_cache_dir();
253        let recorder = ApprovalRecorder::new(temp_dir.path().to_path_buf());
254
255        for _ in 0..5 {
256            let _ = recorder
257                .record_approval(
258                    "cargo test|sandbox_permissions=\"require_escalated\"|additional_permissions=null",
259                    Some("commands starting with `cargo test`"),
260                    true,
261                    None,
262                )
263                .await;
264        }
265
266        let suggestion = recorder
267            .get_auto_approval_suggestion(
268                "cargo test|sandbox_permissions=\"require_escalated\"|additional_permissions=null",
269                "fallback label",
270            )
271            .await
272            .expect("suggestion");
273        assert!(suggestion.contains("commands starting with `cargo test`"));
274    }
275
276    #[tokio::test]
277    async fn test_should_auto_approve_refreshes_patterns_from_disk() {
278        // Simulates a second vtcode session: one ApprovalRecorder records
279        // approvals to disk, then a separately constructed recorder must
280        // observe them on the next auto-approve check without restart.
281        let temp_dir = temp_cache_dir();
282
283        let key =
284            "find src -type f -name '*.rs' '|' sort|sandbox_permissions=\"use_default\"|additional_permissions=null";
285
286        let reader = ApprovalRecorder::new(temp_dir.path().to_path_buf());
287        assert!(!reader.should_auto_approve(key).await);
288
289        let writer = ApprovalRecorder::new(temp_dir.path().to_path_buf());
290        for _ in 0..3 {
291            writer.record_approval(key, Some("find src"), true, None).await.unwrap();
292        }
293
294        // Without the disk refresh in should_auto_approve, the reader's
295        // in-memory map would still be empty and this assertion would fail.
296        assert!(reader.should_auto_approve(key).await);
297    }
298
299    #[tokio::test]
300    async fn recovers_legacy_patterns_and_republishes_to_canonical_cache() {
301        let temp_dir = temp_cache_dir();
302        let legacy_dir = temp_dir.path().join("legacy-cache");
303        let canonical_dir = temp_dir.path().join("cache/approval");
304        std::fs::create_dir_all(&legacy_dir).expect("legacy cache directory");
305
306        let mut patterns = HashMap::new();
307        patterns.insert(
308            "run_command".to_string(),
309            ApprovalPattern {
310                tool_name: "run_command".to_string(),
311                display_name: Some("Run Command".to_string()),
312                approve_count: 3,
313                deny_count: 0,
314                last_decision: Some(true),
315                recent_reason: None,
316            },
317        );
318        std::fs::write(
319            legacy_dir.join("approval_patterns.json"),
320            serde_json::to_vec(&patterns).expect("serialize legacy patterns"),
321        )
322        .expect("write legacy patterns");
323
324        let recorder = ApprovalRecorder::new_with_legacy_cache_dirs(canonical_dir.clone(), [legacy_dir]);
325        assert_eq!(recorder.get_approval_count("run_command").await, 3);
326        assert!(recorder.has_high_approval_rate("run_command").await);
327        assert!(canonical_dir.join("approval_patterns.json").is_file());
328    }
329
330    #[tokio::test]
331    async fn canonical_patterns_take_precedence_over_legacy_patterns() {
332        let temp_dir = temp_cache_dir();
333        let legacy_dir = temp_dir.path().join("legacy-cache");
334        let canonical_dir = temp_dir.path().join("cache/approval");
335        std::fs::create_dir_all(&legacy_dir).expect("legacy cache directory");
336        std::fs::create_dir_all(&canonical_dir).expect("canonical cache directory");
337
338        let pattern = |approve_count| ApprovalPattern {
339            tool_name: "run_command".to_owned(),
340            display_name: Some("Run Command".to_owned()),
341            approve_count,
342            deny_count: 0,
343            last_decision: Some(true),
344            recent_reason: None,
345        };
346        let mut canonical_patterns = HashMap::new();
347        canonical_patterns.insert("run_command".to_owned(), pattern(1));
348        let mut legacy_patterns = HashMap::new();
349        legacy_patterns.insert("run_command".to_owned(), pattern(5));
350
351        std::fs::write(
352            canonical_dir.join("approval_patterns.json"),
353            serde_json::to_vec(&canonical_patterns).expect("serialize canonical patterns"),
354        )
355        .expect("write canonical patterns");
356        std::fs::write(
357            legacy_dir.join("approval_patterns.json"),
358            serde_json::to_vec(&legacy_patterns).expect("serialize legacy patterns"),
359        )
360        .expect("write legacy patterns");
361
362        let recorder = ApprovalRecorder::new_with_legacy_cache_dirs(canonical_dir, [legacy_dir]);
363        assert_eq!(recorder.get_approval_count("run_command").await, 1);
364    }
365
366    #[tokio::test]
367    async fn malformed_canonical_patterns_recover_legacy_without_replacing_them() {
368        let temp_dir = temp_cache_dir();
369        let legacy_dir = temp_dir.path().join("legacy-cache");
370        let canonical_dir = temp_dir.path().join("cache/approval");
371        std::fs::create_dir_all(&legacy_dir).expect("legacy cache directory");
372        std::fs::create_dir_all(&canonical_dir).expect("canonical cache directory");
373        let canonical_file = canonical_dir.join("approval_patterns.json");
374        std::fs::write(&canonical_file, b"not json").expect("malformed canonical patterns");
375
376        let mut patterns = HashMap::new();
377        patterns.insert(
378            "run_command".to_owned(),
379            ApprovalPattern {
380                tool_name: "run_command".to_owned(),
381                display_name: Some("Run Command".to_owned()),
382                approve_count: 3,
383                deny_count: 0,
384                last_decision: Some(true),
385                recent_reason: None,
386            },
387        );
388        std::fs::write(
389            legacy_dir.join("approval_patterns.json"),
390            serde_json::to_vec(&patterns).expect("serialize legacy patterns"),
391        )
392        .expect("write legacy patterns");
393
394        let recorder = ApprovalRecorder::new_with_legacy_cache_dirs(canonical_dir, [legacy_dir]);
395
396        assert_eq!(recorder.get_approval_count("run_command").await, 3);
397        assert_eq!(std::fs::read(canonical_file).expect("read canonical patterns"), b"not json");
398    }
399
400    #[tokio::test]
401    async fn deferred_recorder_skips_disk_read_until_reload() {
402        let temp_dir = temp_cache_dir();
403        let canonical_dir = temp_dir.path().join("cache/approval");
404        std::fs::create_dir_all(&canonical_dir).expect("canonical cache directory");
405
406        let mut patterns = HashMap::new();
407        patterns.insert(
408            "run_command".to_string(),
409            ApprovalPattern {
410                tool_name: "run_command".to_string(),
411                display_name: Some("Run Command".to_string()),
412                approve_count: 3,
413                deny_count: 0,
414                last_decision: Some(true),
415                recent_reason: None,
416            },
417        );
418        std::fs::write(
419            canonical_dir.join("approval_patterns.json"),
420            serde_json::to_vec(&patterns).expect("serialize patterns"),
421        )
422        .expect("write patterns");
423
424        let recorder = ApprovalRecorder::new_deferred(canonical_dir, Vec::<PathBuf>::new());
425        assert_eq!(
426            recorder.get_approval_count("run_command").await,
427            0,
428            "deferred recorder must not touch disk on the paint path"
429        );
430
431        recorder.reload().await;
432        assert_eq!(recorder.get_approval_count("run_command").await, 3);
433        assert!(recorder.has_high_approval_rate("run_command").await);
434    }
435
436    #[tokio::test]
437    async fn test_shell_scoped_history_does_not_reuse_tool_level_key() {
438        let temp_dir = temp_cache_dir();
439        let recorder = ApprovalRecorder::new(temp_dir.path().to_path_buf());
440
441        for _ in 0..5 {
442            let _ = recorder
443                .record_approval("command_session", Some("Unified Exec"), true, None)
444                .await;
445        }
446
447        assert_eq!(
448            recorder
449                .get_approval_count("cargo test|sandbox_permissions=\"require_escalated\"|additional_permissions=null")
450                .await,
451            0
452        );
453        assert!(
454            recorder
455                .get_auto_approval_suggestion(
456                    "cargo test|sandbox_permissions=\"require_escalated\"|additional_permissions=null",
457                    "commands starting with `cargo test`",
458                )
459                .await
460                .is_none()
461        );
462    }
463}