mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
test(mcp): pin credential isolation across server entries sharing an upstream URL
This commit is contained in:
parent
7e0af8fbbf
commit
a786ba9005
3 changed files with 137 additions and 0 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue