Skip to content
Merged
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
149 changes: 143 additions & 6 deletions crates/rmcp/src/transport/auth.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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. Callback errors are logged, and the
/// manager still returns [`AuthError::TokenRefreshRejected`] so callers can reauthorize.
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
Expand Down Expand Up @@ -2279,6 +2293,7 @@ impl AuthorizationManager {
}
let current_credentials = stored_credentials
.token_response
.as_ref()
.ok_or(AuthError::AuthorizationRequired)?;

let refresh_token = current_credentials
Expand All @@ -2291,26 +2306,39 @@ 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())
if let Err(store_error) = self
.credential_store
.on_refresh_token_rejected(&stored_credentials)
.await
{
tracing::warn!(
error = %store_error,
"Failed to handle rejected refresh token"
);
}
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
Expand Down Expand Up @@ -8679,13 +8707,19 @@ 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();

assert!(
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]
Expand Down Expand Up @@ -9196,6 +9230,7 @@ mod tests {
credentials: InMemoryCredentialStore,
lock: Arc<Mutex<()>>,
events: Arc<StdMutex<Vec<&'static str>>>,
rejected_credentials: Arc<StdMutex<Option<StoredCredentials>>>,
guard_requested: Arc<Semaphore>,
save_started: Arc<Semaphore>,
save_gate: Option<Arc<Semaphore>>,
Expand Down Expand Up @@ -9241,6 +9276,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<Option<CredentialRefreshGuard>, AuthError> {
self.events.lock().unwrap().push("acquire");
self.guard_requested.add_permits(1);
Expand Down Expand Up @@ -9269,6 +9317,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,
Expand Down Expand Up @@ -9311,6 +9360,16 @@ mod tests {
})
}

fn refresh_error_http_client(store: &RefreshStore, error: &str) -> Arc<RefreshHttpClient> {
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<RefreshHttpClient>,
Expand Down Expand Up @@ -9370,6 +9429,84 @@ 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();

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 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;
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;
Expand Down
Loading