feat: mcp router disable extensions (#2319)

This commit is contained in:
Wendy Tang
2025-04-23 15:31:50 -07:00
committed by GitHub
parent 93abd4e54b
commit 1b9699cca3
7 changed files with 110 additions and 22 deletions
+30 -5
View File
@@ -22,7 +22,7 @@ use tracing::{debug, error, instrument, warn};
use crate::agents::extension::{ExtensionConfig, ExtensionResult, ToolInfo}; use crate::agents::extension::{ExtensionConfig, ExtensionResult, ToolInfo};
use crate::agents::extension_manager::{get_parameter_names, ExtensionManager}; use crate::agents::extension_manager::{get_parameter_names, ExtensionManager};
use crate::agents::platform_tools::{ use crate::agents::platform_tools::{
PLATFORM_ENABLE_EXTENSION_TOOL_NAME, PLATFORM_LIST_RESOURCES_TOOL_NAME, PLATFORM_LIST_RESOURCES_TOOL_NAME, PLATFORM_MANAGE_EXTENSIONS_TOOL_NAME,
PLATFORM_READ_RESOURCE_TOOL_NAME, PLATFORM_SEARCH_AVAILABLE_EXTENSIONS_TOOL_NAME, PLATFORM_READ_RESOURCE_TOOL_NAME, PLATFORM_SEARCH_AVAILABLE_EXTENSIONS_TOOL_NAME,
}; };
use crate::agents::prompt_manager::PromptManager; use crate::agents::prompt_manager::PromptManager;
@@ -112,14 +112,22 @@ impl Agent {
tool_call: mcp_core::tool::ToolCall, tool_call: mcp_core::tool::ToolCall,
request_id: String, request_id: String,
) -> (String, Result<Vec<Content>, ToolError>) { ) -> (String, Result<Vec<Content>, ToolError>) {
if tool_call.name == PLATFORM_ENABLE_EXTENSION_TOOL_NAME { if tool_call.name == PLATFORM_MANAGE_EXTENSIONS_TOOL_NAME {
let extension_name = tool_call let extension_name = tool_call
.arguments .arguments
.get("extension_name") .get("extension_name")
.and_then(|v| v.as_str()) .and_then(|v| v.as_str())
.unwrap_or("") .unwrap_or("")
.to_string(); .to_string();
return self.enable_extension(extension_name, request_id).await; let action = tool_call
.arguments
.get("action")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
return self
.manage_extensions(action, extension_name, request_id)
.await;
} }
let extension_manager = self.extension_manager.lock().await; let extension_manager = self.extension_manager.lock().await;
@@ -202,12 +210,28 @@ impl Agent {
) )
} }
pub(super) async fn enable_extension( pub(super) async fn manage_extensions(
&self, &self,
action: String,
extension_name: String, extension_name: String,
request_id: String, request_id: String,
) -> (String, Result<Vec<Content>, ToolError>) { ) -> (String, Result<Vec<Content>, ToolError>) {
let mut extension_manager = self.extension_manager.lock().await; let mut extension_manager = self.extension_manager.lock().await;
if action == "disable" {
let result = extension_manager
.remove_extension(&extension_name)
.await
.map(|_| {
vec![Content::text(format!(
"The extension '{}' has been disabled successfully",
extension_name
))]
})
.map_err(|e| ToolError::ExecutionError(e.to_string()));
return (request_id, result);
}
let config = match ExtensionConfigManager::get_config_by_name(&extension_name) { let config = match ExtensionConfigManager::get_config_by_name(&extension_name) {
Ok(Some(config)) => config, Ok(Some(config)) => config,
Ok(None) => { Ok(None) => {
@@ -289,7 +313,7 @@ impl Agent {
if extension_name.is_none() || extension_name.as_deref() == Some("platform") { if extension_name.is_none() || extension_name.as_deref() == Some("platform") {
// Add platform tools // Add platform tools
prefixed_tools.push(platform_tools::search_available_extensions_tool()); prefixed_tools.push(platform_tools::search_available_extensions_tool());
prefixed_tools.push(platform_tools::enable_extension_tool()); prefixed_tools.push(platform_tools::manage_extensions_tool());
// Add resource tools if supported // Add resource tools if supported
if extension_manager.supports_resources() { if extension_manager.supports_resources() {
@@ -616,6 +640,7 @@ impl Agent {
let system_prompt = self.prompt_manager.build_system_prompt( let system_prompt = self.prompt_manager.build_system_prompt(
extensions_info, extensions_info,
self.frontend_instructions.clone(), self.frontend_instructions.clone(),
extension_manager.suggest_disable_extensions_prompt().await,
Some(model_name), Some(model_name),
); );
+51 -3
View File
@@ -288,6 +288,38 @@ impl ExtensionManager {
Ok(()) Ok(())
} }
pub async fn suggest_disable_extensions_prompt(&self) -> Value {
let enabled_extensions_count = self.clients.len();
let total_tools = self
.get_prefixed_tools(None)
.await
.map(|tools| tools.len())
.unwrap_or(0);
// Check if either condition is met
const MIN_EXTENSIONS: usize = 5;
const MIN_TOOLS: usize = 50;
if enabled_extensions_count > MIN_EXTENSIONS || total_tools > MIN_TOOLS {
Value::String(format!(
"The user currently has enabled {} extensions with a total of {} tools. \
Since this exceeds the recommended limits ({} extensions or {} tools), \
you should ask the user if they would like to disable some extensions for this session.\n\n\
Use the search_available_extensions tool to find extensions available to disable. \
You should only disable extensions found from the search_available_extensions tool. \
List all the extensions available to disable in the response. \
Explain that minimizing extensions helps with the recall of the correct tools to use.",
enabled_extensions_count,
total_tools,
MIN_EXTENSIONS,
MIN_TOOLS,
))
} else {
Value::String(String::new()) // Empty string if under limits
}
}
pub async fn list_extensions(&self) -> ExtensionResult<Vec<String>> { pub async fn list_extensions(&self) -> ExtensionResult<Vec<String>> {
Ok(self.clients.keys().cloned().collect()) Ok(self.clients.keys().cloned().collect())
} }
@@ -725,14 +757,30 @@ impl ExtensionManager {
} }
} }
// Get currently enabled extensions that can be disabled
let enabled_extensions: Vec<String> = self.clients.keys().cloned().collect();
// Build output string
if !disabled_extensions.is_empty() { if !disabled_extensions.is_empty() {
output_parts.push(format!( output_parts.push(format!(
"Currently available extensions user can enable:\n{}\n", "Extensions available to enable:\n{}\n",
disabled_extensions.join("\n") disabled_extensions.join("\n")
)); ));
} else { } else {
output_parts output_parts.push("No extensions available to enable.\n".to_string());
.push("No available extensions found in current configuration.\n".to_string()); }
if !enabled_extensions.is_empty() {
output_parts.push(format!(
"\n\nExtensions available to disable:\n{}\n",
enabled_extensions
.iter()
.map(|name| format!("- {}", name))
.collect::<Vec<_>>()
.join("\n")
));
} else {
output_parts.push("No extensions that can be disabled.\n".to_string());
} }
Ok(vec![Content::text(output_parts.join("\n"))]) Ok(vec![Content::text(output_parts.join("\n"))])
+9 -7
View File
@@ -6,7 +6,7 @@ pub const PLATFORM_READ_RESOURCE_TOOL_NAME: &str = "platform__read_resource";
pub const PLATFORM_LIST_RESOURCES_TOOL_NAME: &str = "platform__list_resources"; pub const PLATFORM_LIST_RESOURCES_TOOL_NAME: &str = "platform__list_resources";
pub const PLATFORM_SEARCH_AVAILABLE_EXTENSIONS_TOOL_NAME: &str = pub const PLATFORM_SEARCH_AVAILABLE_EXTENSIONS_TOOL_NAME: &str =
"platform__search_available_extensions"; "platform__search_available_extensions";
pub const PLATFORM_ENABLE_EXTENSION_TOOL_NAME: &str = "platform__enable_extension"; pub const PLATFORM_MANAGE_EXTENSIONS_TOOL_NAME: &str = "platform__manage_extensions";
pub fn read_resource_tool() -> Tool { pub fn read_resource_tool() -> Tool {
Tool::new( Tool::new(
@@ -87,22 +87,24 @@ pub fn search_available_extensions_tool() -> Tool {
) )
} }
pub fn enable_extension_tool() -> Tool { pub fn manage_extensions_tool() -> Tool {
Tool::new( Tool::new(
PLATFORM_ENABLE_EXTENSION_TOOL_NAME.to_string(), PLATFORM_MANAGE_EXTENSIONS_TOOL_NAME.to_string(),
"Enable extensions to help complete tasks. "Tool to manage extensions and tools in goose context.
Enable an extension by providing the extension name. Enable or disable extensions to help complete tasks.
Enable or disable an extension by providing the extension name.
" "
.to_string(), .to_string(),
json!({ json!({
"type": "object", "type": "object",
"required": ["extension_name"], "required": ["action", "extension_name"],
"properties": { "properties": {
"action": {"type": "string", "description": "The action to perform", "enum": ["enable", "disable"]},
"extension_name": {"type": "string", "description": "The name of the extension to enable"} "extension_name": {"type": "string", "description": "The name of the extension to enable"}
} }
}), }),
Some(ToolAnnotations { Some(ToolAnnotations {
title: Some("Enable extensions".to_string()), title: Some("Enable or disable an extension".to_string()),
read_only_hint: false, read_only_hint: false,
destructive_hint: false, destructive_hint: false,
idempotent_hint: false, idempotent_hint: false,
@@ -62,6 +62,7 @@ impl PromptManager {
&self, &self,
extensions_info: Vec<ExtensionInfo>, extensions_info: Vec<ExtensionInfo>,
frontend_instructions: Option<String>, frontend_instructions: Option<String>,
suggest_disable_extensions_prompt: Value,
model_name: Option<&str>, model_name: Option<&str>,
) -> String { ) -> String {
let mut context: HashMap<&str, Value> = HashMap::new(); let mut context: HashMap<&str, Value> = HashMap::new();
@@ -81,6 +82,12 @@ impl PromptManager {
let current_date_time = Utc::now().format("%Y-%m-%d %H:%M:%S").to_string(); let current_date_time = Utc::now().format("%Y-%m-%d %H:%M:%S").to_string();
context.insert("current_date_time", Value::String(current_date_time)); context.insert("current_date_time", Value::String(current_date_time));
// Add the suggestion about disabling extensions if flag is true
context.insert(
"suggest_disable",
Value::String(suggest_disable_extensions_prompt.to_string()),
);
// First check the global store, and only if it's not available, fall back to the provided model_name // First check the global store, and only if it's not available, fall back to the provided model_name
let model_to_use: Option<String> = let model_to_use: Option<String> =
get_current_model().or_else(|| model_name.map(|s| s.to_string())); get_current_model().or_else(|| model_name.map(|s| s.to_string()));
+1
View File
@@ -37,6 +37,7 @@ impl Agent {
let mut system_prompt = self.prompt_manager.build_system_prompt( let mut system_prompt = self.prompt_manager.build_system_prompt(
extensions_info, extensions_info,
self.frontend_instructions.clone(), self.frontend_instructions.clone(),
extension_manager.suggest_disable_extensions_prompt().await,
Some(model_name), Some(model_name),
); );
@@ -1,4 +1,4 @@
use crate::agents::platform_tools::PLATFORM_ENABLE_EXTENSION_TOOL_NAME; use crate::agents::platform_tools::PLATFORM_MANAGE_EXTENSIONS_TOOL_NAME;
use crate::config::permission::PermissionLevel; use crate::config::permission::PermissionLevel;
use crate::config::PermissionManager; use crate::config::PermissionManager;
use crate::message::{Message, MessageContent, ToolRequest}; use crate::message::{Message, MessageContent, ToolRequest};
@@ -170,7 +170,7 @@ pub async fn check_tool_permissions(
let mut needs_approval = vec![]; let mut needs_approval = vec![];
let mut denied = vec![]; let mut denied = vec![];
let mut llm_detect_candidates = vec![]; let mut llm_detect_candidates = vec![];
let mut enable_extension_request_ids = vec![]; let mut extension_request_ids = vec![];
for request in candidate_requests { for request in candidate_requests {
if let Ok(tool_call) = request.tool_call.clone() { if let Ok(tool_call) = request.tool_call.clone() {
@@ -179,8 +179,8 @@ pub async fn check_tool_permissions(
} else if mode == "auto" { } else if mode == "auto" {
approved.push(request.clone()); approved.push(request.clone());
} else { } else {
if tool_call.name == PLATFORM_ENABLE_EXTENSION_TOOL_NAME { if tool_call.name == PLATFORM_MANAGE_EXTENSIONS_TOOL_NAME {
enable_extension_request_ids.push(request.id.clone()); extension_request_ids.push(request.id.clone());
} }
// 1. Check user-defined permission // 1. Check user-defined permission
@@ -255,7 +255,7 @@ pub async fn check_tool_permissions(
needs_approval, needs_approval,
denied, denied,
}, },
enable_extension_request_ids, extension_request_ids,
) )
} }
@@ -430,8 +430,8 @@ mod tests {
let enable_extension = ToolRequest { let enable_extension = ToolRequest {
id: "tool_3".to_string(), id: "tool_3".to_string(),
tool_call: ToolResult::Ok(ToolCall { tool_call: ToolResult::Ok(ToolCall {
name: PLATFORM_ENABLE_EXTENSION_TOOL_NAME.to_string(), name: PLATFORM_MANAGE_EXTENSIONS_TOOL_NAME.to_string(),
arguments: serde_json::json!({"url": "http://example.com"}), arguments: serde_json::json!({"action": "enable", "extension_name": "data_fetcher"}),
}), }),
}; };
+5
View File
@@ -31,6 +31,11 @@ and platform__list_resources on this extension.
No extensions are defined. You should let the user know that they should add extensions. No extensions are defined. You should let the user know that they should add extensions.
{% endif %} {% endif %}
{% if suggest_disable is defined %}
# Suggestion
{{suggest_disable}}
{% endif %}
# Response Guidelines # Response Guidelines
- Use Markdown formatting for all responses. - Use Markdown formatting for all responses.