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
11 changes: 1 addition & 10 deletions codex-rs/core/src/mcp_tool_call.rs
Original file line number Diff line number Diff line change
Expand Up @@ -140,16 +140,7 @@ pub(crate) async fn handle_mcp_tool_call(
arguments: arguments_value.clone(),
};

sess.refresh_mcp_if_dirty().await;
let current_binding = sess
.services
.mcp_runtime
.current_binding_for_call(&server)
.await;
let Some(prepared_call) = current_binding
.as_ref()
.and_then(|binding| binding.prepare_call(&server, &tool_name))
else {
let Some(prepared_call) = sess.prepare_mcp_call(&server, &tool_name).await else {
let item_metadata =
McpToolCallItemMetadata::from_tool_metadata(&server, /*metadata*/ None);
let result = notify_mcp_tool_call_skip(
Expand Down
24 changes: 24 additions & 0 deletions codex-rs/core/src/session/mcp_runtime.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ use super::session::SessionConfiguration;
use super::*;
use crate::mcp::McpRuntimeProjection;
use codex_mcp::ElicitationReviewerHandle;
use codex_mcp::PreparedMcpCall;
use codex_protocol::capabilities::SelectedCapabilityRoot;

pub(super) struct McpDesiredState {
Expand All @@ -30,6 +31,29 @@ impl McpDesiredState {
}

impl Session {
/// Waits on this session's refreshed server before tool execution is admitted.
pub(crate) async fn wait_for_mcp_server(self: &Arc<Self>, server: &str) {
self.refresh_mcp_if_dirty().await;
self.services
.mcp_runtime
.wait_for_server_startup(server)
.await;
}

/// Captures this session's current MCP client and catalog for one tool call.
pub(crate) async fn prepare_mcp_call(
self: &Arc<Self>,
server: &str,
tool: &str,
) -> Option<PreparedMcpCall> {
self.refresh_mcp_if_dirty().await;
self.services
.mcp_runtime
.current_binding_for_call(server)
.await?
.prepare_call(server, tool)
}

pub(super) async fn latest_mcp_desired_state(
&self,
auth: Option<CodexAuth>,
Expand Down
5 changes: 1 addition & 4 deletions codex-rs/core/src/tools/handlers/mcp.rs
Original file line number Diff line number Diff line change
Expand Up @@ -170,11 +170,8 @@ impl McpHandler {
impl CoreToolRuntime for McpHandler {
fn wait_until_ready<'a>(&'a self, session: &'a Arc<Session>) -> Option<BoxFuture<'a, ()>> {
Some(Box::pin(async move {
session.refresh_mcp_if_dirty().await;
session
.services
.mcp_runtime
.wait_for_server_startup(&self.tool_info.server_name)
.wait_for_mcp_server(&self.tool_info.server_name)
.await;
}))
}
Expand Down
143 changes: 143 additions & 0 deletions codex-rs/core/tests/suite/mcp_tool_cache.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,13 +3,15 @@ use std::time::Duration;

use anyhow::Context;
use codex_config::Constrained;
use codex_config::types::McpServerConfig;
use codex_core::NewThread;
use codex_core::StartThreadOptions;
use codex_exec_server::ExecutorFileSystem;
use codex_exec_server::RemoveOptions;
use codex_protocol::models::PermissionProfile;
use codex_protocol::protocol::AskForApproval;
use codex_protocol::protocol::EventMsg;
use codex_protocol::protocol::McpInvocation;
use codex_protocol::protocol::Op;
use codex_protocol::user_input::UserInput;
use codex_utils_path_uri::PathUri;
Expand All @@ -20,6 +22,7 @@ use core_test_support::skip_if_no_network;
use core_test_support::skip_if_wine_exec;
use core_test_support::test_codex::test_codex;
use core_test_support::wait_for_event;
use core_test_support::wait_for_mcp_server;
use pretty_assertions::assert_eq;
use serde_json::Value;
use serde_json::json;
Expand Down Expand Up @@ -94,6 +97,146 @@ async fn wait_for_new_pid(
.context("timed out waiting for a new MCP server process")
}

#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn mcp_calls_stay_bound_to_each_thread() -> anyhow::Result<()> {
skip_if_wine_exec!(
Ok(()),
"requires a Windows test_stdio_server in the Wine-exec environment"
);
skip_if_no_network!(Ok(()));

let responses_server = responses::start_mock_server().await;
let command = remote_aware_stdio_server_bin()?;
let environment_id = remote_aware_environment_id();
let make_server = |marker| {
serde_json::from_value::<McpServerConfig>(json!({
"command": command,
"environment_id": environment_id,
"env": {
"MCP_TEST_DYNAMIC_SERVER_METADATA": "1",
"MCP_TEST_VALUE": marker,
},
"enabled_tools": ["echo"],
"startup_timeout_sec": 10,
}))
};
let first_server = make_server("first-runtime")?;
let second_server = make_server("second-runtime")?;
let fixture = test_codex()
.with_model_info_override("gpt-5.4", |model| model.supports_search_tool = false)
.with_config(move |config| {
config.permissions.approval_policy = Constrained::allow_any(AskForApproval::Never);
config
.permissions
.set_permission_profile(PermissionProfile::Disabled)
.expect("first thread should accept disabled permissions");
let mut servers = config.mcp_servers.get().clone();
servers.insert(SERVER_NAME.to_string(), first_server);
config
.mcp_servers
.set(servers)
.expect("first thread should accept its MCP servers");
})
.build_with_auto_env(&responses_server)
.await?;

let mut second_config = fixture.config.clone();
let mut second_servers = second_config.mcp_servers.get().clone();
second_servers.insert(SERVER_NAME.to_string(), second_server);
second_config.mcp_servers.set(second_servers)?;
let NewThread {
thread: second_thread,
..
} = fixture
.thread_manager
.start_thread(StartThreadOptions::new(second_config))
.await?;

wait_for_mcp_server(&fixture.codex, SERVER_NAME).await?;
wait_for_mcp_server(&second_thread, SERVER_NAME).await?;

let calls = [
(&fixture.codex, "first-call", "first-runtime"),
(&second_thread, "second-call", "second-runtime"),
(&fixture.codex, "first-again", "first-runtime"),
];
let mut processes = Vec::new();
for (thread, call_id, marker) in calls {
let call_response = mount_sse_once(
&responses_server,
responses::sse(vec![
responses::ev_response_created(call_id),
responses::ev_function_call_with_namespace(
call_id,
NAMESPACE,
"echo",
&json!({ "message": call_id }).to_string(),
),
responses::ev_completed(call_id),
]),
)
.await;
let completion_response = mount_sse_once(
&responses_server,
responses::sse(vec![
responses::ev_response_created(&format!("{call_id}-done")),
responses::ev_assistant_message(call_id, "done"),
responses::ev_completed(&format!("{call_id}-done")),
]),
)
.await;
thread
.submit(user_turn(&format!("Call the {SERVER_NAME} echo tool.")))
.await?;
let EventMsg::McpToolCallEnd(end) = wait_for_event(
thread,
|event| matches!(event, EventMsg::McpToolCallEnd(end) if end.call_id == call_id),
)
.await
else {
unreachable!("event predicate guarantees the requested MCP result");
};
assert_eq!(
end.invocation,
McpInvocation {
server: SERVER_NAME.to_string(),
tool: "echo".to_string(),
arguments: Some(json!({ "message": call_id })),
}
);
let content = end
.result
.expect("thread-local MCP call should succeed")
.structured_content
.expect("echo should return structured content");
let process = content
.get("echo")
.and_then(Value::as_str)
.expect("echo should identify its server process")
.to_string();
assert!(process.starts_with("rmcp-test-process-"));
assert_eq!(content, json!({ "echo": process, "env": marker }));
wait_for_event(thread, |event| matches!(event, EventMsg::TurnComplete(_))).await;
let request = call_response.single_request();
assert!(request.tool_by_name(NAMESPACE, "echo").is_some());
let output = completion_response
.single_request()
.function_call_output_text(call_id)
.expect("MCP result should be returned to the model");
assert!(output.contains(&process));
assert!(output.contains(marker));
processes.push(process);
}

assert_ne!(processes[0], processes[1]);
assert_eq!(processes[0], processes[2]);

fixture.codex.shutdown_and_wait().await?;
second_thread.shutdown_and_wait().await?;
responses_server.verify().await;
Ok(())
}

#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn regular_mcp_definition_cache_preserves_live_session_state() -> anyhow::Result<()> {
skip_if_wine_exec!(
Expand Down
Loading