diff --git a/codex-rs/Cargo.lock b/codex-rs/Cargo.lock index 0b38db1eb4..5e0656fbe2 100644 --- a/codex-rs/Cargo.lock +++ b/codex-rs/Cargo.lock @@ -887,6 +887,7 @@ dependencies = [ "codex-file-search", "codex-login", "codex-protocol", + "codex-rmcp-client", "codex-utils-json-to-toml", "core_test_support", "mcp-types", diff --git a/codex-rs/app-server-protocol/src/protocol/common.rs b/codex-rs/app-server-protocol/src/protocol/common.rs index 2858366739..105f97b292 100644 --- a/codex-rs/app-server-protocol/src/protocol/common.rs +++ b/codex-rs/app-server-protocol/src/protocol/common.rs @@ -139,6 +139,11 @@ client_request_definitions! { response: v2::ModelListResponse, }, + McpServerOauthLogin => "mcpServer/oauthLogin" { + params: v2::McpServerOauthLoginParams, + response: v2::McpServerOauthLoginResponse, + }, + McpServersList => "mcpServers/list" { params: v2::ListMcpServersParams, response: v2::ListMcpServersResponse, diff --git a/codex-rs/app-server-protocol/src/protocol/v2.rs b/codex-rs/app-server-protocol/src/protocol/v2.rs index ea70b805b0..468ca834c3 100644 --- a/codex-rs/app-server-protocol/src/protocol/v2.rs +++ b/codex-rs/app-server-protocol/src/protocol/v2.rs @@ -688,6 +688,23 @@ pub struct ListMcpServersResponse { pub next_cursor: Option, } +#[derive(Serialize, Deserialize, Debug, Clone, PartialEq, JsonSchema, TS)] +#[serde(rename_all = "camelCase")] +#[ts(export_to = "v2/")] +pub struct McpServerOauthLoginParams { + pub name: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + #[ts(optional)] + pub scopes: Option>, +} + +#[derive(Serialize, Deserialize, Debug, Clone, PartialEq, JsonSchema, TS)] +#[serde(rename_all = "camelCase")] +#[ts(export_to = "v2/")] +pub struct McpServerOauthLoginResponse { + pub authorization_url: String, +} + #[derive(Serialize, Deserialize, Debug, Clone, PartialEq, JsonSchema, TS)] #[serde(rename_all = "camelCase")] #[ts(export_to = "v2/")] diff --git a/codex-rs/app-server/Cargo.toml b/codex-rs/app-server/Cargo.toml index 99d5a7a141..e4a326a2c3 100644 --- a/codex-rs/app-server/Cargo.toml +++ b/codex-rs/app-server/Cargo.toml @@ -26,6 +26,7 @@ codex-login = { workspace = true } codex-protocol = { workspace = true } codex-app-server-protocol = { workspace = true } codex-feedback = { workspace = true } +codex-rmcp-client = { workspace = true } codex-utils-json-to-toml = { workspace = true } chrono = { workspace = true } serde = { workspace = true, features = ["derive"] } diff --git a/codex-rs/app-server/src/codex_message_processor.rs b/codex-rs/app-server/src/codex_message_processor.rs index 65721a698e..59b19f1c6b 100644 --- a/codex-rs/app-server/src/codex_message_processor.rs +++ b/codex-rs/app-server/src/codex_message_processor.rs @@ -55,6 +55,8 @@ use codex_app_server_protocol::LoginChatGptResponse; use codex_app_server_protocol::LogoutAccountResponse; use codex_app_server_protocol::LogoutChatGptResponse; use codex_app_server_protocol::McpServer; +use codex_app_server_protocol::McpServerOauthLoginParams; +use codex_app_server_protocol::McpServerOauthLoginResponse; use codex_app_server_protocol::ModelListParams; use codex_app_server_protocol::ModelListResponse; use codex_app_server_protocol::NewConversationParams; @@ -115,6 +117,7 @@ use codex_core::config::Config; use codex_core::config::ConfigOverrides; use codex_core::config::ConfigToml; use codex_core::config::edit::ConfigEditsBuilder; +use codex_core::config::types::McpServerTransportConfig; use codex_core::config_loader::load_config_as_toml; use codex_core::default_client::get_codex_user_agent; use codex_core::exec::ExecParams; @@ -147,6 +150,7 @@ use codex_protocol::protocol::RolloutItem; use codex_protocol::protocol::SessionMetaLine; use codex_protocol::protocol::USER_MESSAGE_BEGIN; use codex_protocol::user_input::UserInput as CoreInputItem; +use codex_rmcp_client::perform_oauth_login_return_url; use codex_utils_json_to_toml::json_to_toml; use std::collections::HashMap; use std::collections::HashSet; @@ -369,6 +373,9 @@ impl CodexMessageProcessor { ClientRequest::ModelList { request_id, params } => { self.list_models(request_id, params).await; } + ClientRequest::McpServerOauthLogin { request_id, params } => { + self.mcp_server_oauth_login(request_id, params).await; + } ClientRequest::McpServersList { request_id, params } => { self.list_mcp_servers(request_id, params).await; } @@ -1916,6 +1923,77 @@ impl CodexMessageProcessor { self.outgoing.send_response(request_id, response).await; } + async fn mcp_server_oauth_login( + &self, + request_id: RequestId, + params: McpServerOauthLoginParams, + ) { + if !self.config.features.enabled(Feature::RmcpClient) { + let error = JSONRPCErrorError { + code: INVALID_REQUEST_ERROR_CODE, + message: "OAuth login is only supported when [features].rmcp_client is true in config.toml".to_string(), + data: None, + }; + self.outgoing.send_error(request_id, error).await; + return; + } + + let McpServerOauthLoginParams { name, scopes } = params; + + let Some(server) = self.config.mcp_servers.get(&name) else { + let error = JSONRPCErrorError { + code: INVALID_REQUEST_ERROR_CODE, + message: format!("No MCP server named '{name}' found."), + data: None, + }; + self.outgoing.send_error(request_id, error).await; + return; + }; + + let (url, http_headers, env_http_headers) = match &server.transport { + McpServerTransportConfig::StreamableHttp { + url, + http_headers, + env_http_headers, + .. + } => (url.clone(), http_headers.clone(), env_http_headers.clone()), + _ => { + let error = JSONRPCErrorError { + code: INVALID_REQUEST_ERROR_CODE, + message: "OAuth login is only supported for streamable HTTP servers." + .to_string(), + data: None, + }; + self.outgoing.send_error(request_id, error).await; + return; + } + }; + + match perform_oauth_login_return_url( + &name, + &url, + self.config.mcp_oauth_credentials_store_mode, + http_headers, + env_http_headers, + scopes.as_deref().unwrap_or_default(), + ) + .await + { + Ok(authorization_url) => { + let response = McpServerOauthLoginResponse { authorization_url }; + self.outgoing.send_response(request_id, response).await; + } + Err(err) => { + let error = JSONRPCErrorError { + code: INTERNAL_ERROR_CODE, + message: format!("failed to login to MCP server '{name}': {err}"), + data: None, + }; + self.outgoing.send_error(request_id, error).await; + } + } + } + async fn list_mcp_servers(&self, request_id: RequestId, params: ListMcpServersParams) { let snapshot = collect_mcp_snapshot(self.config.as_ref()).await; diff --git a/codex-rs/rmcp-client/src/lib.rs b/codex-rs/rmcp-client/src/lib.rs index ac617f3d29..6c74ead7fe 100644 --- a/codex-rs/rmcp-client/src/lib.rs +++ b/codex-rs/rmcp-client/src/lib.rs @@ -17,6 +17,7 @@ pub use oauth::delete_oauth_tokens; pub(crate) use oauth::load_oauth_tokens; pub use oauth::save_oauth_tokens; pub use perform_oauth_login::perform_oauth_login; +pub use perform_oauth_login::perform_oauth_login_return_url; pub use rmcp::model::ElicitationAction; pub use rmcp_client::Elicitation; pub use rmcp_client::ElicitationResponse; diff --git a/codex-rs/rmcp-client/src/perform_oauth_login.rs b/codex-rs/rmcp-client/src/perform_oauth_login.rs index d8ffdd3949..2e8d7adaed 100644 --- a/codex-rs/rmcp-client/src/perform_oauth_login.rs +++ b/codex-rs/rmcp-client/src/perform_oauth_login.rs @@ -40,70 +40,44 @@ pub async fn perform_oauth_login( env_http_headers: Option>, scopes: &[String], ) -> Result<()> { - let server = Arc::new(Server::http("127.0.0.1:0").map_err(|err| anyhow!(err))?); - let guard = CallbackServerGuard { - server: Arc::clone(&server), - }; + OauthLoginFlow::new( + server_name, + server_url, + store_mode, + http_headers, + env_http_headers, + scopes, + true, + ) + .await? + .finish() + .await +} - let redirect_uri = match server.server_addr() { - tiny_http::ListenAddr::IP(std::net::SocketAddr::V4(addr)) => { - format!("http://{}:{}/callback", addr.ip(), addr.port()) - } - tiny_http::ListenAddr::IP(std::net::SocketAddr::V6(addr)) => { - format!("http://[{}]:{}/callback", addr.ip(), addr.port()) - } - #[cfg(not(target_os = "windows"))] - _ => return Err(anyhow!("unable to determine callback address")), - }; +pub async fn perform_oauth_login_return_url( + server_name: &str, + server_url: &str, + store_mode: OAuthCredentialsStoreMode, + http_headers: Option>, + env_http_headers: Option>, + scopes: &[String], +) -> Result { + let flow = OauthLoginFlow::new( + server_name, + server_url, + store_mode, + http_headers, + env_http_headers, + scopes, + false, + ) + .await?; - let (tx, rx) = oneshot::channel(); - spawn_callback_server(server, tx); + let auth_url = flow.authorization_url(); + let server_name_for_logging = flow.server_name.clone(); + flow.spawn(server_name_for_logging); - let default_headers = build_default_headers(http_headers, env_http_headers)?; - let http_client = apply_default_headers(ClientBuilder::new(), &default_headers).build()?; - - let mut oauth_state = OAuthState::new(server_url, Some(http_client)).await?; - let scope_refs: Vec<&str> = scopes.iter().map(String::as_str).collect(); - oauth_state - .start_authorization(&scope_refs, &redirect_uri, Some("Codex")) - .await?; - let auth_url = oauth_state.get_authorization_url().await?; - - println!("Authorize `{server_name}` by opening this URL in your browser:\n{auth_url}\n"); - - if webbrowser::open(&auth_url).is_err() { - println!("(Browser launch failed; please copy the URL above manually.)"); - } - - let (code, csrf_state) = timeout(Duration::from_secs(300), rx) - .await - .context("timed out waiting for OAuth callback")? - .context("OAuth callback was cancelled")?; - - oauth_state - .handle_callback(&code, &csrf_state) - .await - .context("failed to handle OAuth callback")?; - - let (client_id, credentials_opt) = oauth_state - .get_credentials() - .await - .context("failed to retrieve OAuth credentials")?; - let credentials = - credentials_opt.ok_or_else(|| anyhow!("OAuth provider did not return credentials"))?; - - let expires_at = compute_expires_at_millis(&credentials); - let stored = StoredOAuthTokens { - server_name: server_name.to_string(), - url: server_url.to_string(), - client_id, - token_response: WrappedOAuthTokenResponse(credentials), - expires_at, - }; - save_oauth_tokens(server_name, &stored, store_mode)?; - - drop(guard); - Ok(()) + Ok(auth_url) } fn spawn_callback_server(server: Arc, tx: oneshot::Sender<(String, String)>) { @@ -160,3 +134,134 @@ fn parse_oauth_callback(path: &str) -> Option { state: state?, }) } + +struct OauthLoginFlow { + auth_url: String, + oauth_state: OAuthState, + rx: oneshot::Receiver<(String, String)>, + guard: CallbackServerGuard, + server_name: String, + server_url: String, + store_mode: OAuthCredentialsStoreMode, + launch_browser: bool, +} + +impl OauthLoginFlow { + async fn new( + server_name: &str, + server_url: &str, + store_mode: OAuthCredentialsStoreMode, + http_headers: Option>, + env_http_headers: Option>, + scopes: &[String], + launch_browser: bool, + ) -> Result { + let server = Arc::new(Server::http("127.0.0.1:0").map_err(|err| anyhow!(err))?); + let guard = CallbackServerGuard { + server: Arc::clone(&server), + }; + + let redirect_uri = match server.server_addr() { + tiny_http::ListenAddr::IP(std::net::SocketAddr::V4(addr)) => { + let ip = addr.ip(); + let port = addr.port(); + format!("http://{ip}:{port}/callback") + } + tiny_http::ListenAddr::IP(std::net::SocketAddr::V6(addr)) => { + let ip = addr.ip(); + let port = addr.port(); + format!("http://[{ip}]:{port}/callback") + } + #[cfg(not(target_os = "windows"))] + _ => return Err(anyhow!("unable to determine callback address")), + }; + + let (tx, rx) = oneshot::channel(); + spawn_callback_server(server, tx); + + let default_headers = build_default_headers(http_headers, env_http_headers)?; + let http_client = apply_default_headers(ClientBuilder::new(), &default_headers).build()?; + + let mut oauth_state = OAuthState::new(server_url, Some(http_client)).await?; + let scope_refs: Vec<&str> = scopes.iter().map(String::as_str).collect(); + oauth_state + .start_authorization(&scope_refs, &redirect_uri, Some("Codex")) + .await?; + let auth_url = oauth_state.get_authorization_url().await?; + + Ok(Self { + auth_url, + oauth_state, + rx, + guard, + server_name: server_name.to_string(), + server_url: server_url.to_string(), + store_mode, + launch_browser, + }) + } + + fn authorization_url(&self) -> String { + self.auth_url.clone() + } + + async fn finish(mut self) -> Result<()> { + if self.launch_browser { + let server_name = &self.server_name; + let auth_url = &self.auth_url; + println!( + "Authorize `{server_name}` by opening this URL in your browser:\n{auth_url}\n" + ); + + if webbrowser::open(auth_url).is_err() { + println!("(Browser launch failed; please copy the URL above manually.)"); + } + } + + let result = async { + let (code, csrf_state) = timeout(Duration::from_secs(300), &mut self.rx) + .await + .context("timed out waiting for OAuth callback")? + .context("OAuth callback was cancelled")?; + + self.oauth_state + .handle_callback(&code, &csrf_state) + .await + .context("failed to handle OAuth callback")?; + + let (client_id, credentials_opt) = self + .oauth_state + .get_credentials() + .await + .context("failed to retrieve OAuth credentials")?; + let credentials = credentials_opt + .ok_or_else(|| anyhow!("OAuth provider did not return credentials"))?; + + let expires_at = compute_expires_at_millis(&credentials); + let stored = StoredOAuthTokens { + server_name: self.server_name.clone(), + url: self.server_url.clone(), + client_id, + token_response: WrappedOAuthTokenResponse(credentials), + expires_at, + }; + save_oauth_tokens(&self.server_name, &stored, self.store_mode)?; + + Ok(()) + } + .await; + + drop(self.guard); + result + } + + fn spawn(self, server_name_for_logging: String) { + tokio::spawn(async move { + if let Err(err) = self.finish().await { + eprintln!( + "Failed to complete OAuth login for '{server_name_for_logging}': {err:#}" + ); + } + }); + } +}