diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py index 64226eea821..c88027abcd4 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py @@ -149,6 +149,29 @@ async def test_authorization_code_isolates_by_subject(): assert isinstance(bob, Error) and bob.error.tag == "unauthorized" +@pytest.mark.asyncio +async def test_authorization_code_isolates_by_server_id_even_when_servers_share_a_url(): + """A token stored for one server must be invisible to a different server_id pointing at the + same upstream URL: credentials bind to the server entry they were authorized for, so a + recreated or duplicated server starts unauthorized instead of inheriting the old grant. Guards + against any future token lookup keyed on the resource URL instead of (user_id, server_id) -- + both the egress resolve and the has_user_token discovery check must agree.""" + shared_url = "https://upstream.example.com" + store = _FakeTokenStore({("alice", "server-a"): OAuthToken(access_token="at-alice")}) + provider = UpstreamCredentialProvider(oauth_token_store=store) + subject = Subject(tenant_id="", subject_id="alice") + spec_a = ServerSpec(server_id="server-a", resource=shared_url, config=AuthorizationCodeConfig()) + spec_b = ServerSpec(server_id="server-b", resource=shared_url, config=AuthorizationCodeConfig()) + + granted = await provider.resolve_credentials(subject, spec_a) + fresh = await provider.resolve_credentials(subject, spec_b) + + assert isinstance(granted, Ok) and _emitted(granted.ok)["Authorization"] == "Bearer at-alice" + assert isinstance(fresh, Error) and fresh.error.tag == "unauthorized" + assert await provider.has_user_token(subject, spec_a) is True + assert await provider.has_user_token(subject, spec_b) is False + + @pytest.mark.asyncio async def test_has_user_token_reflects_the_stored_token(): present = UpstreamCredentialProvider( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index c95e3bbc42a..c8e871b3bc1 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -753,6 +753,90 @@ async def test_register_client_persist_discriminator_oauth2_persists(): assert await _register_persistence_attempted_for_auth_type(MCPAuth.oauth2) is True +@pytest.mark.asyncio +async def test_register_client_persists_only_to_its_own_row_when_another_server_shares_the_url(): + """A fresh server must mint and persist its OWN DCR client even when another server row with + the same upstream URL already holds one: both the reuse lookup and the persist are keyed by + server_id, never by URL, so OAuth client identity is not transferable between server entries. + If either side ever falls back to a URL match, this fails: the fresh server would skip the + upstream registration (adopting the sibling's client) or persist onto the wrong row.""" + from fastapi import Request + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + register_client_with_server, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + shared_url = "https://provider.example/mcp" + fresh_server = MCPServer( + server_id="server-b", + name="server-b", + server_name="server-b", + alias="server-b", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + url=shared_url, + client_id=None, + client_secret=None, + authorization_url="https://provider.example/oauth/authorize", + token_url="https://provider.example/oauth/token", + registration_url="https://provider.example/oauth/register", + ) + sibling_row_with_client = MagicMock(server_id="server-a", url=shared_url) + sibling_row_with_client.credentials = {"client_id": "client-a-do-not-adopt"} + own_row_without_client = MagicMock(server_id="server-b", url=shared_url) + own_row_without_client.credentials = {} + rows_by_server_id = {"server-a": sibling_row_with_client, "server-b": own_row_without_client} + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://proxy.litellm.example/" + mock_request.headers = {} + + mock_response = MagicMock() + mock_response.json.return_value = {"client_id": "fresh-client-b", "token_endpoint_auth_method": "none"} + mock_response.raise_for_status = MagicMock() + mock_async_client = MagicMock() + mock_async_client.post = AsyncMock(return_value=mock_response) + + mock_update = AsyncMock(return_value=MagicMock()) + + async def _get_row(prisma_client, server_id): + return rows_by_server_id.get(server_id) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=mock_async_client, + ), + patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()), + patch("litellm.proxy._experimental.mcp_server.db.get_mcp_server", new=AsyncMock(side_effect=_get_row)), + patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=mock_update), + patch.object(global_mcp_server_manager, "update_server", new=AsyncMock()), + ): + response = await register_client_with_server( + request=mock_request, + mcp_server=fresh_server, + client_name="Litellm Proxy", + grant_types=["authorization_code", "refresh_token"], + response_types=["code"], + token_endpoint_auth_method="none", + persist_credentials=True, + ) + + mock_async_client.post.assert_called_once() + body = json.loads(response.body.decode("utf-8")) + assert body["client_id"] == "fresh-client-b" + + mock_update.assert_called_once() + update_data = mock_update.call_args.kwargs["data"] + assert update_data.server_id == "server-b" + assert update_data.credentials["client_id"] == "fresh-client-b" + + @pytest.mark.asyncio async def test_register_client_does_not_clobber_token_url_when_absent(): """When the in-memory server has no token_url, the DCR persist must omit it from the diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_token_cache.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_token_cache.py index a60dab9148d..7d2cf4442a5 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_token_cache.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_token_cache.py @@ -92,6 +92,36 @@ async def test_token_cached_across_calls(): assert mock_client.post.call_count == 1 +@pytest.mark.asyncio +async def test_m2m_token_not_shared_across_server_ids_with_identical_config(): + """Two servers with byte-identical client_credentials config but different server_ids must not + share a cached M2M token: the cache is keyed by server_id, so a new server entry (even one + recreated with the same URL and credentials) mints its own token instead of inheriting the + sibling's. Guards against the cache key ever collapsing to the URL or the client config.""" + cache = MCPOAuth2TokenCache() + server_a = _server(server_id="srv-a") + server_b = _server(server_id="srv-b") + mock_client = AsyncMock() + mock_client.post.side_effect = [_token_response("tok-for-a"), _token_response("tok-for-b")] + + with ( + patch( + "litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client", + return_value=mock_client, + ), + patch( + "litellm.proxy._experimental.mcp_server.oauth2_token_cache.mcp_oauth2_token_cache", + cache, + ), + ): + token_a = await resolve_mcp_auth(server_a) + token_b = await resolve_mcp_auth(server_b) + + assert token_a == "tok-for-a" + assert token_b == "tok-for-b" + assert mock_client.post.call_count == 2 + + @pytest.mark.asyncio async def test_per_request_header_beats_oauth2(): """An explicit mcp_auth_header takes priority over the OAuth2 token."""