diff --git a/crates/goose/src/agents/agent.rs b/crates/goose/src/agents/agent.rs index 3bee1e59..6a35a4d0 100644 --- a/crates/goose/src/agents/agent.rs +++ b/crates/goose/src/agents/agent.rs @@ -1,17 +1,9 @@ -use std::collections::{HashMap, HashSet}; +use std::collections::HashMap; use std::sync::Arc; use anyhow::{anyhow, Result}; use futures::stream::BoxStream; -use regex::Regex; -use serde_json::Value; -use tokio::sync::{mpsc, Mutex}; -use tracing::{debug, error, instrument, warn}; - -use crate::agents::extension::{ExtensionConfig, ExtensionResult, ToolInfo}; -use crate::agents::extension_manager::{get_parameter_names, ExtensionManager}; -use crate::agents::types::ToolResultReceiver; use crate::config::permission::PermissionLevel; use crate::config::{Config, ExtensionConfigManager, PermissionManager}; use crate::message::{Message, MessageContent, ToolRequest}; @@ -19,43 +11,42 @@ use crate::permission::permission_judge::check_tool_permissions; use crate::permission::{Permission, PermissionConfirmation}; use crate::providers::base::Provider; use crate::providers::errors::ProviderError; -use crate::providers::toolshim::{ - augment_message_with_tool_calls, modify_system_prompt_for_tool_json, OllamaInterpreter, -}; use crate::recipe::{Author, Recipe}; -use crate::session; use crate::token_counter::TokenCounter; use crate::truncate::{truncate_messages, OldestFirstTruncation}; +use regex::Regex; +use serde_json::Value; +use tokio::sync::{mpsc, Mutex}; +use tracing::{debug, error, instrument, warn}; -use mcp_core::{ - prompt::Prompt, protocol::GetPromptResult, tool::Tool, Content, ToolError, ToolResult, -}; - +use crate::agents::extension::{ExtensionConfig, ExtensionResult, ToolInfo}; +use crate::agents::extension_manager::{get_parameter_names, ExtensionManager}; use crate::agents::platform_tools::{ - self, PLATFORM_LIST_RESOURCES_TOOL_NAME, PLATFORM_READ_RESOURCE_TOOL_NAME, - PLATFORM_SEARCH_AVAILABLE_EXTENSIONS_TOOL_NAME, + PLATFORM_ENABLE_EXTENSION_TOOL_NAME, PLATFORM_LIST_RESOURCES_TOOL_NAME, + PLATFORM_READ_RESOURCE_TOOL_NAME, PLATFORM_SEARCH_AVAILABLE_EXTENSIONS_TOOL_NAME, }; use crate::agents::prompt_manager::PromptManager; use crate::agents::types::SessionConfig; - -use super::platform_tools::PLATFORM_ENABLE_EXTENSION_TOOL_NAME; -use super::types::FrontendTool; +use crate::agents::types::{FrontendTool, ToolResultReceiver}; +use mcp_core::{ + prompt::Prompt, protocol::GetPromptResult, tool::Tool, Content, ToolError, ToolResult, +}; const MAX_TRUNCATION_ATTEMPTS: usize = 3; const ESTIMATE_FACTOR_DECAY: f32 = 0.9; /// The main goose Agent pub struct Agent { - provider: Arc, - extension_manager: Mutex, - frontend_tools: HashMap, - frontend_instructions: Option, - prompt_manager: PromptManager, - token_counter: TokenCounter, - confirmation_tx: mpsc::Sender<(String, PermissionConfirmation)>, - confirmation_rx: Mutex>, - tool_result_tx: mpsc::Sender<(String, ToolResult>)>, - tool_result_rx: ToolResultReceiver, + pub(super) provider: Arc, + pub(super) extension_manager: Mutex, + pub(super) frontend_tools: HashMap, + pub(super) frontend_instructions: Option, + pub(super) prompt_manager: PromptManager, + pub(super) token_counter: TokenCounter, + pub(super) confirmation_tx: mpsc::Sender<(String, PermissionConfirmation)>, + pub(super) confirmation_rx: Mutex>, + pub(super) tool_result_tx: mpsc::Sender<(String, ToolResult>)>, + pub(super) tool_result_rx: ToolResultReceiver, } impl Agent { @@ -112,13 +103,13 @@ impl Agent { } /// Dispatch a single tool call to the appropriate client - #[instrument(skip(tool_call, extension_manager, request_id), fields(input, output))] - async fn create_tool_future( - extension_manager: &ExtensionManager, + #[instrument(skip(self, tool_call, request_id), fields(input, output))] + async fn dispatch_tool_call( + &self, tool_call: mcp_core::tool::ToolCall, - is_frontend_tool: bool, request_id: String, ) -> (String, Result, ToolError>) { + let extension_manager = self.extension_manager.lock().await; let result = if tool_call.name == PLATFORM_READ_RESOURCE_TOOL_NAME { // Check if the tool is read_resource and handle it separately extension_manager @@ -130,7 +121,7 @@ impl Agent { .await } else if tool_call.name == PLATFORM_SEARCH_AVAILABLE_EXTENSIONS_TOOL_NAME { extension_manager.search_available_extensions().await - } else if is_frontend_tool { + } else if self.is_frontend_tool(&tool_call.name) { // For frontend tools, return an error indicating we need frontend execution Err(ToolError::ExecutionError( "Frontend tool execution required".to_string(), @@ -199,10 +190,11 @@ impl Agent { } async fn enable_extension( - extension_manager: &mut ExtensionManager, + &self, extension_name: String, request_id: String, ) -> (String, Result, ToolError>) { + let mut extension_manager = self.extension_manager.lock().await; let config = match ExtensionConfigManager::get_config_by_name(&extension_name) { Ok(Some(config)) => config, Ok(None) => { @@ -275,7 +267,7 @@ impl Agent { } pub async fn list_tools(&self) -> Vec { - let mut extension_manager = self.extension_manager.lock().await; + let extension_manager = self.extension_manager.lock().await; extension_manager .get_prefixed_tools() .await @@ -317,56 +309,20 @@ impl Agent { ) -> anyhow::Result>> { let mut messages = messages.to_vec(); let reply_span = tracing::Span::current(); - let mut extension_manager = self.extension_manager.lock().await; - let mut tools = extension_manager.get_prefixed_tools().await?; let mut truncation_attempt: usize = 0; // Load settings from config let config = Config::global(); + + // Setup tools and prompt + let (mut tools, mut toolshim_tools, mut system_prompt) = + self.prepare_tools_and_prompt().await?; + let goose_mode = config.get_param("GOOSE_MODE").unwrap_or("auto".to_string()); - // we add in the 2 resource tools if any extensions support resources - // TODO: make sure there is no collision with another extension's tool name - if extension_manager.supports_resources() { - tools.push(platform_tools::read_resource_tool()); - tools.push(platform_tools::list_resources_tool()); - } - tools.push(platform_tools::search_available_extensions_tool()); - tools.push(platform_tools::enable_extension_tool()); + let (tools_with_readonly_annotation, tools_without_annotation) = + Self::categorize_tools_by_annotation(&tools); - let (tools_with_readonly_annotation, tools_without_annotation): ( - HashSet, - HashSet, - ) = tools - .iter() - .fold((HashSet::new(), HashSet::new()), |mut acc, tool| { - match &tool.annotations { - Some(annotations) if annotations.read_only_hint => { - acc.0.insert(tool.name.clone()); - } - _ => { - acc.1.insert(tool.name.clone()); - } - } - acc - }); - - let config = self.provider.get_model_config(); - let extensions_info = extension_manager.get_extensions_info().await; - let mut system_prompt = self - .prompt_manager - .build_system_prompt(extensions_info, self.frontend_instructions.clone()); - let mut toolshim_tools = vec![]; - if config.toolshim { - // If tool interpretation is enabled, modify the system prompt to instruct to return JSON tool requests - system_prompt = modify_system_prompt_for_tool_json(&system_prompt, &tools); - // make a copy of tools before empty - toolshim_tools = tools.clone(); - // pass empty tools vector to provider completion since toolshim will handle tool calls instead - tools = vec![]; - } - - // Set the user_message field in the span instead of creating a new event if let Some(content) = messages .last() .and_then(|msg| msg.content.first()) @@ -376,47 +332,19 @@ impl Agent { } Ok(Box::pin(async_stream::try_stream! { - let _reply_guard = reply_span.enter(); + let _ = reply_span.enter(); loop { - match self.provider().complete( + match Self::generate_response_from_provider( + self.provider(), &system_prompt, &messages, &tools, + &toolshim_tools, ).await { - Ok((mut response, usage)) => { - // Post-process / structure the response only if tool interpretation is enabled - if config.toolshim { - let interpreter = OllamaInterpreter::new() - .map_err(|e| anyhow::anyhow!("Failed to create OllamaInterpreter: {}", e))?; - - response = augment_message_with_tool_calls(&interpreter, response, &toolshim_tools).await?; - } - + Ok((response, usage)) => { // record usage for the session in the session file - if let Some(session) = session.clone() { - // TODO: track session_id in langfuse tracing - let session_file = session::get_path(session.id); - let mut metadata = session::read_metadata(&session_file)?; - metadata.working_dir = session.working_dir; - metadata.total_tokens = usage.usage.total_tokens; - metadata.input_tokens = usage.usage.input_tokens; - metadata.output_tokens = usage.usage.output_tokens; - - let accumulate = |a: Option, b: Option| -> Option { - match (a, b) { - (Some(x), Some(y)) => Some(x + y), - _ => a.or(b) - } - }; - - metadata.accumulated_total_tokens = accumulate(metadata.accumulated_total_tokens, usage.usage.total_tokens); - metadata.accumulated_input_tokens = accumulate(metadata.accumulated_input_tokens, usage.usage.input_tokens); - metadata.accumulated_output_tokens = accumulate(metadata.accumulated_output_tokens, usage.usage.output_tokens); - - // The message count is the number of messages in the session + 1 for the response - // The message count does not include the tool response till next iteration - metadata.message_count = messages.len() + 1; - session::update_metadata(&session_file, &metadata).await?; + if let Some(session_config) = session.clone() { + Self::update_session_metrics(session_config, &usage, messages.len()).await?; } // Reset truncation attempt @@ -539,7 +467,7 @@ impl Agent { .and_then(|v| v.as_str()) .unwrap_or("") .to_string(); - let install_result = Self::enable_extension(&mut extension_manager, extension_name, request.id.clone()).await; + let install_result = self.enable_extension(extension_name, request.id.clone()).await; install_results.push(install_result); } else { // User declined - add declined response @@ -557,8 +485,7 @@ impl Agent { // Skip the confirmation for approved tools for request in &permission_check_result.approved { if let Ok(tool_call) = request.tool_call.clone() { - let is_frontend_tool = self.is_frontend_tool(&tool_call.name); - let tool_future = Self::create_tool_future(&extension_manager, tool_call, is_frontend_tool, request.id.clone()); + let tool_future = self.dispatch_tool_call(tool_call, request.id.clone()); tool_futures.push(tool_future); } } @@ -573,7 +500,6 @@ impl Agent { // Process read-only tools for request in &permission_check_result.needs_approval { if let Ok(tool_call) = request.tool_call.clone() { - let is_frontend_tool = self.is_frontend_tool(&tool_call.name); let confirmation = Message::user().with_tool_confirmation_request( request.id.clone(), tool_call.name.clone(), @@ -589,7 +515,7 @@ impl Agent { let confirmed = tool_confirmation.permission == Permission::AllowOnce || tool_confirmation.permission == Permission::AlwaysAllow; if confirmed { // Add this tool call to the futures collection - let tool_future = Self::create_tool_future(&extension_manager, tool_call.clone(), is_frontend_tool, request.id.clone()); + let tool_future = self.dispatch_tool_call(tool_call.clone(), request.id.clone()); tool_futures.push(tool_future); if tool_confirmation.permission == Permission::AlwaysAllow { permission_manager.update_user_permission(&tool_call.name, PermissionLevel::AlwaysAllow); @@ -617,8 +543,7 @@ impl Agent { } // Check if any install results had errors before processing them - let all_successful = !install_results.iter().any(|(_, result)| result.is_err()); - + let all_install_successful = !install_results.iter().any(|(_, result)| result.is_err()); for (request_id, output) in install_results { message_tool_response = message_tool_response.with_tool_response( request_id, @@ -626,19 +551,11 @@ impl Agent { ); } - // Update system prompt and tools if all installations were successful - if all_successful { - let extensions_info = extension_manager.get_extensions_info().await; - system_prompt = self.prompt_manager.build_system_prompt(extensions_info, self.frontend_instructions.clone()); - tools = extension_manager.get_prefixed_tools().await?; - if extension_manager.supports_resources() { - tools.push(platform_tools::read_resource_tool()); - tools.push(platform_tools::list_resources_tool()); - } - tools.push(platform_tools::search_available_extensions_tool()); - tools.push(platform_tools::enable_extension_tool()); + // Update system prompt and tools if installations were successful + if all_install_successful { + (tools, toolshim_tools, system_prompt) = self.prepare_tools_and_prompt().await?; } - } + } yield message_tool_response.clone(); @@ -653,26 +570,15 @@ impl Agent { yield Message::assistant().with_text("Error: Context length exceeds limits even after multiple attempts to truncate. Please start a new session with fresh context and try again."); break; } - truncation_attempt += 1; warn!("Context length exceeded. Truncation Attempt: {}/{}.", truncation_attempt, MAX_TRUNCATION_ATTEMPTS); - // Decay the estimate factor as we make more truncation attempts // Estimate factor decays like this over time: 0.9, 0.81, 0.729, ... let estimate_factor: f32 = ESTIMATE_FACTOR_DECAY.powi(truncation_attempt as i32); - - // release the lock before truncation to prevent deadlock - drop(extension_manager); - if let Err(err) = self.truncate_messages(&mut messages, estimate_factor, &system_prompt, &mut tools).await { yield Message::assistant().with_text(format!("Error: Unable to truncate messages to stay within context limit. \n\nRan into this error: {}.\n\nPlease start a new session with fresh context and try again.", err)); break; } - - - // Re-acquire the lock - extension_manager = self.extension_manager.lock().await; - // Retry the loop after truncation continue; }, @@ -732,7 +638,7 @@ impl Agent { } pub async fn get_plan_prompt(&self) -> anyhow::Result { - let mut extension_manager = self.extension_manager.lock().await; + let extension_manager = self.extension_manager.lock().await; let tools = extension_manager.get_prefixed_tools().await?; let tools_info = tools .into_iter() @@ -758,7 +664,7 @@ impl Agent { } pub async fn create_recipe(&self, mut messages: Vec) -> Result { - let mut extension_manager = self.extension_manager.lock().await; + let extension_manager = self.extension_manager.lock().await; let extensions_info = extension_manager.get_extensions_info().await; let system_prompt = self .prompt_manager diff --git a/crates/goose/src/agents/extension_manager.rs b/crates/goose/src/agents/extension_manager.rs index d5eb6365..5694840c 100644 --- a/crates/goose/src/agents/extension_manager.rs +++ b/crates/goose/src/agents/extension_manager.rs @@ -229,7 +229,7 @@ impl ExtensionManager { } /// Get all tools from all clients with proper prefixing - pub async fn get_prefixed_tools(&mut self) -> ExtensionResult> { + pub async fn get_prefixed_tools(&self) -> ExtensionResult> { let mut tools = Vec::new(); // Add tools from MCP extensions with prefixing diff --git a/crates/goose/src/agents/mod.rs b/crates/goose/src/agents/mod.rs index dea6806f..d4c5ee6a 100644 --- a/crates/goose/src/agents/mod.rs +++ b/crates/goose/src/agents/mod.rs @@ -3,6 +3,7 @@ pub mod extension; pub mod extension_manager; pub mod platform_tools; pub mod prompt_manager; +mod reply_parts; mod types; pub use agent::Agent; diff --git a/crates/goose/src/agents/reply_parts.rs b/crates/goose/src/agents/reply_parts.rs new file mode 100644 index 00000000..b0ab2ea6 --- /dev/null +++ b/crates/goose/src/agents/reply_parts.rs @@ -0,0 +1,150 @@ +use std::collections::HashSet; +use std::sync::Arc; + +use anyhow::Result; + +use crate::agents::platform_tools; +use crate::message::Message; +use crate::providers::base::{Provider, ProviderUsage}; +use crate::providers::errors::ProviderError; +use crate::providers::toolshim::{ + augment_message_with_tool_calls, modify_system_prompt_for_tool_json, OllamaInterpreter, +}; +use crate::session; +use mcp_core::tool::Tool; + +use super::super::agents::Agent; + +impl Agent { + /// Prepares tools and system prompt for a provider request + pub(crate) async fn prepare_tools_and_prompt( + &self, + ) -> anyhow::Result<(Vec, Vec, String)> { + let extension_manager = self.extension_manager.lock().await; + // Get tools from extension manager + let mut tools = extension_manager.get_prefixed_tools().await?; + + // Add resource tools if supported + if extension_manager.supports_resources() { + tools.push(platform_tools::read_resource_tool()); + tools.push(platform_tools::list_resources_tool()); + } + + // Add platform tools + tools.push(platform_tools::search_available_extensions_tool()); + tools.push(platform_tools::enable_extension_tool()); + + // Add frontend tools + for frontend_tool in self.frontend_tools.values() { + tools.push(frontend_tool.tool.clone()); + } + + // Prepare system prompt + let extensions_info = extension_manager.get_extensions_info().await; + let mut system_prompt = self + .prompt_manager + .build_system_prompt(extensions_info, self.frontend_instructions.clone()); + + // Handle toolshim if enabled + let mut toolshim_tools = vec![]; + if self.provider.get_model_config().toolshim { + // If tool interpretation is enabled, modify the system prompt + system_prompt = modify_system_prompt_for_tool_json(&system_prompt, &tools); + // Make a copy of tools before emptying + toolshim_tools = tools.clone(); + // Empty the tools vector for provider completion + tools = vec![]; + } + + Ok((tools, toolshim_tools, system_prompt)) + } + + /// Categorize tools based on their annotations + /// Returns: + /// - read_only_tools: Tools with read-only annotations + /// - non_read_tools: Tools without read-only annotations + pub(crate) fn categorize_tools_by_annotation( + tools: &[Tool], + ) -> (HashSet, HashSet) { + tools + .iter() + .fold((HashSet::new(), HashSet::new()), |mut acc, tool| { + match &tool.annotations { + Some(annotations) if annotations.read_only_hint => { + acc.0.insert(tool.name.clone()); + } + _ => { + acc.1.insert(tool.name.clone()); + } + } + acc + }) + } + + /// Generate a response from the LLM provider + /// Handles toolshim transformations if needed + pub(crate) async fn generate_response_from_provider( + provider: Arc, + system_prompt: &str, + messages: &[Message], + tools: &[Tool], + toolshim_tools: &[Tool], + ) -> Result<(Message, ProviderUsage), ProviderError> { + let config = provider.get_model_config(); + + // Call the provider to get a response + let (mut response, usage) = provider.complete(system_prompt, messages, tools).await?; + + // Post-process / structure the response only if tool interpretation is enabled + if config.toolshim { + let interpreter = OllamaInterpreter::new().map_err(|e| { + ProviderError::ExecutionError(format!("Failed to create OllamaInterpreter: {}", e)) + })?; + + response = augment_message_with_tool_calls(&interpreter, response, toolshim_tools) + .await + .map_err(|e| { + ProviderError::ExecutionError(format!("Failed to augment message: {}", e)) + })?; + } + + Ok((response, usage)) + } + + /// Update session metrics after a response + pub(crate) async fn update_session_metrics( + session_config: crate::agents::types::SessionConfig, + usage: &crate::providers::base::ProviderUsage, + messages_length: usize, + ) -> Result<()> { + let session_file = session::get_path(session_config.id); + let mut metadata = session::read_metadata(&session_file)?; + + metadata.working_dir = session_config.working_dir.clone(); + metadata.total_tokens = usage.usage.total_tokens; + metadata.input_tokens = usage.usage.input_tokens; + metadata.output_tokens = usage.usage.output_tokens; + // The message count is the number of messages in the session + 1 for the response + // The message count does not include the tool response till next iteration + metadata.message_count = messages_length + 1; + + // Keep running sum of tokens to track cost over the entire session + let accumulate = |a: Option, b: Option| -> Option { + match (a, b) { + (Some(x), Some(y)) => Some(x + y), + _ => a.or(b), + } + }; + metadata.accumulated_total_tokens = + accumulate(metadata.accumulated_total_tokens, usage.usage.total_tokens); + metadata.accumulated_input_tokens = + accumulate(metadata.accumulated_input_tokens, usage.usage.input_tokens); + metadata.accumulated_output_tokens = accumulate( + metadata.accumulated_output_tokens, + usage.usage.output_tokens, + ); + session::update_metadata(&session_file, &metadata).await?; + + Ok(()) + } +} diff --git a/crates/goose/src/session/storage.rs b/crates/goose/src/session/storage.rs index a6ade795..d0fcb588 100644 --- a/crates/goose/src/session/storage.rs +++ b/crates/goose/src/session/storage.rs @@ -31,7 +31,7 @@ pub struct SessionMetadata { pub input_tokens: Option, /// The number of output tokens used in the session. Retrieved from the provider's last usage. pub output_tokens: Option, - /// The total number of tokens used in the session. Accumulated across all messages. + /// The total number of tokens used in the session. Accumulated across all messages (useful for tracking cost over an entire session). pub accumulated_total_tokens: Option, /// The number of input tokens used in the session. Accumulated across all messages. pub accumulated_input_tokens: Option,