From c6a3b5b4e17d382305fdb0e1cc2da6904784473f Mon Sep 17 00:00:00 2001 From: Matt Van Horn Date: Wed, 10 Jun 2026 22:10:06 -0700 Subject: [PATCH] fix: preserve unparseable extension entries during config refresh (#9439) Co-authored-by: Michael Neale --- crates/goose/src/config/extensions.rs | 430 +++++++++++++++++++++++--- 1 file changed, 391 insertions(+), 39 deletions(-) diff --git a/crates/goose/src/config/extensions.rs b/crates/goose/src/config/extensions.rs index 237ee8b68..f36419972 100644 --- a/crates/goose/src/config/extensions.rs +++ b/crates/goose/src/config/extensions.rs @@ -4,7 +4,7 @@ use crate::agents::ExtensionConfig; use indexmap::IndexMap; use serde::{Deserialize, Serialize}; use serde_yaml::Mapping; -use tracing::warn; +use tracing::{info, warn}; use utoipa::ToSchema; pub const DEFAULT_EXTENSION: &str = "developer"; @@ -35,12 +35,40 @@ pub fn name_to_key(name: &str) -> String { pub(crate) fn is_extension_available(config: &ExtensionConfig) -> bool { match config { ExtensionConfig::Platform { name, .. } => { - PLATFORM_EXTENSIONS.contains_key(name_to_key(name).as_str()) + crate::agents::extension::PLATFORM_EXTENSIONS.contains_key(name_to_key(name).as_str()) } _ => true, } } +fn parse_extensions_map(raw: &Mapping) -> IndexMap { + let mut extensions_map = IndexMap::with_capacity(raw.len()); + for (k, v) in raw { + let Some(key) = k.as_str() else { + warn!(key = ?k, "Skipping malformed extension config entry"); + continue; + }; + + match serde_yaml::from_value::(v.clone()) { + Ok(entry) => { + if !is_extension_available(&entry.config) { + continue; + } + extensions_map.insert(key.to_string(), entry); + } + Err(err) => { + info!( + key = %key, + error = %err, + "Skipping malformed extension config entry" + ); + } + } + } + + extensions_map +} + fn get_extensions_map_with_config(config: &Config) -> IndexMap { let raw: Mapping = config .get_param(EXTENSIONS_CONFIG_KEY) @@ -52,36 +80,49 @@ fn get_extensions_map_with_config(config: &Config) -> IndexMap(v)) { - (serde_yaml::Value::String(key), Ok(entry)) => { - if !is_extension_available(&entry.config) { - continue; - } - extensions_map.insert(key, entry); - } - (k, v) => { - warn!( - key = ?k, - value = ?v, - "Skipping malformed extension config entry" - ); - } - } - } - - extensions_map + parse_extensions_map(&raw) } fn get_extensions_map() -> IndexMap { get_extensions_map_with_config(Config::global()) } -fn save_extensions_map(extensions: IndexMap) { - let config = Config::global(); - if let Err(e) = config.set_param(EXTENSIONS_CONFIG_KEY, &extensions) { - tracing::warn!("Failed to save extensions config: {}", e); +enum ExtensionMutation { + Upsert(String, Box), + Remove(String), + Noop, +} + +fn with_raw_extensions_mapping(config: &Config, mutate: F) +where + F: FnOnce(&mut IndexMap) -> ExtensionMutation, +{ + let mut serialize_error = None; + let result = config.update_param::(EXTENSIONS_CONFIG_KEY, |mut raw| { + let mut extensions = parse_extensions_map(&raw); + + match mutate(&mut extensions) { + ExtensionMutation::Upsert(key, entry) => match serde_yaml::to_value(entry) { + Ok(value) => { + raw.insert(serde_yaml::Value::String(key), value); + } + Err(err) => { + serialize_error = Some(err); + } + }, + ExtensionMutation::Remove(key) => { + raw.shift_remove(key.as_str()); + } + ExtensionMutation::Noop => {} + } + + raw + }); + + if let Some(e) = serialize_error { + warn!("Failed to serialize extensions config entry: {}", e); + } else if let Err(e) = result { + warn!("Failed to save extensions config: {}", e); } } @@ -94,28 +135,40 @@ pub fn get_extension_by_name(name: &str) -> Option { } pub fn set_extension(entry: ExtensionEntry) { - let mut extensions = get_extensions_map(); + set_extension_with_config(Config::global(), entry); +} + +fn set_extension_with_config(config: &Config, entry: ExtensionEntry) { let key = entry.config.key(); - extensions.insert(key, entry); - save_extensions_map(extensions); + with_raw_extensions_mapping(config, |_| ExtensionMutation::Upsert(key, Box::new(entry))); } pub fn remove_extension(key: &str) { - let mut extensions = get_extensions_map(); - extensions.shift_remove(key); - save_extensions_map(extensions); + remove_extension_with_config(Config::global(), key); +} + +fn remove_extension_with_config(config: &Config, key: &str) { + with_raw_extensions_mapping(config, |_| ExtensionMutation::Remove(key.to_string())); } /// Returns true when an existing extension was updated, false when the key was missing. pub fn set_extension_enabled(key: &str, enabled: bool) -> bool { - let mut extensions = get_extensions_map(); - let Some(entry) = extensions.get_mut(key) else { - return false; - }; + set_extension_enabled_with_config(Config::global(), key, enabled) +} - entry.enabled = enabled; - save_extensions_map(extensions); - true +fn set_extension_enabled_with_config(config: &Config, key: &str, enabled: bool) -> bool { + let mut updated = false; + with_raw_extensions_mapping(config, |extensions| { + let Some(entry) = extensions.get_mut(key) else { + return ExtensionMutation::Noop; + }; + + entry.enabled = enabled; + updated = true; + ExtensionMutation::Upsert(key.to_string(), Box::new(entry.clone())) + }); + + updated } pub fn get_all_extensions() -> Vec { @@ -225,6 +278,45 @@ pub fn resolve_extensions_for_new_session( #[cfg(test)] mod tests { use super::*; + use std::fmt; + use std::sync::{Arc, Mutex}; + use tempfile::NamedTempFile; + use tracing::{Event, Level, Subscriber}; + use tracing_subscriber::layer::SubscriberExt; + + fn test_config(content: &str) -> (Config, NamedTempFile, NamedTempFile) { + let config_file = NamedTempFile::new().unwrap(); + let secrets_file = NamedTempFile::new().unwrap(); + std::fs::write(config_file.path(), content).unwrap(); + let config = + Config::new_with_file_secrets(config_file.path(), secrets_file.path()).unwrap(); + (config, config_file, secrets_file) + } + + fn read_extensions(config: &Config) -> Mapping { + let content = std::fs::read_to_string(config.path()).unwrap(); + let values: Mapping = serde_yaml::from_str(&content).unwrap(); + values + .get(EXTENSIONS_CONFIG_KEY) + .unwrap() + .as_mapping() + .unwrap() + .clone() + } + + fn builtin_entry(name: &str, enabled: bool) -> ExtensionEntry { + ExtensionEntry { + enabled, + config: ExtensionConfig::Builtin { + name: name.to_string(), + description: format!("{name} description"), + display_name: Some(name.to_string()), + timeout: None, + bundled: None, + available_tools: Vec::new(), + }, + } + } #[test] fn test_is_extension_available_filters_unknown_platform() { @@ -248,4 +340,264 @@ mod tests { assert!(!is_extension_available(&unknown_platform)); assert!(is_extension_available(&builtin)); } + + #[test] + fn test_set_extension_enabled_preserves_clean_siblings() { + let (config, _config_file, _secrets_file) = test_config( + r#" +extensions: + first: + enabled: true + type: builtin + name: first + description: first description + display_name: First + second: + enabled: true + type: builtin + name: second + description: second description + display_name: Second + extra_field: preserved +"#, + ); + let before = read_extensions(&config); + let second_before = before.get("second").unwrap().clone(); + + set_extension_enabled_with_config(&config, "first", false); + + let extensions = read_extensions(&config); + assert_eq!( + extensions + .get("first") + .unwrap() + .as_mapping() + .unwrap() + .get("enabled") + .unwrap() + .as_bool(), + Some(false) + ); + assert_eq!(extensions.get("second").unwrap(), &second_before); + } + + #[test] + fn test_set_extension_enabled_preserves_unparseable_sibling() { + let (config, _config_file, _secrets_file) = test_config( + r#" +extensions: + valid: + enabled: true + type: builtin + name: valid + description: valid description + display_name: Valid + broken: + enabled: true + type: stdio + name: Broken + description: missing cmd + args: [] +"#, + ); + let before = read_extensions(&config); + let broken_before = before.get("broken").unwrap().clone(); + + set_extension_enabled_with_config(&config, "valid", false); + + let extensions = read_extensions(&config); + assert!(extensions.contains_key("valid")); + assert_eq!(extensions.get("broken").unwrap(), &broken_before); + assert_eq!( + extensions + .get("valid") + .unwrap() + .as_mapping() + .unwrap() + .get("enabled") + .unwrap() + .as_bool(), + Some(false) + ); + } + + #[test] + fn test_set_extension_adds_entry_without_dropping_unparseable_entries() { + let (config, _config_file, _secrets_file) = test_config( + r#" +extensions: + broken: + enabled: true + type: stdio + name: Broken + description: missing cmd + args: [] +"#, + ); + let before = read_extensions(&config); + let broken_before = before.get("broken").unwrap().clone(); + + set_extension_with_config(&config, builtin_entry("new extension", true)); + + let extensions = read_extensions(&config); + assert_eq!(extensions.get("broken").unwrap(), &broken_before); + assert!(extensions.contains_key("newextension")); + } + + #[test] + fn test_remove_extension_preserves_unparseable_sibling() { + let (config, _config_file, _secrets_file) = test_config( + r#" +extensions: + valid: + enabled: true + type: builtin + name: valid + description: valid description + display_name: Valid + broken: + enabled: true + type: stdio + name: Broken + description: missing cmd + args: [] +"#, + ); + let before = read_extensions(&config); + let broken_before = before.get("broken").unwrap().clone(); + + remove_extension_with_config(&config, "valid"); + + let extensions = read_extensions(&config); + assert!(!extensions.contains_key("valid")); + assert_eq!(extensions.get("broken").unwrap(), &broken_before); + } + + #[derive(Clone, Default)] + struct CapturedLogs { + events: Arc>>, + } + + #[derive(Debug)] + struct CapturedEvent { + level: Level, + message: String, + key: Option, + } + + impl tracing_subscriber::Layer for CapturedLogs + where + S: Subscriber, + { + fn on_event(&self, event: &Event<'_>, _ctx: tracing_subscriber::layer::Context<'_, S>) { + let mut visitor = EventVisitor::default(); + event.record(&mut visitor); + self.events.lock().unwrap().push(CapturedEvent { + level: *event.metadata().level(), + message: visitor.message, + key: visitor.key, + }); + } + } + + #[derive(Default)] + struct EventVisitor { + message: String, + key: Option, + } + + impl tracing::field::Visit for EventVisitor { + fn record_str(&mut self, field: &tracing::field::Field, value: &str) { + match field.name() { + "message" => self.message = value.to_string(), + "key" => self.key = Some(value.to_string()), + _ => {} + } + } + + fn record_debug(&mut self, field: &tracing::field::Field, value: &dyn fmt::Debug) { + match field.name() { + "message" => self.message = format!("{value:?}").trim_matches('"').to_string(), + "key" => { + self.key = Some(format!("{value:?}").trim_matches('"').to_string()); + } + _ => {} + } + } + } + + #[test] + fn test_deserialization_failure_logs_offending_key() { + let (config, _config_file, _secrets_file) = test_config( + r#" +extensions: + valid: + enabled: true + type: builtin + name: valid + description: valid description + display_name: Valid + broken: + enabled: true + type: stdio + name: Broken + description: missing cmd + args: [] +"#, + ); + let logs = CapturedLogs::default(); + let subscriber = tracing_subscriber::registry().with(logs.clone()); + + tracing::subscriber::with_default(subscriber, || { + let extensions = get_enabled_extensions_with_config(&config); + // Bundled platform extensions are auto-injected; filter to user-declared entries + // (Builtin or anything with the test YAML's names) for the invariant check. + let user_names: Vec<&str> = extensions + .iter() + .filter_map(|ext| match ext { + ExtensionConfig::Builtin { name, .. } => Some(name.as_str()), + _ => None, + }) + .collect(); + assert_eq!( + user_names, + vec!["valid"], + "expected only the parseable user extension to be enabled, got {:?}", + user_names + ); + }); + + let matching_events: Vec<_> = logs + .events + .lock() + .unwrap() + .iter() + .filter(|event| { + event.level == Level::INFO + && event + .message + .contains("Skipping malformed extension config entry") + }) + .map(|event| event.key.clone()) + .collect(); + + let broken_logs: Vec<_> = matching_events + .iter() + .filter(|k| k.as_deref() == Some("broken")) + .collect(); + assert!( + !broken_logs.is_empty(), + "expected at least one log naming the broken extension key, got {:?}", + matching_events + ); + let other_keys: Vec<_> = matching_events + .iter() + .filter(|k| k.as_deref() != Some("broken")) + .collect(); + assert!( + other_keys.is_empty(), + "expected no logs for other extension keys, got {:?}", + other_keys + ); + } }