vtcode_core/tools/registry/
approval_recorder.rs1use 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#[derive(Clone)]
14pub struct ApprovalRecorder {
15 manager: Arc<RwLock<JustificationManager>>,
16}
17
18impl ApprovalRecorder {
19 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 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 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 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 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 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 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 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 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 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 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 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 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 assert!(recorder.get_auto_approval_suggestion("read_file", "Read File").await.is_none());
221
222 for _ in 0..5 {
224 let _ = recorder.record_approval("read_file", Some("Read File"), true, None).await;
225 }
226
227 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 assert!(!recorder.should_auto_approve("run_command").await);
240
241 for _ in 0..3 {
243 let _ = recorder.record_approval("run_command", Some("Run Command"), true, None).await;
244 }
245
246 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 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 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}