diff --git a/codex-rs/app-server/src/mcp_refresh.rs b/codex-rs/app-server/src/mcp_refresh.rs index 884e73d0a9..84a6b8e83b 100644 --- a/codex-rs/app-server/src/mcp_refresh.rs +++ b/codex-rs/app-server/src/mcp_refresh.rs @@ -93,7 +93,7 @@ async fn queue_refresh( } #[cfg(test)] -pub(crate) mod tests { +mod tests { use super::*; use crate::extensions::ThreadExtensionDependencies; use crate::extensions::guardian_agent_spawner; @@ -177,7 +177,7 @@ pub(crate) mod tests { Ok(()) } - pub(crate) async fn refresh_test_state() -> anyhow::Result<( + async fn refresh_test_state() -> anyhow::Result<( TempDir, Arc, ConfigManager, @@ -275,11 +275,11 @@ pub(crate) mod tests { Ok((temp_dir, thread_manager, config_manager, loader)) } - pub(crate) struct CountingThreadConfigLoader { + struct CountingThreadConfigLoader { good_cwd: AbsolutePathBuf, bad_cwd: AbsolutePathBuf, - pub(crate) good_loads: AtomicUsize, - pub(crate) bad_loads: AtomicUsize, + good_loads: AtomicUsize, + bad_loads: AtomicUsize, } impl CountingThreadConfigLoader { diff --git a/codex-rs/app-server/src/request_processors/mcp_processor.rs b/codex-rs/app-server/src/request_processors/mcp_processor.rs index d9bdf186d0..08d464598b 100644 --- a/codex-rs/app-server/src/request_processors/mcp_processor.rs +++ b/codex-rs/app-server/src/request_processors/mcp_processor.rs @@ -173,14 +173,25 @@ impl McpRequestProcessor { ) .await .map_err(|err| internal_error(format!("failed to login to MCP server '{name}': {err}")))?; - let (authorization_url, completion) = handle.into_parts(); - spawn_mcp_oauth_login_completion( - name.clone(), - completion, - Arc::clone(&self.outgoing), - Arc::clone(&self.thread_manager), - self.config_manager.clone(), - ); + let authorization_url = handle.authorization_url().to_string(); + let notification_name = name.clone(); + let outgoing = Arc::clone(&self.outgoing); + + tokio::spawn(async move { + let (success, error) = match handle.wait().await { + Ok(()) => (true, None), + Err(err) => (false, Some(err.to_string())), + }; + + let notification = ServerNotification::McpServerOauthLoginCompleted( + McpServerOauthLoginCompletedNotification { + name: notification_name, + success, + error, + }, + ); + outgoing.send_server_notification(notification).await; + }); Ok(McpServerOauthLoginResponse { authorization_url }) } @@ -432,54 +443,6 @@ impl McpRequestProcessor { } } -fn spawn_mcp_oauth_login_completion( - notification_name: String, - completion: oneshot::Receiver>, - outgoing: Arc, - thread_manager: Arc, - config_manager: ConfigManager, -) { - tokio::spawn(async move { - finish_mcp_oauth_login_completion( - notification_name, - completion, - outgoing, - thread_manager, - config_manager, - ) - .await; - }); -} - -async fn finish_mcp_oauth_login_completion( - notification_name: String, - completion: oneshot::Receiver>, - outgoing: Arc, - thread_manager: Arc, - config_manager: ConfigManager, -) { - let (success, error) = match completion.await { - Ok(Ok(())) => { - crate::mcp_refresh::queue_best_effort_refresh(&thread_manager, &config_manager).await; - (true, None) - } - Ok(Err(err)) => (false, Some(err.to_string())), - Err(err) => ( - false, - Some(format!("OAuth login task was cancelled: {err}")), - ), - }; - - let notification = ServerNotification::McpServerOauthLoginCompleted( - McpServerOauthLoginCompletedNotification { - name: notification_name, - success, - error, - }, - ); - outgoing.send_server_notification(notification).await; -} - fn with_mcp_tool_call_thread_id_meta( meta: Option, thread_id: &str, @@ -503,117 +466,3 @@ fn with_mcp_tool_call_thread_id_meta( other => other, } } - -#[cfg(test)] -mod tests { - use super::*; - use crate::outgoing_message::OutgoingEnvelope; - use crate::outgoing_message::OutgoingMessage; - use pretty_assertions::assert_eq; - use std::sync::atomic::Ordering; - use tokio::sync::mpsc; - - #[tokio::test] - async fn successful_oauth_login_completion_queues_mcp_refresh() -> anyhow::Result<()> { - let (_temp_dir, thread_manager, config_manager, loader) = - crate::mcp_refresh::tests::refresh_test_state().await?; - let outgoing = outgoing_sender(); - let (completion_tx, completion_rx) = oneshot::channel(); - completion_tx - .send(Ok(())) - .expect("completion receiver should be active"); - - finish_mcp_oauth_login_completion( - "test-server".to_string(), - completion_rx, - Arc::clone(&outgoing.sender), - thread_manager, - config_manager, - ) - .await; - - assert_eq!(loader.good_loads.load(Ordering::Relaxed), 1); - assert_eq!(loader.bad_loads.load(Ordering::Relaxed), 1); - assert_oauth_completion( - outgoing.receiver, - /*success*/ true, - /*error*/ None, - ) - .await; - Ok(()) - } - - #[tokio::test] - async fn failed_oauth_login_completion_does_not_queue_mcp_refresh() -> anyhow::Result<()> { - let (_temp_dir, thread_manager, config_manager, loader) = - crate::mcp_refresh::tests::refresh_test_state().await?; - let outgoing = outgoing_sender(); - let (completion_tx, completion_rx) = oneshot::channel(); - completion_tx - .send(Err(anyhow::anyhow!("login failed"))) - .expect("completion receiver should be active"); - - finish_mcp_oauth_login_completion( - "test-server".to_string(), - completion_rx, - Arc::clone(&outgoing.sender), - thread_manager, - config_manager, - ) - .await; - - assert_eq!(loader.good_loads.load(Ordering::Relaxed), 0); - assert_eq!(loader.bad_loads.load(Ordering::Relaxed), 0); - assert_oauth_completion( - outgoing.receiver, - /*success*/ false, - Some("login failed"), - ) - .await; - Ok(()) - } - - struct TestOutgoing { - sender: Arc, - receiver: mpsc::Receiver, - } - - fn outgoing_sender() -> TestOutgoing { - let (tx, rx) = mpsc::channel(4); - TestOutgoing { - sender: Arc::new(OutgoingMessageSender::new( - tx, - codex_analytics::AnalyticsEventsClient::disabled(), - )), - receiver: rx, - } - } - - async fn assert_oauth_completion( - mut rx: mpsc::Receiver, - success: bool, - error: Option<&str>, - ) { - let envelope = rx - .recv() - .await - .expect("completion notification should send"); - let OutgoingEnvelope::Broadcast { - message: - OutgoingMessage::AppServerNotification( - ServerNotification::McpServerOauthLoginCompleted(notification), - ), - } = envelope - else { - panic!("unexpected outgoing envelope: {envelope:?}"); - }; - assert_eq!( - notification, - McpServerOauthLoginCompletedNotification { - name: "test-server".to_string(), - success, - error: error.map(str::to_string), - } - ); - } -}