Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 37 additions & 3 deletions codex-rs/rmcp-client/src/http_client_adapter.rs
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@ use codex_exec_server::HttpResponseBodyStream;
use futures::StreamExt;
use futures::stream;
use futures::stream::BoxStream;
use oauth2::AccessToken;
use reqwest::StatusCode;
use reqwest::header::ACCEPT;
use reqwest::header::AUTHORIZATION;
Expand Down Expand Up @@ -61,6 +62,10 @@ pub(crate) struct StreamableHttpClientAdapter {
pub(crate) enum StreamableHttpClientAdapterError {
#[error("streamable HTTP session expired with 404 Not Found")]
SessionExpired404,
#[error("MCP server rejected the access token with HTTP 401 Unauthorized")]
AccessTokenRejected { rejected_access_token: AccessToken },
#[error("MCP OAuth operation failed: {0:#}")]
OAuth(#[source] anyhow::Error),
#[error(transparent)]
HttpRequest(#[from] ExecServerError),
#[error("invalid HTTP header: {0}")]
Expand Down Expand Up @@ -109,7 +114,7 @@ impl StreamableHttpClient for StreamableHttpClientAdapter {
JSON_MIME_TYPE.to_string(),
StreamableHttpClientAdapterError::Header,
)?;
if let Some(auth_token) = auth_token {
if let Some(auth_token) = auth_token.as_deref() {
insert_header(
&mut headers,
AUTHORIZATION,
Expand Down Expand Up @@ -162,6 +167,11 @@ impl StreamableHttpClient for StreamableHttpClientAdapter {
StreamableHttpClientAdapterError::SessionExpired404,
));
}
if response.status == StatusCode::UNAUTHORIZED.as_u16()
&& let Some(error) = access_token_rejected(auth_token.as_deref())
{
return Err(error);
}
if response.status == StatusCode::UNAUTHORIZED.as_u16()
&& let Some(header) =
response_header(&response.headers, reqwest::header::WWW_AUTHENTICATE)
Expand Down Expand Up @@ -240,7 +250,7 @@ impl StreamableHttpClient for StreamableHttpClientAdapter {
let mut headers = self.default_headers.clone();
headers.extend(custom_headers);
self.add_auth_headers(&mut headers);
if let Some(auth_token) = auth_token {
if let Some(auth_token) = auth_token.as_deref() {
insert_header(
&mut headers,
AUTHORIZATION,
Expand Down Expand Up @@ -274,6 +284,11 @@ impl StreamableHttpClient for StreamableHttpClientAdapter {
if response.status == StatusCode::METHOD_NOT_ALLOWED.as_u16() {
return Ok(());
}
if response.status == StatusCode::UNAUTHORIZED.as_u16()
&& let Some(error) = access_token_rejected(auth_token.as_deref())
{
return Err(error);
}
if !status_is_success(response.status) {
return Err(StreamableHttpError::UnexpectedServerResponse(
format!("DELETE returned HTTP {}", response.status).into(),
Expand Down Expand Up @@ -316,7 +331,7 @@ impl StreamableHttpClient for StreamableHttpClientAdapter {
StreamableHttpClientAdapterError::Header,
)?;
}
if let Some(auth_token) = auth_token {
if let Some(auth_token) = auth_token.as_deref() {
insert_header(
&mut headers,
AUTHORIZATION,
Expand Down Expand Up @@ -349,6 +364,11 @@ impl StreamableHttpClient for StreamableHttpClientAdapter {
StreamableHttpClientAdapterError::SessionExpired404,
));
}
if response.status == StatusCode::UNAUTHORIZED.as_u16()
&& let Some(error) = access_token_rejected(auth_token.as_deref())
{
return Err(error);
}
if !status_is_success(response.status) {
return Err(StreamableHttpError::UnexpectedServerResponse(
format!("GET returned HTTP {}", response.status).into(),
Expand All @@ -371,6 +391,20 @@ impl StreamableHttpClient for StreamableHttpClientAdapter {
}
}

fn access_token_rejected(
auth_token: Option<&str>,
) -> Option<StreamableHttpError<StreamableHttpClientAdapterError>> {
// Preserve the token associated with this response. Reading the current credential after a
// delayed 401 is racy: another concurrent request may already have refreshed A to B, in which
// case recovery must retry B rather than refresh B a second time. AccessToken's Debug
// implementation redacts the secret if this error is logged.
auth_token.map(|rejected_access_token| {
StreamableHttpError::Client(StreamableHttpClientAdapterError::AccessTokenRejected {
rejected_access_token: AccessToken::new(rejected_access_token.to_string()),
})
})
}

impl StreamableHttpClientAdapter {
fn add_auth_headers(&self, headers: &mut HeaderMap) {
if let Some(auth_provider) = &self.auth_provider {
Expand Down
1 change: 1 addition & 0 deletions codex-rs/rmcp-client/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ mod in_process_transport;
mod logging_client_handler;
mod oauth;
mod oauth_http_client;
mod oauth_transport;
mod perform_oauth_login;
mod program_resolver;
mod rmcp_client;
Expand Down
4 changes: 3 additions & 1 deletion codex-rs/rmcp-client/src/oauth.rs
Original file line number Diff line number Diff line change
Expand Up @@ -65,6 +65,7 @@ use codex_keyring_store::KeyringStore;
use codex_utils_home_dir::find_codex_home;

pub(crate) use self::persistor::OAuthPersistor;
pub(crate) use self::persistor::request_oauth_token_response;
use self::refresh_lock::RefreshCredentialLock;
pub(crate) use self::resolved_store::ResolvedOAuthCredentialStore;
pub(crate) use self::resolved_store::ResolvedOAuthTokens;
Expand All @@ -73,7 +74,8 @@ pub(crate) use self::resolved_store::resolve_oauth_tokens;

const KEYRING_SERVICE: &str = "Codex MCP Credentials";
const MCP_OAUTH_SECRET_PREFIX: &str = "MCP_OAUTH";
const REFRESH_SKEW_MILLIS: u64 = 30_000;
// Refresh proactively so ordinary requests do not race token expiry.
const REFRESH_SKEW_MILLIS: u64 = 60_000;

#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct StoredOAuthTokens {
Expand Down
1 change: 0 additions & 1 deletion codex-rs/rmcp-client/src/oauth/persistor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -71,7 +71,6 @@ impl OAuthPersistor {
.await
}

#[expect(dead_code, reason = "wired by part 4 of this stack")]
pub(crate) async fn refresh_after_unauthorized(
&self,
rejected_access_token: AccessToken,
Expand Down
217 changes: 217 additions & 0 deletions codex-rs/rmcp-client/src/oauth_transport.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,217 @@
//! Codex-owned OAuth policy for RMCP Streamable HTTP traffic.
//!
//! RMCP remains responsible for transport mechanics and bearer-token injection. Codex owns the
//! credential lifecycle: every POST, SSE GET/reconnect, and session DELETE receives proactive
//! refresh from its owning Codex layer, and each path has at most one 401 recovery. The
//! authorization manager only receives request-safe credentials, so it cannot independently
//! refresh outside Codex's serialized transaction.
//!
//! POST recovery is split at an intentional ownership boundary. Client-originated requests and
//! notifications retain their outer `RmcpClient` recovery, which knows the startup/tool deadline
//! and can avoid replaying a request after its caller timed out. RMCP-owned responses to
//! server-initiated requests have no such outer operation, so they recover here. GET/reconnect and
//! DELETE are always RMCP-owned and also recover here.

use std::collections::HashMap;
use std::sync::Arc;

use reqwest::header::HeaderName;
use reqwest::header::HeaderValue;
use rmcp::model::ClientJsonRpcMessage;
use rmcp::model::JsonRpcMessage;
use rmcp::transport::auth::AuthClient;
use rmcp::transport::streamable_http_client::StreamableHttpClient;
use rmcp::transport::streamable_http_client::StreamableHttpError;
use rmcp::transport::streamable_http_client::StreamableHttpPostResponse;
use tracing::debug;

use crate::http_client_adapter::StreamableHttpClientAdapter;
use crate::http_client_adapter::StreamableHttpClientAdapterError;
use crate::oauth::OAuthPersistor;

type TransportResult<T> =
std::result::Result<T, StreamableHttpError<StreamableHttpClientAdapterError>>;

#[derive(Clone)]
pub(crate) struct OAuthTransportClient {
auth_client: AuthClient<StreamableHttpClientAdapter>,
persistor: OAuthPersistor,
}

impl OAuthTransportClient {
pub(crate) fn new(
auth_client: AuthClient<StreamableHttpClientAdapter>,
persistor: OAuthPersistor,
) -> Self {
Self {
auth_client,
persistor,
}
}

pub(crate) fn persistor(&self) -> OAuthPersistor {
self.persistor.clone()
}

async fn preflight(&self, operation: &'static str) -> TransportResult<()> {
debug!(
operation,
"checking MCP OAuth credentials before transport request"
);
self.persistor
.refresh_if_needed()
.await
.map_err(oauth_transport_error)
}

async fn recover_after_unauthorized(
&self,
operation: &'static str,
rejected_access_token: Option<oauth2::AccessToken>,
) -> TransportResult<bool> {
let Some(rejected_access_token) = rejected_access_token else {
return Ok(false);
};

debug!(
operation,
"recovering once after MCP transport rejected an OAuth access token"
);
self.persistor
.refresh_after_unauthorized(rejected_access_token)
.await
.map_err(oauth_transport_error)?;
Ok(true)
}
}

impl StreamableHttpClient for OAuthTransportClient {
type Error = StreamableHttpClientAdapterError;

async fn post_message(
&self,
uri: Arc<str>,
message: ClientJsonRpcMessage,
session_id: Option<Arc<str>>,
auth_token: Option<String>,
custom_headers: HashMap<HeaderName, HeaderValue>,
) -> TransportResult<StreamableHttpPostResponse> {
let is_rmcp_owned_response = matches!(
message,
JsonRpcMessage::Response(_) | JsonRpcMessage::Error(_)
);
if is_rmcp_owned_response {
self.preflight("post_message").await?;
}
let result = self
.auth_client
.post_message(
Arc::clone(&uri),
message.clone(),
session_id.clone(),
auth_token.clone(),
custom_headers.clone(),
)
.await;

// RMCP queues client-originated requests independently of the caller waiting on them. If
// recovery happened here, a timed-out public tool call could still be replayed after its
// refresh finished. The outer RmcpClient path owns those deadlines. Responses to
// server-initiated requests have no outer operation and therefore recover here.
if !is_rmcp_owned_response {
return result;
}
let rejected_access_token = result.as_ref().err().and_then(rejected_access_token);
if self
.recover_after_unauthorized("post_message", rejected_access_token)
.await?
{
self.auth_client
.post_message(uri, message, session_id, auth_token, custom_headers)
.await
} else {
result
}
}

async fn delete_session(
&self,
uri: Arc<str>,
session_id: Arc<str>,
auth_token: Option<String>,
custom_headers: HashMap<HeaderName, HeaderValue>,
) -> TransportResult<()> {
self.preflight("delete_session").await?;
let result = self
.auth_client
.delete_session(
Arc::clone(&uri),
Arc::clone(&session_id),
auth_token.clone(),
custom_headers.clone(),
)
.await;
let rejected_access_token = result.as_ref().err().and_then(rejected_access_token);
if self
.recover_after_unauthorized("delete_session", rejected_access_token)
.await?
{
self.auth_client
.delete_session(uri, session_id, auth_token, custom_headers)
.await
} else {
result
}
}

async fn get_stream(
&self,
uri: Arc<str>,
session_id: Arc<str>,
last_event_id: Option<String>,
auth_token: Option<String>,
custom_headers: HashMap<HeaderName, HeaderValue>,
) -> TransportResult<
futures::stream::BoxStream<'static, Result<sse_stream::Sse, sse_stream::Error>>,
> {
self.preflight("get_stream").await?;
let result = self
.auth_client
.get_stream(
Arc::clone(&uri),
Arc::clone(&session_id),
last_event_id.clone(),
auth_token.clone(),
custom_headers.clone(),
)
.await;
let rejected_access_token = result.as_ref().err().and_then(rejected_access_token);
if self
.recover_after_unauthorized("get_stream", rejected_access_token)
.await?
{
self.auth_client
.get_stream(uri, session_id, last_event_id, auth_token, custom_headers)
.await
} else {
result
}
}
}

fn rejected_access_token(
error: &StreamableHttpError<StreamableHttpClientAdapterError>,
) -> Option<oauth2::AccessToken> {
match error {
StreamableHttpError::Client(StreamableHttpClientAdapterError::AccessTokenRejected {
rejected_access_token,
}) => Some(rejected_access_token.clone()),
_ => None,
}
}

fn oauth_transport_error(
error: anyhow::Error,
) -> StreamableHttpError<StreamableHttpClientAdapterError> {
StreamableHttpError::Client(StreamableHttpClientAdapterError::OAuth(error))
}
Loading
Loading