From 76ae383685bb8b5797cfeb33b8e7f8364e1bf9d0 Mon Sep 17 00:00:00 2001 From: Matthew Zeng Date: Fri, 18 Sep 2026 14:19:31 -0700 Subject: [PATCH 1/2] fix(auth): notify credential stores when refresh tokens are rejected --- crates/rmcp/src/transport/auth.rs | 117 ++++++++++++++++++++++++++++-- 1 file changed, 111 insertions(+), 6 deletions(-) diff --git a/crates/rmcp/src/transport/auth.rs b/crates/rmcp/src/transport/auth.rs index 0727368cf..24cd330df 100644 --- a/crates/rmcp/src/transport/auth.rs +++ b/crates/rmcp/src/transport/auth.rs @@ -305,6 +305,20 @@ pub trait CredentialStore: Send + Sync { async fn clear(&self) -> Result<(), AuthError>; + /// Optionally handle the provider's definitive rejection of a refresh token. + /// + /// Called only for a definitive `invalid_grant` response, with the credentials used + /// for that exchange, before its refresh guard (if any) is released. Implementations + /// must not reacquire that guard or change credentials that have since been replaced. + /// The default leaves credentials unchanged. On success, the manager returns + /// [`AuthError::TokenRefreshRejected`]; a callback error is propagated instead. + async fn on_refresh_token_rejected( + &self, + _credentials: &StoredCredentials, + ) -> Result<(), AuthError> { + Ok(()) + } + /// Optionally coordinate refreshes that share these credentials. /// /// The manager acquires this guard before loading credentials and retains it @@ -2279,6 +2293,7 @@ impl AuthorizationManager { } let current_credentials = stored_credentials .token_response + .as_ref() .ok_or(AuthError::AuthorizationRequired)?; let refresh_token = current_credentials @@ -2291,26 +2306,32 @@ impl AuthorizationManager { .exchange_refresh_token(&refresh_token_value) // RFC 8707: the resource indicator is required on token requests, including refreshes .add_extra_param("resource", self.oauth_resource().await); - let mut refresh_scopes = stored_credentials.granted_scopes; + let mut refresh_scopes = stored_credentials.granted_scopes.clone(); self.add_offline_access_if_supported(&mut refresh_scopes); let requested_scopes = refresh_scopes.clone(); for scope in refresh_scopes { refresh_request = refresh_request.add_scope(Scope::new(scope)); } - let mut token_result = refresh_request + let mut token_result = match refresh_request .request_async(&OAuth2HttpClient { client: self.http_client.as_ref(), redirect_policy: self.refresh_redirect_policy, }) .await - .map_err(|error| match &error { + { + Ok(token_result) => token_result, + Err(error) => match &error { RequestTokenError::ServerResponse(response) if response.error() == &BasicErrorResponseType::InvalidGrant => { - AuthError::TokenRefreshRejected(error.to_string()) + self.credential_store + .on_refresh_token_rejected(&stored_credentials) + .await?; + return Err(AuthError::TokenRefreshRejected(error.to_string())); } - _ => AuthError::TokenRefreshFailed(error.to_string()), - })?; + _ => return Err(AuthError::TokenRefreshFailed(error.to_string())), + }, + }; // RFC 6749 section 6: issuing a new refresh token on refresh is optional. // When the response omits one, keep the existing refresh token rather than @@ -8679,6 +8700,7 @@ mod tests { #[tokio::test] async fn invalid_grant_refresh_requires_reauthorization() { let manager = manager_with_refresh_error("invalid_grant").await; + let before = manager.credential_store.load().await.unwrap(); let err = manager.try_refresh_or_reauth().await.unwrap_err(); @@ -8686,6 +8708,11 @@ mod tests { matches!(err, AuthError::AuthorizationRequired), "expected AuthorizationRequired when the refresh token is rejected, got: {err:?}" ); + assert_eq!( + serde_json::to_value(manager.credential_store.load().await.unwrap()).unwrap(), + serde_json::to_value(before).unwrap(), + "the default rejection callback must leave credentials unchanged" + ); } #[tokio::test] @@ -9196,6 +9223,7 @@ mod tests { credentials: InMemoryCredentialStore, lock: Arc>, events: Arc>>, + rejected_credentials: Arc>>, guard_requested: Arc, save_started: Arc, save_gate: Option>, @@ -9241,6 +9269,19 @@ mod tests { self.credentials.clear().await } + async fn on_refresh_token_rejected( + &self, + credentials: &StoredCredentials, + ) -> Result<(), AuthError> { + assert!(self.lock.try_lock().is_err()); + self.events.lock().unwrap().push("rejected"); + *self.rejected_credentials.lock().unwrap() = Some(credentials.clone()); + if self.fail_at == Some("rejected") { + return Err(AuthError::CredentialStoreError("rejection failed".into())); + } + Ok(()) + } + async fn acquire_refresh_guard(&self) -> Result, AuthError> { self.events.lock().unwrap().push("acquire"); self.guard_requested.add_permits(1); @@ -9269,6 +9310,7 @@ mod tests { credentials: credential_store, lock: Arc::new(Mutex::new(())), events: Arc::new(StdMutex::new(Vec::new())), + rejected_credentials: Arc::new(StdMutex::new(None)), guard_requested: Arc::new(Semaphore::new(0)), save_started: Arc::new(Semaphore::new(0)), save_gate: None, @@ -9311,6 +9353,16 @@ mod tests { }) } + fn refresh_error_http_client(store: &RefreshStore, error: &str) -> Arc { + Arc::new(RefreshHttpClient { + recording: RecordingOAuthHttpClient::with_responses(vec![http_response( + 400, + serde_json::json!({"error": error}), + )]), + events: store.events.clone(), + }) + } + async fn refresh_manager( store: RefreshStore, http_client: Arc, @@ -9370,6 +9422,59 @@ mod tests { assert!(store.lock.try_lock().is_ok()); } + #[rstest] + #[case(None)] + #[case(Some("rejected"))] + #[tokio::test] + async fn rejected_refresh_notifies_store_with_attempted_credentials_under_guard( + #[case] fail_at: Option<&'static str>, + ) { + let mut store = refresh_store().await; + store.fail_at = fail_at; + let attempted = store.credentials.load().await.unwrap(); + let http_client = refresh_error_http_client(&store, "invalid_grant"); + let manager = refresh_manager(store.clone(), http_client.clone()).await; + + let error = manager.refresh_token().await.unwrap_err(); + + match fail_at { + Some(_) => assert!(matches!(error, + AuthError::CredentialStoreError(message) if message == "rejection failed")), + None => assert!(matches!(error, AuthError::TokenRefreshRejected(_))), + } + assert_eq!( + serde_json::to_value(&*store.rejected_credentials.lock().unwrap()).unwrap(), + serde_json::to_value(attempted).unwrap() + ); + assert_eq!( + *store.events.lock().unwrap(), + [ + "acquire", "acquired", "load", "provider", "rejected", "release" + ] + ); + assert_eq!(http_client.recording.requests().len(), 1); + assert!(store.lock.try_lock().is_ok()); + } + + #[tokio::test] + async fn transient_refresh_failure_does_not_notify_store_of_rejection() { + let store = refresh_store().await; + let http_client = refresh_error_http_client(&store, "temporarily_unavailable"); + let manager = refresh_manager(store.clone(), http_client).await; + + assert!(matches!( + manager.refresh_token().await, + Err(AuthError::TokenRefreshFailed(_)) + )); + + assert!(store.rejected_credentials.lock().unwrap().is_none()); + assert_eq!( + *store.events.lock().unwrap(), + ["acquire", "acquired", "load", "provider", "release"] + ); + assert!(store.lock.try_lock().is_ok()); + } + #[tokio::test] async fn concurrent_refreshes_wait_for_save_and_use_the_latest_token() { let mut store = refresh_store().await; From 285f59fe7b26f2224b92308a1ddcdb2a76574399 Mon Sep 17 00:00:00 2001 From: Matthew Zeng Date: Wed, 23 Sep 2026 00:26:52 +0000 Subject: [PATCH 2/2] fix(auth): preserve refresh rejection when credential callback fails --- crates/rmcp/src/transport/auth.rs | 50 +++++++++++++++++++++++++------ 1 file changed, 41 insertions(+), 9 deletions(-) diff --git a/crates/rmcp/src/transport/auth.rs b/crates/rmcp/src/transport/auth.rs index 24cd330df..ab11142ae 100644 --- a/crates/rmcp/src/transport/auth.rs +++ b/crates/rmcp/src/transport/auth.rs @@ -310,8 +310,8 @@ pub trait CredentialStore: Send + Sync { /// Called only for a definitive `invalid_grant` response, with the credentials used /// for that exchange, before its refresh guard (if any) is released. Implementations /// must not reacquire that guard or change credentials that have since been replaced. - /// The default leaves credentials unchanged. On success, the manager returns - /// [`AuthError::TokenRefreshRejected`]; a callback error is propagated instead. + /// The default leaves credentials unchanged. Callback errors are logged, and the + /// manager still returns [`AuthError::TokenRefreshRejected`] so callers can reauthorize. async fn on_refresh_token_rejected( &self, _credentials: &StoredCredentials, @@ -2324,9 +2324,16 @@ impl AuthorizationManager { RequestTokenError::ServerResponse(response) if response.error() == &BasicErrorResponseType::InvalidGrant => { - self.credential_store + if let Err(store_error) = self + .credential_store .on_refresh_token_rejected(&stored_credentials) - .await?; + .await + { + tracing::warn!( + error = %store_error, + "Failed to handle rejected refresh token" + ); + } return Err(AuthError::TokenRefreshRejected(error.to_string())); } _ => return Err(AuthError::TokenRefreshFailed(error.to_string())), @@ -9437,11 +9444,7 @@ mod tests { let error = manager.refresh_token().await.unwrap_err(); - match fail_at { - Some(_) => assert!(matches!(error, - AuthError::CredentialStoreError(message) if message == "rejection failed")), - None => assert!(matches!(error, AuthError::TokenRefreshRejected(_))), - } + assert!(matches!(error, AuthError::TokenRefreshRejected(_))); assert_eq!( serde_json::to_value(&*store.rejected_credentials.lock().unwrap()).unwrap(), serde_json::to_value(attempted).unwrap() @@ -9456,6 +9459,35 @@ mod tests { assert!(store.lock.try_lock().is_ok()); } + #[tokio::test] + async fn get_access_token_requires_reauth_when_rejection_callback_fails() { + let mut store = refresh_store().await; + store.fail_at = Some("rejected"); + let mut credentials = store.credentials.load().await.unwrap().unwrap(); + credentials + .token_response + .as_mut() + .unwrap() + .set_expires_in(Some(&std::time::Duration::ZERO)); + store.credentials.save(credentials).await.unwrap(); + let http_client = refresh_error_http_client(&store, "invalid_grant"); + let manager = refresh_manager(store.clone(), http_client.clone()).await; + + assert!(matches!( + manager.get_access_token().await, + Err(AuthError::AuthorizationRequired) + )); + + assert_eq!( + *store.events.lock().unwrap(), + [ + "load", "acquire", "acquired", "load", "provider", "rejected", "release" + ] + ); + assert_eq!(http_client.recording.requests().len(), 1); + assert!(store.lock.try_lock().is_ok()); + } + #[tokio::test] async fn transient_refresh_failure_does_not_notify_store_of_rejection() { let store = refresh_store().await;