fix: preserve unparseable extension entries during config refresh (#9439)

Co-authored-by: Michael Neale <michael.neale@gmail.com>
This commit is contained in:
Matt Van Horn
2026-06-10 22:10:06 -07:00
committed by GitHub
parent db1f0fcc63
commit c6a3b5b4e1
+391 -39
View File
@@ -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<String, ExtensionEntry> {
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::<ExtensionEntry>(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<String, ExtensionEntry> {
let raw: Mapping = config
.get_param(EXTENSIONS_CONFIG_KEY)
@@ -52,36 +80,49 @@ fn get_extensions_map_with_config(config: &Config) -> IndexMap<String, Extension
Default::default()
});
let mut extensions_map = IndexMap::with_capacity(raw.len());
for (k, v) in raw {
match (k, serde_yaml::from_value::<ExtensionEntry>(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<String, ExtensionEntry> {
get_extensions_map_with_config(Config::global())
}
fn save_extensions_map(extensions: IndexMap<String, ExtensionEntry>) {
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<ExtensionEntry>),
Remove(String),
Noop,
}
fn with_raw_extensions_mapping<F>(config: &Config, mutate: F)
where
F: FnOnce(&mut IndexMap<String, ExtensionEntry>) -> ExtensionMutation,
{
let mut serialize_error = None;
let result = config.update_param::<Mapping, Mapping, _>(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<ExtensionConfig> {
}
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<ExtensionEntry> {
@@ -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<Mutex<Vec<CapturedEvent>>>,
}
#[derive(Debug)]
struct CapturedEvent {
level: Level,
message: String,
key: Option<String>,
}
impl<S> tracing_subscriber::Layer<S> 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<String>,
}
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
);
}
}