fix: preserve unparseable extension entries during config refresh (#9439)
Co-authored-by: Michael Neale <michael.neale@gmail.com>
This commit is contained in:
@@ -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
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user