Fix case-insensitive matching for builtin extension names (#6825)
Co-authored-by: Douwe Osinga <douwe@squareup.com> Co-authored-by: Claude Sonnet 4.5 <noreply@anthropic.com>
This commit is contained in:
@@ -286,7 +286,11 @@ async fn add_builtins(agent: &Agent, builtins: Vec<String>) {
|
|||||||
async fn add_extensions(agent: &Agent, extensions: Vec<ExtensionConfig>) {
|
async fn add_extensions(agent: &Agent, extensions: Vec<ExtensionConfig>) {
|
||||||
for extension in extensions {
|
for extension in extensions {
|
||||||
let name = extension.name().to_string();
|
let name = extension.name().to_string();
|
||||||
match agent.add_extension(extension).await {
|
match agent
|
||||||
|
.extension_manager
|
||||||
|
.add_extension(extension, None, None)
|
||||||
|
.await
|
||||||
|
{
|
||||||
Ok(_) => info!(extension = %name, "extension loaded"),
|
Ok(_) => info!(extension = %name, "extension loaded"),
|
||||||
Err(e) => warn!(extension = %name, error = %e, "extension load failed"),
|
Err(e) => warn!(extension = %name, error = %e, "extension load failed"),
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -35,6 +35,7 @@ use super::types::SharedProvider;
|
|||||||
use crate::agents::extension::{Envs, ProcessExit};
|
use crate::agents::extension::{Envs, ProcessExit};
|
||||||
use crate::agents::extension_malware_check;
|
use crate::agents::extension_malware_check;
|
||||||
use crate::agents::mcp_client::{McpClient, McpClientTrait};
|
use crate::agents::mcp_client::{McpClient, McpClientTrait};
|
||||||
|
use crate::config::extensions::name_to_key;
|
||||||
use crate::config::search_path::SearchPaths;
|
use crate::config::search_path::SearchPaths;
|
||||||
use crate::config::{get_all_extensions, Config};
|
use crate::config::{get_all_extensions, Config};
|
||||||
use crate::oauth::oauth_flow;
|
use crate::oauth::oauth_flow;
|
||||||
@@ -133,20 +134,6 @@ impl ResourceItem {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Sanitizes a string by replacing invalid characters with underscores.
|
|
||||||
/// Valid characters match [a-zA-Z0-9_-]
|
|
||||||
pub fn normalize(input: &str) -> String {
|
|
||||||
let mut result = String::with_capacity(input.len());
|
|
||||||
for c in input.chars() {
|
|
||||||
result.push(match c {
|
|
||||||
c if c.is_ascii_alphanumeric() || c == '_' || c == '-' => c,
|
|
||||||
c if c.is_whitespace() => continue, // effectively "strip" whitespace
|
|
||||||
_ => '_', // Replace any other non-ASCII character with '_'
|
|
||||||
});
|
|
||||||
}
|
|
||||||
result.to_lowercase()
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Generates extension name from server info; adds random suffix on collision.
|
/// Generates extension name from server info; adds random suffix on collision.
|
||||||
fn generate_extension_name(
|
fn generate_extension_name(
|
||||||
server_info: Option<&ServerInfo>,
|
server_info: Option<&ServerInfo>,
|
||||||
@@ -155,7 +142,7 @@ fn generate_extension_name(
|
|||||||
let base = server_info
|
let base = server_info
|
||||||
.and_then(|info| {
|
.and_then(|info| {
|
||||||
let name = info.server_info.name.as_str();
|
let name = info.server_info.name.as_str();
|
||||||
(!name.is_empty()).then(|| normalize(name))
|
(!name.is_empty()).then(|| name_to_key(name))
|
||||||
})
|
})
|
||||||
.unwrap_or_else(|| "unnamed".to_string());
|
.unwrap_or_else(|| "unnamed".to_string());
|
||||||
|
|
||||||
@@ -491,7 +478,7 @@ impl ExtensionManager {
|
|||||||
container: Option<&Container>,
|
container: Option<&Container>,
|
||||||
) -> ExtensionResult<()> {
|
) -> ExtensionResult<()> {
|
||||||
let config_name = config.key().to_string();
|
let config_name = config.key().to_string();
|
||||||
let sanitized_name = normalize(&config_name);
|
let sanitized_name = name_to_key(&config_name);
|
||||||
|
|
||||||
if self.extensions.lock().await.contains_key(&sanitized_name) {
|
if self.extensions.lock().await.contains_key(&sanitized_name) {
|
||||||
return Ok(());
|
return Ok(());
|
||||||
@@ -577,8 +564,9 @@ impl ExtensionManager {
|
|||||||
}
|
}
|
||||||
ExtensionConfig::Builtin { name, timeout, .. } => {
|
ExtensionConfig::Builtin { name, timeout, .. } => {
|
||||||
let timeout_duration = Duration::from_secs(timeout.unwrap_or(300));
|
let timeout_duration = Duration::from_secs(timeout.unwrap_or(300));
|
||||||
|
let normalized_name = name_to_key(name);
|
||||||
|
|
||||||
if !goose_mcp::BUILTIN_EXTENSIONS.contains_key(name.as_str()) {
|
if !goose_mcp::BUILTIN_EXTENSIONS.contains_key(normalized_name.as_str()) {
|
||||||
return Err(ExtensionError::ConfigError(format!(
|
return Err(ExtensionError::ConfigError(format!(
|
||||||
"Unknown builtin extension: {}",
|
"Unknown builtin extension: {}",
|
||||||
name
|
name
|
||||||
@@ -599,7 +587,7 @@ impl ExtensionManager {
|
|||||||
.arg(container_id)
|
.arg(container_id)
|
||||||
.arg("goose")
|
.arg("goose")
|
||||||
.arg("mcp")
|
.arg("mcp")
|
||||||
.arg(name);
|
.arg(&normalized_name);
|
||||||
});
|
});
|
||||||
|
|
||||||
let client = child_process_client(
|
let client = child_process_client(
|
||||||
@@ -612,7 +600,9 @@ impl ExtensionManager {
|
|||||||
.await?;
|
.await?;
|
||||||
Box::new(client)
|
Box::new(client)
|
||||||
} else {
|
} else {
|
||||||
let def = goose_mcp::BUILTIN_EXTENSIONS.get(name.as_str()).unwrap();
|
let def = goose_mcp::BUILTIN_EXTENSIONS
|
||||||
|
.get(normalized_name.as_str())
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
// Set GOOSE_WORKING_DIR in the current process for builtin extensions
|
// Set GOOSE_WORKING_DIR in the current process for builtin extensions
|
||||||
// since they run in-process and read from std::env::var
|
// since they run in-process and read from std::env::var
|
||||||
@@ -638,7 +628,7 @@ impl ExtensionManager {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
ExtensionConfig::Platform { name, .. } => {
|
ExtensionConfig::Platform { name, .. } => {
|
||||||
let normalized_key = normalize(name);
|
let normalized_key = name_to_key(name);
|
||||||
let def = PLATFORM_EXTENSIONS
|
let def = PLATFORM_EXTENSIONS
|
||||||
.get(normalized_key.as_str())
|
.get(normalized_key.as_str())
|
||||||
.ok_or_else(|| {
|
.ok_or_else(|| {
|
||||||
@@ -713,7 +703,7 @@ impl ExtensionManager {
|
|||||||
info: Option<ServerInfo>,
|
info: Option<ServerInfo>,
|
||||||
temp_dir: Option<TempDir>,
|
temp_dir: Option<TempDir>,
|
||||||
) {
|
) {
|
||||||
let normalized = normalize(&name);
|
let normalized = name_to_key(&name);
|
||||||
self.extensions
|
self.extensions
|
||||||
.lock()
|
.lock()
|
||||||
.await
|
.await
|
||||||
@@ -739,7 +729,7 @@ impl ExtensionManager {
|
|||||||
|
|
||||||
/// Get aggregated usage statistics
|
/// Get aggregated usage statistics
|
||||||
pub async fn remove_extension(&self, name: &str) -> ExtensionResult<()> {
|
pub async fn remove_extension(&self, name: &str) -> ExtensionResult<()> {
|
||||||
let sanitized_name = normalize(name);
|
let sanitized_name = name_to_key(name);
|
||||||
self.extensions.lock().await.remove(&sanitized_name);
|
self.extensions.lock().await.remove(&sanitized_name);
|
||||||
self.invalidate_tools_cache_and_bump_version().await;
|
self.invalidate_tools_cache_and_bump_version().await;
|
||||||
Ok(())
|
Ok(())
|
||||||
@@ -762,7 +752,7 @@ impl ExtensionManager {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub async fn is_extension_enabled(&self, name: &str) -> bool {
|
pub async fn is_extension_enabled(&self, name: &str) -> bool {
|
||||||
let normalized = normalize(name);
|
let normalized = name_to_key(name);
|
||||||
self.extensions.lock().await.contains_key(&normalized)
|
self.extensions.lock().await.contains_key(&normalized)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -800,8 +790,8 @@ impl ExtensionManager {
|
|||||||
extension_name: Option<&str>,
|
extension_name: Option<&str>,
|
||||||
exclude: Option<&str>,
|
exclude: Option<&str>,
|
||||||
) -> Vec<Tool> {
|
) -> Vec<Tool> {
|
||||||
let extension_name_normalized = extension_name.map(normalize);
|
let extension_name_normalized = extension_name.map(name_to_key);
|
||||||
let exclude_normalized = exclude.map(normalize);
|
let exclude_normalized = exclude.map(name_to_key);
|
||||||
|
|
||||||
tools
|
tools
|
||||||
.iter()
|
.iter()
|
||||||
@@ -1449,7 +1439,7 @@ impl ExtensionManager {
|
|||||||
}
|
}
|
||||||
|
|
||||||
async fn get_server_client(&self, name: impl Into<String>) -> Option<McpClientBox> {
|
async fn get_server_client(&self, name: impl Into<String>) -> Option<McpClientBox> {
|
||||||
let normalized = normalize(&name.into());
|
let normalized = name_to_key(&name.into());
|
||||||
self.extensions
|
self.extensions
|
||||||
.lock()
|
.lock()
|
||||||
.await
|
.await
|
||||||
@@ -1526,7 +1516,7 @@ mod tests {
|
|||||||
client: McpClientBox,
|
client: McpClientBox,
|
||||||
available_tools: Vec<String>,
|
available_tools: Vec<String>,
|
||||||
) {
|
) {
|
||||||
let sanitized_name = normalize(&name);
|
let sanitized_name = name_to_key(&name);
|
||||||
let config = ExtensionConfig::Builtin {
|
let config = ExtensionConfig::Builtin {
|
||||||
name: name.clone(),
|
name: name.clone(),
|
||||||
display_name: Some(name.clone()),
|
display_name: Some(name.clone()),
|
||||||
|
|||||||
@@ -31,7 +31,7 @@ pub use agent::{Agent, AgentConfig, AgentEvent, ExtensionLoadResult};
|
|||||||
pub use container::Container;
|
pub use container::Container;
|
||||||
pub use execute_commands::COMPACT_TRIGGERS;
|
pub use execute_commands::COMPACT_TRIGGERS;
|
||||||
pub use extension::ExtensionConfig;
|
pub use extension::ExtensionConfig;
|
||||||
pub use extension_manager::{normalize, ExtensionManager};
|
pub use extension_manager::ExtensionManager;
|
||||||
pub use prompt_manager::PromptManager;
|
pub use prompt_manager::PromptManager;
|
||||||
pub use subagent_task_config::TaskConfig;
|
pub use subagent_task_config::TaskConfig;
|
||||||
pub use types::{FrontendTool, RetryConfig, SessionConfig, SuccessCheck};
|
pub use types::{FrontendTool, RetryConfig, SessionConfig, SuccessCheck};
|
||||||
|
|||||||
@@ -20,10 +20,15 @@ pub struct ExtensionEntry {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn name_to_key(name: &str) -> String {
|
pub fn name_to_key(name: &str) -> String {
|
||||||
name.chars()
|
let mut result = String::with_capacity(name.len());
|
||||||
.filter(|c| !c.is_whitespace())
|
for c in name.chars() {
|
||||||
.collect::<String>()
|
result.push(match c {
|
||||||
.to_lowercase()
|
c if c.is_ascii_alphanumeric() || c == '_' || c == '-' => c,
|
||||||
|
c if c.is_whitespace() => continue,
|
||||||
|
_ => '_',
|
||||||
|
});
|
||||||
|
}
|
||||||
|
result.to_lowercase()
|
||||||
}
|
}
|
||||||
|
|
||||||
fn get_extensions_map_with_config(config: &Config) -> IndexMap<String, ExtensionEntry> {
|
fn get_extensions_map_with_config(config: &Config) -> IndexMap<String, ExtensionEntry> {
|
||||||
|
|||||||
Reference in New Issue
Block a user