Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
420 changes: 368 additions & 52 deletions codex-rs/app-server/tests/suite/v2/selected_capability_stack.rs

Large diffs are not rendered by default.

7 changes: 6 additions & 1 deletion codex-rs/codex-mcp/src/connection_manager/tool_catalog.rs
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
use std::collections::HashSet;
use std::sync::Arc;
use std::sync::atomic::Ordering;
use std::time::Instant;
Expand Down Expand Up @@ -170,6 +171,7 @@ impl McpConnectionSet {
config: Arc<crate::McpConfig>,
plugins_available: bool,
required_servers: &[String],
required_plugins: &HashSet<String>,
) -> McpBinding {
let revision = self.tool_catalog_revision.read().await;
let mut listed_tools = Vec::new();
Expand All @@ -190,10 +192,13 @@ impl McpConnectionSet {
let has_cached_tools = cached_tools.is_some();
let must_wait_for_startup = (required
&& (!view.connection.startup_is_dormant() || !has_cached_tools))
|| self.is_selected_plugin_mcp_server(server_name)
|| required_servers
.iter()
.any(|required| required == server_name)
|| (self.is_selected_plugin_mcp_server(server_name)
&& self
.plugin_id_for_mcp_server_name(server_name)
.is_some_and(|plugin_id| required_plugins.contains(plugin_id)))
|| (server_name == CODEX_APPS_MCP_SERVER_NAME && !has_cached_tools);
if !must_wait_for_startup && has_cached_tools {
return (server_name, view, cached_tools);
Expand Down
85 changes: 73 additions & 12 deletions codex-rs/codex-mcp/src/connection_manager_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -231,6 +231,7 @@ async fn capture_binding(manager: &Arc<McpConnectionSet>) -> McpBinding {
Arc::new(config),
/*plugins_available*/ false,
/*required_servers*/ &[],
/*required_plugins*/ &HashSet::new(),
)
.await
}
Expand Down Expand Up @@ -2592,10 +2593,17 @@ async fn capture_binding_skips_pending_optional_servers_after_configured_shared_
serde_json::from_value(serde_json::json!({ "command": "optional-plugin" }))
.expect("optional plugin MCP config"),
));
catalog.register(crate::McpServerRegistration::from_selected_plugin(
"pending-selected".to_string(),
crate::McpPluginAttribution::new("selected-plugin".to_string(), "Selected".to_string()),
/*selection_order*/ 0,
serde_json::from_value(serde_json::json!({ "command": "selected-plugin" }))
.expect("selected plugin MCP config"),
));
plugin_config.mcp_server_catalog = catalog.build();
plugin_config.optional_mcp_startup_grace = Duration::from_millis(250);
manager.tool_plugin_provenance = Arc::new(crate::tool_plugin_provenance(&plugin_config));
for server_name in ["pending-one", "pending-two"] {
for server_name in ["pending-one", "pending-two", "pending-selected"] {
manager.insert_test_client(
server_name.to_string(),
AsyncManagedClient {
Expand All @@ -2613,19 +2621,34 @@ async fn capture_binding_skips_pending_optional_servers_after_configured_shared_
);
}

let mut required_manager = McpConnectionSet::new_uninitialized(
&approval_policy,
&permission_profile,
/*prefix_mcp_tool_names*/ true,
);
required_manager.tool_plugin_provenance = Arc::clone(&manager.tool_plugin_provenance);
required_manager.insert_test_client(
"pending-selected",
manager.test_client("pending-selected").clone(),
);
required_manager.required_servers = vec!["pending-selected".to_string()];

let manager = Arc::new(manager);
assert_eq!(manager.stable_catalog_revision().await, None);
let started = tokio::time::Instant::now();
let binding = tokio::time::timeout(
Duration::from_millis(500),
manager.capture_binding_with_metadata(
Arc::new(plugin_config),
/*plugins_available*/ false,
/*required_servers*/ &[],
/*required_plugins*/ &HashSet::new(),
),
)
.await
.expect("all optional servers should share the configured startup grace");
assert!(binding.tools().is_empty());
assert_eq!(started.elapsed(), Duration::from_millis(250));

let binding = tokio::time::timeout(Duration::from_millis(1), capture_binding(&manager))
.await
Expand All @@ -2651,17 +2674,52 @@ async fn capture_binding_skips_pending_optional_servers_after_configured_shared_
"resource discovery must not wait for an omitted optional server"
);

let required_servers = vec!["pending-one".to_string()];
let binding = tokio::time::timeout(
Duration::from_millis(1),
manager.capture_binding_with_metadata(
Arc::new(crate::mcp::tests::test_mcp_config(std::env::temp_dir())),
/*plugins_available*/ false,
&required_servers,
),
)
.await;
assert!(binding.is_err(), "explicitly requested servers must wait");
for server_name in ["pending-one", "pending-selected"] {
let required_servers = vec![server_name.to_string()];
let binding = tokio::time::timeout(
Duration::from_millis(1),
manager.capture_binding_with_metadata(
Arc::new(crate::mcp::tests::test_mcp_config(std::env::temp_dir())),
/*plugins_available*/ false,
&required_servers,
/*required_plugins*/ &HashSet::new(),
),
)
.await;
assert!(binding.is_err(), "explicitly requested servers must wait");
}
// A plugin mention must still require startup after the optional grace has elapsed.
for (plugin_id, must_wait) in [
("selected-plugin", true),
("optional-plugin", false),
("selected-plugin-other", false),
] {
let required_plugins = HashSet::from([plugin_id.to_string()]);
let binding = tokio::time::timeout(
Duration::from_millis(1),
manager.capture_binding_with_metadata(
Arc::new(crate::mcp::tests::test_mcp_config(std::env::temp_dir())),
/*plugins_available*/ false,
/*required_servers*/ &[],
&required_plugins,
),
)
.await;
assert_eq!(
binding.is_err(),
must_wait,
"plugin requirement {plugin_id}"
);
}
assert!(
tokio::time::timeout(
Duration::from_millis(1500),
capture_binding(&Arc::new(required_manager)),
)
.await
.is_err(),
"configured-required selected plugin servers must wait beyond the optional grace"
);
}

#[tokio::test(start_paused = true)]
Expand Down Expand Up @@ -2690,6 +2748,7 @@ async fn capture_binding_waits_for_optional_startup_when_shared_grace_is_disable
Arc::new(config),
/*plugins_available*/ false,
/*required_servers*/ &[],
/*required_plugins*/ &HashSet::new(),
)
.await
});
Expand Down Expand Up @@ -2818,6 +2877,7 @@ async fn capture_binding_shares_optional_startup_grace_across_connection_sets()
Arc::new(disabled_config),
/*plugins_available*/ false,
/*required_servers*/ &[],
/*required_plugins*/ &HashSet::new(),
),
)
.await
Expand Down Expand Up @@ -2845,6 +2905,7 @@ async fn capture_binding_shares_optional_startup_grace_across_connection_sets()
Arc::new(updated_config),
/*plugins_available*/ false,
/*required_servers*/ &[],
/*required_plugins*/ &HashSet::new(),
),
)
.await
Expand Down
45 changes: 35 additions & 10 deletions codex-rs/codex-mcp/src/runtime.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
//! [`crate::connection_manager`].

use std::collections::HashMap;
use std::collections::HashSet;
use std::path::PathBuf;
use std::sync::Arc;
use std::sync::Mutex;
Expand Down Expand Up @@ -335,20 +336,29 @@ impl McpRuntime {

/// Captures the latest published configuration and live client handles.
pub async fn current_binding(&self) -> Option<Arc<McpBinding>> {
self.current_binding_with_required_servers(&[]).await
self.current_binding_with_requirements(&[], &HashSet::new())
.await
}

/// Captures the latest runtime, waiting for servers explicitly required by this turn.
pub async fn current_binding_with_required_servers(
/// Captures one runtime, waiting for explicitly required servers and selected plugins.
/// Plugin IDs are resolved by the captured connection set, even if a refresh publishes later.
pub async fn current_binding_with_requirements(
&self,
required_servers: &[String],
required_plugins: &HashSet<String>,
) -> Option<Arc<McpBinding>> {
Self::binding_from_published_runtime(self.current.load_full(), required_servers).await
Self::binding_from_published_runtime(
self.current.load_full(),
required_servers,
required_plugins,
)
.await
}

async fn binding_from_published_runtime(
current: Arc<PublishedMcpRuntime>,
required_servers: &[String],
required_plugins: &HashSet<String>,
) -> Option<Arc<McpBinding>> {
let config = Arc::clone(current.config.as_ref()?);
let stable_catalog_revision = current.connections.stable_catalog_revision().await;
Expand All @@ -367,7 +377,12 @@ impl McpRuntime {
let binding = Arc::new(
current
.connections
.capture_binding_with_metadata(config, current.plugins_available, required_servers)
.capture_binding_with_metadata(
config,
current.plugins_available,
required_servers,
required_plugins,
)
.await,
);
if let Some(catalog_revision) = stable_catalog_revision
Expand Down Expand Up @@ -445,7 +460,12 @@ impl McpRuntime {
if !current.connections.wait_for_server_startup(server).await {
return None;
}
Self::binding_from_published_runtime(current, /*required_servers*/ &[]).await
Self::binding_from_published_runtime(
current,
/*required_servers*/ &[],
/*required_plugins*/ &HashSet::new(),
)
.await
}

/// Returns the latest published configuration without waiting for clients.
Expand Down Expand Up @@ -877,12 +897,14 @@ mod tests {
let first = McpRuntime::binding_from_published_runtime(
Arc::clone(&published),
/*required_servers*/ &[],
/*required_plugins*/ &HashSet::new(),
)
.await
.expect("first binding");
let repeated = McpRuntime::binding_from_published_runtime(
Arc::clone(&published),
/*required_servers*/ &[],
/*required_plugins*/ &HashSet::new(),
)
.await
.expect("repeated binding");
Expand All @@ -893,10 +915,13 @@ mod tests {
cached_binding: Mutex::new(None),
..previous
});
let refreshed =
McpRuntime::binding_from_published_runtime(republished, /*required_servers*/ &[])
.await
.expect("republished binding");
let refreshed = McpRuntime::binding_from_published_runtime(
republished,
/*required_servers*/ &[],
/*required_plugins*/ &HashSet::new(),
)
.await
.expect("republished binding");
assert!(!Arc::ptr_eq(&first, &refreshed));
}

Expand Down
25 changes: 14 additions & 11 deletions codex-rs/core/src/plugins/mentions.rs
Original file line number Diff line number Diff line change
Expand Up @@ -68,6 +68,19 @@ pub(crate) fn collect_explicit_plugin_mentions(
return Vec::new();
}

// `config_name` stores the full plugin ID, not its display or mention name.
let mentioned_plugin_ids = collect_explicit_plugin_ids(input);
plugins
.iter()
.filter(|plugin| mentioned_plugin_ids.contains(plugin.config_name.as_str()))
.cloned()
.collect()
}

/// Collect exact IDs from explicit `plugin://` references, independently of display names.
///
/// Host plugins use `<plugin>@<marketplace>` IDs; selected roots may supply opaque IDs.
pub(crate) fn collect_explicit_plugin_ids(input: &[UserInput]) -> HashSet<String> {
let messages = input
.iter()
.filter_map(|item| match item {
Expand All @@ -76,7 +89,7 @@ pub(crate) fn collect_explicit_plugin_mentions(
})
.collect::<Vec<String>>();

let mentioned_config_names: HashSet<String> = input
input
.iter()
.filter_map(|item| match item {
UserInput::Mention { path, .. } => Some(path.clone()),
Expand All @@ -89,16 +102,6 @@ pub(crate) fn collect_explicit_plugin_mentions(
)
.filter(|path| tool_kind_for_path(path.as_str()) == ToolMentionKind::Plugin)
.filter_map(|path| plugin_config_name_from_path(path.as_str()).map(str::to_string))
.collect();

if mentioned_config_names.is_empty() {
return Vec::new();
}

plugins
.iter()
.filter(|plugin| mentioned_config_names.contains(plugin.config_name.as_str()))
.cloned()
.collect()
}

Expand Down
Loading
Loading