From eaf7932d950d20cc9d2c757fd06a481372a40763 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Thu, 25 Jun 2026 21:43:33 -0700 Subject: [PATCH] refactor(mcp): thread user_id/server_id through the TokenRefresher seam The refresh seam took only the OAuthToken, but a refresher needs the server's config (token endpoint, client credentials, scopes) to run the grant and the (user_id, server_id) key to persist the minted token, neither of which is derivable from the token. Widen TokenRefresher.refresh to (user_id, server_id, token) and pass them through from RefreshingTokenStore so each stacked mode PR plugs into the final seam rather than forcing a later signature change across the stack. --- .../outbound_credentials/oauth_token_store.py | 12 ++++++++++-- .../test_oauth_token_store.py | 17 ++++++++++++++--- 2 files changed, 24 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py index 656ffff5c47..e909c99685f 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py @@ -71,9 +71,15 @@ class TokenRefresher(Protocol): ``client_credentials`` grant, or an RFC 8693 re-exchange. Returns ``None`` when it cannot refresh (e.g. no ``refresh_token``), which the caller turns into a 401 challenge. It must persist the new token so later requests (and the surrounding cache) read it without refreshing. + + ``server_id`` selects the upstream's config (token endpoint, client credentials, scopes) the + grant runs against; ``(user_id, server_id)`` is the key the new token is persisted under. They + are not derivable from ``token``, so the seam threads them alongside it. """ - async def refresh(self, token: OAuthToken) -> OAuthToken | None: ... + async def refresh( + self, user_id: str, server_id: str, token: OAuthToken + ) -> OAuthToken | None: ... class CachedOAuthTokenStore: @@ -190,7 +196,9 @@ class RefreshingTokenStore: # result (or exception). The done-callback removes the entry, so the map self-cleans and # is bounded by in-flight refreshes, not by the number of distinct users/servers. The # task is detached from the caller, so a cancelled caller does not abort the refresh. - task = asyncio.ensure_future(self._refresher.refresh(token)) + task = asyncio.ensure_future( + self._refresher.refresh(user_id, server_id, token) + ) self._inflight[key] = task task.add_done_callback(lambda _t, k=key: self._inflight.pop(k, None)) return await task diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_oauth_token_store.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_oauth_token_store.py index c54a027f340..eb7c31e8dc6 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_oauth_token_store.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_oauth_token_store.py @@ -176,13 +176,17 @@ class _RefreshablePair: self._current = initial self.fetch_calls = 0 self.refresh_calls = 0 + self.refresh_args: List[Tuple[str, str]] = [] async def fetch(self, user_id: str, server_id: str) -> Optional[OAuthToken]: self.fetch_calls += 1 return self._current - async def refresh(self, token: OAuthToken) -> Optional[OAuthToken]: + async def refresh( + self, user_id: str, server_id: str, token: OAuthToken + ) -> Optional[OAuthToken]: self.refresh_calls += 1 + self.refresh_args.append((user_id, server_id)) await asyncio.sleep( 0 ) # yield so other concurrent callers reach the lock and wait @@ -210,6 +214,9 @@ async def test_refreshing_mints_a_fresh_token_when_expired(): token = await store.fetch("u", "s") assert token is not None and token.access_token == "refreshed" assert pair.refresh_calls == 1 + assert pair.refresh_args == [ + ("u", "s") + ] # the seam threads the grant/persist key through async def test_refreshing_returns_none_when_it_cannot_refresh(): @@ -217,7 +224,9 @@ async def test_refreshing_returns_none_when_it_cannot_refresh(): async def fetch(self, user_id: str, server_id: str) -> Optional[OAuthToken]: return OAuthToken(access_token="old", expires_at=900.0) - async def refresh(self, token: OAuthToken) -> Optional[OAuthToken]: + async def refresh( + self, user_id: str, server_id: str, token: OAuthToken + ) -> Optional[OAuthToken]: return None # e.g. no refresh_token src = _NoRefresh() @@ -245,7 +254,9 @@ async def test_refresh_failure_is_shared_by_joiners_not_re_run(): async def fetch(self, user_id: str, server_id: str) -> Optional[OAuthToken]: return OAuthToken(access_token="old", expires_at=900.0) - async def refresh(self, token: OAuthToken) -> Optional[OAuthToken]: + async def refresh( + self, user_id: str, server_id: str, token: OAuthToken + ) -> Optional[OAuthToken]: self.calls += 1 await asyncio.sleep(0) # let the concurrent callers join the same task raise RuntimeError("refresh boom")