refactor: smaller pieces of agent.reply() (#2153)

This commit is contained in:
Salman Mohammed
2025-04-11 12:49:25 -04:00
committed by GitHub
parent 31bb81e2c8
commit 0b800e05b7
5 changed files with 208 additions and 151 deletions
+55 -149
View File
@@ -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<dyn Provider>,
extension_manager: Mutex<ExtensionManager>,
frontend_tools: HashMap<String, FrontendTool>,
frontend_instructions: Option<String>,
prompt_manager: PromptManager,
token_counter: TokenCounter,
confirmation_tx: mpsc::Sender<(String, PermissionConfirmation)>,
confirmation_rx: Mutex<mpsc::Receiver<(String, PermissionConfirmation)>>,
tool_result_tx: mpsc::Sender<(String, ToolResult<Vec<Content>>)>,
tool_result_rx: ToolResultReceiver,
pub(super) provider: Arc<dyn Provider>,
pub(super) extension_manager: Mutex<ExtensionManager>,
pub(super) frontend_tools: HashMap<String, FrontendTool>,
pub(super) frontend_instructions: Option<String>,
pub(super) prompt_manager: PromptManager,
pub(super) token_counter: TokenCounter,
pub(super) confirmation_tx: mpsc::Sender<(String, PermissionConfirmation)>,
pub(super) confirmation_rx: Mutex<mpsc::Receiver<(String, PermissionConfirmation)>>,
pub(super) tool_result_tx: mpsc::Sender<(String, ToolResult<Vec<Content>>)>,
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<Vec<Content>, 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<Vec<Content>, 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<Tool> {
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<BoxStream<'_, anyhow::Result<Message>>> {
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<String>,
HashSet<String>,
) = 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<i32>, b: Option<i32>| -> Option<i32> {
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<String> {
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<Message>) -> Result<Recipe> {
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
+1 -1
View File
@@ -229,7 +229,7 @@ impl ExtensionManager {
}
/// Get all tools from all clients with proper prefixing
pub async fn get_prefixed_tools(&mut self) -> ExtensionResult<Vec<Tool>> {
pub async fn get_prefixed_tools(&self) -> ExtensionResult<Vec<Tool>> {
let mut tools = Vec::new();
// Add tools from MCP extensions with prefixing
+1
View File
@@ -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;
+150
View File
@@ -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<Tool>, Vec<Tool>, 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<String>, HashSet<String>) {
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<dyn Provider>,
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<i32>, b: Option<i32>| -> Option<i32> {
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(())
}
}
+1 -1
View File
@@ -31,7 +31,7 @@ pub struct SessionMetadata {
pub input_tokens: Option<i32>,
/// The number of output tokens used in the session. Retrieved from the provider's last usage.
pub output_tokens: Option<i32>,
/// 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<i32>,
/// The number of input tokens used in the session. Accumulated across all messages.
pub accumulated_input_tokens: Option<i32>,