Skip to content

Commit 04c2f3f

Browse files
authored
fix(auth): unify refresh checks and error handling (#1236)
* fix(auth): unify refresh checks and error handling Claude-Session: https://claude.ai/code/session_011pHFfoTygeG84mCDXzCcmw * test(auth): add live authorization-server checks Claude-Session: https://claude.ai/code/session_011pHFfoTygeG84mCDXzCcmw
1 parent add9cbe commit 04c2f3f

4 files changed

Lines changed: 502 additions & 23 deletions

File tree

‎crates/rmcp/src/transport/auth.rs‎

Lines changed: 100 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -294,6 +294,7 @@ impl CredentialRefreshGuard {
294294
/// Implementations of this trait can provide custom storage backends
295295
/// for OAuth2 credentials, such as file-based storage, keychain integration,
296296
/// or database storage.
297+
///
297298
/// Return [`AuthError::CredentialStoreError`] for backend or locking failures
298299
/// so they remain distinct from errors requiring reauthorization.
299300
#[async_trait]
@@ -2248,12 +2249,19 @@ impl AuthorizationManager {
22482249
.as_ref()
22492250
.ok_or_else(|| AuthError::InternalError("OAuth client not configured".to_string()))?;
22502251

2251-
let refresh_guard = self.credential_store.acquire_refresh_guard().await?;
2252+
// Held for the rest of this function so the load, the exchange, and the
2253+
// save stay inside one guarded section.
2254+
let _refresh_guard = self.credential_store.acquire_refresh_guard().await?;
22522255
let stored = self.credential_store.load().await?;
22532256
let stored_credentials = stored.ok_or(AuthError::AuthorizationRequired)?;
2254-
if refresh_guard.is_some()
2255-
&& stored_credentials.client_id != oauth_client.client_id().as_str()
2256-
{
2257+
// Refreshing with another client's stored token would put that token on a
2258+
// request authenticated as this client.
2259+
if stored_credentials.client_id != oauth_client.client_id().as_str() {
2260+
tracing::warn!(
2261+
stored_client_id = stored_credentials.client_id.as_str(),
2262+
configured_client_id = oauth_client.client_id().as_str(),
2263+
"stored credentials belong to a different client; reauthorization required"
2264+
);
22572265
return Err(AuthError::AuthorizationRequired);
22582266
}
22592267
let current_credentials = stored_credentials
@@ -2271,8 +2279,8 @@ impl AuthorizationManager {
22712279
// RFC 8707: the resource indicator is required on token requests, including refreshes
22722280
.add_extra_param("resource", self.oauth_resource().await);
22732281
let mut refresh_scopes = stored_credentials.granted_scopes;
2274-
let authoritative_scopes = refresh_guard.is_some().then(|| refresh_scopes.clone());
22752282
self.add_offline_access_if_supported(&mut refresh_scopes);
2283+
let requested_scopes = refresh_scopes.clone();
22762284
for scope in refresh_scopes {
22772285
refresh_request = refresh_request.add_scope(Scope::new(scope));
22782286
}
@@ -2298,10 +2306,12 @@ impl AuthorizationManager {
22982306
token_result.set_refresh_token(Some(refresh_token_value));
22992307
}
23002308

2301-
let granted_scopes: Vec<String> = match (token_result.scopes(), authoritative_scopes) {
2302-
(Some(scopes), _) => scopes.iter().map(|s| s.to_string()).collect(),
2303-
(None, Some(scopes)) => scopes,
2304-
(None, None) => self.current_scopes.read().await.clone(),
2309+
let response_scopes = token_result
2310+
.scopes()
2311+
.map(|scopes| scopes.iter().map(|s| s.to_string()).collect());
2312+
let granted_scopes = {
2313+
let current = self.current_scopes.read().await;
2314+
Self::resolve_granted_scopes(response_scopes, &requested_scopes, &current)
23052315
};
23062316

23072317
*self.current_scopes.write().await = granted_scopes.clone();
@@ -8515,17 +8525,17 @@ mod tests {
85158525
}
85168526
}
85178527

8518-
fn refresh_store() -> RefreshStore {
8528+
async fn refresh_store() -> RefreshStore {
85198529
let credentials = StoredCredentials::new(
85208530
"my-client".into(),
85218531
Some(make_token_response_with_refresh("old-token", "old-refresh")),
85228532
vec!["read".into()],
85238533
Some(AuthorizationManager::now_epoch_secs()),
85248534
);
8535+
let credential_store = InMemoryCredentialStore::new();
8536+
credential_store.save(credentials).await.unwrap();
85258537
RefreshStore {
8526-
credentials: InMemoryCredentialStore {
8527-
credentials: Arc::new(tokio::sync::RwLock::new(Some(credentials))),
8528-
},
8538+
credentials: credential_store,
85298539
lock: Arc::new(Mutex::new(())),
85308540
events: Arc::new(StdMutex::new(Vec::new())),
85318541
guard_requested: Arc::new(Semaphore::new(0)),
@@ -8604,7 +8614,7 @@ mod tests {
86048614

86058615
#[tokio::test]
86068616
async fn refresh_guard_spans_load_exchange_and_completed_save() {
8607-
let store = refresh_store();
8617+
let store = refresh_store().await;
86088618
let manager = refresh_manager(store.clone(), refresh_http_client(&store)).await;
86098619

86108620
manager.refresh_token().await.unwrap();
@@ -8631,7 +8641,7 @@ mod tests {
86318641

86328642
#[tokio::test]
86338643
async fn concurrent_refreshes_wait_for_save_and_use_the_latest_token() {
8634-
let mut store = refresh_store();
8644+
let mut store = refresh_store().await;
86358645
let save_gate = Arc::new(Semaphore::new(0));
86368646
store.save_gate = Some(save_gate.clone());
86378647
let http_client = refresh_http_client(&store);
@@ -8677,8 +8687,8 @@ mod tests {
86778687
}
86788688

86798689
#[tokio::test]
8680-
async fn guarded_refresh_rejects_credentials_for_another_client() {
8681-
let store = refresh_store();
8690+
async fn refresh_rejects_credentials_for_another_client() {
8691+
let store = refresh_store().await;
86828692
let mut credentials = store.credentials.load().await.unwrap().unwrap();
86838693
credentials.client_id = "other-client".into();
86848694
store.credentials.save(credentials).await.unwrap();
@@ -8693,6 +8703,78 @@ mod tests {
86938703
assert!(store.lock.try_lock().is_ok());
86948704
}
86958705

8706+
#[tokio::test]
8707+
async fn refresh_rejects_credentials_for_another_client_without_a_guard() {
8708+
let (base_url, captured) = start_token_server().await;
8709+
let mut manager = manager_with_metadata(Some(AuthorizationMetadata {
8710+
authorization_endpoint: format!("{base_url}/authorize"),
8711+
token_endpoint: format!("{base_url}/token"),
8712+
..Default::default()
8713+
}))
8714+
.await;
8715+
manager.configure_client(test_client_config()).unwrap();
8716+
manager
8717+
.credential_store
8718+
.save(StoredCredentials::new(
8719+
"other-client".into(),
8720+
Some(make_token_response_with_refresh("old-token", "old-refresh")),
8721+
vec!["read".into()],
8722+
Some(AuthorizationManager::now_epoch_secs()),
8723+
))
8724+
.await
8725+
.unwrap();
8726+
8727+
let error = manager.refresh_token().await.unwrap_err();
8728+
8729+
assert!(
8730+
matches!(error, AuthError::AuthorizationRequired),
8731+
"a client mismatch must require reauthorization, got: {error:?}"
8732+
);
8733+
assert!(
8734+
captured.lock().unwrap().is_none(),
8735+
"a client mismatch must be caught before the refresh token leaves the process"
8736+
);
8737+
}
8738+
8739+
#[tokio::test]
8740+
async fn refresh_without_a_guard_keeps_stored_scopes_when_response_omits_them() {
8741+
// start_token_server answers without a `scope`, matching a provider that
8742+
// grants the request in full.
8743+
let (base_url, _captured) = start_token_server().await;
8744+
let mut manager = manager_with_metadata(Some(AuthorizationMetadata {
8745+
authorization_endpoint: format!("{base_url}/authorize"),
8746+
token_endpoint: format!("{base_url}/token"),
8747+
..Default::default()
8748+
}))
8749+
.await;
8750+
manager.configure_client(test_client_config()).unwrap();
8751+
manager
8752+
.credential_store
8753+
.save(StoredCredentials::new(
8754+
"my-client".into(),
8755+
Some(make_token_response_with_refresh("old-token", "old-refresh")),
8756+
vec!["read".into()],
8757+
Some(AuthorizationManager::now_epoch_secs()),
8758+
))
8759+
.await
8760+
.unwrap();
8761+
*manager.current_scopes.write().await = vec!["stale".into()];
8762+
8763+
manager.refresh_token().await.unwrap();
8764+
8765+
let saved = manager.credential_store.load().await.unwrap().unwrap();
8766+
assert_eq!(
8767+
saved.granted_scopes,
8768+
["read"],
8769+
"the stored grant outranks the per-process scope cache"
8770+
);
8771+
assert_eq!(
8772+
manager.get_current_scopes().await,
8773+
["read"],
8774+
"the refreshed grant must replace the stale scope cache"
8775+
);
8776+
}
8777+
86968778
#[rstest]
86978779
#[case("guard", 0)]
86988780
#[case("load", 0)]
@@ -8702,7 +8784,7 @@ mod tests {
87028784
#[case] phase: &'static str,
87038785
#[case] provider_requests: usize,
87048786
) {
8705-
let mut store = refresh_store();
8787+
let mut store = refresh_store().await;
87068788
store.fail_at = Some(phase);
87078789
let http_client = refresh_http_client(&store);
87088790
let manager = refresh_manager(store.clone(), http_client.clone()).await;

‎crates/rmcp/src/transport/common/auth/streamable_http_client.rs‎

Lines changed: 113 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,10 @@ where
1919
/// 401 propagates as [`StreamableHttpError::AuthRequired`] carrying the
2020
/// `WWW-Authenticate` challenge for the caller to authorize with;
2121
/// - a token the server rejects (e.g. revoked) → one silent refresh, one
22-
/// retry, then the challenge propagates.
22+
/// retry, then the challenge propagates;
23+
/// - a refresh that fails for any other reason (credential store, network,
24+
/// provider) → that error propagates so the caller can retry instead of
25+
/// being sent through a new authorization.
2326
async fn call_reacting_to_challenges<T, F, Fut>(
2427
&self,
2528
auth_token: Option<String>,
@@ -54,11 +57,13 @@ where
5457
match refreshed {
5558
Ok(fresh_token) if fresh_token != sent_token => call(Some(fresh_token)).await,
5659
Ok(_) => Err(StreamableHttpError::AuthRequired(challenge)),
57-
Err(error @ AuthError::CredentialStoreError(_)) => Err(error.into()),
58-
Err(error) => {
59-
debug!("token refresh after server rejection failed: {error}");
60+
// `try_refresh_or_reauth` already reports the cases that need a
61+
// new authorization; anything else is retryable or infrastructural.
62+
Err(AuthError::AuthorizationRequired) => {
63+
debug!("token refresh after server rejection requires authorization");
6064
Err(StreamableHttpError::AuthRequired(challenge))
6165
}
66+
Err(error) => Err(error.into()),
6267
}
6368
}
6469
result => result,
@@ -212,11 +217,16 @@ where
212217

213218
#[cfg(all(test, feature = "transport-streamable-http-client-reqwest"))]
214219
mod tests {
220+
use std::sync::Arc;
221+
222+
use oauth2::{AccessToken, RefreshToken, basic::BasicTokenType};
223+
215224
use super::*;
216225
use crate::transport::{
217226
auth::{
218227
AuthorizationManager, AuthorizationMetadata, CredentialRefreshGuard, CredentialStore,
219-
StoredCredentials,
228+
InMemoryCredentialStore, OAuthHttpClient, OAuthHttpClientFuture, OAuthHttpRequest,
229+
OAuthTokenResponse, StoredCredentials, VendorExtraTokenFields,
220230
},
221231
streamable_http_client::AuthRequiredError,
222232
};
@@ -269,4 +279,102 @@ mod tests {
269279
StreamableHttpError::Auth(AuthError::CredentialStoreError(message))
270280
if message == "guard unavailable"));
271281
}
282+
283+
struct UnreachableTokenEndpoint;
284+
285+
impl OAuthHttpClient for UnreachableTokenEndpoint {
286+
fn execute(&self, _: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> {
287+
Box::pin(async { Err("token endpoint unreachable".into()) })
288+
}
289+
}
290+
291+
struct RejectingTokenEndpoint;
292+
293+
impl OAuthHttpClient for RejectingTokenEndpoint {
294+
fn execute(&self, _: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> {
295+
Box::pin(async {
296+
Ok(oauth2::http::Response::builder()
297+
.status(400)
298+
.header("content-type", "application/json")
299+
.body(br#"{"error":"invalid_grant"}"#.to_vec())
300+
.unwrap())
301+
})
302+
}
303+
}
304+
305+
/// A manager holding a refresh token the given token endpoint will answer for.
306+
async fn manager_with_stored_refresh_token(
307+
token_endpoint: Arc<dyn OAuthHttpClient>,
308+
) -> AuthorizationManager {
309+
let mut manager = AuthorizationManager::new_with_oauth_http_client(
310+
"https://mcp.example.com/mcp",
311+
token_endpoint,
312+
)
313+
.await
314+
.unwrap();
315+
manager.set_metadata(AuthorizationMetadata {
316+
authorization_endpoint: "https://auth.example.com/authorize".into(),
317+
token_endpoint: "https://auth.example.com/token".into(),
318+
..Default::default()
319+
});
320+
manager.configure_client_id("client").unwrap();
321+
322+
let mut token_response = OAuthTokenResponse::new(
323+
AccessToken::new("old-token".into()),
324+
BasicTokenType::Bearer,
325+
VendorExtraTokenFields::default(),
326+
);
327+
token_response.set_refresh_token(Some(RefreshToken::new("stored-refresh".into())));
328+
let store = InMemoryCredentialStore::new();
329+
store
330+
.save(StoredCredentials::new(
331+
"client".into(),
332+
Some(token_response),
333+
vec![],
334+
None,
335+
))
336+
.await
337+
.unwrap();
338+
manager.set_credential_store(store);
339+
manager
340+
}
341+
342+
/// Drive one call whose server answer is a 401 challenge.
343+
async fn challenge_once(manager: AuthorizationManager) -> StreamableHttpError<reqwest::Error> {
344+
AuthClient::new(reqwest::Client::new(), manager)
345+
.call_reacting_to_challenges(Some("old-token".into()), |_| async {
346+
Err::<(), _>(StreamableHttpError::AuthRequired(AuthRequiredError::new(
347+
"Bearer".into(),
348+
)))
349+
})
350+
.await
351+
.unwrap_err()
352+
}
353+
354+
#[tokio::test]
355+
async fn reactive_refresh_propagates_retryable_refresh_failure() {
356+
let manager = manager_with_stored_refresh_token(Arc::new(UnreachableTokenEndpoint)).await;
357+
358+
let error = challenge_once(manager).await;
359+
360+
assert!(
361+
matches!(
362+
error,
363+
StreamableHttpError::Auth(AuthError::TokenRefreshFailed(_))
364+
),
365+
"a retryable refresh failure must reach the caller instead of asking for a new authorization, got: {error:?}"
366+
);
367+
}
368+
369+
#[tokio::test]
370+
async fn reactive_refresh_reports_a_rejected_refresh_token_as_a_challenge() {
371+
let manager = manager_with_stored_refresh_token(Arc::new(RejectingTokenEndpoint)).await;
372+
373+
let error = challenge_once(manager).await;
374+
375+
assert!(
376+
matches!(error, StreamableHttpError::AuthRequired(_)),
377+
"a definitively rejected refresh token must surface the challenge, got: {error:?}"
378+
);
379+
}
272380
}

0 commit comments

Comments
 (0)