vtcode_core/tools/registry/
justification.rs1use crate::tools::registry::risk_scorer::RiskLevel;
6use anyhow::{Context, Result};
7use hashbrown::HashMap;
8use serde::{Deserialize, Serialize};
9use std::path::{Path, PathBuf};
10use vtcode_commons::VtCodePaths;
11
12#[derive(Debug, Clone, Serialize, Deserialize)]
14pub struct ToolJustification {
15 pub tool_name: String,
17 pub reason: String,
19 pub expected_outcome: Option<String>,
21 pub risk_level: String,
23 pub timestamp: String,
25}
26
27impl ToolJustification {
28 pub fn new(tool_name: impl Into<String>, reason: impl Into<String>, risk_level: &RiskLevel) -> Self {
30 Self {
31 tool_name: tool_name.into(),
32 reason: reason.into(),
33 expected_outcome: None,
34 risk_level: format!("{risk_level:?}"),
35 timestamp: chrono::Local::now().to_rfc3339(),
36 }
37 }
38
39 pub fn with_outcome(mut self, outcome: impl Into<String>) -> Self {
41 self.expected_outcome = Some(outcome.into());
42 self
43 }
44
45 pub fn format_for_dialog(&self) -> Vec<String> {
47 let mut lines = vec![];
48
49 lines.push(String::new());
50 lines.push("Agent Reasoning:".to_owned());
51
52 for line in self.reason.lines() {
54 let wrapped = textwrap::fill(&format!(" {line}"), 78);
55 for wrapped_line in wrapped.lines() {
56 lines.push(wrapped_line.to_owned());
57 }
58 }
59
60 if let Some(outcome) = &self.expected_outcome {
61 lines.push(String::new());
62 lines.push("Expected Outcome:".to_owned());
63 let wrapped = textwrap::fill(&format!(" {outcome}"), 78);
64 for wrapped_line in wrapped.lines() {
65 lines.push(wrapped_line.to_owned());
66 }
67 }
68
69 lines.push(String::new());
70 lines.push(format!("Risk Level: {}", self.risk_level));
71
72 lines
73 }
74}
75
76#[derive(Debug, Clone, Serialize, Deserialize, Default)]
78pub struct ApprovalPattern {
79 pub tool_name: String,
81 #[serde(default)]
83 pub display_name: Option<String>,
84 pub approve_count: u32,
86 pub deny_count: u32,
88 pub last_decision: Option<bool>,
90 pub recent_reason: Option<String>,
92}
93
94impl ApprovalPattern {
95 pub fn approval_rate(&self) -> f32 {
97 let total = self.approve_count + self.deny_count;
98 if total == 0 {
99 0.0
100 } else {
101 self.approve_count as f32 / total as f32
102 }
103 }
104
105 pub fn has_high_approval_rate(&self) -> bool {
107 self.approval_count() >= 3 && self.approval_rate() > 0.8
108 }
109
110 pub fn approval_count(&self) -> u32 {
112 self.approve_count
113 }
114
115 pub fn display_name<'a>(&'a self, fallback: &'a str) -> &'a str {
116 self.display_name.as_deref().unwrap_or(fallback)
117 }
118}
119
120fn merge_pattern_from_disk(local: &mut ApprovalPattern, disk: &ApprovalPattern) {
124 local.approve_count = local.approve_count.max(disk.approve_count);
125 local.deny_count = local.deny_count.max(disk.deny_count);
126 if disk.display_name.is_some() {
127 local.display_name = disk.display_name.clone();
128 }
129 if disk.last_decision.is_some() {
130 local.last_decision = disk.last_decision;
131 }
132 if disk.recent_reason.is_some() {
133 local.recent_reason = disk.recent_reason.clone();
134 }
135}
136
137fn merge_pattern_map(local: &mut HashMap<String, ApprovalPattern>, disk_patterns: HashMap<String, ApprovalPattern>) {
138 for (key, disk) in disk_patterns {
139 local
140 .entry(key)
141 .and_modify(|local| merge_pattern_from_disk(local, &disk))
142 .or_insert(disk);
143 }
144}
145
146pub struct JustificationManager {
148 cache_dir: PathBuf,
149 legacy_pattern_files: Vec<PathBuf>,
150 patterns: std::sync::Arc<std::sync::Mutex<HashMap<String, ApprovalPattern>>>,
151}
152
153impl JustificationManager {
154 pub fn new(cache_dir: PathBuf) -> Self {
156 Self::new_with_legacy_pattern_files(cache_dir, Vec::new())
157 }
158
159 pub(crate) fn new_with_legacy_pattern_files(
161 cache_dir: PathBuf,
162 legacy_pattern_files: impl IntoIterator<Item = PathBuf>,
163 ) -> Self {
164 let manager = Self::new_without_load(cache_dir, legacy_pattern_files);
165 let _ = manager.load_patterns();
167
168 manager
169 }
170
171 pub(crate) fn new_without_load(
176 cache_dir: PathBuf,
177 legacy_pattern_files: impl IntoIterator<Item = PathBuf>,
178 ) -> Self {
179 let canonical_pattern_file = cache_dir.join("approval_patterns.json");
180 let mut legacy_pattern_files = legacy_pattern_files
181 .into_iter()
182 .filter(|path| path != &canonical_pattern_file)
183 .collect::<Vec<_>>();
184 legacy_pattern_files.dedup();
185
186 let patterns = std::sync::Arc::new(std::sync::Mutex::new(HashMap::new()));
187 Self { cache_dir, legacy_pattern_files, patterns }
188 }
189
190 fn load_patterns(&self) -> Result<()> {
197 let canonical_pattern_file = self.cache_dir.join("approval_patterns.json");
198 let canonical_state = read_patterns_file(&canonical_pattern_file)?;
199 let canonical_missing = matches!(&canonical_state, PatternFileState::Missing);
200 let mut load_error = None;
201 match canonical_state {
202 PatternFileState::Valid(patterns) => {
203 self.merge_loaded_patterns([patterns])?;
204 return Ok(());
205 }
206 PatternFileState::Missing => {}
207 PatternFileState::Malformed(error) => load_error = Some(error),
208 }
209
210 let mut loaded_patterns = Vec::new();
211 for patterns_file in &self.legacy_pattern_files {
212 match read_patterns_file(patterns_file) {
213 Ok(PatternFileState::Valid(patterns)) => loaded_patterns.push(patterns),
214 Ok(PatternFileState::Missing) => {}
215 Ok(PatternFileState::Malformed(error)) | Err(error) => load_error = Some(error),
216 }
217 }
218
219 if loaded_patterns.is_empty() {
220 if let Some(error) = load_error {
221 return Err(error);
222 }
223 return Ok(());
224 }
225
226 self.merge_loaded_patterns(loaded_patterns)?;
227
228 if canonical_missing {
232 self.persist_patterns_if_absent()
233 } else {
234 Ok(())
235 }
236 }
237
238 fn merge_loaded_patterns<I>(&self, loaded_patterns: I) -> Result<()>
239 where
240 I: IntoIterator<Item = HashMap<String, ApprovalPattern>>,
241 {
242 let mut patterns = self
243 .patterns
244 .lock()
245 .map_err(|e| anyhow::anyhow!("Failed to lock patterns: {e}"))?;
246
247 for loaded in loaded_patterns {
248 merge_pattern_map(&mut patterns, loaded);
249 }
250 drop(patterns);
251
252 Ok(())
253 }
254
255 pub fn refresh_patterns(&self) -> Result<()> {
257 self.load_patterns()
258 }
259
260 pub fn get_pattern(&self, approval_key: &str) -> Option<ApprovalPattern> {
262 if let Ok(patterns) = self.patterns.lock() {
263 patterns.get(approval_key).cloned()
264 } else {
265 None
266 }
267 }
268
269 pub fn record_decision(
271 &self,
272 approval_key: &str,
273 display_name: Option<&str>,
274 approved: bool,
275 reason: Option<String>,
276 ) {
277 let should_persist = if let Ok(mut patterns) = self.patterns.lock() {
278 let pattern = patterns.entry(approval_key.to_owned()).or_insert_with(|| ApprovalPattern {
279 tool_name: approval_key.to_owned(),
280 display_name: display_name.map(str::to_owned),
281 approve_count: 0,
282 deny_count: 0,
283 last_decision: None,
284 recent_reason: None,
285 });
286
287 if let Some(display_name) = display_name {
288 pattern.display_name = Some(display_name.to_owned());
289 }
290
291 if approved {
292 pattern.approve_count += 1;
293 } else {
294 pattern.deny_count += 1;
295 }
296
297 pattern.last_decision = Some(approved);
298 pattern.recent_reason = reason;
299 true
300 } else {
301 false
302 };
303
304 if should_persist {
306 let _ = self.persist_patterns();
307 }
308 }
309
310 fn persist_patterns(&self) -> Result<()> {
315 let (patterns_file, mut patterns_snapshot) = self.current_patterns_snapshot()?;
316 VtCodePaths::with_private_file_lock(&patterns_file, || {
317 if let PatternFileState::Valid(disk_patterns) = read_patterns_file(&patterns_file)? {
318 merge_pattern_map(&mut patterns_snapshot, disk_patterns);
319 }
320 let content = serde_json::to_vec_pretty(&patterns_snapshot)?;
321 VtCodePaths::write_private_file_atomic(&patterns_file, &content)
322 .context("failed to write approval patterns cache")?;
323 Ok(())
324 })
325 }
326
327 fn persist_patterns_if_absent(&self) -> Result<()> {
328 let (patterns_file, content) = self.serialized_patterns()?;
329 VtCodePaths::with_private_file_lock(&patterns_file, || {
330 VtCodePaths::write_private_file_atomic_if_absent(&patterns_file, &content).map(|_| ())
331 })
332 .context("failed to publish approval patterns cache")
333 }
334
335 fn current_patterns_snapshot(&self) -> Result<(PathBuf, HashMap<String, ApprovalPattern>)> {
336 VtCodePaths::ensure_user_dir(&self.cache_dir)?;
337 let patterns_snapshot = {
338 let patterns = self
339 .patterns
340 .lock()
341 .map_err(|e| anyhow::anyhow!("Failed to lock patterns: {e}"))?;
342 patterns.clone()
343 };
344 Ok((self.cache_dir.join("approval_patterns.json"), patterns_snapshot))
345 }
346
347 fn serialized_patterns(&self) -> Result<(PathBuf, Vec<u8>)> {
348 let (patterns_file, patterns_snapshot) = self.current_patterns_snapshot()?;
349 let content = serde_json::to_vec_pretty(&patterns_snapshot)?;
350 Ok((patterns_file, content))
351 }
352
353 pub fn get_learning_summary(&self, approval_key: &str) -> Option<String> {
355 let pattern = self.get_pattern(approval_key)?;
356
357 if pattern.approval_count() == 0 {
358 return None;
359 }
360
361 Some(format!(
362 "Approved {} of {} times ({:.0}%)",
363 pattern.approve_count,
364 pattern.approve_count + pattern.deny_count,
365 pattern.approval_rate() * 100.0
366 ))
367 }
368}
369
370enum PatternFileState {
371 Missing,
372 Valid(HashMap<String, ApprovalPattern>),
373 Malformed(anyhow::Error),
374}
375
376fn read_patterns_file(path: &Path) -> Result<PatternFileState> {
377 let metadata = match std::fs::symlink_metadata(path) {
378 Ok(metadata) => metadata,
379 Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(PatternFileState::Missing),
380 Err(error) => {
381 return Err(error).with_context(|| format!("failed to inspect approval patterns cache {}", path.display()));
382 }
383 };
384 if !metadata.is_file() || metadata.file_type().is_symlink() {
385 anyhow::bail!("approval patterns cache is not a regular file: {}", path.display());
386 }
387
388 let content = match String::from_utf8(VtCodePaths::read_file_no_follow(path)?) {
389 Ok(content) => content,
390 Err(error) => {
391 return Ok(PatternFileState::Malformed(anyhow::anyhow!(
392 "failed to read approval patterns cache {} as UTF-8: {error}",
393 path.display()
394 )));
395 }
396 };
397 match serde_json::from_str::<HashMap<String, ApprovalPattern>>(&content) {
398 Ok(patterns) => Ok(PatternFileState::Valid(patterns)),
399 Err(error) => Ok(PatternFileState::Malformed(anyhow::anyhow!(
400 "failed to parse approval patterns cache {}: {error}",
401 path.display()
402 ))),
403 }
404}
405
406#[cfg(test)]
407mod tests {
408 use super::*;
409
410 #[test]
411 fn test_tool_justification_creation() {
412 let just = ToolJustification::new("read_file", "Need to understand code structure", &RiskLevel::Low)
413 .with_outcome("Will analyze the AST to provide better context");
414
415 assert_eq!(just.tool_name, "read_file");
416 assert!(just.reason.contains("understand"));
417 assert!(just.expected_outcome.is_some());
418 }
419
420 #[test]
421 fn test_justification_formatting() {
422 let just =
423 ToolJustification::new("run_command", "Execute build to check for compilation errors", &RiskLevel::High)
424 .with_outcome("Will produce build output for analysis");
425
426 let formatted = just.format_for_dialog();
427 assert!(formatted.iter().any(|l| l.contains("Agent Reasoning")));
428 assert!(formatted.iter().any(|l| l.contains("Expected Outcome")));
429 assert!(formatted.iter().any(|l| l.contains("Risk Level")));
430 }
431
432 #[test]
433 fn test_approval_pattern_calculation() {
434 let mut pattern = ApprovalPattern {
435 tool_name: "read_file".to_owned(),
436 display_name: None,
437 approve_count: 9,
438 deny_count: 1,
439 last_decision: Some(true),
440 recent_reason: None,
441 };
442
443 assert!((pattern.approval_rate() - 0.9).abs() < f32::EPSILON);
444 assert!(pattern.has_high_approval_rate());
445
446 pattern.approve_count = 3;
447 pattern.deny_count = 7;
448 assert!(!pattern.has_high_approval_rate()); }
450
451 #[test]
452 fn test_justification_manager_basic() {
453 let temp_dir = std::env::temp_dir().join(format!("vtcode_test_{}", std::process::id()));
454 let manager = JustificationManager::new(temp_dir.clone());
455
456 manager.record_decision("read_file", Some("Read File"), true, None);
457 manager.record_decision("read_file", Some("Read File"), true, None);
458 manager.record_decision("read_file", Some("Read File"), false, None);
459
460 let pattern = manager.get_pattern("read_file").unwrap();
461 assert_eq!(pattern.approve_count, 2);
462 assert_eq!(pattern.deny_count, 1);
463 assert!((pattern.approval_rate() - 2.0 / 3.0).abs() < f32::EPSILON);
464 assert_eq!(pattern.display_name.as_deref(), Some("Read File"));
465
466 let _ = std::fs::remove_dir_all(&temp_dir);
468 }
469
470 #[test]
471 fn test_justification_manager_preserves_new_display_name() {
472 let temp_dir = std::env::temp_dir().join(format!("vtcode_test_{}", std::process::id()));
473 let manager = JustificationManager::new(temp_dir.clone());
474
475 manager.record_decision("shell:key", Some("command `cargo test`"), true, None);
476 manager.record_decision("shell:key", Some("commands starting with `cargo`"), true, None);
477
478 let pattern = manager.get_pattern("shell:key").unwrap();
479 assert_eq!(pattern.display_name.as_deref(), Some("commands starting with `cargo`"));
480
481 let _ = std::fs::remove_dir_all(&temp_dir);
483 }
484}