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,
        }
    }
}