Skip to main content

vtcode_core/tools/registry/
justification.rs

1/// Tool Justification System
2///
3/// Captures agent reasoning before high-risk tool execution to improve approval UX
4/// and enable learning of approval patterns.
5use 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/// Justification provided by the agent for executing a high-risk tool
13#[derive(Debug, Clone, Serialize, Deserialize)]
14pub struct ToolJustification {
15    /// Tool being justified
16    pub tool_name: String,
17    /// Brief explanation from the agent
18    pub reason: String,
19    /// Expected outcome of tool execution
20    pub expected_outcome: Option<String>,
21    /// Risk level that triggered justification
22    pub risk_level: String,
23    /// Timestamp when justification was provided
24    pub timestamp: String,
25}
26
27impl ToolJustification {
28    /// Create a new tool justification
29    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    /// Add expected outcome to justification
40    pub fn with_outcome(mut self, outcome: impl Into<String>) -> Self {
41        self.expected_outcome = Some(outcome.into());
42        self
43    }
44
45    /// Format justification for display in approval dialog
46    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        // Wrap reason text if needed - iterate directly without collecting
53        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/// Tracks approval patterns to learn from user decisions
77#[derive(Debug, Clone, Serialize, Deserialize, Default)]
78pub struct ApprovalPattern {
79    /// Stable approval key used for lookup and persistence
80    pub tool_name: String,
81    /// Human-readable label for prompts and summaries
82    #[serde(default)]
83    pub display_name: Option<String>,
84    /// Number of times user approved
85    pub approve_count: u32,
86    /// Number of times user denied
87    pub deny_count: u32,
88    /// Last decision (true = approve, false = deny)
89    pub last_decision: Option<bool>,
90    /// Most recent reason (if available)
91    pub recent_reason: Option<String>,
92}
93
94impl ApprovalPattern {
95    /// Compute approval rate (0.0 to 1.0)
96    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    /// Check if this tool has high approval rate (>80%)
106    pub fn has_high_approval_rate(&self) -> bool {
107        self.approval_count() >= 3 && self.approval_rate() > 0.8
108    }
109
110    /// Return approval count
111    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
120/// Merge an on-disk pattern into the in-memory entry by taking the max of
121/// counters and preferring any non-`None` metadata from disk. Conservative:
122/// undercount is safer (more prompts) than overcount (fewer prompts → risk).
123fn 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
146/// Manager for approval pattern learning and justifications
147pub 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    /// Create a new justification manager
155    pub fn new(cache_dir: PathBuf) -> Self {
156        Self::new_with_legacy_pattern_files(cache_dir, Vec::new())
157    }
158
159    /// Create a manager that can recover approval patterns from older cache locations.
160    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        // Try to load existing patterns
166        let _ = manager.load_patterns();
167
168        manager
169    }
170
171    /// Create a manager without touching disk.
172    ///
173    /// First-paint path uses this so approval-pattern file reads move to
174    /// hydration; call [`Self::refresh_patterns`] there before the first turn.
175    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    /// Load approval patterns from disk and merge into the in-memory map.
191    ///
192    /// Merging (rather than replacing) keeps in-memory increments that have not
193    /// yet been flushed to disk — important when refreshing right before an
194    /// auto-approval check while a concurrent vtcode session may have written
195    /// newer counts to the same file.
196    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        // The old file remains in place as a rollback source, but make the
229        // recovered data available at the canonical path immediately. This
230        // also prevents every fresh session from re-reading the legacy file.
231        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    /// Re-read patterns from disk, merging with any in-memory state.
256    pub fn refresh_patterns(&self) -> Result<()> {
257        self.load_patterns()
258    }
259
260    /// Get approval pattern for a key
261    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    /// Record user approval decision
270    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        // Persist to disk after releasing the lock.
305        if should_persist {
306            let _ = self.persist_patterns();
307        }
308    }
309
310    /// Persist patterns to disk
311    ///
312    /// Clone the patterns under the mutex, then merge and write under the
313    /// process-shared file lock outside the mutex.
314    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    /// Get learning summary for a key
354    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()); // < 0.8 rate
449    }
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        // Cleanup
467        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        // Cleanup
482        let _ = std::fs::remove_dir_all(&temp_dir);
483    }
484}