diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index ae790575015..bfa786f4bce 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1558,13 +1558,12 @@ class MCPServerManager: openapi_key_prefix = prefix_root + MCP_TOOL_PREFIX_SEPARATOR global_mcp_tool_registry.unregister_tools_with_prefix(openapi_key_prefix) - surviving = { - tool_name: remaining - for tool_name, owner_ids in self.tool_name_to_mcp_server_ids_mapping.items() - if (remaining := owner_ids - {server.server_id}) - } - self.tool_name_to_mcp_server_ids_mapping.clear() - self.tool_name_to_mcp_server_ids_mapping.update(surviving) + for tool_name in list(self.tool_name_to_mcp_server_ids_mapping): + remaining = self.tool_name_to_mcp_server_ids_mapping[tool_name] - {server.server_id} + if remaining: + self.tool_name_to_mcp_server_ids_mapping[tool_name] = remaining + else: + del self.tool_name_to_mcp_server_ids_mapping[tool_name] def remove_server(self, mcp_server: LiteLLM_MCPServerTable): """ diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 99b1e6b4207..23e853a32d3 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2690,8 +2690,13 @@ if MCP_AVAILABLE: litellm_logging_obj.model_call_details["model"] = f"MCP: {name}" # Resolve the MCP server early so BYOK checks and credential injection # apply to ALL dispatch paths (local tool registry AND managed MCP server). + # Stay within the caller's scope: a scope-blind lookup here would undo the + # scope decision already made above and dispatch to an out-of-scope server. if mcp_server is None: - mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name) + allowed_server_ids = {server.server_id for server in allowed_mcp_servers} + fallback_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name) + if fallback_server is not None and fallback_server.server_id in allowed_server_ids: + mcp_server = fallback_server if mcp_server: standard_logging_mcp_tool_call["mcp_server_cost_info"] = (mcp_server.mcp_info or {}).get( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 1576a4f4270..2d0bba5efc5 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -6161,6 +6161,71 @@ async def test_execute_mcp_tool_jsonrpc_unprefixed_ambiguous_tool_is_rejected(): assert "server" not in injected +@pytest.mark.asyncio +async def test_execute_mcp_tool_jsonrpc_never_dispatches_to_an_out_of_scope_server(): + """End to end: a tool whose only owner is outside the caller's scope must not be dispatched. + + resolve_tool_route already fails closed, but a later scope-blind fallback that only runs + when no server resolved would re-resolve the out-of-scope owner and dispatch to it, + bypassing the permission gate because no server_name was set for it to check. + """ + from litellm.proxy._experimental.mcp_server import server as mcp_module + + in_scope = MCPServer( + server_id="in-scope-id", + name="echo_alpha", + server_name="echo_alpha", + url="http://127.0.0.1:5115/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + authentication_token="in-scope-secret", + ) + out_of_scope = MCPServer( + server_id="out-of-scope-id", + name="echo_zulu", + server_name="echo_zulu", + url="http://127.0.0.1:5115/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + authentication_token="out-of-scope-secret", + ) + + dispatched: dict = {} + + async def fake_create_mcp_client(server, **kwargs): + dispatched["server"] = server + raise AssertionError("must not dispatch to an out-of-scope server") + + with ( + patch.dict( + mcp_module.global_mcp_server_manager.tool_name_to_mcp_server_ids_mapping, + {"secret_tool": frozenset({out_of_scope.server_id})}, + ), + patch.object( + mcp_module.global_mcp_server_manager, + "get_registry", + return_value={in_scope.server_id: in_scope, out_of_scope.server_id: out_of_scope}, + ), + patch.object( + mcp_module.global_mcp_server_manager, + "_create_mcp_client", + new=fake_create_mcp_client, + ), + patch.object(mcp_module.MCPRequestHandler, "is_tool_allowed", return_value=True), + patch.object(mcp_module.global_mcp_tool_registry, "get_tool", return_value=None), + patch("litellm.proxy.proxy_server.proxy_logging_obj", None), + pytest.raises(HTTPException), + ): + await mcp_module.execute_mcp_tool( + name="secret_tool", + arguments={"message": "hello"}, + allowed_mcp_servers=[in_scope], + start_time=datetime.now(), + ) + + assert "server" not in dispatched + + @pytest.mark.asyncio async def test_execute_mcp_tool_jsonrpc_unprefixed_resolves_when_caller_reaches_one_server(): """A caller scoped to one server is served unprefixed names and must still route. 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 fd2e5046c13..cc6cf64b023 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 @@ -4006,6 +4006,28 @@ class TestMCPServerManager: assert manager.tool_name_to_mcp_server_ids_mapping["echo"] == frozenset({"id-alpha"}) assert manager._get_mcp_server_from_tool_name("echo") is alpha + def test_cleanup_mutates_the_mapping_in_place_without_emptying_it(self): + """Cleanup must keep the same dict object and never clear it wholesale. + + The initialize task is dispatched without being awaited and holds a reference to + this dict, so rebinding it would drop that task's writes; clearing it would let a + concurrent reader briefly see an empty map and mis-route. + """ + manager = MCPServerManager() + alpha = MCPServer(server_id="id-alpha", name="echo_alpha", transport=MCPTransport.http) + zulu = MCPServer(server_id="id-zulu", name="echo_zulu", transport=MCPTransport.http) + manager.registry = {"id-alpha": alpha, "id-zulu": zulu} + manager._register_tool_route("shared", "id-alpha") + manager._register_tool_route("shared", "id-zulu") + manager._register_tool_route("alpha_only", "id-alpha") + original_map = manager.tool_name_to_mcp_server_ids_mapping + + manager._cleanup_server_tool_routing_artifacts(zulu) + + assert manager.tool_name_to_mcp_server_ids_mapping is original_map + assert manager.tool_name_to_mcp_server_ids_mapping["shared"] == frozenset({"id-alpha"}) + assert manager.tool_name_to_mcp_server_ids_mapping["alpha_only"] == frozenset({"id-alpha"}) + def test_cleanup_drops_the_route_when_its_last_owner_leaves(self): manager = MCPServerManager() alpha = MCPServer(server_id="id-alpha", name="echo_alpha", transport=MCPTransport.http)