mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): resolve the caller's granted server before the registry pick at connect
The connect preflight runs the scoped router's grant-first lookup for every single-server connect, so a granted plain server named like an ungranted OBO server's alias is served instead of intercepted by that server's 401, and obo_without_subject is read off the server the lookup picks Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
0b262d6594
commit
ee7c8a20d0
2 changed files with 98 additions and 6 deletions
|
|
@ -1625,21 +1625,19 @@ if MCP_AVAILABLE:
|
|||
registry_pick = operations.global_mcp_server_manager.get_mcp_server_answering_to(
|
||||
server_name, client_ip=client_ip
|
||||
)
|
||||
obo_without_subject = (
|
||||
registry_pick is not None
|
||||
and registry_pick.auth_type == MCPAuth.oauth2_token_exchange
|
||||
and not oauth2_headers
|
||||
)
|
||||
allowed_single = (
|
||||
await operations._get_allowed_mcp_servers(
|
||||
user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip
|
||||
)
|
||||
if registry_pick and not obo_without_subject and mcp_servers is not None and len(mcp_servers) == 1
|
||||
if registry_pick and mcp_servers is not None and len(mcp_servers) == 1
|
||||
else ()
|
||||
)
|
||||
granted = next(iter(allowed_single), None)
|
||||
server = granted if granted is not None else registry_pick
|
||||
granted_single = granted is not None
|
||||
obo_without_subject = (
|
||||
server is not None and server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers
|
||||
)
|
||||
if server is not None and allowed_server_ids is not None and server.server_id not in allowed_server_ids:
|
||||
# Caller's narrowed scope excludes this server — skip the
|
||||
# preemptive challenge and let downstream authorization
|
||||
|
|
|
|||
|
|
@ -8985,6 +8985,100 @@ class TestConnectPreflightRoutesLikeTheScopedRouter:
|
|||
)
|
||||
assert guardrail.asked_about == ["d-id"]
|
||||
|
||||
@staticmethod
|
||||
def _register_obo_alias_collision() -> tuple[MCPServer, MCPServer]:
|
||||
obo: Final = MCPServer(
|
||||
server_id="o-id",
|
||||
name="obo_server",
|
||||
server_name="obo_server",
|
||||
alias="obo",
|
||||
url="https://obo.test/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2_token_exchange,
|
||||
token_exchange_endpoint="https://idp.test/token",
|
||||
client_id="cid",
|
||||
client_secret="csecret",
|
||||
)
|
||||
plain: Final = MCPServer(
|
||||
server_id="p-id",
|
||||
name="obo",
|
||||
server_name="obo",
|
||||
url="https://plain.test/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
)
|
||||
mcp_operations.global_mcp_server_manager.registry.update({"o-id": obo, "p-id": plain})
|
||||
return obo, plain
|
||||
|
||||
@staticmethod
|
||||
async def _connect_to(route_name: str) -> None:
|
||||
from litellm.proxy._experimental.mcp_server import server as server_module
|
||||
|
||||
await server_module._raise_preemptive_401_for_unauthenticated_servers(
|
||||
scope={"type": "http", "method": "POST", "path": f"/mcp/{route_name}", "headers": []},
|
||||
mcp_servers=[route_name],
|
||||
oauth2_headers=None,
|
||||
mcp_server_auth_headers=None,
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-collision", user_id="u-1"),
|
||||
client_ip=None,
|
||||
raw_headers={"x-litellm-api-key": "sk-collision"},
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("route_name", ["obo", "OBO"], ids=["exact_alias", "case_variant"])
|
||||
async def test_granted_plain_server_connects_past_an_ungranted_obo_alias_holder(self, route_name):
|
||||
_, plain = self._register_obo_alias_collision()
|
||||
preflight: Final = AsyncMock()
|
||||
with (
|
||||
patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it
|
||||
mcp_operations.global_mcp_server_manager, "get_allowed_mcp_servers", AsyncMock(return_value=["p-id"])
|
||||
),
|
||||
patch.object( # test-quality-ok: a real exchanger would call an IdP; this one records which server ran
|
||||
mcp_operations.global_mcp_server_manager, "preflight_token_exchange", preflight
|
||||
),
|
||||
):
|
||||
await self._connect_to(route_name)
|
||||
|
||||
assert preflight.await_args is not None, "the granted plain server must reach the preflight, not a 401"
|
||||
assert preflight.await_args.kwargs["server"] is plain
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("route_name", ["obo", "obo_server"], ids=["alias", "server_name"])
|
||||
async def test_granted_obo_server_still_challenges_without_a_subject(self, route_name, monkeypatch):
|
||||
monkeypatch.delenv("SERVER_ROOT_PATH", raising=False)
|
||||
self._register_obo_alias_collision()
|
||||
with (
|
||||
patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it
|
||||
mcp_operations.global_mcp_server_manager, "get_allowed_mcp_servers", AsyncMock(return_value=["o-id"])
|
||||
),
|
||||
pytest.raises(HTTPException) as exc,
|
||||
):
|
||||
await self._connect_to(route_name)
|
||||
|
||||
assert exc.value.status_code == 401
|
||||
assert ((exc.value.headers or {}).get("WWW-Authenticate") or "").startswith(
|
||||
'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/obo"'
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("route_name", ["obo", "OBO"], ids=["exact_alias", "case_variant"])
|
||||
async def test_key_granting_neither_server_is_not_preflighted_for_the_plain_one(self, route_name):
|
||||
self._register_obo_alias_collision()
|
||||
preflight: Final = AsyncMock()
|
||||
with (
|
||||
patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it
|
||||
mcp_operations.global_mcp_server_manager, "get_allowed_mcp_servers", AsyncMock(return_value=[])
|
||||
),
|
||||
patch.object( # test-quality-ok: a real exchanger would call an IdP; this one records which server ran
|
||||
mcp_operations.global_mcp_server_manager, "preflight_token_exchange", preflight
|
||||
),
|
||||
pytest.raises(HTTPException) as exc,
|
||||
):
|
||||
await self._connect_to(route_name)
|
||||
|
||||
assert exc.value.status_code == 401
|
||||
assert preflight.await_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_allowed_mcp_servers_from_mcp_server_names_mixed_known_and_unknown():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue