diff --git a/crates/goose-acp/src/fs.rs b/crates/goose-acp/src/fs.rs index aff7173b..a751aa34 100644 --- a/crates/goose-acp/src/fs.rs +++ b/crates/goose-acp/src/fs.rs @@ -1,16 +1,26 @@ +use crate::tools::AcpAwareToolMeta; +use agent_client_protocol_schema::TerminalId; use async_trait::async_trait; use fs_err as fs; use goose::agents::mcp_client::{Error as McpError, McpClientTrait}; use goose::agents::platform_extensions::developer::edit::{ resolve_path, string_replace, FileEditParams, FileReadParams, FileWriteParams, }; +use goose::agents::platform_extensions::developer::shell::{ShellParams, OUTPUT_LIMIT_BYTES}; use goose::agents::platform_extensions::developer::DeveloperClient; use rmcp::model::{CallToolResult, Content as RmcpContent, Tool, ToolAnnotations}; -use sacp::schema::{ReadTextFileRequest, SessionId, WriteTextFileRequest}; +use sacp::schema::{ + CreateTerminalRequest, Diff, KillTerminalCommandRequest, ReadTextFileRequest, + ReleaseTerminalRequest, SessionId, SessionNotification, SessionUpdate, Terminal, + TerminalOutputRequest, ToolCallContent, ToolCallId, ToolCallLocation, ToolCallUpdate, + ToolCallUpdateFields, ToolKind, WaitForTerminalExitRequest, WriteTextFileRequest, +}; use sacp::{AgentToClient, JrConnectionCx}; use schemars::schema_for; use std::path::Path; use std::sync::Arc; +use std::time::Duration; +use tokio::time::timeout; use tokio_util::sync::CancellationToken; async fn acp_read_text_file( @@ -56,6 +66,7 @@ pub(crate) struct AcpTools { pub(crate) session_id: SessionId, pub(crate) fs_read: bool, pub(crate) fs_write: bool, + pub(crate) terminal: bool, } fn error_result(msg: impl std::fmt::Display) -> CallToolResult { @@ -81,19 +92,22 @@ fn read_tool() -> Tool { ) } -pub(crate) fn with_location_meta( - mut result: CallToolResult, - path: &Path, - line: Option, -) -> CallToolResult { - let location = serde_json::json!({ - "tool_locations": [{"path": path.to_string_lossy(), "line": line}] - }); - result.meta = Some(serde_json::from_value(location).unwrap()); - result -} - impl AcpTools { + fn update_tool_call(&self, ctx: &goose::agents::ToolCallContext, fields: ToolCallUpdateFields) { + if let Some(ref req_id) = ctx.tool_call_request_id { + let _ = self + .cx + .send_notification(SessionNotification::new( + self.session_id.clone(), + SessionUpdate::ToolCallUpdate(ToolCallUpdate::new( + ToolCallId::new(req_id.clone()), + fields, + )), + )) + .inspect_err(|e| tracing::error!("error updating tool call with client: {}", e)); + } + } + fn parse_args( arguments: Option, ) -> Result { @@ -111,20 +125,24 @@ impl AcpTools { async fn acp_read( &self, arguments: Option, - working_dir: Option<&str>, + ctx: &goose::agents::ToolCallContext, ) -> Result { let params: FileReadParams = match Self::parse_args(arguments) { Ok(p) => p, Err(e) => return Ok(error_result(e)), }; - let path = resolve_path(¶ms.path, working_dir.map(Path::new)); + let path = resolve_path(¶ms.path, ctx.working_dir.as_deref()); + self.update_tool_call( + ctx, + ToolCallUpdateFields::new() + .kind(ToolKind::Read) + .locations(vec![ToolCallLocation::new(&path)]), + ); match acp_read_text_file(&self.cx, &self.session_id, &path, params.line, params.limit).await { - Ok(content) => Ok(with_location_meta( - CallToolResult::success(vec![RmcpContent::text(content).with_priority(0.0)]), - &path, - params.line, - )), + Ok(content) => Ok(CallToolResult::success(vec![ + RmcpContent::text(content).with_priority(0.0) + ])), Err(e) => Ok(fail("read", ¶ms.path, e)), } } @@ -132,26 +150,35 @@ impl AcpTools { async fn acp_write( &self, arguments: Option, - working_dir: Option<&str>, + ctx: &goose::agents::ToolCallContext, ) -> Result { let params: FileWriteParams = match Self::parse_args(arguments) { Ok(p) => p, Err(e) => return Ok(error_result(e)), }; - let path = resolve_path(¶ms.path, working_dir.map(Path::new)); + let path = resolve_path(¶ms.path, ctx.working_dir.as_deref()); + self.update_tool_call( + ctx, + ToolCallUpdateFields::new() + .kind(ToolKind::Edit) + .locations(vec![ToolCallLocation::new(&path)]), + ); match acp_write_text_file(&self.cx, &self.session_id, &path, ¶ms.content).await { Ok(()) => { + self.update_tool_call( + ctx, + ToolCallUpdateFields::new().content(vec![ToolCallContent::Diff(Diff::new( + &path, + ¶ms.content, + ))]), + ); let line_count = params.content.lines().count(); let action = if path.exists() { "Wrote" } else { "Created" }; - Ok(with_location_meta( - CallToolResult::success(vec![RmcpContent::text(format!( - "{action} {} ({line_count} lines)", - params.path - )) - .with_priority(0.0)]), - &path, - Some(1), + Ok(CallToolResult::success(vec![RmcpContent::text(format!( + "{action} {} ({line_count} lines)", + params.path )) + .with_priority(0.0)])) } Err(e) => Ok(fail("write", ¶ms.path, e)), } @@ -160,13 +187,19 @@ impl AcpTools { async fn acp_edit( &self, arguments: Option, - working_dir: Option<&str>, + ctx: &goose::agents::ToolCallContext, ) -> Result { let params: FileEditParams = match Self::parse_args(arguments) { Ok(p) => p, Err(e) => return Ok(error_result(e)), }; - let path = resolve_path(¶ms.path, working_dir.map(Path::new)); + let path = resolve_path(¶ms.path, ctx.working_dir.as_deref()); + self.update_tool_call( + ctx, + ToolCallUpdateFields::new() + .kind(ToolKind::Edit) + .locations(vec![ToolCallLocation::new(&path)]), + ); let content = match self.read_content(&path).await { Ok(c) => c, @@ -186,21 +219,169 @@ impl AcpTools { match write_result { Ok(()) => { + self.update_tool_call( + ctx, + ToolCallUpdateFields::new().content(vec![ToolCallContent::Diff( + Diff::new(&path, &new_content).old_text(&content), + )]), + ); let old_lines = params.before.lines().count(); let new_lines = params.after.lines().count(); - Ok(with_location_meta( - CallToolResult::success(vec![RmcpContent::text(format!( - "Edited {} ({old_lines} lines -> {new_lines} lines)", - params.path - )) - .with_priority(0.0)]), - &path, - Some(1), + Ok(CallToolResult::success(vec![RmcpContent::text(format!( + "Edited {} ({old_lines} lines -> {new_lines} lines)", + params.path )) + .with_priority(0.0)])) } Err(e) => Ok(fail("write", ¶ms.path, e)), } } + + async fn acp_shell( + &self, + arguments: Option, + ctx: &goose::agents::ToolCallContext, + ) -> Result { + let params: ShellParams = match Self::parse_args(arguments) { + Ok(p) => p, + Err(e) => return Ok(error_result(e)), + }; + self.update_tool_call(ctx, ToolCallUpdateFields::new().kind(ToolKind::Execute)); + + let create_res = self + .cx + .send_request( + CreateTerminalRequest::new(self.session_id.clone(), ¶ms.command) + .cwd(ctx.working_dir.clone()) + .output_byte_limit(OUTPUT_LIMIT_BYTES as u64), + ) + .block_task() + .await + .map_err(|e| { + McpError::McpError(rmcp::model::ErrorData::new( + rmcp::model::ErrorCode::INTERNAL_ERROR, + format!("failed to create terminal: {e:?}"), + None, + )) + })?; + let terminal_id = create_res.terminal_id; + + self.update_tool_call( + ctx, + ToolCallUpdateFields::new().content(vec![ToolCallContent::Terminal(Terminal::new( + terminal_id.clone(), + ))]), + ); + + let result = self + .run_terminal_to_completion(&terminal_id, params.timeout_secs) + .await; + + // Always release the terminal, even if we hit errors above. + let _ = self + .cx + .send_request(ReleaseTerminalRequest::new( + self.session_id.clone(), + terminal_id.clone(), + )) + .block_task() + .await + .inspect_err(|e| tracing::error!("failed to release terminal: {e:?}")); + + let output_res = result?; + + let exit_code = output_res + .exit_status + .and_then(|s| s.exit_code) + .unwrap_or_default(); + + let content = vec![ + RmcpContent::text(format!("exit code: {exit_code}")).with_priority(0.0), + RmcpContent::text(output_res.output).with_priority(0.0), + ]; + + if exit_code != 0 { + Ok(CallToolResult::error(content)) + } else { + Ok(CallToolResult::success(content)) + } + } + + async fn run_terminal_to_completion( + &self, + terminal_id: &TerminalId, + timeout_secs: Option, + ) -> Result { + let wait_fut = self + .cx + .send_request(WaitForTerminalExitRequest::new( + self.session_id.clone(), + terminal_id.clone(), + )) + .block_task(); + + let timed_out = match timeout_secs { + Some(secs) if secs > 0 => match timeout(Duration::from_secs(secs), wait_fut).await { + Ok(res) => { + res.map_err(|e| { + McpError::McpError(rmcp::model::ErrorData::new( + rmcp::model::ErrorCode::INTERNAL_ERROR, + format!("failed to wait for terminal exit: {e:?}"), + None, + )) + })?; + false + } + Err(_) => { + let _ = self + .cx + .send_request(KillTerminalCommandRequest::new( + self.session_id.clone(), + terminal_id.clone(), + )) + .block_task() + .await + .inspect_err(|e| tracing::error!("failed to kill terminal: {e:?}")); + true + } + }, + _ => { + wait_fut.await.map_err(|e| { + McpError::McpError(rmcp::model::ErrorData::new( + rmcp::model::ErrorCode::INTERNAL_ERROR, + format!("failed to wait for terminal exit: {e:?}"), + None, + )) + })?; + false + } + }; + + let mut output_res = self + .cx + .send_request(TerminalOutputRequest::new( + self.session_id.clone(), + terminal_id.clone(), + )) + .block_task() + .await + .map_err(|e| { + McpError::McpError(rmcp::model::ErrorData::new( + rmcp::model::ErrorCode::INTERNAL_ERROR, + format!("failed to get terminal output: {e:?}"), + None, + )) + })?; + + if timed_out { + output_res.output.push_str(&format!( + "\n\nCommand timed out after {} seconds", + timeout_secs.unwrap_or(0) + )); + } + + Ok(output_res) + } } #[async_trait] @@ -223,20 +404,31 @@ impl McpClientTrait for AcpTools { async fn call_tool( &self, - session_id: &str, + ctx: &goose::agents::ToolCallContext, name: &str, arguments: Option, - working_dir: Option<&str>, cancellation_token: CancellationToken, ) -> Result { match name { - "read" if self.fs_read => self.acp_read(arguments, working_dir).await, - "write" if self.fs_write => self.acp_write(arguments, working_dir).await, - // edit reads then writes: require both caps so we don't mix editor buffer with local disk - "edit" if self.fs_read && self.fs_write => self.acp_edit(arguments, working_dir).await, + "read" if self.fs_read => self + .acp_read(arguments, ctx) + .await + .map(|r| r.with_acp_aware_meta()), + "write" if self.fs_write => self + .acp_write(arguments, ctx) + .await + .map(|r| r.with_acp_aware_meta()), + "edit" if self.fs_read && self.fs_write => self + .acp_edit(arguments, ctx) + .await + .map(|r| r.with_acp_aware_meta()), + "shell" if self.terminal => self + .acp_shell(arguments, ctx) + .await + .map(|r| r.with_acp_aware_meta()), _ => { self.inner - .call_tool(session_id, name, arguments, working_dir, cancellation_token) + .call_tool(ctx, name, arguments, cancellation_token) .await } } diff --git a/crates/goose-acp/src/lib.rs b/crates/goose-acp/src/lib.rs index 84c17d24..97c405e2 100644 --- a/crates/goose-acp/src/lib.rs +++ b/crates/goose-acp/src/lib.rs @@ -5,4 +5,5 @@ pub mod custom_requests; mod fs; pub mod server; pub mod server_factory; +pub(crate) mod tools; pub mod transport; diff --git a/crates/goose-acp/src/server.rs b/crates/goose-acp/src/server.rs index 260d919b..ef8413aa 100644 --- a/crates/goose-acp/src/server.rs +++ b/crates/goose-acp/src/server.rs @@ -1,5 +1,6 @@ use crate::custom_requests::*; use crate::fs::AcpTools; +use crate::tools::AcpAwareToolMeta; use anyhow::Result; use fs_err as fs; use goose::acp::PermissionDecision; @@ -40,7 +41,7 @@ use sacp::schema::{ use sacp::{AgentToClient, ByteStreams, Handled, JrConnectionCx, JrMessageHandler, MessageCx}; use std::collections::HashMap; use std::sync::Arc; -use tokio::sync::Mutex; +use tokio::sync::{Mutex, OnceCell}; use tokio_util::compat::{TokioAsyncReadCompatExt as _, TokioAsyncWriteCompatExt as _}; use tokio_util::sync::CancellationToken; use tracing::{debug, error, info, warn}; @@ -59,7 +60,8 @@ pub struct GooseAcpAgent { sessions: Arc>>, provider_factory: ProviderConstructor, builtins: Vec, - client_fs_capabilities: Mutex, + client_fs_capabilities: OnceCell, + client_terminal: OnceCell, config_dir: std::path::PathBuf, session_manager: Arc, permission_manager: Arc, @@ -340,7 +342,8 @@ impl GooseAcpAgent { sessions: Arc::new(Mutex::new(HashMap::new())), provider_factory, builtins, - client_fs_capabilities: Mutex::new(FileSystemCapability::new()), + client_fs_capabilities: OnceCell::new(), + client_terminal: OnceCell::new(), config_dir, session_manager, permission_manager, @@ -371,10 +374,15 @@ impl GooseAcpAgent { .unwrap_or_default(); extensions.extend(self.builtins.iter().map(|b| builtin_to_extension_config(b))); - let caps = self.client_fs_capabilities.lock().await.clone(); + let caps = self + .client_fs_capabilities + .get() + .cloned() + .unwrap_or_default(); + let terminal = self.client_terminal.get().copied().unwrap_or(false); let acp_developer = match (cx, session_id) { (Some(cx), Some(sid)) - if (caps.read_text_file || caps.write_text_file) + if (caps.read_text_file || caps.write_text_file || terminal) && extensions.iter().any(|e| e.name() == "developer") => { let context = agent.extension_manager.get_context().clone(); @@ -384,6 +392,7 @@ impl GooseAcpAgent { session_id: sid.clone(), fs_read: caps.read_text_file, fs_write: caps.write_text_file, + terminal, }); let dev_ext = extensions.iter().find(|e| e.name() == "developer"); let available_tools = dev_ext @@ -569,20 +578,27 @@ impl GooseAcpAgent { Err(_) => ToolCallStatus::Failed, }; - let content = build_tool_call_content(&tool_response.tool_result); + let mut fields = ToolCallUpdateFields::new().status(status); + if !tool_response + .tool_result + .as_ref() + .is_ok_and(|r| r.is_acp_aware()) + { + let content = build_tool_call_content(&tool_response.tool_result); + fields = fields.content(content); - let locations = extract_locations_from_meta(tool_response).unwrap_or_else(|| { - if let Some(tool_request) = session.tool_requests.get(&tool_response.id) { - extract_tool_locations(tool_request, tool_response) - } else { - Vec::new() + let locations = extract_locations_from_meta(tool_response).unwrap_or_else(|| { + if let Some(tool_request) = session.tool_requests.get(&tool_response.id) { + extract_tool_locations(tool_request, tool_response) + } else { + Vec::new() + } + }); + if !locations.is_empty() { + fields = fields.locations(locations); } - }); - - let mut fields = ToolCallUpdateFields::new().status(status).content(content); - if !locations.is_empty() { - fields = fields.locations(locations); } + cx.send_notification(SessionNotification::new( session_id.clone(), SessionUpdate::ToolCallUpdate(ToolCallUpdate::new( @@ -731,7 +747,10 @@ impl GooseAcpAgent { ) -> Result { debug!(?args, "initialize request"); - *self.client_fs_capabilities.lock().await = args.client_capabilities.fs.clone(); + let _ = self + .client_fs_capabilities + .set(args.client_capabilities.fs.clone()); + let _ = self.client_terminal.set(args.client_capabilities.terminal); let capabilities = AgentCapabilities::new() .load_session(true) diff --git a/crates/goose-acp/src/tools.rs b/crates/goose-acp/src/tools.rs new file mode 100644 index 00000000..a5c2082c --- /dev/null +++ b/crates/goose-acp/src/tools.rs @@ -0,0 +1,25 @@ +use rmcp::{ + model::{CallToolResult, Meta}, + object, +}; + +const ACP_AWARE_META_KEY: &str = "_goose/acp-aware"; + +pub trait AcpAwareToolMeta { + fn with_acp_aware_meta(self) -> Self; + fn is_acp_aware(&self) -> bool; +} + +impl AcpAwareToolMeta for CallToolResult { + fn with_acp_aware_meta(self) -> Self { + self.with_meta(Some(Meta(object!({ACP_AWARE_META_KEY: true})))) + } + + fn is_acp_aware(&self) -> bool { + self.meta + .as_ref() + .and_then(|meta| meta.get(ACP_AWARE_META_KEY)) + .and_then(|v| v.as_bool()) + .unwrap_or(false) + } +} diff --git a/crates/goose-cli/src/scenario_tests/mock_client.rs b/crates/goose-cli/src/scenario_tests/mock_client.rs index f1ea8d25..31932e5b 100644 --- a/crates/goose-cli/src/scenario_tests/mock_client.rs +++ b/crates/goose-cli/src/scenario_tests/mock_client.rs @@ -95,10 +95,9 @@ impl McpClientTrait for MockClient { async fn call_tool( &self, - _session_id: &str, + _ctx: &goose::agents::ToolCallContext, name: &str, arguments: Option>, - _working_dir: Option<&str>, _cancel_token: CancellationToken, ) -> Result { if let Some(handler) = self.handlers.get(name) { diff --git a/crates/goose-server/src/routes/agent.rs b/crates/goose-server/src/routes/agent.rs index 73e59423..0250a1bd 100644 --- a/crates/goose-server/src/routes/agent.rs +++ b/crates/goose-server/src/routes/agent.rs @@ -994,14 +994,10 @@ async fn call_tool( params }; + let ctx = goose::agents::ToolCallContext::new(payload.session_id.clone(), None, None); let tool_result = agent .extension_manager - .dispatch_tool_call( - &payload.session_id, - tool_call, - None, - CancellationToken::default(), - ) + .dispatch_tool_call(&ctx, tool_call, CancellationToken::default()) .await .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; diff --git a/crates/goose/src/agents/agent.rs b/crates/goose/src/agents/agent.rs index fc1ace44..62b40d6e 100644 --- a/crates/goose/src/agents/agent.rs +++ b/crates/goose/src/agents/agent.rs @@ -542,22 +542,25 @@ impl Agent { }; } + let ctx = super::tool_execution::ToolCallContext::new( + session.id.clone(), + Some(session.working_dir.clone()), + Some(request_id.clone()), + ); + debug!("WAITING_TOOL_START: {}", tool_call.name); let result: ToolCallResult = if self.is_frontend_tool(&tool_call.name).await { - // For frontend tools, return an error indicating we need frontend execution ToolCallResult::from(Err(ErrorData::new( ErrorCode::INTERNAL_ERROR, "Frontend tool execution required".to_string(), None, ))) } else { - // Clone the result to ensure no references to extension_manager are returned let result = self .extension_manager .dispatch_tool_call( - &session.id, + &ctx, tool_call.clone(), - Some(session.working_dir.as_path()), cancellation_token.unwrap_or_default(), ) .await; @@ -566,7 +569,6 @@ impl Agent { "tool_execution_failed", &format!("{}: {}", tool_call.name, e), ); - // Try to downcast to ErrorData to avoid double wrapping let error_data = e.downcast::().unwrap_or_else(|e| { ErrorData::new(ErrorCode::INTERNAL_ERROR, e.to_string(), None) }); diff --git a/crates/goose/src/agents/extension_manager.rs b/crates/goose/src/agents/extension_manager.rs index 954221e8..30476210 100644 --- a/crates/goose/src/agents/extension_manager.rs +++ b/crates/goose/src/agents/extension_manager.rs @@ -30,7 +30,7 @@ use super::extension::{ ExtensionConfig, ExtensionError, ExtensionInfo, ExtensionResult, PlatformExtensionContext, ToolInfo, PLATFORM_EXTENSIONS, }; -use super::tool_execution::ToolCallResult; +use super::tool_execution::{ToolCallContext, ToolCallResult}; use super::types::SharedProvider; use crate::agents::extension::{Envs, ProcessExit}; use crate::agents::extension_malware_check; @@ -1385,13 +1385,12 @@ impl ExtensionManager { pub async fn dispatch_tool_call( &self, - session_id: &str, + ctx: &super::tool_execution::ToolCallContext, tool_call: CallToolRequestParams, - working_dir: Option<&std::path::Path>, cancellation_token: CancellationToken, ) -> Result { let tool_name_str = tool_call.name.to_string(); - let resolved = self.resolve_tool(session_id, &tool_name_str).await?; + let resolved = self.resolve_tool(&ctx.session_id, &tool_name_str).await?; if let Some(extension) = self.extensions.lock().await.get(&resolved.extension_name) { if !extension @@ -1413,25 +1412,22 @@ impl ExtensionManager { let arguments = tool_call.arguments.clone(); let client = resolved.client.clone(); let notifications_receiver = client.subscribe().await; - let session_id = session_id.to_string(); let actual_tool_name = resolved.actual_tool_name; - let working_dir_str = working_dir.map(|p| p.to_string_lossy().to_string()); + let owned_ctx = ToolCallContext::new( + ctx.session_id.clone(), + ctx.working_dir.clone(), + ctx.tool_call_request_id.clone(), + ); let fut = async move { tracing::debug!( "dispatch_tool_call: calling client.call_tool tool={} session_id={} working_dir={:?}", actual_tool_name, - session_id, - working_dir_str + owned_ctx.session_id, + owned_ctx.working_dir, ); client - .call_tool( - &session_id, - &actual_tool_name, - arguments, - working_dir_str.as_deref(), - cancellation_token, - ) + .call_tool(&owned_ctx, &actual_tool_name, arguments, cancellation_token) .await .map_err(|e| match e { ServiceError::McpError(error_data) => error_data, @@ -1798,10 +1794,9 @@ mod tests { async fn call_tool( &self, - _session_id: &str, + _ctx: &ToolCallContext, name: &str, _arguments: Option, - _working_dir: Option<&str>, _cancellation_token: CancellationToken, ) -> Result { match name { @@ -1838,6 +1833,8 @@ mod tests { #[tokio::test] async fn test_dispatch_tool_call() { + use super::super::tool_execution::ToolCallContext; + let temp_dir = tempfile::tempdir().unwrap(); let extension_manager = ExtensionManager::new_without_provider(temp_dir.path().to_path_buf()); @@ -1855,16 +1852,17 @@ mod tests { .add_mock_extension("client 🚀".to_string(), Arc::new(MockClient {})) .await; + let ctx = ToolCallContext::new( + "test-session-id".to_string(), + None, + Some("test-req-id".to_string()), + ); + let tool_call = CallToolRequestParams::new("test_client__tool".to_string()).with_arguments(object!({})); let result = extension_manager - .dispatch_tool_call( - "test-session-id", - tool_call, - None, - CancellationToken::default(), - ) + .dispatch_tool_call(&ctx, tool_call, CancellationToken::default()) .await; assert!(result.is_ok()); @@ -1872,12 +1870,7 @@ mod tests { .with_arguments(object!({})); let result = extension_manager - .dispatch_tool_call( - "test-session-id", - tool_call, - None, - CancellationToken::default(), - ) + .dispatch_tool_call(&ctx, tool_call, CancellationToken::default()) .await; assert!(result.is_ok()); @@ -1885,12 +1878,7 @@ mod tests { .with_arguments(object!({})); let result = extension_manager - .dispatch_tool_call( - "test-session-id", - tool_call, - None, - CancellationToken::default(), - ) + .dispatch_tool_call(&ctx, tool_call, CancellationToken::default()) .await; assert!(result.is_ok()); @@ -1898,12 +1886,7 @@ mod tests { CallToolRequestParams::new("client___tool".to_string()).with_arguments(object!({})); let result = extension_manager - .dispatch_tool_call( - "test-session-id", - tool_call, - None, - CancellationToken::default(), - ) + .dispatch_tool_call(&ctx, tool_call, CancellationToken::default()) .await; assert!(result.is_ok()); @@ -1911,12 +1894,7 @@ mod tests { CallToolRequestParams::new("client___tools".to_string()).with_arguments(object!({})); let result = extension_manager - .dispatch_tool_call( - "test-session-id", - invalid_tool_call, - None, - CancellationToken::default(), - ) + .dispatch_tool_call(&ctx, invalid_tool_call, CancellationToken::default()) .await; if let Err(err) = result { let tool_err = err.downcast_ref::().expect("Expected ErrorData"); @@ -1929,12 +1907,7 @@ mod tests { CallToolRequestParams::new("_client__tools".to_string()).with_arguments(object!({})); let result = extension_manager - .dispatch_tool_call( - "test-session-id", - invalid_tool_call, - None, - CancellationToken::default(), - ) + .dispatch_tool_call(&ctx, invalid_tool_call, CancellationToken::default()) .await; if let Err(err) = result { let tool_err = err.downcast_ref::().expect("Expected ErrorData"); @@ -2009,6 +1982,8 @@ mod tests { #[tokio::test] async fn test_dispatch_unavailable_tool_returns_error() { + use super::super::tool_execution::ToolCallContext; + let temp_dir = tempfile::tempdir().unwrap(); let extension_manager = ExtensionManager::new_without_provider(temp_dir.path().to_path_buf()); @@ -2023,16 +1998,17 @@ mod tests { ) .await; + let ctx = ToolCallContext::new( + "test-session-id".to_string(), + None, + Some("test-req-id".to_string()), + ); + let unavailable_tool_call = CallToolRequestParams::new("test_extension__tool".to_string()) .with_arguments(object!({})); let result = extension_manager - .dispatch_tool_call( - "test-session-id", - unavailable_tool_call, - None, - CancellationToken::default(), - ) + .dispatch_tool_call(&ctx, unavailable_tool_call, CancellationToken::default()) .await; if let Err(err) = result { @@ -2048,12 +2024,7 @@ mod tests { .with_arguments(object!({})); let result = extension_manager - .dispatch_tool_call( - "test-session-id", - available_tool_call, - None, - CancellationToken::default(), - ) + .dispatch_tool_call(&ctx, available_tool_call, CancellationToken::default()) .await; assert!(result.is_ok()); diff --git a/crates/goose/src/agents/mcp_client.rs b/crates/goose/src/agents/mcp_client.rs index 28a84c94..bae871de 100644 --- a/crates/goose/src/agents/mcp_client.rs +++ b/crates/goose/src/agents/mcp_client.rs @@ -1,4 +1,5 @@ use crate::action_required_manager::ActionRequiredManager; +use crate::agents::tool_execution::ToolCallContext; use crate::agents::types::SharedProvider; use crate::session_context::{SESSION_ID_HEADER, WORKING_DIR_HEADER}; use rmcp::model::{ @@ -47,10 +48,9 @@ pub trait McpClientTrait: Send + Sync { async fn call_tool( &self, - session_id: &str, + ctx: &ToolCallContext, name: &str, arguments: Option, - working_dir: Option<&str>, cancel_token: CancellationToken, ) -> Result; @@ -594,10 +594,9 @@ impl McpClientTrait for McpClient { async fn call_tool( &self, - session_id: &str, + ctx: &ToolCallContext, name: &str, arguments: Option, - working_dir: Option<&str>, cancel_token: CancellationToken, ) -> Result { let mut params = CallToolRequestParams::new(name.to_string()); @@ -607,7 +606,12 @@ impl McpClientTrait for McpClient { let request = ClientRequest::CallToolRequest(Request::new(params)); let result = self - .send_request_with_context(session_id, working_dir, request, cancel_token) + .send_request_with_context( + &ctx.session_id, + ctx.working_dir_str(), + request, + cancel_token, + ) .await; match result? { diff --git a/crates/goose/src/agents/mod.rs b/crates/goose/src/agents/mod.rs index 9f083b88..b3bc386a 100644 --- a/crates/goose/src/agents/mod.rs +++ b/crates/goose/src/agents/mod.rs @@ -30,4 +30,5 @@ pub use extension_manager::ExtensionManager; pub use prompt_manager::PromptManager; pub use subagent_handler::SUBAGENT_TOOL_REQUEST_TYPE; pub use subagent_task_config::TaskConfig; +pub use tool_execution::ToolCallContext; pub use types::{FrontendTool, RetryConfig, SessionConfig, SuccessCheck}; diff --git a/crates/goose/src/agents/platform_extensions/analyze/mod.rs b/crates/goose/src/agents/platform_extensions/analyze/mod.rs index 0b24fca5..ec98eded 100644 --- a/crates/goose/src/agents/platform_extensions/analyze/mod.rs +++ b/crates/goose/src/agents/platform_extensions/analyze/mod.rs @@ -5,6 +5,7 @@ pub mod parser; use crate::agents::extension::PlatformExtensionContext; use crate::agents::mcp_client::{Error, McpClientTrait}; +use crate::agents::tool_execution::ToolCallContext; use anyhow::Result; use async_trait::async_trait; use ignore::WalkBuilder; @@ -234,13 +235,12 @@ impl McpClientTrait for AnalyzeClient { async fn call_tool( &self, - _session_id: &str, + ctx: &ToolCallContext, name: &str, arguments: Option, - working_dir: Option<&str>, _cancellation_token: CancellationToken, ) -> Result { - let working_dir = working_dir.map(Path::new); + let working_dir = ctx.working_dir.as_deref(); match name { "analyze" => match Self::parse_args::(arguments) { Ok(params) => { diff --git a/crates/goose/src/agents/platform_extensions/apps.rs b/crates/goose/src/agents/platform_extensions/apps.rs index 140d84f1..9301b114 100644 --- a/crates/goose/src/agents/platform_extensions/apps.rs +++ b/crates/goose/src/agents/platform_extensions/apps.rs @@ -1,5 +1,6 @@ use crate::agents::extension::PlatformExtensionContext; use crate::agents::mcp_client::{Error, McpClientTrait}; +use crate::agents::tool_execution::ToolCallContext; use crate::config::paths::Paths; use crate::conversation::message::Message; use crate::goose_apps::McpAppResource; @@ -526,12 +527,12 @@ impl McpClientTrait for AppsManagerClient { async fn call_tool( &self, - session_id: &str, + ctx: &ToolCallContext, name: &str, arguments: Option, - _working_dir: Option<&str>, _cancel_token: CancellationToken, ) -> Result { + let session_id = &ctx.session_id; let result = match name { "list_apps" => self.handle_list_apps(arguments).await, "create_app" => self.handle_create_app(session_id, arguments).await, diff --git a/crates/goose/src/agents/platform_extensions/chatrecall.rs b/crates/goose/src/agents/platform_extensions/chatrecall.rs index 564d8602..aa135417 100644 --- a/crates/goose/src/agents/platform_extensions/chatrecall.rs +++ b/crates/goose/src/agents/platform_extensions/chatrecall.rs @@ -1,5 +1,6 @@ use crate::agents::extension::PlatformExtensionContext; use crate::agents::mcp_client::{Error, McpClientTrait}; +use crate::agents::tool_execution::ToolCallContext; use anyhow::Result; use async_trait::async_trait; use indoc::indoc; @@ -273,12 +274,12 @@ impl McpClientTrait for ChatRecallClient { async fn call_tool( &self, - session_id: &str, + ctx: &ToolCallContext, name: &str, arguments: Option, - _working_dir: Option<&str>, _cancellation_token: CancellationToken, ) -> Result { + let session_id = &ctx.session_id; let content = match name { "chatrecall" => self.handle_chatrecall(session_id, arguments).await, _ => Err(format!("Unknown tool: {}", name)), diff --git a/crates/goose/src/agents/platform_extensions/code_execution.rs b/crates/goose/src/agents/platform_extensions/code_execution.rs index db64dea4..98fc0fb2 100644 --- a/crates/goose/src/agents/platform_extensions/code_execution.rs +++ b/crates/goose/src/agents/platform_extensions/code_execution.rs @@ -1,6 +1,7 @@ use crate::agents::extension::PlatformExtensionContext; use crate::agents::extension_manager::get_tool_owner; use crate::agents::mcp_client::{Error, McpClientTrait}; +use crate::agents::tool_execution::ToolCallContext; use anyhow::Result; use async_trait::async_trait; use indoc::indoc; @@ -252,8 +253,13 @@ fn create_tool_callback( } params }; + let ctx = crate::agents::ToolCallContext::new( + session_id, + None, + Some("tool-request-id".to_string()), + ); match manager - .dispatch_tool_call(&session_id, tool_call, None, CancellationToken::new()) + .dispatch_tool_call(&ctx, tool_call, CancellationToken::new()) .await { Ok(dispatch_result) => match dispatch_result.result.await { @@ -318,10 +324,10 @@ impl McpClientTrait for CodeExecutionClient { "list_functions".to_string(), indoc! {r#" List all available functions across all namespaces. - + This will not return function input and output types. After determining which functions are needed use - get_function_details to get input and output type + get_function_details to get input and output type information about specific functions. "#} .to_string(), @@ -420,12 +426,12 @@ impl McpClientTrait for CodeExecutionClient { async fn call_tool( &self, - session_id: &str, + ctx: &ToolCallContext, name: &str, arguments: Option, - _working_dir: Option<&str>, _cancellation_token: CancellationToken, ) -> Result { + let session_id = &ctx.session_id; let result = match name { "list_functions" => self.handle_list_functions(session_id).await, "get_function_details" => { diff --git a/crates/goose/src/agents/platform_extensions/developer/mod.rs b/crates/goose/src/agents/platform_extensions/developer/mod.rs index 95297bba..9633d5ed 100644 --- a/crates/goose/src/agents/platform_extensions/developer/mod.rs +++ b/crates/goose/src/agents/platform_extensions/developer/mod.rs @@ -4,6 +4,7 @@ pub mod tree; use crate::agents::extension::PlatformExtensionContext; use crate::agents::mcp_client::{Error, McpClientTrait}; +use crate::agents::ToolCallContext; use anyhow::Result; use async_trait::async_trait; use edit::{EditTools, FileEditParams, FileWriteParams}; @@ -15,7 +16,6 @@ use rmcp::model::{ use schemars::{schema_for, JsonSchema}; use serde_json::Value; use shell::{ShellOutput, ShellParams, ShellTool}; -use std::path::Path; use std::sync::Arc; use tokio_util::sync::CancellationToken; use tree::{TreeParams, TreeTool}; @@ -146,13 +146,12 @@ impl McpClientTrait for DeveloperClient { async fn call_tool( &self, - _session_id: &str, + ctx: &ToolCallContext, name: &str, arguments: Option, - working_dir: Option<&str>, - _cancellation_token: CancellationToken, + _cancel_token: CancellationToken, ) -> Result { - let working_dir = working_dir.map(Path::new); + let working_dir = ctx.working_dir.as_deref(); match name { "shell" => match Self::parse_args::(arguments) { Ok(params) => Ok(self.shell_tool.shell_with_cwd(params, working_dir).await), @@ -231,15 +230,15 @@ mod tests { let cwd = temp.path().join("workspace"); fs::create_dir_all(&cwd).unwrap(); + let ctx = ToolCallContext::new("session".to_owned(), Some(cwd.clone()), None); let write = client .call_tool( - "session", + &ctx, "write", Some(object!({ "path": "notes.txt", "content": "first line" })), - Some(cwd.to_str().unwrap()), CancellationToken::new(), ) .await @@ -252,14 +251,13 @@ mod tests { let edit = client .call_tool( - "session", + &ctx, "edit", Some(object!({ "path": "notes.txt", "before": "first", "after": "updated" })), - Some(cwd.to_str().unwrap()), CancellationToken::new(), ) .await @@ -279,14 +277,14 @@ mod tests { let cwd = temp.path().join("workspace"); fs::create_dir_all(&cwd).unwrap(); + let ctx = ToolCallContext::new("session".to_owned(), Some(cwd.clone()), None); let result = client .call_tool( - "session", + &ctx, "shell", Some(object!({ "command": "pwd" })), - Some(cwd.to_str().unwrap()), CancellationToken::new(), ) .await diff --git a/crates/goose/src/agents/platform_extensions/developer/shell.rs b/crates/goose/src/agents/platform_extensions/developer/shell.rs index 56ce4ce2..ea7a5c62 100644 --- a/crates/goose/src/agents/platform_extensions/developer/shell.rs +++ b/crates/goose/src/agents/platform_extensions/developer/shell.rs @@ -13,7 +13,7 @@ use tokio_stream::{wrappers::SplitStream, StreamExt}; use crate::subprocess::SubprocessExt; const OUTPUT_LIMIT_LINES: usize = 2000; -const OUTPUT_LIMIT_BYTES: usize = 50_000; +pub const OUTPUT_LIMIT_BYTES: usize = 50_000; const OUTPUT_PREVIEW_LINES: usize = 50; const OUTPUT_SLOTS: usize = 8; diff --git a/crates/goose/src/agents/platform_extensions/ext_manager.rs b/crates/goose/src/agents/platform_extensions/ext_manager.rs index 1dde4309..58221b69 100644 --- a/crates/goose/src/agents/platform_extensions/ext_manager.rs +++ b/crates/goose/src/agents/platform_extensions/ext_manager.rs @@ -1,5 +1,6 @@ use crate::agents::extension::PlatformExtensionContext; use crate::agents::mcp_client::{Error, McpClientTrait}; +use crate::agents::tool_execution::ToolCallContext; use crate::config::get_extension_by_name; use anyhow::Result; use async_trait::async_trait; @@ -408,12 +409,12 @@ impl McpClientTrait for ExtensionManagerClient { async fn call_tool( &self, - session_id: &str, + ctx: &ToolCallContext, name: &str, arguments: Option, - _working_dir: Option<&str>, _cancellation_token: CancellationToken, ) -> Result { + let session_id = &ctx.session_id; let result = match name { SEARCH_AVAILABLE_EXTENSIONS_TOOL_NAME => { self.handle_search_available_extensions().await diff --git a/crates/goose/src/agents/platform_extensions/summarize.rs b/crates/goose/src/agents/platform_extensions/summarize.rs index 214cc513..ea974f64 100644 --- a/crates/goose/src/agents/platform_extensions/summarize.rs +++ b/crates/goose/src/agents/platform_extensions/summarize.rs @@ -1,6 +1,7 @@ use std::path::{Path, PathBuf}; use std::sync::Arc; +use crate::agents::tool_execution::ToolCallContext; use async_trait::async_trait; use ignore::gitignore::{Gitignore, GitignoreBuilder}; use rmcp::model::{ @@ -98,10 +99,9 @@ impl McpClientTrait for SummarizeClient { async fn call_tool( &self, - session_id: &str, + ctx: &ToolCallContext, name: &str, arguments: Option, - working_dir: Option<&str>, _cancellation_token: CancellationToken, ) -> Result { if name != "summarize" { @@ -111,7 +111,7 @@ impl McpClientTrait for SummarizeClient { ))])); } - let Some(working_dir) = working_dir else { + let Some(working_dir) = ctx.working_dir_str() else { return Ok(CallToolResult::error(vec![Content::text( "Error: working_dir is required for summarize", )])); @@ -148,6 +148,7 @@ impl McpClientTrait for SummarizeClient { } }; + let session_id = &ctx.session_id; match execute_summarize(provider, session_id, params, &working_dir).await { Ok(result) => Ok(result), Err(msg) => Ok(CallToolResult::error(vec![Content::text(format!( diff --git a/crates/goose/src/agents/platform_extensions/summon.rs b/crates/goose/src/agents/platform_extensions/summon.rs index 395290ae..84345b1e 100644 --- a/crates/goose/src/agents/platform_extensions/summon.rs +++ b/crates/goose/src/agents/platform_extensions/summon.rs @@ -9,6 +9,7 @@ use crate::agents::extension::PlatformExtensionContext; use crate::agents::mcp_client::{Error, McpClientTrait}; use crate::agents::subagent_handler::{run_subagent_task, OnMessageCallback, SubagentRunParams}; use crate::agents::subagent_task_config::{TaskConfig, DEFAULT_SUBAGENT_MAX_TURNS}; +use crate::agents::tool_execution::ToolCallContext; use crate::agents::AgentConfig; use crate::config::paths::Paths; use crate::config::Config; @@ -1815,12 +1816,12 @@ impl McpClientTrait for SummonClient { async fn call_tool( &self, - session_id: &str, + ctx: &ToolCallContext, name: &str, arguments: Option, - _working_dir: Option<&str>, cancellation_token: CancellationToken, ) -> Result { + let session_id = &ctx.session_id; let content = match name { "load" => self.handle_load(session_id, arguments).await, "delegate" => { @@ -2252,8 +2253,9 @@ You review code."#; let names: Vec<_> = result.tools.iter().map(|t| t.name.as_ref()).collect(); assert!(names.contains(&"load") && names.contains(&"delegate")); + let ctx = ToolCallContext::new("test".to_string(), None, None); let result = client - .call_tool("test", "unknown", None, None, CancellationToken::new()) + .call_tool(&ctx, "unknown", None, CancellationToken::new()) .await .unwrap(); assert!(result.is_error.unwrap_or(false)); diff --git a/crates/goose/src/agents/platform_extensions/todo.rs b/crates/goose/src/agents/platform_extensions/todo.rs index ea47dc38..5eb0d67c 100644 --- a/crates/goose/src/agents/platform_extensions/todo.rs +++ b/crates/goose/src/agents/platform_extensions/todo.rs @@ -1,5 +1,6 @@ use crate::agents::extension::PlatformExtensionContext; use crate::agents::mcp_client::{Error, McpClientTrait}; +use crate::agents::tool_execution::ToolCallContext; use crate::session::extension_data; use crate::session::extension_data::ExtensionState; use anyhow::Result; @@ -155,12 +156,12 @@ impl McpClientTrait for TodoClient { async fn call_tool( &self, - session_id: &str, + ctx: &ToolCallContext, name: &str, arguments: Option, - _working_dir: Option<&str>, _cancellation_token: CancellationToken, ) -> Result { + let session_id = &ctx.session_id; let content = match name { "todo_write" => self.handle_write_todo(session_id, arguments).await, _ => Err(format!("Unknown tool: {}", name)), diff --git a/crates/goose/src/agents/platform_extensions/tom.rs b/crates/goose/src/agents/platform_extensions/tom.rs index c4aa3a4e..6d5d5de4 100644 --- a/crates/goose/src/agents/platform_extensions/tom.rs +++ b/crates/goose/src/agents/platform_extensions/tom.rs @@ -1,5 +1,6 @@ use crate::agents::extension::PlatformExtensionContext; use crate::agents::mcp_client::{Error, McpClientTrait}; +use crate::agents::tool_execution::ToolCallContext; use anyhow::Result; use async_trait::async_trait; use rmcp::model::{ @@ -45,10 +46,9 @@ impl McpClientTrait for TomClient { async fn call_tool( &self, - _session_id: &str, + _ctx: &ToolCallContext, name: &str, _arguments: Option, - _working_dir: Option<&str>, _cancellation_token: CancellationToken, ) -> Result { Ok(CallToolResult::error(vec![Content::text(format!( diff --git a/crates/goose/src/agents/tool_execution.rs b/crates/goose/src/agents/tool_execution.rs index fbc68fa7..2811d81c 100644 --- a/crates/goose/src/agents/tool_execution.rs +++ b/crates/goose/src/agents/tool_execution.rs @@ -8,11 +8,38 @@ use futures::{Stream, StreamExt}; use tokio::sync::Mutex; use tokio_util::sync::CancellationToken; +use std::path::PathBuf; + use crate::config::permission::PermissionLevel; use crate::mcp_utils::ToolResult; use crate::permission::Permission; use rmcp::model::{Content, ServerNotification}; +/// Context passed through the tool call dispatch chain. +pub struct ToolCallContext { + pub session_id: String, + pub working_dir: Option, + pub tool_call_request_id: Option, +} + +impl ToolCallContext { + pub fn new( + session_id: String, + working_dir: Option, + tool_call_request_id: Option, + ) -> Self { + Self { + session_id, + working_dir, + tool_call_request_id, + } + } + + pub fn working_dir_str(&self) -> Option<&str> { + self.working_dir.as_ref().and_then(|p| p.to_str()) + } +} + // ToolCallResult combines the result of a tool call with an optional notification stream that // can be used to receive notifications from the tool. pub struct ToolCallResult { diff --git a/crates/goose/tests/mcp_integration_test.rs b/crates/goose/tests/mcp_integration_test.rs index 3996a355..693f31f9 100644 --- a/crates/goose/tests/mcp_integration_test.rs +++ b/crates/goose/tests/mcp_integration_test.rs @@ -276,13 +276,13 @@ async fn test_replayed_session( new_call = new_call.with_arguments(args); } let tool_call = new_call; + let ctx = goose::agents::ToolCallContext::new( + "test-session-id".to_string(), + None, + Some("test-id".to_string()), + ); let result = extension_manager - .dispatch_tool_call( - "test-session-id", - tool_call, - None, - CancellationToken::default(), - ) + .dispatch_tool_call(&ctx, tool_call, CancellationToken::default()) .await; let tool_result = result?; diff --git a/crates/goose/tests/providers.rs b/crates/goose/tests/providers.rs index c612bf06..b85bc9fe 100644 --- a/crates/goose/tests/providers.rs +++ b/crates/goose/tests/providers.rs @@ -328,10 +328,15 @@ impl ProviderFixture { }; let params = tool_req.tool_call.as_ref().unwrap().clone(); + let ctx = goose::agents::ToolCallContext::new( + self.session_id.to_string(), + None, + Some("test-id".to_string()), + ); let result = self .agent .extension_manager - .dispatch_tool_call(&self.session_id, params, None, CancellationToken::new()) + .dispatch_tool_call(&ctx, params, CancellationToken::new()) .await .unwrap() .result