fix(rmcp): propagate proactive OAuth refresh failures

This commit is contained in:
Adam Perry
2026-06-05 04:48:28 +00:00
parent 45c6f67596
commit ed1ef2a09d
3 changed files with 366 additions and 151 deletions

View File

@@ -439,7 +439,7 @@ impl RmcpClient {
params: Option<PaginatedRequestParams>,
timeout: Option<Duration>,
) -> Result<ListToolsResult> {
self.refresh_oauth_if_needed().await;
self.refresh_oauth_if_needed().await?;
let result = self
.run_service_operation("tools/list", timeout, move |service| {
let params = params.clone();
@@ -455,7 +455,7 @@ impl RmcpClient {
params: Option<PaginatedRequestParams>,
timeout: Option<Duration>,
) -> Result<ListToolsWithConnectorIdResult> {
self.refresh_oauth_if_needed().await;
self.refresh_oauth_if_needed().await?;
let result = self
.run_service_operation("tools/list", timeout, move |service| {
let params = params.clone();
@@ -500,7 +500,7 @@ impl RmcpClient {
params: Option<PaginatedRequestParams>,
timeout: Option<Duration>,
) -> Result<ListResourcesResult> {
self.refresh_oauth_if_needed().await;
self.refresh_oauth_if_needed().await?;
let result = self
.run_service_operation("resources/list", timeout, move |service| {
let params = params.clone();
@@ -516,7 +516,7 @@ impl RmcpClient {
params: Option<PaginatedRequestParams>,
timeout: Option<Duration>,
) -> Result<ListResourceTemplatesResult> {
self.refresh_oauth_if_needed().await;
self.refresh_oauth_if_needed().await?;
let result = self
.run_service_operation("resources/templates/list", timeout, move |service| {
let params = params.clone();
@@ -532,7 +532,7 @@ impl RmcpClient {
params: ReadResourceRequestParams,
timeout: Option<Duration>,
) -> Result<ReadResourceResult> {
self.refresh_oauth_if_needed().await;
self.refresh_oauth_if_needed().await?;
let result = self
.run_service_operation("resources/read", timeout, move |service| {
let params = params.clone();
@@ -550,7 +550,7 @@ impl RmcpClient {
meta: Option<serde_json::Value>,
timeout: Option<Duration>,
) -> Result<CallToolResult> {
self.refresh_oauth_if_needed().await;
self.refresh_oauth_if_needed().await?;
let arguments = match arguments {
Some(Value::Object(map)) => Some(map),
Some(other) => {
@@ -606,7 +606,7 @@ impl RmcpClient {
method: &str,
params: Option<serde_json::Value>,
) -> Result<()> {
self.refresh_oauth_if_needed().await;
self.refresh_oauth_if_needed().await?;
self.run_service_operation(
"notifications/custom",
/*timeout*/ None,
@@ -636,7 +636,7 @@ impl RmcpClient {
method: &str,
params: Option<serde_json::Value>,
) -> Result<ServerResult> {
self.refresh_oauth_if_needed().await;
self.refresh_oauth_if_needed().await?;
let response = self
.run_service_operation("requests/custom", /*timeout*/ None, move |service| {
let params = params.clone();
@@ -700,12 +700,11 @@ impl RmcpClient {
}
}
async fn refresh_oauth_if_needed(&self) {
if let Some(runtime) = self.oauth_persistor().await
&& let Err(error) = runtime.refresh_if_needed().await
{
warn!("failed to refresh OAuth tokens: {error}");
async fn refresh_oauth_if_needed(&self) -> Result<()> {
if let Some(runtime) = self.oauth_persistor().await {
runtime.refresh_if_needed().await?;
}
Ok(())
}
async fn create_pending_transport(
@@ -1061,6 +1060,10 @@ async fn create_oauth_transport_and_runtime(
Ok((transport, runtime))
}
#[cfg(test)]
#[path = "rmcp_client_oauth_tests.rs"]
mod oauth_tests;
#[cfg(test)]
mod tests {
use std::time::Duration;

View File

@@ -0,0 +1,300 @@
use std::io;
use std::sync::Arc;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::time::Duration;
use codex_config::types::OAuthCredentialsStoreMode;
use codex_exec_server::Environment;
use futures::FutureExt;
use oauth2::AccessToken;
use oauth2::RefreshToken;
use oauth2::basic::BasicTokenType;
use pretty_assertions::assert_eq;
use rmcp::ErrorData as McpError;
use rmcp::handler::server::ServerHandler;
use rmcp::model::ClientCapabilities;
use rmcp::model::Implementation;
use rmcp::model::ListToolsResult;
use rmcp::model::ProtocolVersion;
use rmcp::model::ServerCapabilities;
use rmcp::model::ServerInfo;
use rmcp::service::RequestContext;
use rmcp::service::RoleServer;
use rmcp::transport::auth::OAuthTokenResponse;
use rmcp::transport::auth::VendorExtraTokenFields;
use serde_json::json;
use tempfile::TempDir;
use tokio::io::DuplexStream;
use tokio::process::Command;
use wiremock::Mock;
use wiremock::MockServer;
use wiremock::Request;
use wiremock::Respond;
use wiremock::ResponseTemplate;
use wiremock::matchers::method;
use wiremock::matchers::path;
use super::*;
use crate::oauth::WrappedOAuthTokenResponse;
const SERVER_NAME: &str = "request-time-refresh-test";
const INITIAL_ACCESS_TOKEN: &str = "initial-access-token";
const INITIAL_REFRESH_TOKEN: &str = "initial-refresh-token";
const REFRESHED_ACCESS_TOKEN: &str = "refreshed-access-token";
const REFRESHED_REFRESH_TOKEN: &str = "refreshed-refresh-token";
const CHILD_SERVER_URL_ENV: &str = "MCP_TEST_REQUEST_TIME_REFRESH_SERVER_URL";
const CHILD_SCENARIO_ENV: &str = "MCP_TEST_REQUEST_TIME_REFRESH_SCENARIO";
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Scenario {
RefreshSucceeds,
ProviderRejectsRefresh,
}
impl Scenario {
fn as_env(self) -> &'static str {
match self {
Self::RefreshSucceeds => "refresh-succeeds",
Self::ProviderRejectsRefresh => "provider-rejects-refresh",
}
}
fn from_env(value: &str) -> anyhow::Result<Self> {
match value {
"refresh-succeeds" => Ok(Self::RefreshSucceeds),
"provider-rejects-refresh" => Ok(Self::ProviderRejectsRefresh),
_ => anyhow::bail!("unknown request-time refresh scenario: {value}"),
}
}
}
#[derive(Clone, Copy)]
struct TokenResponder {
scenario: Scenario,
}
impl Respond for TokenResponder {
fn respond(&self, request: &Request) -> ResponseTemplate {
let body = String::from_utf8_lossy(&request.body);
assert!(
body.contains(INITIAL_REFRESH_TOKEN),
"unexpected refresh request body: {body}"
);
match self.scenario {
Scenario::RefreshSucceeds => ResponseTemplate::new(200).set_body_json(json!({
"access_token": REFRESHED_ACCESS_TOKEN,
"token_type": "Bearer",
"expires_in": 7200,
"refresh_token": REFRESHED_REFRESH_TOKEN,
})),
Scenario::ProviderRejectsRefresh => ResponseTemplate::new(400).set_body_json(json!({
"error": "invalid_grant",
"error_description": "provider rejected refresh",
})),
}
}
}
#[derive(Clone)]
struct CountingServer {
list_tools_calls: Arc<AtomicUsize>,
}
impl ServerHandler for CountingServer {
fn get_info(&self) -> ServerInfo {
ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
}
fn list_tools(
&self,
_request: Option<PaginatedRequestParams>,
_context: RequestContext<RoleServer>,
) -> impl Future<Output = std::result::Result<ListToolsResult, McpError>> + Send + '_ {
self.list_tools_calls.fetch_add(1, Ordering::SeqCst);
async {
Ok(ListToolsResult {
tools: Vec::new(),
next_cursor: None,
meta: None,
})
}
}
}
#[derive(Clone)]
struct TestTransportFactory {
server: CountingServer,
}
impl InProcessTransportFactory for TestTransportFactory {
fn open(&self) -> BoxFuture<'static, io::Result<DuplexStream>> {
let server = self.server.clone();
async move {
let (client_transport, server_transport) = tokio::io::duplex(64 * 1024);
let _server_task = tokio::spawn(async move {
if let Ok(service) = rmcp::serve_server(server, server_transport).await {
let _result = service.waiting().await;
}
});
Ok(client_transport)
}
.boxed()
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 1)]
async fn proactive_refresh_succeeds_before_request() -> anyhow::Result<()> {
run_parent_scenario(Scenario::RefreshSucceeds).await
}
#[tokio::test(flavor = "multi_thread", worker_threads = 1)]
async fn provider_rejection_of_proactive_refresh_stops_before_request() -> anyhow::Result<()> {
run_parent_scenario(Scenario::ProviderRejectsRefresh).await
}
async fn run_parent_scenario(scenario: Scenario) -> anyhow::Result<()> {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/.well-known/oauth-authorization-server/mcp"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"authorization_endpoint": format!("{}/oauth/authorize", server.uri()),
"token_endpoint": format!("{}/oauth/token", server.uri()),
"scopes_supported": [""],
})))
.expect(1)
.mount(&server)
.await;
Mock::given(method("POST"))
.and(path("/oauth/token"))
.respond_with(TokenResponder { scenario })
.expect(1)
.mount(&server)
.await;
let codex_home = TempDir::new()?;
let status = Command::new(std::env::current_exe()?)
.args([
"rmcp_client::oauth_tests::oauth_request_time_child",
"--exact",
"--ignored",
"--nocapture",
])
.env("CODEX_HOME", codex_home.path())
.env(CHILD_SERVER_URL_ENV, server.uri())
.env(CHILD_SCENARIO_ENV, scenario.as_env())
.status()
.await?;
assert!(
status.success(),
"request-time refresh child failed: {status}"
);
server.verify().await;
Ok(())
}
#[tokio::test(flavor = "multi_thread", worker_threads = 1)]
#[ignore = "spawned by request-time refresh parent tests"]
async fn oauth_request_time_child() -> anyhow::Result<()> {
let server_url = std::env::var(CHILD_SERVER_URL_ENV)?;
let scenario = Scenario::from_env(&std::env::var(CHILD_SCENARIO_ENV)?)?;
let list_tools_calls = Arc::new(AtomicUsize::new(0));
let client = RmcpClient::new_in_process_client(Arc::new(TestTransportFactory {
server: CountingServer {
list_tools_calls: Arc::clone(&list_tools_calls),
},
}))
.await?;
initialize_in_process_client(&client).await?;
let mut response = OAuthTokenResponse::new(
AccessToken::new(INITIAL_ACCESS_TOKEN.to_string()),
BasicTokenType::Bearer,
VendorExtraTokenFields::default(),
);
response.set_refresh_token(Some(RefreshToken::new(INITIAL_REFRESH_TOKEN.to_string())));
response.set_expires_in(Some(&Duration::from_secs(7200)));
let mcp_url = format!("{server_url}/mcp");
let initial_tokens = StoredOAuthTokens {
server_name: SERVER_NAME.to_string(),
url: mcp_url.clone(),
client_id: "test-client-id".to_string(),
token_response: WrappedOAuthTokenResponse(response),
expires_at: Some(0),
};
let (_, runtime) = create_oauth_transport_and_runtime(
SERVER_NAME,
&mcp_url,
initial_tokens,
OAuthCredentialsStoreMode::File,
HeaderMap::new(),
Environment::default_for_tests().get_http_client(),
)
.await?;
{
let mut state = client.state.lock().await;
match &mut *state {
ClientState::Ready { oauth, .. } => *oauth = Some(runtime),
ClientState::Connecting { .. } => panic!("client was not initialized"),
ClientState::Closed => panic!("client was unexpectedly closed"),
}
}
let result = client
.list_tools(/*params*/ None, Some(Duration::from_secs(5)))
.await;
match scenario {
Scenario::RefreshSucceeds => {
assert_eq!(result?.tools, Vec::new());
assert_eq!(list_tools_calls.load(Ordering::SeqCst), 1);
}
Scenario::ProviderRejectsRefresh => {
let error = match result {
Ok(_) => panic!("provider rejection should fail before tools/list"),
Err(error) => error,
};
let error_chain = format!("{error:#}");
assert!(
error_chain.contains(
"failed to refresh OAuth tokens for server request-time-refresh-test"
),
"unexpected provider rejection error: {error_chain}"
);
assert!(
error_chain.contains("invalid_grant"),
"provider error was not preserved: {error_chain}"
);
assert_eq!(list_tools_calls.load(Ordering::SeqCst), 0);
}
}
client.shutdown().await;
Ok(())
}
async fn initialize_in_process_client(client: &RmcpClient) -> anyhow::Result<()> {
let params = InitializeRequestParams::new(
ClientCapabilities::default(),
Implementation::new("codex-test", "0.0.0-test"),
)
.with_protocol_version(ProtocolVersion::V_2025_06_18);
client
.initialize(
params,
Some(Duration::from_secs(5)),
Box::new(|_, _| {
async {
Ok(ElicitationResponse {
action: ElicitationAction::Accept,
content: Some(json!({})),
meta: None,
})
}
.boxed()
}),
)
.await?;
Ok(())
}

View File

@@ -2,8 +2,6 @@ mod streamable_http_test_support;
use std::sync::Arc;
use std::sync::Mutex;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::time::Duration;
use codex_config::types::OAuthCredentialsStoreMode;
@@ -35,36 +33,10 @@ use streamable_http_test_support::initialize_client;
const SERVER_NAME: &str = "test-streamable-http-oauth-lifecycle";
const INITIAL_ACCESS_TOKEN: &str = "initial-expired-access-token";
const INITIAL_REFRESH_TOKEN: &str = "initial-refresh-token";
const SHORT_LIVED_ACCESS_TOKEN: &str = "short-lived-access-token";
const SHORT_LIVED_REFRESH_TOKEN: &str = "short-lived-refresh-token";
const LONG_LIVED_ACCESS_TOKEN: &str = "long-lived-access-token";
const LONG_LIVED_REFRESH_TOKEN: &str = "long-lived-refresh-token";
const CHILD_SERVER_URL_ENV: &str = "MCP_TEST_OAUTH_LIFECYCLE_SERVER_URL";
const CHILD_CHECKPOINT_URL_ENV: &str = "MCP_TEST_OAUTH_LIFECYCLE_CHECKPOINT_URL";
const CHILD_SCENARIO_ENV: &str = "MCP_TEST_OAUTH_LIFECYCLE_SCENARIO";
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Scenario {
RefreshSucceeds,
RefreshTokenExpired,
}
impl Scenario {
fn as_env(self) -> &'static str {
match self {
Self::RefreshSucceeds => "refresh-succeeds",
Self::RefreshTokenExpired => "refresh-token-expired",
}
}
fn from_env(value: &str) -> anyhow::Result<Self> {
match value {
"refresh-succeeds" => Ok(Self::RefreshSucceeds),
"refresh-token-expired" => Ok(Self::RefreshTokenExpired),
_ => anyhow::bail!("unknown OAuth lifecycle scenario: {value}"),
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
enum TimelineEvent {
@@ -104,46 +76,27 @@ impl Timeline {
}
}
#[derive(Clone)]
struct TokenResponder {
scenario: Scenario,
timeline: Timeline,
calls: AtomicUsize,
}
impl Respond for TokenResponder {
fn respond(&self, request: &Request) -> ResponseTemplate {
let body = String::from_utf8_lossy(&request.body);
let refresh_token = if body.contains(INITIAL_REFRESH_TOKEN) {
INITIAL_REFRESH_TOKEN
} else if body.contains(SHORT_LIVED_REFRESH_TOKEN) {
SHORT_LIVED_REFRESH_TOKEN
} else {
panic!("unexpected refresh request body: {body}");
};
assert!(
body.contains(INITIAL_REFRESH_TOKEN),
"unexpected refresh request body: {body}"
);
self.timeline
.push(TimelineEvent::Refresh(refresh_token.to_string()));
.push(TimelineEvent::Refresh(INITIAL_REFRESH_TOKEN.to_string()));
match self.calls.fetch_add(1, Ordering::SeqCst) {
0 => ResponseTemplate::new(200).set_body_json(json!({
"access_token": SHORT_LIVED_ACCESS_TOKEN,
"token_type": "Bearer",
"expires_in": 32,
"refresh_token": SHORT_LIVED_REFRESH_TOKEN,
})),
1 if self.scenario == Scenario::RefreshSucceeds => ResponseTemplate::new(200)
.set_body_json(json!({
"access_token": LONG_LIVED_ACCESS_TOKEN,
"token_type": "Bearer",
"expires_in": 7200,
"refresh_token": LONG_LIVED_REFRESH_TOKEN,
})),
1 | 2 if self.scenario == Scenario::RefreshTokenExpired => ResponseTemplate::new(400)
.set_body_json(json!({
"error": "invalid_grant",
"error_description": "refresh token expired",
})),
call => panic!("unexpected OAuth token request #{call}"),
}
ResponseTemplate::new(200).set_body_json(json!({
"access_token": LONG_LIVED_ACCESS_TOKEN,
"token_type": "Bearer",
"expires_in": 7200,
"refresh_token": LONG_LIVED_REFRESH_TOKEN,
}))
}
}
@@ -206,60 +159,7 @@ impl Respond for McpResponder {
}
#[tokio::test(flavor = "multi_thread", worker_threads = 1)]
async fn refreshes_only_when_the_next_operation_starts_after_idle() -> anyhow::Result<()> {
let timeline = run_scenario(Scenario::RefreshSucceeds).await?;
assert_eq!(
timeline,
vec![
TimelineEvent::Refresh(INITIAL_REFRESH_TOKEN.to_string()),
TimelineEvent::Mcp {
method: "initialize".to_string(),
authorization: format!("Bearer {SHORT_LIVED_ACCESS_TOKEN}"),
},
TimelineEvent::Mcp {
method: "notifications/initialized".to_string(),
authorization: format!("Bearer {SHORT_LIVED_ACCESS_TOKEN}"),
},
TimelineEvent::IdleStarted,
TimelineEvent::IdleFinished,
TimelineEvent::Refresh(SHORT_LIVED_REFRESH_TOKEN.to_string()),
TimelineEvent::Mcp {
method: "tools/list".to_string(),
authorization: format!("Bearer {LONG_LIVED_ACCESS_TOKEN}"),
},
]
);
Ok(())
}
#[tokio::test(flavor = "multi_thread", worker_threads = 1)]
async fn expired_refresh_token_before_next_operation_requires_reauthorization() -> anyhow::Result<()>
{
let timeline = run_scenario(Scenario::RefreshTokenExpired).await?;
assert_eq!(
timeline,
vec![
TimelineEvent::Refresh(INITIAL_REFRESH_TOKEN.to_string()),
TimelineEvent::Mcp {
method: "initialize".to_string(),
authorization: format!("Bearer {SHORT_LIVED_ACCESS_TOKEN}"),
},
TimelineEvent::Mcp {
method: "notifications/initialized".to_string(),
authorization: format!("Bearer {SHORT_LIVED_ACCESS_TOKEN}"),
},
TimelineEvent::IdleStarted,
TimelineEvent::IdleFinished,
TimelineEvent::Refresh(SHORT_LIVED_REFRESH_TOKEN.to_string()),
TimelineEvent::Refresh(SHORT_LIVED_REFRESH_TOKEN.to_string()),
]
);
Ok(())
}
async fn run_scenario(scenario: Scenario) -> anyhow::Result<Vec<TimelineEvent>> {
async fn does_not_refresh_in_background_while_idle() -> anyhow::Result<()> {
let server = MockServer::start().await;
let timeline = Timeline::new();
@@ -270,15 +170,15 @@ async fn run_scenario(scenario: Scenario) -> anyhow::Result<Vec<TimelineEvent>>
"token_endpoint": format!("{}/oauth/token", server.uri()),
"scopes_supported": [""],
})))
.expect(1)
.mount(&server)
.await;
Mock::given(method("POST"))
.and(path("/oauth/token"))
.respond_with(TokenResponder {
scenario,
timeline: timeline.clone(),
calls: AtomicUsize::new(0),
})
.expect(1)
.mount(&server)
.await;
Mock::given(method("POST"))
@@ -286,20 +186,22 @@ async fn run_scenario(scenario: Scenario) -> anyhow::Result<Vec<TimelineEvent>>
.respond_with(McpResponder {
timeline: timeline.clone(),
})
.expect(3)
.mount(&server)
.await;
for (path, event) in [
for (checkpoint_path, event) in [
("/checkpoint/idle-started", TimelineEvent::IdleStarted),
("/checkpoint/idle-finished", TimelineEvent::IdleFinished),
] {
let timeline = timeline.clone();
Mock::given(method("POST"))
.and(wiremock::matchers::path(path))
.and(path(checkpoint_path))
.respond_with(move |_: &Request| {
timeline.push(event.clone());
ResponseTemplate::new(204)
})
.expect(1)
.mount(&server)
.await;
}
@@ -315,20 +217,39 @@ async fn run_scenario(scenario: Scenario) -> anyhow::Result<Vec<TimelineEvent>>
.env("CODEX_HOME", codex_home.path())
.env(CHILD_SERVER_URL_ENV, format!("{}/mcp", server.uri()))
.env(CHILD_CHECKPOINT_URL_ENV, server.uri())
.env(CHILD_SCENARIO_ENV, scenario.as_env())
.status()
.await?;
assert!(status.success(), "OAuth lifecycle child failed: {status}");
Ok(timeline.snapshot())
assert_eq!(
timeline.snapshot(),
vec![
TimelineEvent::Refresh(INITIAL_REFRESH_TOKEN.to_string()),
TimelineEvent::Mcp {
method: "initialize".to_string(),
authorization: format!("Bearer {LONG_LIVED_ACCESS_TOKEN}"),
},
TimelineEvent::Mcp {
method: "notifications/initialized".to_string(),
authorization: format!("Bearer {LONG_LIVED_ACCESS_TOKEN}"),
},
TimelineEvent::IdleStarted,
TimelineEvent::IdleFinished,
TimelineEvent::Mcp {
method: "tools/list".to_string(),
authorization: format!("Bearer {LONG_LIVED_ACCESS_TOKEN}"),
},
]
);
server.verify().await;
Ok(())
}
#[tokio::test(flavor = "multi_thread", worker_threads = 1)]
#[ignore = "spawned by OAuth lifecycle parent tests"]
#[ignore = "spawned by OAuth lifecycle parent test"]
async fn oauth_lifecycle_child() -> anyhow::Result<()> {
let server_url = std::env::var(CHILD_SERVER_URL_ENV)?;
let checkpoint_url = std::env::var(CHILD_CHECKPOINT_URL_ENV)?;
let scenario = Scenario::from_env(&std::env::var(CHILD_SCENARIO_ENV)?)?;
let mut response = OAuthTokenResponse::new(
AccessToken::new(INITIAL_ACCESS_TOKEN.to_string()),
@@ -363,25 +284,16 @@ async fn oauth_lifecycle_child() -> anyhow::Result<()> {
initialize_client(&client).await?;
post_checkpoint(&checkpoint_url, "idle-started").await?;
tokio::time::sleep(Duration::from_secs(3)).await;
tokio::time::sleep(Duration::from_millis(100)).await;
post_checkpoint(&checkpoint_url, "idle-finished").await?;
let result = client
.list_tools(/*params*/ None, Some(Duration::from_secs(5)))
.await;
match scenario {
Scenario::RefreshSucceeds => {
assert_eq!(result?.tools, Vec::new());
}
Scenario::RefreshTokenExpired => {
let error = result.expect_err("expired refresh token should fail the operation");
assert!(
error.to_string().contains("authorization required"),
"unexpected expired refresh token error: {error:#}"
);
}
}
assert_eq!(
client
.list_tools(/*params*/ None, Some(Duration::from_secs(5)))
.await?
.tools,
Vec::new()
);
Ok(())
}