Merge platform/builtin extensions (#7630)

This commit is contained in:
Jack Amadeo
2026-03-04 12:29:45 -05:00
committed by GitHub
parent 7dca516fbf
commit 3152191418
5 changed files with 143 additions and 165 deletions
+93 -83
View File
@@ -197,12 +197,13 @@ pub fn get_tool_owner(tool: &Tool) -> Option<String> {
}
fn is_unprefixed_extension(config: &ExtensionConfig) -> bool {
if let ExtensionConfig::Platform { name, .. } = config {
PLATFORM_EXTENSIONS
.get(name_to_key(name).as_str())
.is_some_and(|def| def.unprefixed_tools)
} else {
false
match config {
ExtensionConfig::Platform { name, .. } | ExtensionConfig::Builtin { name, .. } => {
PLATFORM_EXTENSIONS
.get(name_to_key(name).as_str())
.is_some_and(|def| def.unprefixed_tools)
}
_ => false,
}
}
@@ -579,72 +580,91 @@ impl ExtensionManager {
)
.await?
}
ExtensionConfig::Builtin { name, timeout, .. } => {
let timeout_duration = Duration::from_secs(timeout.unwrap_or(300));
let normalized_name = name_to_key(name);
let extension_fn =
get_builtin_extension(normalized_name.as_str()).ok_or_else(|| {
ExtensionError::ConfigError(format!("Unknown builtin extension: {}", name))
})?;
if let Some(container) = container {
let container_id = container.id();
tracing::info!(
container = %container_id,
builtin = %name,
"Starting builtin extension inside Docker container"
);
let normalized_name = name_to_key(name);
let command = Command::new("docker").configure(|command| {
command
.arg("exec")
.arg("-i")
.arg(container_id)
.arg("goose")
.arg("mcp")
.arg(&normalized_name);
});
let effective_working_dir = working_dir
.clone()
.unwrap_or_else(|| std::env::current_dir().unwrap_or_default());
let capabilities = GooseMcpClientCapabilities {
mcpui: self.capabilities.mcpui,
};
let client = child_process_client(
command,
timeout,
self.provider.clone(),
Some(&effective_working_dir),
Some(container_id.to_string()),
self.client_name.clone(),
capabilities,
)
.await?;
Box::new(client)
ExtensionConfig::Builtin { ref name, .. }
| ExtensionConfig::Platform { ref name, .. } => {
let timeout = if let ExtensionConfig::Builtin { timeout, .. } = &config {
*timeout
} else {
// Non-containerized builtin runs in-process via duplex channels.
// Working directory is passed per-request via call_tool metadata, not here.
let (server_read, client_write) = tokio::io::duplex(65536);
let (client_read, server_write) = tokio::io::duplex(65536);
extension_fn(server_read, server_write);
None
};
let normalized_name = name_to_key(name);
let capabilities = GooseMcpClientCapabilities {
mcpui: self.capabilities.mcpui,
};
if let Some(def) = PLATFORM_EXTENSIONS.get(normalized_name.as_str()) {
// Platform extension: create via in-process client factory
let mut context = self.context.clone();
context.extension_manager = Some(Arc::downgrade(self));
if let Some(id) = session_id {
if let Ok(session) =
self.context.session_manager.get_session(id, false).await
{
context.session = Some(Arc::new(session));
}
}
(def.client_factory)(context)
} else {
// Builtin MCP server extension
let timeout_secs = timeout.unwrap_or(crate::config::DEFAULT_EXTENSION_TIMEOUT);
let extension_fn =
get_builtin_extension(normalized_name.as_str()).ok_or_else(|| {
ExtensionError::ConfigError(format!("Unknown extension: {}", name))
})?;
Box::new(
McpClient::connect(
(client_read, client_write),
timeout_duration,
if let Some(container) = container {
let container_id = container.id();
tracing::info!(
container = %container_id,
builtin = %name,
"Starting builtin extension inside Docker container"
);
let command = Command::new("docker").configure(|command| {
command
.arg("exec")
.arg("-i")
.arg(container_id)
.arg("goose")
.arg("mcp")
.arg(&normalized_name);
});
let effective_working_dir = working_dir
.clone()
.unwrap_or_else(|| std::env::current_dir().unwrap_or_default());
let capabilities = GooseMcpClientCapabilities {
mcpui: self.capabilities.mcpui,
};
let client = child_process_client(
command,
&Some(timeout_secs),
self.provider.clone(),
Some(&effective_working_dir),
Some(container_id.to_string()),
self.client_name.clone(),
capabilities,
)
.await?,
)
.await?;
Box::new(client)
} else {
let (server_read, client_write) = tokio::io::duplex(65536);
let (client_read, server_write) = tokio::io::duplex(65536);
extension_fn(server_read, server_write);
let capabilities = GooseMcpClientCapabilities {
mcpui: self.capabilities.mcpui,
};
Box::new(
McpClient::connect(
(client_read, client_write),
Duration::from_secs(timeout_secs),
self.provider.clone(),
self.client_name.clone(),
capabilities,
)
.await?,
)
}
}
}
ExtensionConfig::Stdio {
@@ -707,23 +727,6 @@ impl ExtensionManager {
.await?;
Box::new(client)
}
ExtensionConfig::Platform { name, .. } => {
let normalized_key = name_to_key(name);
let def = PLATFORM_EXTENSIONS
.get(normalized_key.as_str())
.ok_or_else(|| {
ExtensionError::ConfigError(format!("Unknown platform extension: {}", name))
})?;
let mut context = self.context.clone();
context.extension_manager = Some(Arc::downgrade(self));
if let Some(id) = session_id {
if let Ok(session) = self.context.session_manager.get_session(id, false).await {
context.session = Some(Arc::new(session));
}
}
(def.client_factory)(context)
}
ExtensionConfig::InlinePython {
name,
code,
@@ -1640,7 +1643,14 @@ impl ExtensionManager {
extensions
.iter()
.filter_map(|(name, extension)| {
if let ExtensionConfig::Platform { .. } = &extension.config {
let is_platform = match &extension.config {
ExtensionConfig::Platform { .. } => true,
ExtensionConfig::Builtin { name: ext_name, .. } => {
PLATFORM_EXTENSIONS.contains_key(name_to_key(ext_name).as_str())
}
_ => false,
};
if is_platform {
Some((name.clone(), extension.get_client()))
} else {
None
+10 -10
View File
@@ -321,7 +321,7 @@ impl Config {
// Run migrations on the loaded config
if crate::config::migrations::run_migrations(&mut values) {
if let Err(e) = self.save_values(values.clone()) {
if let Err(e) = self.save_values(&values) {
tracing::warn!("Failed to save migrated config: {}", e);
}
}
@@ -352,7 +352,7 @@ impl Config {
default_config: Mapping,
) -> Result<Mapping, ConfigError> {
// Try to write the default config to disk
match self.save_values(default_config.clone()) {
match self.save_values(&default_config) {
Ok(_) => {
if default_config.is_empty() {
tracing::info!("Created fresh empty config file");
@@ -407,7 +407,7 @@ impl Config {
match parse_yaml_content(&backup_content) {
Ok(values) => {
// Successfully parsed backup, restore it as the main config
if let Err(e) = self.save_values(values.clone()) {
if let Err(e) = self.save_values(&values) {
tracing::warn!(
"Failed to restore backup as main config: {}",
e
@@ -468,12 +468,12 @@ impl Config {
load_init_config_from_workspace()
}
fn save_values(&self, values: Mapping) -> Result<(), ConfigError> {
fn save_values(&self, values: &Mapping) -> Result<(), ConfigError> {
// Create backup before writing new config
self.create_backup_if_needed()?;
// Convert to YAML for storage
let yaml_value = serde_yaml::to_string(&values)?;
let yaml_value = serde_yaml::to_string(values)?;
if let Some(parent) = self.config_path.parent() {
std::fs::create_dir_all(parent)
@@ -510,7 +510,7 @@ impl Config {
pub fn initialize_if_empty(&self, values: Mapping) -> Result<(), ConfigError> {
let _guard = self.guard.lock().unwrap();
if !self.exists() {
self.save_values(values)
self.save_values(&values)
} else {
Ok(())
}
@@ -743,7 +743,7 @@ impl Config {
let _guard = self.guard.lock().unwrap();
let mut values = self.load_raw()?;
values.insert(serde_yaml::to_value(key)?, serde_yaml::to_value(value)?);
self.save_values(values)
self.save_values(&values)
}
/// Delete a configuration value in the config file.
@@ -766,7 +766,7 @@ impl Config {
let mut values = self.load_raw()?;
values.shift_remove(key);
self.save_values(values)
self.save_values(&values)
}
/// Get a secret value.
@@ -1240,7 +1240,7 @@ mod tests {
let mut handles = vec![];
// Initialize with empty values
config.save_values(Default::default())?;
config.save_values(&Default::default())?;
// Spawn 3 threads that will try to write simultaneously
for i in 0..3 {
@@ -1259,7 +1259,7 @@ mod tests {
);
// Write all values
config.save_values(values.clone())?;
config.save_values(&values)?;
Ok(())
});
handles.push(handle);
+2 -43
View File
@@ -41,40 +41,6 @@ pub(crate) fn is_extension_available(config: &ExtensionConfig) -> bool {
}
}
pub(crate) fn normalize_platform_extension(config: ExtensionConfig) -> ExtensionConfig {
match config {
ExtensionConfig::Builtin {
name,
description,
display_name,
timeout,
bundled,
available_tools,
} => {
let normalized = name_to_key(&name);
if let Some(def) = PLATFORM_EXTENSIONS.get(normalized.as_str()) {
ExtensionConfig::Platform {
name: def.name.to_string(),
description: def.description.to_string(),
display_name: Some(def.display_name.to_string()),
bundled: bundled.or(Some(true)),
available_tools,
}
} else {
ExtensionConfig::Builtin {
name,
description,
display_name,
timeout,
bundled,
available_tools,
}
}
}
other => other,
}
}
fn get_extensions_map_with_config(config: &Config) -> IndexMap<String, ExtensionEntry> {
let raw: Mapping = config
.get_param(EXTENSIONS_CONFIG_KEY)
@@ -90,17 +56,10 @@ fn get_extensions_map_with_config(config: &Config) -> IndexMap<String, Extension
for (k, v) in raw {
match (k, serde_yaml::from_value::<ExtensionEntry>(v)) {
(serde_yaml::Value::String(key), Ok(entry)) => {
let config = normalize_platform_extension(entry.config);
if !is_extension_available(&config) {
if !is_extension_available(&entry.config) {
continue;
}
extensions_map.insert(
key,
ExtensionEntry {
enabled: entry.enabled,
config,
},
);
extensions_map.insert(key, entry);
}
(k, v) => {
warn!(
+34 -13
View File
@@ -33,40 +33,61 @@ fn migrate_platform_extensions(config: &mut Mapping) -> bool {
let needs_migration = match existing {
None => true,
Some(value) => match serde_yaml::from_value::<ExtensionEntry>(value.clone()) {
Ok(entry) => {
if let ExtensionConfig::Platform {
Ok(entry) => match &entry.config {
ExtensionConfig::Platform {
description,
display_name,
..
} = &entry.config
{
}
| ExtensionConfig::Builtin {
description,
display_name,
..
} => {
description != def.description
|| display_name.as_deref() != Some(def.display_name)
} else {
true
}
}
_ => true,
},
Err(_) => true,
},
};
if needs_migration {
let enabled = existing
.and_then(|v| serde_yaml::from_value::<ExtensionEntry>(v.clone()).ok())
let existing_entry =
existing.and_then(|v| serde_yaml::from_value::<ExtensionEntry>(v.clone()).ok());
let enabled = existing_entry
.as_ref()
.map(|e| e.enabled)
.unwrap_or(def.default_enabled);
let new_entry = ExtensionEntry {
config: ExtensionConfig::Platform {
// If the extension already exists as type 'builtin', preserve that type
let is_existing_builtin = existing_entry
.as_ref()
.is_some_and(|e| matches!(e.config, ExtensionConfig::Builtin { .. }));
let config = if is_existing_builtin {
ExtensionConfig::Builtin {
name: def.name.to_string(),
description: def.description.to_string(),
display_name: Some(def.display_name.to_string()),
timeout: None,
bundled: Some(true),
available_tools: Vec::new(),
}
} else {
ExtensionConfig::Platform {
name: def.name.to_string(),
description: def.description.to_string(),
display_name: Some(def.display_name.to_string()),
bundled: Some(true),
available_tools: Vec::new(),
},
enabled,
}
};
let new_entry = ExtensionEntry { config, enabled };
if let Ok(value) = serde_yaml::to_value(&new_entry) {
extensions_map.insert(ext_key, value);
needs_save = true;
+4 -16
View File
@@ -2,7 +2,7 @@
// Provides a simple way to store extension-specific data with versioned keys
use crate::config::base::Config;
use crate::config::extensions::{is_extension_available, normalize_platform_extension};
use crate::config::extensions::is_extension_available;
use crate::config::ExtensionConfig;
use crate::session::SessionManager;
use anyhow::Result;
@@ -117,11 +117,6 @@ impl EnabledExtensionsState {
pub fn from_extension_data(extension_data: &ExtensionData) -> Option<Self> {
let mut state = <Self as ExtensionState>::from_extension_data(extension_data)?;
state.extensions = state
.extensions
.into_iter()
.map(normalize_platform_extension)
.collect();
state.extensions.retain(is_extension_available);
Some(state)
}
@@ -161,7 +156,7 @@ mod tests {
Config::new_with_file_secrets(config_file.path(), secrets_file.path()).unwrap()
}
fn legacy_test_extension() -> ExtensionConfig {
fn test_extension() -> ExtensionConfig {
ExtensionConfig::Builtin {
name: "developer".into(),
description: "dev".into(),
@@ -172,10 +167,6 @@ mod tests {
}
}
fn normalized_test_extension() -> ExtensionConfig {
normalize_platform_extension(legacy_test_extension())
}
fn extension_data_with(extensions: Vec<ExtensionConfig>) -> ExtensionData {
let mut data = ExtensionData::new();
EnabledExtensionsState::new(extensions)
@@ -185,8 +176,8 @@ mod tests {
}
#[test_case(
Some(extension_data_with(vec![legacy_test_extension()])),
Some(vec![normalized_test_extension()])
Some(extension_data_with(vec![test_extension()])),
Some(vec![test_extension()])
; "prefers_session_data"
)]
#[test_case(None, None ; "no_session_falls_back_to_config")]
@@ -304,9 +295,6 @@ mod tests {
let names: Vec<String> = loaded.extensions.iter().map(|ext| ext.name()).collect();
assert!(names.iter().any(|name| name == "developer"));
assert!(loaded.extensions.iter().any(
|ext| matches!(ext, ExtensionConfig::Platform { name, .. } if name == "developer")
));
assert!(!names
.iter()
.any(|name| name == "definitely_not_real_platform_extension"));