diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 613f78e0c17..d9533714a81 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -5850,10 +5850,10 @@ class MCPServerManager: if auth_failure is None and not isinstance(exc, UpstreamSessionClosedError): return MCPClient.error_tool_result(exc) if auth_failure is not None: - self._drop_upstream_session(persistent_session) spec: Final = to_server_spec(mcp_server) if spec is not None: await self._cred_provider.invalidate_credentials(to_subject(user_api_key_auth, subject_token), spec) + self._drop_upstream_session(persistent_session) retry_client: Final = await self._create_mcp_client( server=mcp_server, mcp_auth_header=server_auth_header, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index cbc3f288a60..2db8f946475 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -10243,6 +10243,42 @@ class TestOBOCallToolRetry: manager._cred_provider.invalidate_credentials.assert_not_awaited() assert first.attempts == 1 and retry.attempts == 1 + @pytest.mark.asyncio + async def test_a_401_evicts_the_cached_token_before_releasing_peers_on_the_shared_session(self): + """Peers woken by the shared session closing rebuild their client at once, so the stale token + must already be gone from the cache when the session closes or they re-exchange the same token.""" + manager = self._manager() + order: list[str] = [] + + async def slow_invalidate(*_: object) -> None: + await asyncio.sleep(0.01) + order.append("invalidate") + + manager._cred_provider.invalidate_credentials = AsyncMock(side_effect=slow_invalidate) + shared = MagicMock() + shared.close = MagicMock(side_effect=lambda: order.append("close")) + manager._upstream_sessions[("gw", "obo-srv", "fp")] = shared + first = _RetryFakeClient(raises=_UpstreamAuthError(401)) + retry = _RetryFakeClient(result=CallToolResult(content=[], isError=False)) + manager._create_mcp_client = AsyncMock(return_value=retry) + manager._upstream_session_for = AsyncMock(return_value=None) + + await manager._obo_call_tool_with_retry( + client=first, + call_tool_params=MagicMock(), + host_progress_callback=None, + mcp_server=_obo_server(), + server_auth_header=None, + extra_headers=None, + stdio_env=None, + subject_token="caller-jwt", + user_api_key_auth=None, + persistent_session=shared, + ) + + assert order == ["invalidate", "close"], order + assert ("gw", "obo-srv", "fp") not in manager._upstream_sessions + class TestOBOConcurrencyLimit: """OBO (token_exchange) tool calls must honor the server's max_concurrent_requests.