diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 5282b228f92..4ffa2ba5534 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -7199,12 +7199,18 @@ class MCPServerManager: return server return None - def get_mcp_server_answering_to(self, name: str, client_ip: str | None = None) -> MCPServer | None: + def get_mcp_server_answering_to( + self, name: str, client_ip: str | None = None, *, among: Sequence[MCPServer] | None = None + ) -> MCPServer | None: """The one server a ``/mcp/{name}`` segment denotes, shared by the connect preflight, the scoped router, and RFC 9728 discovery so all three name the same server: the exact ``get_mcp_server_by_name`` priority first, then the exact ``server_id``, then the name priority case-insensitively, then any prefix form routing accepts. A name that denotes a server hidden from ``client_ip`` resolves to ``None`` at the - pass that found it: it never falls through to a looser pass that could name another server.""" + pass that found it: it never falls through to a looser pass that could name another server. ``among`` + runs the same passes over those servers alone instead of the registry, which is how the scoped router + picks the caller's granted server answering to ``name``.""" + if among is not None: + return self._server_among_answering_to(name, tuple(among), client_ip) exact: Final = self.get_mcp_server_by_name(name) if exact is not None: return exact if self._is_server_accessible_from_ip(exact, client_ip) else None @@ -7230,6 +7236,33 @@ class MCPServerManager: None, ) + def _server_among_answering_to( + self, name: str, servers: Sequence[MCPServer], client_ip: str | None + ) -> MCPServer | None: + """``get_mcp_server_answering_to`` over ``servers`` instead of the registry: the same passes in the + same order, with a server hidden from ``client_ip`` resolving to ``None`` at the pass that found it.""" + requested: Final = name.lower() + passes: Final[tuple[Callable[[MCPServer], bool], ...]] = ( + lambda server: server.alias == name, + lambda server: server.server_name == name, + lambda server: server.name == name, + lambda server: server.server_id == name, + lambda server: (server.alias or "").lower() == requested, + lambda server: (server.server_name or "").lower() == requested, + lambda server: (server.name or "").lower() == requested, + ) + for matches in passes: + if (found := next((server for server in servers if matches(server)), None)) is not None: + return found if self._is_server_accessible_from_ip(found, client_ip) else None + return next( + ( + server + for server in servers + if self._is_server_accessible_from_ip(server, client_ip) and server_answers_to_name(server, name) + ), + None, + ) + def get_filtered_registry(self, client_ip: str | None = None) -> dict[str, MCPServer]: """ Get registry filtered by client IP access control. diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 15a9183ef7a..8636afbebef 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -501,17 +501,22 @@ def _server_answers_to(server: MCPServer, name: str) -> bool: def _scoped_server( name: str, allowed_mcp_servers: Sequence[MCPServer], client_ip: str | None ) -> MCPServer | Literal["denied"] | None: - """The granted server a scoped ``name`` selects: the registry's ``get_mcp_server_answering_to`` pick, made - with the same ``client_ip`` the connect preflight and discovery use, when the caller holds it. ``"denied"`` - when the registry names a server the caller does not hold, or one hidden from ``client_ip``, so the name - is neither rerouted to another granted server nor retried as an access group. ``None`` when the registry - cannot place the name for any caller, after trying the granted servers answering to it.""" + """The server a scoped ``name`` selects for the caller, in this order. ``"denied"`` when the registry's + ``get_mcp_server_answering_to`` pick, made with the same ``client_ip`` the connect preflight and discovery + use, is a server hidden from that IP. Otherwise the granted server answering to ``name``: the registry's + own pass order run over ``allowed_mcp_servers`` alone, so a granted server wins over an ungranted alias or + case variant the registry would pick. With no granted server answering: ``"denied"`` when the registry + places the name on a server the caller does not hold, so it is not retried as an access group; ``None`` + when the registry cannot place the name for any caller.""" registry_pick: Final = global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip) - if registry_pick is not None: - return next((s for s in allowed_mcp_servers if s.server_id == registry_pick.server_id), "denied") - if global_mcp_server_manager.get_mcp_server_answering_to(name) is not None: + if registry_pick is None and global_mcp_server_manager.get_mcp_server_answering_to(name) is not None: return "denied" - return next((s for s in allowed_mcp_servers if s and _server_answers_to(s, name)), None) + granted: Final = global_mcp_server_manager.get_mcp_server_answering_to( + name, client_ip=client_ip, among=allowed_mcp_servers + ) + if granted is not None: + return granted + return "denied" if registry_pick is not None else None async def raise_denied_scoped_mcp_access( diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 988f77ea98b..49fa2f898ed 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1603,7 +1603,24 @@ if MCP_AVAILABLE: a server it will be 403'd on immediately after authentication. """ for server_name in mcp_servers or []: - server = operations.global_mcp_server_manager.get_mcp_server_answering_to(server_name, client_ip=client_ip) + 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 + else () + ) + granted = next(iter(allowed_single), None) + server = granted if granted is not None else registry_pick + granted_single = granted is not None 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 @@ -1703,7 +1720,7 @@ if MCP_AVAILABLE: # JSON-RPC error and the WWW-Authenticate header is lost. OBO keeps its connect gate; # guardrail-only gates fire only on a single-server connect the key's grant admits, so a # key without access gets the grant's 403 instead of a sign-in it could not use. The one - # admission lookup below serves the challenge, the sign-in preflight and the exchange. + # admission lookup above serves the challenge, the sign-in preflight and the exchange. sign_in = caller_sign_in_for(server, user_api_key_auth) if server is not None else None subject_token = ( operations.global_mcp_server_manager._extract_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight @@ -1712,19 +1729,6 @@ if MCP_AVAILABLE: if server is not None else None ) - obo_without_subject = ( - server is not None and server.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 server and not obo_without_subject and mcp_servers is not None and len(mcp_servers) == 1 - else () - ) - granted_single = server is not None and any( - allowed.server_id == server.server_id for allowed in allowed_single - ) if server and sign_in is not None and subject_token is None and (obo_without_subject or granted_single): from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph raise_token_exchange_challenge, diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py index cb4c690e18c..d7ba2a3cc93 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py @@ -924,7 +924,7 @@ async def test_get_tools_from_mcp_servers(): mock_manager.get_mcp_server_by_id = lambda server_id: ( mock_server_1 if server_id == "server1_id" else mock_server_2 ) - mock_manager.get_mcp_server_answering_to = MagicMock(return_value=None) + mock_manager.get_mcp_server_answering_to = MCPServerManager().get_mcp_server_answering_to mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1]) # Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test) mock_manager.filter_server_ids_by_ip_with_info = MagicMock( diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index 14f284389e0..c733ba58b1c 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -32,7 +32,7 @@ import litellm from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.proxy._experimental.mcp_server import operations as mcp_operations from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var -from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller +from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller, MCPServerManager from litellm.proxy._types import ( LiteLLM_MCPServerTable, MCPTransport, @@ -1279,7 +1279,7 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails(): mock_manager.get_mcp_server_by_id = lambda server_id: ( working_server if server_id == "working_server" else failing_server ) - mock_manager.get_mcp_server_answering_to = lambda name, client_ip=None: None + mock_manager.get_mcp_server_answering_to = MCPServerManager().get_mcp_server_answering_to # Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering) mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: ( server_ids, @@ -6765,7 +6765,7 @@ async def test_list_tools_with_legacy_db_m2m_server_resolves_oauth2_flow(): ): mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["legacy-m2m-id"]) mock_manager.get_mcp_server_by_id = MagicMock(return_value=legacy_server) - mock_manager.get_mcp_server_answering_to = MagicMock(return_value=None) + mock_manager.get_mcp_server_answering_to = MCPServerManager().get_mcp_server_answering_to mock_manager.filter_server_ids_by_ip_with_info = MagicMock(return_value=(["legacy-m2m-id"], 0)) mock_manager._get_tools_from_server = AsyncMock(side_effect=capture_extra_headers) @@ -8620,7 +8620,9 @@ async def test_scoped_router_selects_the_server_the_connect_preflight_resolves(a finally: global_mcp_server_manager.registry.clear() - assert only_b == [], "a name the registry gives to an ungranted server must not fall through to another" + assert [s.server_id for s in only_b] == ["b-id"], ( + "the granted server answering to the name wins over the registry's ungranted alias holder" + ) @pytest.mark.asyncio @@ -8691,10 +8693,202 @@ async def test_scoped_router_hides_a_private_server_from_an_external_ip_like_the global_mcp_server_manager.registry.clear() assert external == [], "a name the preflight hides from this IP must not reroute to a case variant" - assert internal == [] + assert [s.server_id for s in internal] == ["u-id"], "with no IP hiding in play the granted case variant wins" assert [s.server_id for s in by_own_name] == ["u-id"] +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("registry", "scope", "granted", "expected"), + [ + pytest.param(("a-id", "d-id"), "docs", ("d-id",), ["d-id"], id="alias-collision-alias-holder-listed-first"), + pytest.param(("d-id", "a-id"), "docs", ("d-id",), ["d-id"], id="alias-collision-exact-name-listed-first"), + pytest.param(("g1", "g2"), "GITHUB", ("g2",), ["g2"], id="case-variant-collision"), + pytest.param(("p-id", "m-id"), "shared", ("m-id",), [], id="ungranted-only-name-stays-denied"), + ], +) +async def test_scoped_router_prefers_the_granted_server_answering_to_the_name(registry, scope, granted, expected): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + from litellm.proxy._experimental.mcp_server.server import _get_allowed_mcp_servers_from_mcp_server_names + + servers: Final = { + "a-id": MCPServer( + server_id="a-id", name="a_docs", server_name="a_docs", alias="docs", transport=MCPTransport.http + ), + "d-id": MCPServer(server_id="d-id", name="docs", server_name="docs", transport=MCPTransport.http), + "g1": MCPServer(server_id="g1", name="GitHub", server_name="GitHub", transport=MCPTransport.http), + "g2": MCPServer(server_id="g2", name="github", server_name="github", transport=MCPTransport.http), + "p-id": MCPServer(server_id="p-id", name="p", server_name="p", alias="shared", transport=MCPTransport.http), + "m-id": MCPServer(server_id="m-id", name="m", server_name="m", transport=MCPTransport.http), + } + global_mcp_server_manager.registry.clear() + global_mcp_server_manager.registry.update({server_id: servers[server_id] for server_id in registry}) + try: + with patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." + "MCPRequestHandler._get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + return_value=["m-id"], + ) as groups: + selected: Final = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=[scope], allowed_mcp_servers=[servers[server_id] for server_id in granted] + ) + finally: + global_mcp_server_manager.registry.clear() + + assert [s.server_id for s in selected] == expected + groups.assert_not_awaited() + + +def test_get_mcp_server_answering_to_among_applies_the_registry_pass_order_and_ip_hiding(): + manager: Final = MCPServerManager() + by_alias: Final = MCPServer(server_id="a-id", name="a", server_name="a", alias="svc", transport=MCPTransport.http) + by_server_name: Final = MCPServer(server_id="b-id", name="b", server_name="svc", transport=MCPTransport.http) + by_name: Final = MCPServer(server_id="c-id", name="svc", server_name="c", transport=MCPTransport.http) + by_id: Final = MCPServer(server_id="svc", name="d", server_name="d", transport=MCPTransport.http) + folded: Final = MCPServer(server_id="e-id", name="e", server_name="SVC", transport=MCPTransport.http) + hidden: Final = MCPServer( + server_id="h-id", + name="h", + server_name="h", + alias="svc", + transport=MCPTransport.http, + available_on_public_internet=False, + ) + + assert manager.get_mcp_server_answering_to("svc", among=[by_name, by_server_name, by_alias]) is by_alias + assert manager.get_mcp_server_answering_to("svc", among=[by_name, by_server_name]) is by_server_name + assert manager.get_mcp_server_answering_to("svc", among=[by_id, by_name]) is by_name + assert manager.get_mcp_server_answering_to("svc", among=[folded, by_id]) is by_id + assert manager.get_mcp_server_answering_to("svc", among=[folded]) is folded + assert manager.get_mcp_server_answering_to("E-ID", among=[folded]) is folded + assert manager.get_mcp_server_answering_to("svc", among=()) is None + assert manager.get_mcp_server_answering_to("svc", client_ip="203.0.113.7", among=[hidden, folded]) is None + assert manager.get_mcp_server_answering_to("svc", client_ip="10.0.0.7", among=[hidden, folded]) is hidden + assert manager.get_mcp_server_answering_to("svc", client_ip="203.0.113.7", among=[folded]) is folded + assert manager.get_mcp_server_answering_to("svc") is None + + manager.registry = {"e-id": folded, "b-id": by_server_name, "a-id": by_alias} + + assert manager.get_mcp_server_answering_to("svc") is by_alias + assert manager.get_mcp_server_answering_to("svc") is manager.get_mcp_server_answering_to( + "svc", among=tuple(manager.registry.values()) + ) + + +class _GrantedServerSignInGuardrail(CustomGuardrail): + """Requires caller sign-in on one server only and records every server it is asked about.""" + + def __init__(self, *args, gated_server_id: str, **kwargs): + super().__init__(*args, **kwargs) + self._gated_server_id = gated_server_id + self.asked_about = [] # mutable-ok: call recorder + + def caller_sign_in(self, server, user_api_key_auth): + from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn + + self.asked_about.append(server.server_id) + if server.server_id != self._gated_server_id: + return None + return CallerSignIn(issuers=("https://idp.test",), scopes=("scope-a",)) + + async def preflight_caller_sign_in(self, server, user_api_key_auth, subject_token): + from litellm.proxy._experimental.mcp_server.caller_sign_in import SignedIn + + return SignedIn() + + +class TestConnectPreflightRoutesLikeTheScopedRouter: + """A granted key connecting to ``/mcp/{name}`` is pre-flighted for the server the scoped router routes + it to, so an ungranted server holding the name as an alias neither hides the granted server's sign-in + challenge nor skips its connect-time exchange.""" + + @staticmethod + def _register_alias_collision() -> None: + alias_holder: Final = MCPServer( + server_id="a-id", + name="a_docs", + server_name="a_docs", + alias="docs", + url="https://a.test/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + ) + granted: Final = MCPServer( + server_id="d-id", + name="docs", + server_name="docs", + url="https://d.test/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + ) + mcp_operations.global_mcp_server_manager.registry.update({"a-id": alias_holder, "d-id": granted}) + + @staticmethod + async def _connect_to_docs() -> 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": "/mcp/docs", "headers": []}, + mcp_servers=["docs"], + oauth2_headers=None, + mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-granted-docs", user_id="u-1"), + client_ip=None, + raw_headers={"x-litellm-api-key": "sk-granted-docs"}, + ) + + @pytest.mark.asyncio + async def test_exchange_runs_for_the_granted_server_not_the_alias_holder(self): + self._register_alias_collision() + + async def refuse_exchange(server, **kwargs): + raise HTTPException( + status_code=401, detail=f"exchange refused for {server.server_id} as {kwargs['connected_as']}" + ) + + 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=["d-id"]) + ), + patch.object( # test-quality-ok: a real exchanger would call an IdP; this one reports which server it ran for + mcp_operations.global_mcp_server_manager, "preflight_token_exchange", refuse_exchange + ), + pytest.raises(HTTPException) as exc, + ): + await self._connect_to_docs() + + assert exc.value.status_code == 401 + assert exc.value.detail == "exchange refused for d-id as docs" + + @pytest.mark.asyncio + async def test_sign_in_challenge_names_the_granted_server_not_the_alias_holder(self, monkeypatch): + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + self._register_alias_collision() + guardrail: Final = _GrantedServerSignInGuardrail(guardrail_name="sign-in-stub", gated_server_id="d-id") + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + 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=["d-id"]), + ), + pytest.raises(HTTPException) as exc, + ): + await self._connect_to_docs() + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + 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/docs"' + ) + assert guardrail.asked_about == ["d-id"] + + @pytest.mark.asyncio async def test_get_allowed_mcp_servers_from_mcp_server_names_mixed_known_and_unknown(): """ @@ -9836,7 +10030,7 @@ async def test_aggregate_listing_reports_per_server_outcomes(): mock_manager.get_mcp_server_by_id = lambda server_id: ( working_server if server_id == "working_server" else broken_server ) - mock_manager.get_mcp_server_answering_to = lambda name, client_ip=None: None + mock_manager.get_mcp_server_answering_to = MCPServerManager().get_mcp_server_answering_to mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0) async def mock_get_tools_from_server(server, **kwargs): @@ -10351,9 +10545,7 @@ class TestOboPreflightScopedToAllowedServers: requested = _make_obo_server("obo_tools") key = UserAPIKeyAuth(api_key="sk-plain-only") - allowed_lookup, preflight = await self._run( - requested, allowed=[_make_obo_server("plain_tools")], user_api_key_auth=key - ) + allowed_lookup, preflight = await self._run(requested, allowed=[], user_api_key_auth=key) preflight.assert_not_awaited() allowed_lookup.assert_awaited_once_with(