Persist dynamic extension config so we can resume recipe sessions w/ extensions (#4331)

This commit is contained in:
Will Pfleger
2025-10-04 17:17:12 -04:00
committed by GitHub
parent 51651550e0
commit a59fbab69c
14 changed files with 357 additions and 305 deletions
+39 -21
View File
@@ -31,7 +31,7 @@ use crate::agents::tool_route_manager::ToolRouteManager;
use crate::agents::tool_router_index_manager::ToolRouterIndexManager;
use crate::agents::types::SessionConfig;
use crate::agents::types::{FrontendTool, ToolResultReceiver};
use crate::config::{Config, ExtensionConfigManager};
use crate::config::{get_enabled_extensions, get_extension_by_name, Config};
use crate::context_mgmt::auto_compact;
use crate::conversation::{debug_conversation_fix, fix_conversation, Conversation};
use crate::mcp_utils::ToolResult;
@@ -62,6 +62,7 @@ use super::platform_tools;
use super::tool_execution::{ToolCallResult, CHAT_MODE_TOOL_SKIPPED_RESPONSE, DECLINED_RESPONSE};
use crate::agents::subagent_task_config::TaskConfig;
use crate::conversation::message::{Message, ToolRequest};
use crate::session::extension_data::{EnabledExtensionsState, ExtensionState};
use crate::session::SessionManager;
const DEFAULT_MAX_TURNS: u32 = 1000;
@@ -549,6 +550,28 @@ impl Agent {
)
}
/// Save current extension state to session metadata
/// Should be called after any extension add/remove operation
pub async fn save_extension_state(&self, session: &SessionConfig) -> Result<()> {
let extension_configs = self.extension_manager.get_extension_configs().await;
let extensions_state = EnabledExtensionsState::new(extension_configs);
let mut session_data = SessionManager::get_session(&session.id, false).await?;
if let Err(e) = extensions_state.to_extension_data(&mut session_data.extension_data) {
warn!("Failed to serialize extension state: {}", e);
return Err(anyhow!("Extension state serialization failed: {}", e));
}
SessionManager::update_session(&session.id)
.extension_data(session_data.extension_data)
.apply()
.await?;
Ok(())
}
#[allow(clippy::too_many_lines)]
pub(super) async fn manage_extensions(
&self,
@@ -595,9 +618,9 @@ impl Agent {
return (request_id, result);
}
let config = match ExtensionConfigManager::get_config_by_name(&extension_name) {
Ok(Some(config)) => config,
Ok(None) => {
let config = match get_extension_by_name(&extension_name) {
Some(config) => config,
None => {
return (
request_id,
Err(ErrorData::new(
@@ -610,16 +633,6 @@ impl Agent {
)),
)
}
Err(e) => {
return (
request_id,
Err(ErrorData::new(
ErrorCode::INTERNAL_ERROR,
format!("Failed to get extension config: {}", e),
None,
)),
)
}
};
let result = self
.extension_manager
@@ -658,6 +671,7 @@ impl Agent {
}
}
}
(request_id, result)
}
@@ -792,6 +806,10 @@ impl Agent {
.expect("Failed to list extensions")
}
pub async fn get_extension_configs(&self) -> Vec<ExtensionConfig> {
self.extension_manager.get_extension_configs().await
}
/// Handle a confirmation response for a tool request
pub async fn handle_confirmation(
&self,
@@ -1199,7 +1217,12 @@ impl Agent {
}
}
if all_install_successful {
if all_install_successful && !enable_extension_request_ids.is_empty() {
if let Some(ref session_config) = session {
if let Err(e) = self.save_extension_state(session_config).await {
warn!("Failed to save extension state after runtime changes: {}", e);
}
}
tools_updated = true;
}
}
@@ -1558,12 +1581,7 @@ impl Agent {
(instructions, activities)
};
let extensions = ExtensionConfigManager::get_all().unwrap_or_default();
let extension_configs: Vec<_> = extensions
.iter()
.filter(|e| e.enabled)
.map(|e| e.config.clone())
.collect();
let extension_configs = get_enabled_extensions();
let author = Author {
contact: std::env::var("USER")
+11 -2
View File
@@ -32,7 +32,7 @@ use super::tool_execution::ToolCallResult;
use crate::agents::extension::{Envs, ProcessExit};
use crate::agents::extension_malware_check;
use crate::agents::mcp_client::{McpClient, McpClientTrait};
use crate::config::{Config, ExtensionConfigManager};
use crate::config::{get_all_extensions, Config};
use crate::oauth::oauth_flow;
use crate::prompt_template;
use rmcp::model::{
@@ -576,6 +576,15 @@ impl ExtensionManager {
Ok(self.extensions.lock().await.keys().cloned().collect())
}
pub async fn get_extension_configs(&self) -> Vec<ExtensionConfig> {
self.extensions
.lock()
.await
.values()
.map(|ext| ext.config.clone())
.collect()
}
/// Get all tools from all clients with proper prefixing
pub async fn get_prefixed_tools(
&self,
@@ -1035,7 +1044,7 @@ impl ExtensionManager {
// First get disabled extensions from current config
let mut disabled_extensions: Vec<String> = vec![];
for extension in ExtensionConfigManager::get_all().expect("should load extensions") {
for extension in get_all_extensions() {
if !extension.enabled {
let config = extension.config.clone();
let description = match &config {
@@ -108,24 +108,14 @@ fn process_extensions(
for ext in arr {
if let Some(name_str) = ext.as_str() {
// Look up the full extension config by name
match crate::config::ExtensionConfigManager::get_config_by_name(name_str) {
Ok(Some(config)) => {
// Check if the extension is enabled
if crate::config::ExtensionConfigManager::is_enabled(&config.key())
.unwrap_or(false)
{
converted_extensions.push(config);
} else {
tracing::warn!("Extension '{}' is disabled, skipping", name_str);
}
}
Ok(None) => {
tracing::warn!("Extension '{}' not found in configuration", name_str);
}
Err(e) => {
tracing::warn!("Error looking up extension '{}': {}", name_str, e);
if let Some(config) = crate::config::get_extension_by_name(name_str) {
if crate::config::is_extension_enabled(&config.key()) {
converted_extensions.push(config);
} else {
tracing::warn!("Extension '{}' is disabled, skipping", name_str);
}
} else {
tracing::warn!("Extension '{}' not found in configuration", name_str);
}
} else if let Ok(ext_config) = serde_json::from_value::<ExtensionConfig>(ext.clone()) {
converted_extensions.push(ext_config);
+3 -5
View File
@@ -1,8 +1,7 @@
use crate::agents::subagent_task_config::DEFAULT_SUBAGENT_MAX_TURNS;
use crate::{
agents::extension::ExtensionConfig,
agents::{extension_manager::ExtensionManager, Agent, TaskConfig},
config::ExtensionConfigManager,
config::get_all_extensions,
prompt_template::render_global_file,
providers::errors::ProviderError,
};
@@ -68,12 +67,11 @@ impl SubAgent {
extensions.clone()
} else {
// Default behavior: use all enabled extensions
ExtensionConfigManager::get_all()
.unwrap_or_default()
get_all_extensions()
.into_iter()
.filter(|ext| ext.enabled)
.map(|ext| ext.config)
.collect::<Vec<ExtensionConfig>>()
.collect()
};
// Add the determined extensions to the subagent's extension manager
+123 -116
View File
@@ -1,7 +1,6 @@
use super::base::Config;
use crate::agents::extension::PLATFORM_EXTENSIONS;
use crate::agents::ExtensionConfig;
use anyhow::Result;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::collections::HashMap;
@@ -28,132 +27,140 @@ pub fn name_to_key(name: &str) -> String {
.to_lowercase()
}
pub struct ExtensionConfigManager;
fn get_extensions_map() -> HashMap<String, ExtensionEntry> {
let raw: Value = Config::global()
.get_param::<Value>(EXTENSIONS_CONFIG_KEY)
.unwrap_or_else(|err| {
warn!(
"Failed to load {}: {err}. Falling back to empty object.",
EXTENSIONS_CONFIG_KEY
);
Value::Object(serde_json::Map::new())
});
impl ExtensionConfigManager {
fn get_extensions_map() -> Result<HashMap<String, ExtensionEntry>> {
let raw: Value = Config::global()
.get_param::<Value>(EXTENSIONS_CONFIG_KEY)
.unwrap_or_else(|err| {
warn!(
"Failed to load {}: {err}. Falling back to empty object.",
EXTENSIONS_CONFIG_KEY
);
Value::Object(serde_json::Map::new())
});
let mut extensions_map: HashMap<String, ExtensionEntry> = match raw {
Value::Object(obj) => {
let mut m = HashMap::with_capacity(obj.len());
for (k, mut v) in obj {
if let Value::Object(ref mut inner) = v {
match inner.get("description") {
Some(Value::Null) | None => {
inner.insert(
"description".to_string(),
Value::String(String::new()),
);
}
_ => {}
}
}
match serde_json::from_value::<ExtensionEntry>(v.clone()) {
Ok(entry) => {
m.insert(k, entry);
}
Err(err) => {
let bad_json = serde_json::to_string(&v).unwrap_or_else(|e| {
format!("<failed to serialize malformed value: {e}>")
});
warn!(
extension = %k,
error = %err,
bad_json = %bad_json,
"Skipping malformed extension"
);
let mut extensions_map: HashMap<String, ExtensionEntry> = match raw {
Value::Object(obj) => {
let mut m = HashMap::with_capacity(obj.len());
for (k, mut v) in obj {
if let Value::Object(ref mut inner) = v {
match inner.get("description") {
Some(Value::Null) | None => {
inner.insert("description".to_string(), Value::String(String::new()));
}
_ => {}
}
}
match serde_json::from_value::<ExtensionEntry>(v.clone()) {
Ok(entry) => {
m.insert(k, entry);
}
Err(err) => {
let bad_json = serde_json::to_string(&v).unwrap_or_else(|e| {
format!("<failed to serialize malformed value: {e}>")
});
warn!(
extension = %k,
error = %err,
bad_json = %bad_json,
"Skipping malformed extension"
);
}
}
m
}
other => {
warn!(
"Expected object for {}, got {}. Using empty map.",
EXTENSIONS_CONFIG_KEY, other
);
HashMap::new()
}
};
m
}
other => {
warn!(
"Expected object for {}, got {}. Using empty map.",
EXTENSIONS_CONFIG_KEY, other
);
HashMap::new()
}
};
if !extensions_map.is_empty() {
for (name, def) in PLATFORM_EXTENSIONS.iter() {
if !extensions_map.contains_key(*name) {
extensions_map.insert(
name.to_string(),
ExtensionEntry {
config: ExtensionConfig::Platform {
name: def.name.to_string(),
description: def.description.to_string(),
bundled: Some(true),
available_tools: Vec::new(),
},
enabled: true,
if !extensions_map.is_empty() {
for (name, def) in PLATFORM_EXTENSIONS.iter() {
if !extensions_map.contains_key(*name) {
extensions_map.insert(
name.to_string(),
ExtensionEntry {
config: ExtensionConfig::Platform {
name: def.name.to_string(),
description: def.description.to_string(),
bundled: Some(true),
available_tools: Vec::new(),
},
);
}
enabled: true,
},
);
}
}
Ok(extensions_map)
}
extensions_map
}
fn save_extensions_map(extensions: HashMap<String, ExtensionEntry>) -> Result<()> {
let config = Config::global();
config.set_param(EXTENSIONS_CONFIG_KEY, serde_json::to_value(extensions)?)?;
Ok(())
}
pub fn get_config_by_name(name: &str) -> Result<Option<ExtensionConfig>> {
let extensions = Self::get_extensions_map()?;
Ok(extensions
.values()
.find(|entry| entry.config.name() == name)
.map(|entry| entry.config.clone()))
}
pub fn set(entry: ExtensionEntry) -> Result<()> {
let mut extensions = Self::get_extensions_map()?;
let key = entry.config.key();
extensions.insert(key, entry);
Self::save_extensions_map(extensions)
}
pub fn remove(key: &str) -> Result<()> {
let mut extensions = Self::get_extensions_map()?;
extensions.remove(key);
Self::save_extensions_map(extensions)
}
pub fn set_enabled(key: &str, enabled: bool) -> Result<()> {
let mut extensions = Self::get_extensions_map()?;
if let Some(entry) = extensions.get_mut(key) {
entry.enabled = enabled;
Self::save_extensions_map(extensions)?;
fn save_extensions_map(extensions: HashMap<String, ExtensionEntry>) {
let config = Config::global();
match serde_json::to_value(extensions) {
Ok(value) => {
if let Err(e) = config.set_param(EXTENSIONS_CONFIG_KEY, value) {
tracing::debug!("Failed to save extensions config: {}", e);
}
}
Err(e) => {
tracing::debug!("Failed to serialize extensions: {}", e);
}
Ok(())
}
pub fn get_all() -> Result<Vec<ExtensionEntry>> {
let extensions = Self::get_extensions_map()?;
Ok(extensions.into_values().collect())
}
pub fn get_all_names() -> Result<Vec<String>> {
let extensions = Self::get_extensions_map()?;
Ok(extensions.keys().cloned().collect())
}
pub fn is_enabled(key: &str) -> Result<bool> {
let extensions = Self::get_extensions_map()?;
Ok(extensions.get(key).map(|e| e.enabled).unwrap_or(false))
}
}
pub fn get_extension_by_name(name: &str) -> Option<ExtensionConfig> {
let extensions = get_extensions_map();
extensions
.values()
.find(|entry| entry.config.name() == name)
.map(|entry| entry.config.clone())
}
pub fn set_extension(entry: ExtensionEntry) {
let mut extensions = get_extensions_map();
let key = entry.config.key();
extensions.insert(key, entry);
save_extensions_map(extensions);
}
pub fn remove_extension(key: &str) {
let mut extensions = get_extensions_map();
extensions.remove(key);
save_extensions_map(extensions);
}
pub fn set_extension_enabled(key: &str, enabled: bool) {
let mut extensions = get_extensions_map();
if let Some(entry) = extensions.get_mut(key) {
entry.enabled = enabled;
save_extensions_map(extensions);
}
}
pub fn get_all_extensions() -> Vec<ExtensionEntry> {
let extensions = get_extensions_map();
extensions.into_values().collect()
}
pub fn get_all_extension_names() -> Vec<String> {
let extensions = get_extensions_map();
extensions.keys().cloned().collect()
}
pub fn is_extension_enabled(key: &str) -> bool {
let extensions = get_extensions_map();
extensions.get(key).map(|e| e.enabled).unwrap_or(false)
}
pub fn get_enabled_extensions() -> Vec<ExtensionConfig> {
get_all_extensions()
.into_iter()
.filter(|ext| ext.enabled)
.map(|ext| ext.config)
.collect()
}
+4 -1
View File
@@ -10,7 +10,10 @@ pub use crate::agents::ExtensionConfig;
pub use base::{get_config_dir, Config, ConfigError, APP_STRATEGY};
pub use custom_providers::CustomProviderConfig;
pub use experiments::ExperimentManager;
pub use extensions::{ExtensionConfigManager, ExtensionEntry};
pub use extensions::{
get_all_extension_names, get_all_extensions, get_enabled_extensions, get_extension_by_name,
is_extension_enabled, remove_extension, set_extension, set_extension_enabled, ExtensionEntry,
};
pub use permission::PermissionManager;
pub use signup_openrouter::configure_openrouter;
pub use signup_tetrate::configure_tetrate;
@@ -1,6 +1,7 @@
// Extension data management for sessions
// Provides a simple way to store extension-specific data with versioned keys
use crate::config::ExtensionConfig;
use anyhow::Result;
use serde::{Deserialize, Serialize};
use serde_json::Value;
@@ -95,6 +96,23 @@ impl TodoState {
}
}
/// Enabled extensions state implementation for storing which extensions are active
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EnabledExtensionsState {
pub extensions: Vec<ExtensionConfig>,
}
impl ExtensionState for EnabledExtensionsState {
const EXTENSION_NAME: &'static str = "enabled_extensions";
const VERSION: &'static str = "v0";
}
impl EnabledExtensionsState {
pub fn new(extensions: Vec<ExtensionConfig>) -> Self {
Self { extensions }
}
}
#[cfg(test)]
mod tests {
use super::*;
+1
View File
@@ -2,4 +2,5 @@ pub mod extension_data;
mod legacy;
pub mod session_manager;
pub use extension_data::{EnabledExtensionsState, ExtensionData, ExtensionState, TodoState};
pub use session_manager::{Session, SessionInsights, SessionManager};