vtcode_core/tools/registry/
risk_scorer.rs1use crate::config::constants::tools;
4use crate::utils::ansi_codes::{FG_GREEN, FG_MAGENTA, FG_RED, FG_YELLOW};
5use serde::{Deserialize, Serialize};
6
7#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
9#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
10pub enum RiskLevel {
11 Low,
13
14 Medium,
16
17 High,
19
20 Critical,
22}
23
24impl RiskLevel {
25 pub fn as_str(self) -> &'static str {
26 match self {
27 Self::Low => "low",
28 Self::Medium => "medium",
29 Self::High => "high",
30 Self::Critical => "critical",
31 }
32 }
33
34 pub fn color_code(self) -> &'static str {
35 match self {
36 Self::Low => FG_GREEN,
37 Self::Medium => FG_YELLOW,
38 Self::High => FG_RED,
39 Self::Critical => FG_MAGENTA,
40 }
41 }
42}
43
44impl std::fmt::Display for RiskLevel {
45 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
46 f.write_str(self.as_str())
47 }
48}
49
50#[derive(Debug, Clone, Copy, PartialEq, Eq)]
52pub enum ToolSource {
53 Internal,
55
56 Mcp,
58
59 Acp,
61
62 External,
64}
65
66impl ToolSource {
67 pub fn risk_multiplier(self) -> f32 {
70 match self {
71 Self::Internal => 1.0,
72 Self::Mcp => 1.5,
73 Self::Acp => 1.2,
74 Self::External => 2.0,
75 }
76 }
77}
78
79#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
81pub enum WorkspaceTrust {
82 Untrusted,
83 Partial,
84 Trusted,
85 FullAuto,
86}
87
88impl WorkspaceTrust {
89 pub fn risk_reduction(self) -> f32 {
91 match self {
92 Self::Untrusted => 1.0,
93 Self::Partial => 0.8,
94 Self::Trusted => 0.6,
95 Self::FullAuto => 0.3,
96 }
97 }
98}
99
100#[derive(Debug, Clone)]
102pub struct ToolRiskContext {
103 pub tool_name: String,
105
106 pub source: ToolSource,
108
109 pub workspace_trust: WorkspaceTrust,
111
112 pub recent_approvals: usize,
114
115 pub command_args: Vec<String>,
117
118 pub is_write: bool,
120
121 pub is_destructive: bool,
123
124 pub accesses_network: bool,
126}
127
128impl ToolRiskContext {
129 pub fn new(tool_name: String, source: ToolSource, workspace_trust: WorkspaceTrust) -> Self {
131 Self {
132 tool_name,
133 source,
134 workspace_trust,
135 recent_approvals: 0,
136 command_args: Vec::new(),
137 is_write: false,
138 is_destructive: false,
139 accesses_network: false,
140 }
141 }
142
143 pub fn with_args(mut self, args: Vec<String>) -> Self {
145 self.command_args = args;
146 self
147 }
148
149 pub fn as_write(mut self) -> Self {
151 self.is_write = true;
152 self
153 }
154
155 pub fn as_destructive(mut self) -> Self {
157 self.is_destructive = true;
158 self
159 }
160
161 pub fn accesses_network(mut self) -> Self {
163 self.accesses_network = true;
164 self
165 }
166}
167
168pub struct ToolRiskScorer;
170
171impl ToolRiskScorer {
172 pub fn calculate_risk(ctx: &ToolRiskContext) -> RiskLevel {
174 let mut base_score = Self::base_risk_for_tool(&ctx.tool_name);
175
176 if ctx.is_destructive {
178 base_score += 30;
179 }
180 if ctx.is_write {
181 base_score += 15;
182 }
183 if ctx.accesses_network {
184 base_score += 10;
185 }
186
187 #[allow(
189 clippy::cast_sign_loss,
190 clippy::let_and_return,
191 reason = "Intentional compatibility, platform, or test-only suppression."
192 )]
193 let adjusted = ((base_score as f32 * ctx.source.risk_multiplier()).max(0.0)) as u32;
194 base_score = adjusted;
195
196 #[allow(
198 clippy::cast_sign_loss,
199 clippy::let_and_return,
200 reason = "Intentional compatibility, platform, or test-only suppression."
201 )]
202 let adjusted = ((base_score as f32 * ctx.workspace_trust.risk_reduction()).max(0.0)) as u32;
203 base_score = adjusted;
204
205 let approval_reduction = ctx.recent_approvals.min(3) as u32 * 5;
207 base_score = base_score.saturating_sub(approval_reduction);
208
209 match base_score {
211 0..=25 => RiskLevel::Low,
212 26..=50 => RiskLevel::Medium,
213 51..=75 => RiskLevel::High,
214 _ => RiskLevel::Critical,
215 }
216 }
217
218 pub fn is_network_tool(tool_name: &str) -> bool {
223 matches!(
224 tool_name,
225 tools::WEB_SEARCH
226 | tools::WEB_FETCH
227 | tools::FETCH_URL
228 | tools::DEFUDDLE_FETCH
229 | tools::MCP_CONNECT_SERVER
230 | tools::MCP_DISCONNECT_SERVER
231 ) || matches!(tool_name, "mcp:connect" | "mcp:disconnect")
232 }
233
234 pub fn requires_justification(risk: RiskLevel, threshold: RiskLevel) -> bool {
236 risk >= threshold
237 }
238
239 fn base_risk_for_tool(tool_name: &str) -> u32 {
241 match tool_name {
242 tools::READ_FILE
244 | tools::CODE_SEARCH
245 | tools::MCP_SEARCH_TOOLS
246 | tools::MCP_GET_TOOL_DETAILS
247 | tools::MCP_LIST_SERVERS => 0,
248
249 "file_info" | "status" | "logs" => 5,
251
252 tools::WRITE_FILE | tools::EDIT_FILE | tools::CREATE_FILE => 20,
254
255 tools::APPLY_PATCH | tools::DELETE_FILE => 25,
257
258 tools::LOAD_SKILL => 60,
262
263 tools::CREATE_PTY_SESSION | tools::RUN_PTY_CMD | tools::SEND_PTY_INPUT | tools::UNIFIED_EXEC => 35,
265
266 tools::WEB_SEARCH
268 | tools::WEB_FETCH
269 | tools::FETCH_URL
270 | tools::DEFUDDLE_FETCH
271 | tools::MCP_CONNECT_SERVER
272 | tools::MCP_DISCONNECT_SERVER
273 | "mcp:connect"
274 | "mcp:disconnect" => 40,
275
276 _ if tool_name.starts_with("mcp_") => 30,
278
279 _ => 35,
281 }
282 }
283}
284
285#[cfg(test)]
286mod tests {
287 use super::*;
288
289 #[test]
290 fn test_risk_level_ordering() {
291 assert!(RiskLevel::Low < RiskLevel::Medium);
292 assert!(RiskLevel::Medium < RiskLevel::High);
293 assert!(RiskLevel::High < RiskLevel::Critical);
294 }
295
296 #[test]
297 fn test_risk_calculation() {
298 let ctx = ToolRiskContext::new(tools::READ_FILE.to_string(), ToolSource::Internal, WorkspaceTrust::Trusted);
300 let risk = ToolRiskScorer::calculate_risk(&ctx);
301 assert_eq!(risk, RiskLevel::Low);
302
303 let ctx = ToolRiskContext::new(tools::WRITE_FILE.to_string(), ToolSource::External, WorkspaceTrust::Untrusted)
305 .as_write();
306 let risk = ToolRiskScorer::calculate_risk(&ctx);
307 assert!(risk >= RiskLevel::High);
308 }
309
310 #[test]
311 fn test_network_tools_are_classified() {
312 assert!(ToolRiskScorer::is_network_tool(tools::WEB_FETCH));
313 assert!(ToolRiskScorer::is_network_tool(tools::WEB_SEARCH));
314 assert!(ToolRiskScorer::is_network_tool(tools::FETCH_URL));
315 assert!(ToolRiskScorer::is_network_tool(tools::DEFUDDLE_FETCH));
316 assert!(ToolRiskScorer::is_network_tool("mcp:connect"));
317 assert!(ToolRiskScorer::is_network_tool("mcp:disconnect"));
318 assert!(!ToolRiskScorer::is_network_tool(tools::CODE_SEARCH));
319 assert!(!ToolRiskScorer::is_network_tool(tools::READ_FILE));
320 }
321
322 #[test]
323 fn test_network_fetch_not_low_risk_even_when_trusted() {
324 for tool in [tools::WEB_FETCH, "mcp:connect", "mcp:disconnect"] {
327 let ctx = ToolRiskContext::new(tool.to_string(), ToolSource::Internal, WorkspaceTrust::Trusted)
328 .accesses_network();
329 let risk = ToolRiskScorer::calculate_risk(&ctx);
330 assert!(
331 risk > RiskLevel::Low,
332 "network tool '{tool}' should not be auto-approved as low risk, got {risk:?}"
333 );
334 }
335 }
336
337 #[test]
338 fn test_skill_loading_is_not_low_risk_even_when_trusted() {
339 let ctx = ToolRiskContext::new(tools::LOAD_SKILL.to_string(), ToolSource::Internal, WorkspaceTrust::Trusted);
340 let risk = ToolRiskScorer::calculate_risk(&ctx);
341
342 assert!(risk > RiskLevel::Low, "skill loading should require approval, got {risk:?}");
343 }
344
345 #[test]
346 fn test_approval_history_reduces_risk() {
347 let mut ctx =
348 ToolRiskContext::new(tools::RUN_PTY_CMD.to_string(), ToolSource::Internal, WorkspaceTrust::Untrusted);
349
350 let risk_before = ToolRiskScorer::calculate_risk(&ctx);
351
352 ctx.recent_approvals = 3;
353 let risk_after = ToolRiskScorer::calculate_risk(&ctx);
354
355 assert!(risk_after <= risk_before);
356 }
357
358 #[test]
359 fn test_source_multiplier() {
360 let base = ToolRiskContext::new("mcp_tool".to_string(), ToolSource::Internal, WorkspaceTrust::Trusted);
361 let base_risk = ToolRiskScorer::calculate_risk(&base);
362
363 let mcp = ToolRiskContext::new("mcp_tool".to_string(), ToolSource::Mcp, WorkspaceTrust::Trusted);
364 let mcp_risk = ToolRiskScorer::calculate_risk(&mcp);
365
366 assert!(mcp_risk > base_risk || mcp_risk == RiskLevel::Critical);
368 }
369
370 #[test]
371 fn test_requires_justification() {
372 assert!(ToolRiskScorer::requires_justification(RiskLevel::High, RiskLevel::High));
373 assert!(!ToolRiskScorer::requires_justification(RiskLevel::Medium, RiskLevel::High));
374 }
375}