Avoid full MCP reload after OAuth login

This commit is contained in:
Steven Lee
2026-06-18 21:00:43 +00:00
parent f089e6ef96
commit d5b2af01cb
2 changed files with 24 additions and 175 deletions

View File

@@ -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<ThreadManager>,
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 {

View File

@@ -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<anyhow::Result<()>>,
outgoing: Arc<OutgoingMessageSender>,
thread_manager: Arc<ThreadManager>,
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<anyhow::Result<()>>,
outgoing: Arc<OutgoingMessageSender>,
thread_manager: Arc<ThreadManager>,
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<serde_json::Value>,
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<OutgoingMessageSender>,
receiver: mpsc::Receiver<OutgoingEnvelope>,
}
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<OutgoingEnvelope>,
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),
}
);
}
}