From aa93b0ad1cdac09525249b86b97d707249aac41f Mon Sep 17 00:00:00 2001 From: Michael Bolin Date: Thu, 19 Feb 2026 10:40:59 -0800 Subject: [PATCH] chore: consolidate new() and initialize() for McpConnectionManager --- codex-rs/core/src/codex.rs | 94 ++++++++++++--------- codex-rs/core/src/connectors.rs | 21 ++--- codex-rs/core/src/mcp/mod.rs | 21 ++--- codex-rs/core/src/mcp_connection_manager.rs | 29 +++++-- 4 files changed, 90 insertions(+), 75 deletions(-) diff --git a/codex-rs/core/src/codex.rs b/codex-rs/core/src/codex.rs index 6fa5ae3b80..e909cdb3cd 100644 --- a/codex-rs/core/src/codex.rs +++ b/codex-rs/core/src/codex.rs @@ -1263,7 +1263,9 @@ impl Session { }; let services = SessionServices { - mcp_connection_manager: Arc::new(RwLock::new(McpConnectionManager::default())), + mcp_connection_manager: Arc::new( + RwLock::new(McpConnectionManager::new_uninitialized()), + ), mcp_startup_cancellation_token: Mutex::new(CancellationToken::new()), unified_exec_manager: UnifiedExecProcessManager::default(), analytics_events_client: AnalyticsEventsClient::new( @@ -1364,7 +1366,7 @@ impl Session { // Start the watcher after SessionConfigured so it cannot emit earlier events. sess.start_file_watcher_listener(); - // Construct sandbox_state before initialize() so it can be sent to each + // Construct sandbox_state before MCP startup so it can be sent to each // MCP server immediately after it becomes ready (avoiding blocking). let sandbox_state = SandboxState { sandbox_policy: session_configuration.sandbox_policy.get().clone(), @@ -1378,21 +1380,30 @@ impl Session { .map(|(name, _)| name.clone()) .collect(); required_mcp_servers.sort(); - let cancel_token = sess.mcp_startup_cancellation_token().await; - - sess.services - .mcp_connection_manager - .write() - .await - .initialize( - &mcp_servers, - config.mcp_oauth_credentials_store_mode, - auth_statuses.clone(), - tx_event.clone(), - cancel_token, - sandbox_state, - ) - .await; + { + let mut cancel_guard = sess.services.mcp_startup_cancellation_token.lock().await; + cancel_guard.cancel(); + *cancel_guard = CancellationToken::new(); + } + let (mcp_connection_manager, cancel_token) = McpConnectionManager::new( + &mcp_servers, + config.mcp_oauth_credentials_store_mode, + auth_statuses.clone(), + tx_event.clone(), + sandbox_state, + ) + .await; + { + let mut manager_guard = sess.services.mcp_connection_manager.write().await; + *manager_guard = mcp_connection_manager; + } + { + let mut cancel_guard = sess.services.mcp_startup_cancellation_token.lock().await; + if cancel_guard.is_cancelled() { + cancel_token.cancel(); + } + *cancel_guard = cancel_token; + } if !required_mcp_servers.is_empty() { let failures = sess .services @@ -2995,19 +3006,26 @@ impl Session { sandbox_cwd: turn_context.cwd.clone(), use_linux_sandbox_bwrap: turn_context.features.enabled(Feature::UseLinuxSandboxBwrap), }; - let cancel_token = self.reset_mcp_startup_cancellation_token().await; - - let mut refreshed_manager = McpConnectionManager::default(); - refreshed_manager - .initialize( - &mcp_servers, - store_mode, - auth_statuses, - self.get_tx_event(), - cancel_token, - sandbox_state, - ) - .await; + { + let mut cancel_guard = self.services.mcp_startup_cancellation_token.lock().await; + cancel_guard.cancel(); + *cancel_guard = CancellationToken::new(); + } + let (refreshed_manager, cancel_token) = McpConnectionManager::new( + &mcp_servers, + store_mode, + auth_statuses, + self.get_tx_event(), + sandbox_state, + ) + .await; + { + let mut cancel_guard = self.services.mcp_startup_cancellation_token.lock().await; + if cancel_guard.is_cancelled() { + cancel_token.cancel(); + } + *cancel_guard = cancel_token; + } let mut manager = self.services.mcp_connection_manager.write().await; *manager = refreshed_manager; @@ -3064,14 +3082,6 @@ impl Session { .clone() } - async fn reset_mcp_startup_cancellation_token(&self) -> CancellationToken { - let mut guard = self.services.mcp_startup_cancellation_token.lock().await; - guard.cancel(); - let cancel_token = CancellationToken::new(); - *guard = cancel_token.clone(); - cancel_token - } - fn show_raw_agent_reasoning(&self) -> bool { self.services.show_raw_agent_reasoning } @@ -7265,7 +7275,9 @@ mod tests { let file_watcher = Arc::new(FileWatcher::noop()); let services = SessionServices { - mcp_connection_manager: Arc::new(RwLock::new(McpConnectionManager::default())), + mcp_connection_manager: Arc::new(RwLock::new( + McpConnectionManager::new_mcp_connection_manager_for_tests(), + )), mcp_startup_cancellation_token: Mutex::new(CancellationToken::new()), unified_exec_manager: UnifiedExecProcessManager::default(), analytics_events_client: AnalyticsEventsClient::new( @@ -7413,7 +7425,9 @@ mod tests { let file_watcher = Arc::new(FileWatcher::noop()); let services = SessionServices { - mcp_connection_manager: Arc::new(RwLock::new(McpConnectionManager::default())), + mcp_connection_manager: Arc::new(RwLock::new( + McpConnectionManager::new_mcp_connection_manager_for_tests(), + )), mcp_startup_cancellation_token: Mutex::new(CancellationToken::new()), unified_exec_manager: UnifiedExecProcessManager::default(), analytics_events_client: AnalyticsEventsClient::new( diff --git a/codex-rs/core/src/connectors.rs b/codex-rs/core/src/connectors.rs index 5a8930231b..ebad8cd3b7 100644 --- a/codex-rs/core/src/connectors.rs +++ b/codex-rs/core/src/connectors.rs @@ -12,7 +12,6 @@ pub use codex_app_server_protocol::AppInfo; pub use codex_app_server_protocol::AppMetadata; use codex_protocol::protocol::SandboxPolicy; use serde::Deserialize; -use tokio_util::sync::CancellationToken; use tracing::warn; use crate::AuthManager; @@ -91,10 +90,8 @@ pub async fn list_accessible_connectors_from_mcp_tools_with_options( let auth_status_entries = compute_auth_statuses(mcp_servers.iter(), config.mcp_oauth_credentials_store_mode).await; - let mut mcp_connection_manager = McpConnectionManager::default(); let (tx_event, rx_event) = unbounded(); drop(rx_event); - let cancel_token = CancellationToken::new(); let sandbox_state = SandboxState { sandbox_policy: SandboxPolicy::new_read_only_policy(), @@ -103,16 +100,14 @@ pub async fn list_accessible_connectors_from_mcp_tools_with_options( use_linux_sandbox_bwrap: config.features.enabled(Feature::UseLinuxSandboxBwrap), }; - mcp_connection_manager - .initialize( - &mcp_servers, - config.mcp_oauth_credentials_store_mode, - auth_status_entries, - tx_event, - cancel_token.clone(), - sandbox_state, - ) - .await; + let (mcp_connection_manager, cancel_token) = McpConnectionManager::new( + &mcp_servers, + config.mcp_oauth_credentials_store_mode, + auth_status_entries, + tx_event, + sandbox_state, + ) + .await; if force_refetch && let Err(err) = mcp_connection_manager diff --git a/codex-rs/core/src/mcp/mod.rs b/codex-rs/core/src/mcp/mod.rs index 1365b5da8b..84b979499e 100644 --- a/codex-rs/core/src/mcp/mod.rs +++ b/codex-rs/core/src/mcp/mod.rs @@ -14,7 +14,6 @@ use codex_protocol::mcp::Tool; use codex_protocol::protocol::McpListToolsResponseEvent; use codex_protocol::protocol::SandboxPolicy; use serde_json::Value; -use tokio_util::sync::CancellationToken; use crate::AuthManager; use crate::CodexAuth; @@ -191,10 +190,8 @@ pub async fn collect_mcp_snapshot(config: &Config) -> McpListToolsResponseEvent let auth_status_entries = compute_auth_statuses(mcp_servers.iter(), config.mcp_oauth_credentials_store_mode).await; - let mut mcp_connection_manager = McpConnectionManager::default(); let (tx_event, rx_event) = unbounded(); drop(rx_event); - let cancel_token = CancellationToken::new(); // Use ReadOnly sandbox policy for MCP snapshot collection (safest default) let sandbox_state = SandboxState { @@ -204,16 +201,14 @@ pub async fn collect_mcp_snapshot(config: &Config) -> McpListToolsResponseEvent use_linux_sandbox_bwrap: config.features.enabled(Feature::UseLinuxSandboxBwrap), }; - mcp_connection_manager - .initialize( - &mcp_servers, - config.mcp_oauth_credentials_store_mode, - auth_status_entries.clone(), - tx_event, - cancel_token.clone(), - sandbox_state, - ) - .await; + let (mcp_connection_manager, cancel_token) = McpConnectionManager::new( + &mcp_servers, + config.mcp_oauth_credentials_store_mode, + auth_status_entries.clone(), + tx_event, + sandbox_state, + ) + .await; let snapshot = collect_mcp_snapshot_from_manager(&mcp_connection_manager, auth_status_entries).await; diff --git a/codex-rs/core/src/mcp_connection_manager.rs b/codex-rs/core/src/mcp_connection_manager.rs index 589be7aac2..7af11bf5aa 100644 --- a/codex-rs/core/src/mcp_connection_manager.rs +++ b/codex-rs/core/src/mcp_connection_manager.rs @@ -353,22 +353,30 @@ pub(crate) struct McpConnectionManager { } impl McpConnectionManager { + pub(crate) fn new_uninitialized() -> Self { + Self { + clients: HashMap::new(), + elicitation_requests: ElicitationRequestManager::default(), + } + } + + #[cfg(test)] + pub(crate) fn new_mcp_connection_manager_for_tests() -> Self { + Self::new_uninitialized() + } pub(crate) fn has_servers(&self) -> bool { !self.clients.is_empty() } - pub async fn initialize( - &mut self, + #[allow(clippy::new_ret_no_self)] + pub async fn new( mcp_servers: &HashMap, store_mode: OAuthCredentialsStoreMode, auth_entries: HashMap, tx_event: Sender, - cancel_token: CancellationToken, initial_sandbox_state: SandboxState, - ) { - if cancel_token.is_cancelled() { - return; - } + ) -> (Self, CancellationToken) { + let cancel_token = CancellationToken::new(); let mut clients = HashMap::new(); let mut join_set = JoinSet::new(); let elicitation_requests = ElicitationRequestManager::default(); @@ -435,8 +443,10 @@ impl McpConnectionManager { (server_name, outcome) }); } - self.clients = clients; - self.elicitation_requests = elicitation_requests.clone(); + let manager = Self { + clients, + elicitation_requests: elicitation_requests.clone(), + }; tokio::spawn(async move { let outcomes = join_set.join_all().await; let mut summary = McpStartupCompleteEvent::default(); @@ -459,6 +469,7 @@ impl McpConnectionManager { }) .await; }); + (manager, cancel_token) } async fn client_by_name(&self, name: &str) -> Result {