diff --git a/codex-rs/core/src/session/mod.rs b/codex-rs/core/src/session/mod.rs index 39f7a0627e..4c5c7b1134 100644 --- a/codex-rs/core/src/session/mod.rs +++ b/codex-rs/core/src/session/mod.rs @@ -3044,16 +3044,21 @@ impl Session { .await; let extension_data = codex_extension_api::ExtensionData::new(turn_context.sub_id.clone()); extension_data.insert(selected_capability_roots.clone()); - let mcp = self - .mcp_runtime_for_step(turn_context.as_ref(), &selected_capability_roots) - .or_cancel(cancellation_token) - .await?; + let (mcp, prepared_recommendations) = async { + tokio::join!( + self.mcp_runtime_for_step(turn_context.as_ref(), &selected_capability_roots), + turn::prepare_tool_recommendations(self.as_ref(), turn_context.as_ref()), + ) + } + .or_cancel(cancellation_token) + .await?; let (mcp_tools, tool_router) = turn::built_tools( self.as_ref(), turn_context.as_ref(), &environments, mcp.as_ref(), &extension_data, + prepared_recommendations, ) .or_cancel(cancellation_token) .await?; diff --git a/codex-rs/core/src/session/turn.rs b/codex-rs/core/src/session/turn.rs index cb38541de2..6c76a46bfa 100644 --- a/codex-rs/core/src/session/turn.rs +++ b/codex-rs/core/src/session/turn.rs @@ -82,6 +82,7 @@ use codex_extension_api::TurnInputEnvironment; use codex_features::Feature; use codex_file_system::FindUpErrorPolicy; use codex_file_system::find_nearest_ancestor_with_markers; +use codex_login::CodexAuth; use codex_mcp::ToolInfo; use codex_protocol::ResponseItemId; use codex_protocol::config_types::AutoCompactTokenLimitScope; @@ -111,6 +112,7 @@ use codex_protocol::protocol::SafetyBufferingEvent; use codex_protocol::protocol::TurnDiffEvent; use codex_protocol::protocol::WarningEvent; use codex_protocol::user_input::UserInput; +use codex_tools::DiscoverableTool; use codex_tools::ToolName; use codex_tools::filter_request_plugin_install_discoverable_tools_for_client; use codex_utils_path_uri::PathUri; @@ -1270,6 +1272,50 @@ async fn run_sampling_request( } } +pub(crate) struct PreparedToolRecommendations { + auth: Option, + endpoint_candidates: Option>, +} + +#[instrument(level = "trace", skip_all)] +pub(crate) async fn prepare_tool_recommendations( + sess: &Session, + turn_context: &TurnContext, +) -> PreparedToolRecommendations { + let loaded_plugins = sess + .services + .plugins_manager + .plugins_for_config(&turn_context.config.plugins_config_input()) + .instrument(trace_span!("built_tools.load_plugins")) + .await; + let tool_suggest_is_enabled = tool_suggest_enabled(turn_context); + let auth = if tool_suggest_is_enabled { + sess.services.auth_manager.auth().await + } else { + None + }; + let endpoint_candidates = if tool_suggest_is_enabled { + let plugins_config = turn_context.config.plugins_config_input(); + sess.services + .plugins_manager + .recommended_plugin_candidates_for_config(RecommendedPluginCandidatesInput { + plugins_config: &plugins_config, + loaded_plugins: &loaded_plugins, + auth: auth.as_ref(), + disabled_tools: &turn_context.config.tool_suggest.disabled_tools, + app_server_client_name: turn_context.app_server_client_name.as_deref(), + }) + .await + } else { + None + }; + + PreparedToolRecommendations { + auth, + endpoint_candidates, + } +} + #[instrument(level = "trace", skip_all, fields( @@ -1284,14 +1330,9 @@ pub(crate) async fn built_tools( environments: &TurnEnvironmentSnapshot, mcp: &codex_mcp::McpBinding, step_store: &ExtensionData, + prepared_recommendations: PreparedToolRecommendations, ) -> (Vec, Arc) { let all_mcp_tools = mcp.tools().to_vec(); - let loaded_plugins = sess - .services - .plugins_manager - .plugins_for_config(&turn_context.config.plugins_config_input()) - .instrument(trace_span!("built_tools.load_plugins")) - .await; let connector_snapshot = mcp.config().connector_snapshot.clone(); let apps_enabled = turn_context.apps_enabled(); @@ -1317,26 +1358,10 @@ pub(crate) async fn built_tools( None }; let tool_suggest_is_enabled = tool_suggest_enabled(turn_context); - let auth = if tool_suggest_is_enabled { - sess.services.auth_manager.auth().await - } else { - None - }; - let endpoint_recommended_plugin_candidates = if tool_suggest_is_enabled { - let plugins_config = turn_context.config.plugins_config_input(); - sess.services - .plugins_manager - .recommended_plugin_candidates_for_config(RecommendedPluginCandidatesInput { - plugins_config: &plugins_config, - loaded_plugins: &loaded_plugins, - auth: auth.as_ref(), - disabled_tools: &turn_context.config.tool_suggest.disabled_tools, - app_server_client_name: turn_context.app_server_client_name.as_deref(), - }) - .await - } else { - None - }; + let PreparedToolRecommendations { + auth, + endpoint_candidates: endpoint_recommended_plugin_candidates, + } = prepared_recommendations; let tool_suggest_candidates = if let Some(recommended_plugin_candidates) = endpoint_recommended_plugin_candidates { Some(ToolSuggestCandidates { diff --git a/codex-rs/core/tests/suite/request_plugin_install.rs b/codex-rs/core/tests/suite/request_plugin_install.rs index 1592577107..735d1b20c5 100644 --- a/codex-rs/core/tests/suite/request_plugin_install.rs +++ b/codex-rs/core/tests/suite/request_plugin_install.rs @@ -1,6 +1,7 @@ #![cfg(not(target_os = "windows"))] #![allow(clippy::unwrap_used)] +use anyhow::Context; use anyhow::Result; use codex_config::types::ToolSuggestDisabledTool; use codex_config::types::ToolSuggestDiscoverable; @@ -21,6 +22,7 @@ use codex_protocol::protocol::EventMsg; use codex_protocol::protocol::Op; use codex_protocol::protocol::ThreadSettingsOverrides; use codex_protocol::user_input::UserInput; +use codex_utils_path_uri::PathUri; use core_test_support::apps_test_server::AppsTestServer; use core_test_support::responses::ev_assistant_message; use core_test_support::responses::ev_completed; @@ -31,6 +33,7 @@ use core_test_support::responses::mount_sse_sequence; use core_test_support::responses::sse; use core_test_support::responses::start_mock_server; use core_test_support::skip_if_no_network; +use core_test_support::skip_if_wine_exec; use core_test_support::test_codex::TestCodex; use core_test_support::test_codex::test_codex; use core_test_support::test_codex::turn_permission_fields; @@ -45,11 +48,15 @@ use std::time::Duration; use std::time::Instant; use wiremock::Mock; use wiremock::MockGuard; +use wiremock::MockServer; use wiremock::ResponseTemplate; use wiremock::matchers::method; use wiremock::matchers::path; use wiremock::matchers::query_param; +use super::rmcp_client::remote_aware_environment_id; +use super::rmcp_client::remote_aware_stdio_server_bin; + const TOOL_SEARCH_TOOL_NAME: &str = "tool_search"; const LIST_AVAILABLE_PLUGINS_TO_INSTALL_TOOL_NAME: &str = "list_available_plugins_to_install"; const REQUEST_PLUGIN_INSTALL_TOOL_NAME: &str = "request_plugin_install"; @@ -59,6 +66,7 @@ const REMOTE_CALENDAR_PLUGIN_ID: &str = "plugin_calendar"; const CALENDAR_CONNECTOR_ID: &str = "calendar"; const CALENDAR_NAMESPACE: &str = "mcp__codex_apps__calendar"; const CALENDAR_CREATE_EVENT_TOOL: &str = "_create_event"; +const STEP_PREPARATION_MCP_SERVER: &str = "step_preparation"; fn tool_names(body: &Value) -> Vec { body.get("tools") @@ -151,6 +159,103 @@ async fn build_test( builder.build(server).await } +async fn build_gated_step_preparation_test( + server: &MockServer, + apps_server: &AppsTestServer, +) -> Result { + let command = remote_aware_stdio_server_bin()?; + let environment_id = remote_aware_environment_id(); + let apps_base_url = apps_server.chatgpt_base_url.clone(); + let mut builder = test_codex() + .with_auth(CodexAuth::create_dummy_chatgpt_auth_for_testing()) + .with_config(move |config| { + config + .permissions + .set_permission_profile(PermissionProfile::Disabled) + .expect("test config should allow disabled permissions"); + configure_apps_without_search_tool(config, apps_base_url.as_str()); + + let barrier_file = config.cwd.join("allow-step-preparation-initialize"); + let pid_file = config.cwd.join("step-preparation.pid"); + let mut servers = config.mcp_servers.get().clone(); + servers.insert( + STEP_PREPARATION_MCP_SERVER.to_string(), + serde_json::from_value(json!({ + "command": command.clone(), + "environment_id": environment_id.clone(), + "env": { + "MCP_TEST_INITIALIZE_BARRIER_FILE": barrier_file, + "MCP_TEST_PID_FILE": pid_file, + }, + "enabled_tools": ["echo"], + "startup_timeout_sec": 10, + })) + .expect("test MCP server configuration"), + ); + config + .mcp_servers + .set(servers) + .expect("test config should allow MCP servers"); + }); + builder.build_with_auto_env(server).await +} + +async fn start_gated_step_preparation(test: &TestCodex, server: &MockServer) -> Result { + let prior_recommendation_count = server + .received_requests() + .await + .unwrap_or_default() + .into_iter() + .filter(|request| request.url.path() == "/ps/plugins/suggested") + .count(); + let (sandbox_policy, permission_profile) = + turn_permission_fields(PermissionProfile::Disabled, test.config.cwd.as_path()); + test.codex + .submit(Op::UserInput { + items: vec![UserInput::Text { + text: "prepare MCP and plugin recommendations".to_string(), + text_elements: Vec::new(), + }], + final_output_json_schema: None, + responsesapi_client_metadata: None, + additional_context: Default::default(), + thread_settings: ThreadSettingsOverrides { + approval_policy: Some(AskForApproval::Never), + sandbox_policy: Some(sandbox_policy), + permission_profile, + ..Default::default() + }, + }) + .await?; + + let fs = test.fs(); + let pid_file = PathUri::from_host_native_path(test.config.cwd.join("step-preparation.pid"))?; + tokio::time::timeout(Duration::from_secs(5), async { + loop { + let mcp_started = fs + .read_file_text(&pid_file, /*sandbox*/ None) + .await + .is_ok_and(|pid| !pid.trim().is_empty()); + let recommendation_count = server + .received_requests() + .await + .unwrap_or_default() + .into_iter() + .filter(|request| request.url.path() == "/ps/plugins/suggested") + .count(); + if mcp_started && recommendation_count > prior_recommendation_count { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .context("MCP startup and plugin recommendations should begin before MCP is released")?; + + PathUri::from_host_native_path(test.config.cwd.join("allow-step-preparation-initialize")) + .map_err(Into::into) +} + async fn start_install_turn(test: &TestCodex, prompt: &str) -> Result { let (sandbox_policy, permission_profile) = turn_permission_fields(PermissionProfile::Disabled, test.config.cwd.as_path()); @@ -279,6 +384,122 @@ async fn mount_remote_calendar_installed_plugins(server: &wiremock::MockServer) } } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn mcp_discovery_overlaps_endpoint_plugin_recommendations() -> Result<()> { + skip_if_wine_exec!( + Ok(()), + "requires a Windows test_stdio_server in the Wine-exec environment" + ); + skip_if_no_network!(Ok(())); + + let server = start_mock_server().await; + let apps_server = AppsTestServer::mount(&server).await?; + mount_recommendations( + &server, + ResponseTemplate::new(200).set_body_json(json!({ + "enabled": true, + "plugins": [{ + "id": "plugin_github", + "name": "github", + "status": "ENABLED", + "installation_policy": "AVAILABLE", + "release": {"display_name": "GitHub"} + }] + })), + ) + .await; + let response = mount_sse_once( + &server, + sse(vec![ + ev_response_created("concurrent-endpoint"), + ev_assistant_message("endpoint-message", "done"), + ev_completed("concurrent-endpoint"), + ]), + ) + .await; + let test = build_gated_step_preparation_test(&server, &apps_server).await?; + + let barrier = start_gated_step_preparation(&test, &server).await?; + assert!( + response.requests().is_empty(), + "sampling should wait for the complete MCP catalog" + ); + test.fs() + .write_file(&barrier, b"ready".to_vec(), /*sandbox*/ None) + .await?; + wait_for_event(&test.codex, |event| { + matches!(event, EventMsg::TurnComplete(_)) + }) + .await; + + let request = response.single_request(); + assert!( + request + .message_input_texts("user") + .join("\n") + .contains("github@openai-curated-remote"), + "the completed request should preserve endpoint recommendations" + ); + assert!( + request + .tool_by_name("mcp__step_preparation", "echo") + .is_some(), + "the completed request should expose the live gated MCP tool" + ); + + test.codex.shutdown_and_wait().await?; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn interrupting_concurrent_step_preparation_prevents_sampling() -> Result<()> { + skip_if_wine_exec!( + Ok(()), + "requires a Windows test_stdio_server in the Wine-exec environment" + ); + skip_if_no_network!(Ok(())); + + let server = start_mock_server().await; + let apps_server = AppsTestServer::mount(&server).await?; + mount_recommendations( + &server, + ResponseTemplate::new(200).set_body_json(json!({"enabled": true, "plugins": []})), + ) + .await; + let response = mount_sse_once( + &server, + sse(vec![ + ev_response_created("cancelled-step-preparation"), + ev_assistant_message("cancelled-message", "unexpected"), + ev_completed("cancelled-step-preparation"), + ]), + ) + .await; + let test = build_gated_step_preparation_test(&server, &apps_server).await?; + + let barrier = start_gated_step_preparation(&test, &server).await?; + assert!(response.requests().is_empty()); + test.codex.submit(Op::Interrupt).await?; + wait_for_event(&test.codex, |event| { + matches!(event, EventMsg::TurnAborted(_)) + }) + .await; + assert!( + response.requests().is_empty(), + "cancelling concurrent step preparation must prevent model sampling" + ); + + test.fs() + .write_file(&barrier, b"ready".to_vec(), /*sandbox*/ None) + .await?; + test.codex.shutdown_and_wait().await?; + assert!( + response.requests().is_empty(), + "releasing the cancelled MCP startup must not revive the aborted turn" + ); + Ok(()) +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn explicit_false_preserves_legacy_workflow() -> Result<()> { skip_if_no_network!(Ok(()));