File size: 6,363 Bytes
17f328f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 | //! Opens MCP event streams with clients that remain connected after task unloading.
use std::sync::Arc;
use anyhow::Result;
use anyhow::bail;
use codex_api::SharedAuthProvider;
use codex_config::types::AuthKeyringBackendKind;
use codex_config::types::OAuthCredentialsStoreMode;
use codex_exec_server::Environment;
use codex_login::AuthManager;
use codex_login::CodexAuth;
use codex_protocol::mcp::ClientMcpExtensions;
use codex_rmcp_client::ElicitationResponse;
use codex_rmcp_client::McpOAuthRefreshMode;
use rmcp::model::ElicitationAction;
use rmcp::model::ElicitationCapability;
use serde_json::Map;
use serde_json::Value;
use tokio::sync::watch;
use crate::CODEX_APPS_MCP_SERVER_NAME;
use crate::EffectiveMcpServer;
use crate::McpEventStream;
use crate::McpProtocolMode;
use crate::McpRuntimeContext;
use crate::rmcp_client::DEFAULT_STARTUP_TIMEOUT;
use crate::rmcp_client::make_rmcp_client;
use crate::rmcp_client::mcp_initialize_request_params;
pub(crate) struct EventStreamConnectionSettings {
pub server: EffectiveMcpServer,
pub store_mode: OAuthCredentialsStoreMode,
pub keyring_backend_kind: AuthKeyringBackendKind,
pub oauth_refresh_mode: McpOAuthRefreshMode,
pub runtime_context: McpRuntimeContext,
pub resolved_environment: std::result::Result<Option<Arc<Environment>>, String>,
pub auth_provider: Option<SharedAuthProvider>,
pub auth_manager: Option<Arc<AuthManager>>,
pub auth: Option<CodexAuth>,
pub protocol_mode: McpProtocolMode,
pub client_mcp_extensions: ClientMcpExtensions,
}
/// Opens each event stream with its own MCP client.
/// Owners watch `wait_for_access_change` to cancel streams when access changes.
#[derive(Clone)]
pub struct McpEventStreamOpener {
pub(crate) connection: Arc<EventStreamConnectionSettings>,
pub(crate) cancellation_receiver: watch::Receiver<()>,
pub(crate) cancel_event_streams_on_server_removal: watch::Sender<()>,
}
impl McpEventStreamOpener {
/// Retains cancellation for this task's subscriptions across MCP runtime replacement.
pub fn event_stream_cancellation_sender(&self) -> watch::Sender<()> {
self.cancel_event_streams_on_server_removal.clone()
}
/// Creates an MCP client and opens an event stream.
pub async fn open(
&self,
event_name: &str,
arguments: &Value,
request_meta: Option<&Map<String, Value>>,
) -> Result<McpEventStream> {
tokio::select! {
biased;
() = self.wait_for_access_change() => bail!("event subscription access changed"),
result = async {
let connection = &self.connection;
if let Some(manager) = &connection.auth_manager {
let auth = manager.auth().await;
if !self.matches_auth(auth.as_ref()) {
bail!("event subscription account changed");
}
}
let startup_timeout = connection.server.config().startup_timeout_sec
.unwrap_or(DEFAULT_STARTUP_TIMEOUT);
let client = Arc::new(tokio::time::timeout(startup_timeout, make_rmcp_client(
CODEX_APPS_MCP_SERVER_NAME,
connection.server.clone(),
connection.store_mode,
connection.keyring_backend_kind,
connection.oauth_refresh_mode,
connection.runtime_context.clone(),
connection.resolved_environment.clone(),
connection.auth_provider.clone(),
connection.protocol_mode,
)).await??);
client.initialize(
mcp_initialize_request_params(
ElicitationCapability::default(),
connection.client_mcp_extensions.clone(),
),
Some(startup_timeout),
Box::new(|_, _| Box::pin(async {
Ok(ElicitationResponse {
action: ElicitationAction::Decline,
content: None,
meta: None,
})
})),
).await?;
McpEventStream::open(
client,
self.cancellation_receiver.clone(),
event_name,
arguments,
request_meta,
).await
} => result,
}
}
/// Waits for an account change or removal of the event server from the task.
pub async fn wait_for_access_change(&self) {
let mut cancellation_receiver = self.cancellation_receiver.clone();
let auth_change = async {
let Some(manager) = &self.connection.auth_manager else {
return std::future::pending().await;
};
let mut changes = manager.auth_change_receiver();
loop {
if !self.matches_auth(manager.auth_cached().as_ref()) {
return;
}
if changes.changed().await.is_err() {
return;
}
}
};
tokio::select! {
Ok(()) = cancellation_receiver.changed() => {},
() = auth_change => {},
}
}
fn matches_auth(&self, current: Option<&CodexAuth>) -> bool {
match (self.connection.auth.as_ref(), current) {
(Some(CodexAuth::AgentIdentity(expected)), Some(CodexAuth::AgentIdentity(current))) => {
expected.record() == current.record()
}
(Some(CodexAuth::AgentIdentity(_)), _) | (_, Some(CodexAuth::AgentIdentity(_))) => {
false
}
(Some(expected), Some(current)) => {
expected.get_account_id() == current.get_account_id()
&& expected.get_chatgpt_user_id() == current.get_chatgpt_user_id()
&& expected.is_workspace_account() == current.is_workspace_account()
&& expected.is_fedramp_account() == current.is_fedramp_account()
&& (expected.get_account_id().is_some() || expected == current)
}
(None, None) => true,
_ => false,
}
}
}
|