1use crate::config::constants::tools;
14use hashbrown::HashMap;
15use std::path::PathBuf;
16use std::sync::Arc;
17
18use anyhow::Result;
19use async_trait::async_trait;
20use serde::{Deserialize, Serialize};
21use serde_json::Value;
22use vtcode_macros::DebugNoInline;
23pub use vtcode_utility_tool_specs::{
24 AdditionalProperties, FreeformTool, FreeformToolFormat, JsonSchema, ResponsesApiTool,
25};
26
27#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
29pub enum ToolKind {
30 Function,
32 Mcp,
34 Custom,
36}
37
38#[derive(Clone, Debug)]
40pub enum ToolPayload {
41 Function { arguments: String },
43 Custom { input: String },
45 Mcp { arguments: Option<Value> },
47 LocalShell { params: ShellToolCallParams },
49}
50
51#[derive(Clone, Debug, Deserialize, Serialize)]
53pub struct ShellToolCallParams {
54 pub command: Vec<String>,
55 pub workdir: Option<String>,
56 pub timeout_ms: Option<u64>,
57 pub sandbox_permissions: Option<SandboxPermissions>,
58 pub justification: Option<String>,
59}
60
61use crate::sandboxing::LinuxSandboxLauncher;
63pub use crate::sandboxing::SandboxPermissions;
64
65#[derive(Clone, Debug)]
67pub enum ToolOutput {
68 Function {
70 content: String,
71 content_items: Option<Vec<ContentItem>>,
72 success: Option<bool>,
73 },
74 Mcp { result: McpToolResult },
76}
77
78impl ToolOutput {
79 pub fn simple(content: impl Into<String>) -> Self {
81 Self::Function {
82 content: content.into(),
83 content_items: None,
84 success: Some(true),
85 }
86 }
87
88 pub fn with_success(content: impl Into<String>, success: bool) -> Self {
90 Self::Function {
91 content: content.into(),
92 content_items: None,
93 success: Some(success),
94 }
95 }
96
97 pub fn error(message: impl Into<String>) -> Self {
99 Self::Function {
100 content: message.into(),
101 content_items: None,
102 success: Some(false),
103 }
104 }
105
106 pub fn content(&self) -> Option<&str> {
108 match self {
109 Self::Function { content, .. } => Some(content),
110 Self::Mcp { result } => result.content.first().and_then(|c| c.as_text()),
111 }
112 }
113
114 pub fn is_success(&self) -> bool {
116 match self {
117 Self::Function { success, .. } => success.unwrap_or(true),
118 Self::Mcp { result } => !result.is_error.unwrap_or(false),
119 }
120 }
121}
122
123#[derive(Clone, Debug, Serialize, Deserialize)]
125#[serde(tag = "type", rename_all = "snake_case")]
126pub enum ContentItem {
127 Text { text: String },
128 Image { data: String, mime_type: String },
129 Resource { uri: String, mime_type: Option<String> },
130}
131
132impl ContentItem {
133 pub fn as_text(&self) -> Option<&str> {
134 match self {
135 ContentItem::Text { text } => Some(text),
136 _ => None,
137 }
138 }
139}
140
141#[derive(Clone, Debug, Serialize, Deserialize)]
143pub struct McpToolResult {
144 pub content: Vec<ContentItem>,
145 pub is_error: Option<bool>,
146}
147
148pub struct ToolInvocation {
150 pub session: Arc<dyn ToolSession>,
151 pub turn: Arc<TurnContext>,
152 pub tracker: Option<SharedDiffTracker>,
153 pub call_id: String,
154 pub tool_name: String,
155 pub payload: ToolPayload,
156}
157
158pub type SharedDiffTracker = Arc<tokio::sync::Mutex<DiffTracker>>;
160
161#[derive(Clone, Debug, PartialEq, Eq)]
163pub struct Constrained<T> {
164 value: T,
165}
166
167impl<T> Constrained<T> {
168 pub fn allow_any(initial_value: T) -> Self {
169 Self { value: initial_value }
170 }
171
172 pub fn get(&self) -> &T {
173 &self.value
174 }
175}
176
177impl<T> std::ops::Deref for Constrained<T> {
182 type Target = T;
183
184 fn deref(&self) -> &Self::Target {
185 &self.value
186 }
187}
188
189impl<T: Copy> Constrained<T> {
190 pub fn value(&self) -> T {
191 self.value
192 }
193}
194
195impl<T: Default> Default for Constrained<T> {
196 fn default() -> Self {
197 Self::allow_any(T::default())
198 }
199}
200
201#[async_trait]
203pub trait ToolSession: Send + Sync {
204 fn cwd(&self) -> &PathBuf;
206
207 fn workspace_root(&self) -> &PathBuf;
209
210 async fn record_warning(&self, message: String);
212
213 fn user_shell(&self) -> &str;
215}
216
217#[derive(Clone, Debug)]
219pub struct TurnContext {
220 pub cwd: PathBuf,
221 pub turn_id: String,
222 pub sub_id: Option<String>,
223 pub shell_environment_policy: ShellEnvironmentPolicy,
224 pub approval_policy: Constrained<ApprovalPolicy>,
225 pub linux_sandbox_launcher: Option<LinuxSandboxLauncher>,
226 pub sandbox_policy: Constrained<super::sandboxing::SandboxConfig>,
228}
229
230impl TurnContext {
231 pub fn resolve_path(&self, path: Option<String>) -> PathBuf {
233 self.resolve_path_ref(path.as_deref())
234 }
235
236 pub fn resolve_path_ref(&self, path: Option<&str>) -> PathBuf {
238 match path {
239 Some(p) => {
240 let path = PathBuf::from(p);
241 if path.is_absolute() { path } else { self.cwd.join(path) }
242 }
243 None => self.cwd.clone(),
244 }
245 }
246}
247
248#[derive(Clone, Debug, Default)]
250pub enum ShellEnvironmentPolicy {
251 #[default]
252 Inherit,
253 Clean,
254 Custom(HashMap<String, String>),
255}
256
257#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
259pub enum ApprovalPolicy {
260 #[default]
261 Never,
262 OnMutation,
263 Always,
264}
265
266#[derive(Default, Debug)]
268pub struct DiffTracker {
269 pub changes: HashMap<PathBuf, FileChange>,
270}
271
272impl DiffTracker {
273 pub fn on_patch_begin(&mut self, changes: &HashMap<PathBuf, FileChange>) {
274 self.changes.extend(changes.clone());
275 }
276
277 pub fn on_patch_end(&mut self, success: bool) {
278 if !success {
279 self.changes.clear();
280 }
281 }
282}
283
284#[derive(Clone, Debug, Serialize, Deserialize)]
286#[serde(tag = "type", rename_all = "snake_case")]
287pub enum FileChange {
288 Add { content: String },
289 Delete,
290 Update { old_content: String, new_content: String },
291 Rename { new_path: PathBuf, content: Option<String> },
292}
293
294#[derive(DebugNoInline, thiserror::Error)]
299pub enum ToolCallError {
300 #[error("Tool error: {0}")]
302 RespondToModel(String),
303
304 #[error("Internal error: {0}")]
306 Internal(#[from] anyhow::Error),
307
308 #[error("Tool rejected: {0}")]
310 Rejected(String),
311
312 #[error("Tool timed out after {0}ms")]
314 Timeout(u64),
315}
316
317impl ToolCallError {
318 pub fn respond(message: impl Into<String>) -> Self {
320 Self::RespondToModel(message.into())
321 }
322}
323
324impl From<super::sandboxing::ToolError> for ToolCallError {
325 fn from(err: super::sandboxing::ToolError) -> Self {
326 match err {
327 super::sandboxing::ToolError::Rejected(msg) => ToolCallError::Rejected(msg),
328 super::sandboxing::ToolError::Codex(e) => ToolCallError::Internal(e),
329 super::sandboxing::ToolError::SandboxDenied(msg) => {
330 ToolCallError::Rejected(format!("Sandbox denied: {msg}"))
331 }
332 super::sandboxing::ToolError::Timeout(ms) => ToolCallError::Timeout(ms),
333 }
334 }
335}
336
337#[async_trait]
342pub trait ToolHandler: Send + Sync {
343 fn kind(&self) -> ToolKind;
345
346 fn matches_kind(&self, payload: &ToolPayload) -> bool {
348 matches!(
349 (self.kind(), payload),
350 (ToolKind::Function, ToolPayload::Function { .. })
351 | (ToolKind::Mcp, ToolPayload::Mcp { .. })
352 | (ToolKind::Custom, ToolPayload::Custom { .. })
353 )
354 }
355
356 async fn is_mutating(&self, _invocation: &ToolInvocation) -> bool {
360 false
361 }
362
363 async fn handle(&self, invocation: ToolInvocation) -> Result<ToolOutput, ToolCallError>;
365}
366
367#[derive(Clone, Debug, Serialize, Deserialize)]
369#[serde(tag = "type", rename_all = "snake_case")]
370pub enum ToolSpec {
371 Function(ResponsesApiTool),
372 Freeform(FreeformTool),
373 WebSearch {},
374 LocalShell {},
375}
376
377impl ToolSpec {
378 pub fn name(&self) -> &str {
379 match self {
380 ToolSpec::Function(tool) => &tool.name,
381 ToolSpec::Freeform(tool) => &tool.name,
382 ToolSpec::WebSearch {} => tools::WEB_SEARCH,
383 ToolSpec::LocalShell {} => "local_shell",
384 }
385 }
386}
387
388#[derive(Clone, Debug)]
390pub struct ConfiguredToolSpec {
391 pub spec: ToolSpec,
392 pub supports_parallel_tool_calls: bool,
393}
394
395impl ConfiguredToolSpec {
396 pub fn new(spec: ToolSpec, supports_parallel: bool) -> Self {
397 Self {
398 spec,
399 supports_parallel_tool_calls: supports_parallel,
400 }
401 }
402}
403
404#[cfg(test)]
405mod tests {
406 use super::*;
407
408 #[test]
409 fn test_tool_output_simple() {
410 let output = ToolOutput::simple("Hello, world!");
411 assert!(output.is_success());
412 assert_eq!(output.content(), Some("Hello, world!"));
413 }
414
415 #[test]
416 fn test_tool_output_error() {
417 let output = ToolOutput::error("Something went wrong");
418 assert!(!output.is_success());
419 assert_eq!(output.content(), Some("Something went wrong"));
420 }
421
422 #[test]
423 fn test_sandbox_permissions_default() {
424 let perms = SandboxPermissions::default();
425 assert_eq!(perms, SandboxPermissions::UseDefault);
426 }
427
428 #[test]
429 fn test_turn_context_resolve_path_absolute() {
430 let ctx = TurnContext {
431 cwd: PathBuf::from("/workspace"),
432 turn_id: "test".to_string(),
433 sub_id: None,
434 shell_environment_policy: ShellEnvironmentPolicy::default(),
435 approval_policy: Constrained::allow_any(ApprovalPolicy::default()),
436 linux_sandbox_launcher: None,
437 sandbox_policy: Constrained::allow_any(Default::default()),
438 };
439
440 let resolved = ctx.resolve_path(Some("/absolute/path".to_string()));
441 assert_eq!(resolved, PathBuf::from("/absolute/path"));
442 }
443
444 #[test]
445 fn test_turn_context_resolve_path_relative() {
446 let ctx = TurnContext {
447 cwd: PathBuf::from("/workspace"),
448 turn_id: "test".to_string(),
449 sub_id: None,
450 shell_environment_policy: ShellEnvironmentPolicy::default(),
451 approval_policy: Constrained::allow_any(ApprovalPolicy::default()),
452 linux_sandbox_launcher: None,
453 sandbox_policy: Constrained::allow_any(Default::default()),
454 };
455
456 let resolved = ctx.resolve_path(Some("relative/path".to_string()));
457 assert_eq!(resolved, PathBuf::from("/workspace/relative/path"));
458 }
459}