From 001fcad93076ed27b651e3221882549456a9b8fc Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 09:31:42 +0000 Subject: [PATCH] test(mcp): align listed-tool slot tests with per-caller keying Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/test_mcp_server_manager.py | 65 ++++++++++--------- 1 file changed, 36 insertions(+), 29 deletions(-) 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 7000497769f..26d4263eac1 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 @@ -7007,10 +7007,8 @@ class TestMCPServerManager: # Mock _create_mcp_client to return our mock client manager._create_mcp_client = AsyncMock(return_value=mock_client) - # Mock user auth with no restrictions - user_api_key_auth = MagicMock() - user_api_key_auth.object_permission = None - user_api_key_auth.object_permission_id = None + # Real auth: the listed-tool slot identity is hashed from these fields + user_api_key_auth = UserAPIKeyAuth(api_key="sk-test") # Mock proxy logging proxy_logging_obj = MagicMock() @@ -7037,7 +7035,9 @@ class TestMCPServerManager: assert mock_client.call_tool.await_count == 1 @staticmethod - def _manager_ready_for_call_tool(listed_tools: list[MCPTool]) -> tuple[MCPServerManager, MagicMock]: + def _manager_ready_for_call_tool( + listed_tools: list[MCPTool], caller: ListedToolsCaller | None = None + ) -> tuple[MCPServerManager, MagicMock]: from mcp.types import CallToolResult manager = MCPServerManager() @@ -7050,7 +7050,7 @@ class TestMCPServerManager: manager.registry = {"test-server": server} manager.tool_name_to_mcp_server_name_mapping["test_tool"] = "test-server" manager.tool_name_to_mcp_server_name_mapping["test-server-test_tool"] = "test-server" - manager._create_prefixed_tools(listed_tools, server) + manager._create_prefixed_tools(listed_tools, server, caller=caller) mock_client = AsyncMock() mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) @@ -7064,23 +7064,23 @@ class TestMCPServerManager: return manager, proxy_logging_obj @staticmethod - def _unrestricted_auth() -> MagicMock: - user_api_key_auth = MagicMock() - user_api_key_auth.object_permission = None - user_api_key_auth.object_permission_id = None - return user_api_key_auth + def _unrestricted_auth() -> UserAPIKeyAuth: + return UserAPIKeyAuth(api_key="sk-test") @pytest.mark.asyncio async def test_call_tool_hands_listed_tool_description_and_schema_to_pre_call_hooks(self): schema = {"type": "object", "properties": {"param": {"type": "string"}}, "required": ["param"]} listed = [MCPTool(name="test_tool", description="Runs the test tool", inputSchema=schema)] - manager, proxy_logging_obj = self._manager_ready_for_call_tool(listed) + auth = self._unrestricted_auth() + manager, proxy_logging_obj = self._manager_ready_for_call_tool( + listed, caller=ListedToolsCaller(user_api_key_auth=auth) + ) await manager.call_tool( server_name="test-server", name="test_tool", arguments={"param": "value"}, - user_api_key_auth=self._unrestricted_auth(), + user_api_key_auth=auth, proxy_logging_obj=proxy_logging_obj, ) @@ -7091,7 +7091,8 @@ class TestMCPServerManager: async def test_call_tool_hands_listed_tool_metadata_to_during_call_hooks_through_real_conversion(self): schema = {"type": "object", "properties": {"param": {"type": "string"}}} listed = [MCPTool(name="test_tool", description="Runs the test tool", inputSchema=schema)] - manager, _ = self._manager_ready_for_call_tool(listed) + auth = UserAPIKeyAuth(api_key="sk-test") + manager, _ = self._manager_ready_for_call_tool(listed, caller=ListedToolsCaller(user_api_key_auth=auth)) proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) @@ -7100,7 +7101,7 @@ class TestMCPServerManager: server_name="test-server", name="test_tool", arguments={"param": "value"}, - user_api_key_auth=UserAPIKeyAuth(api_key="sk-test"), + user_api_key_auth=auth, proxy_logging_obj=proxy_logging_obj, ) @@ -7112,15 +7113,17 @@ class TestMCPServerManager: @pytest.mark.asyncio async def test_call_tool_passes_no_tool_metadata_when_tool_was_never_listed(self): + auth = self._unrestricted_auth() manager, proxy_logging_obj = self._manager_ready_for_call_tool( - [MCPTool(name="other_tool", description="Unrelated", inputSchema={"type": "object"})] + [MCPTool(name="other_tool", description="Unrelated", inputSchema={"type": "object"})], + caller=ListedToolsCaller(user_api_key_auth=auth), ) await manager.call_tool( server_name="test-server", name="test_tool", arguments={"param": "value"}, - user_api_key_auth=self._unrestricted_auth(), + user_api_key_auth=auth, proxy_logging_obj=proxy_logging_obj, ) @@ -7244,7 +7247,9 @@ class TestMCPServerManager: caller=ListedToolsCaller(user_api_key_auth=alice), ) for_bob = manager.get_listed_tool(shared, "echo", ListedToolsCaller(user_api_key_auth=bob)) - assert for_bob is not None and for_bob.description == "everyone" + assert for_bob is None, "keyed callers get their own slot even on servers without upstream per-user auth" + anonymous = manager.get_listed_tool(shared, "echo") + assert anonymous is None @pytest.mark.parametrize( ("server_kwargs", "caller_a", "caller_b"), @@ -7371,20 +7376,24 @@ class TestMCPServerManager: user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm"), ) - caller: Final = ListedToolsCaller(mcp_auth_header="Bearer hdr") + caller: Final = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm"), mcp_auth_header="Bearer hdr" + ) listed = manager.get_listed_tool(server, "turn", caller) assert listed is not None and listed.description == "t" assert manager._create_mcp_client.await_args.kwargs["mcp_auth_header"] == "Bearer hdr" @pytest.mark.parametrize( - ("signer", "static_headers", "shared"), + ("signer", "static_headers"), [ - pytest.param(MagicMock(), None, False, id="signer-mints-per-caller-authorization"), - pytest.param(MagicMock(), {"Authorization": "Bearer admin-token"}, True, id="static-authorization-wins"), - pytest.param(None, None, True, id="no-signer-stays-shared"), + pytest.param(MagicMock(), None, id="signer"), + pytest.param(MagicMock(), {"Authorization": "Bearer admin-token"}, id="static-authorization"), + pytest.param(None, None, id="no-signer"), ], ) - def test_jwt_signer_makes_a_shared_server_list_per_caller(self, signer, static_headers, shared): + def test_keyed_callers_always_list_into_their_own_slot(self, signer, static_headers): + """The catalog is guardrail-shaped per key, so a keyed caller never reads another caller's + listing regardless of the signer or static authorization configuration.""" manager = MCPServerManager() server = MCPServer( server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", static_headers=static_headers @@ -7399,12 +7408,10 @@ class TestMCPServerManager: manager._create_prefixed_tools( [MCPTool(name="turn", description="alice view", inputSchema={})], server, caller=alice ) - for_bob = manager.get_listed_tool(server, "srv-turn", bob) + assert manager.get_listed_tool(server, "srv-turn", bob) is None - if shared: - assert for_bob is not None and for_bob.description == "alice view" - else: - assert for_bob is None + for_alice = manager.get_listed_tool(server, "srv-turn", alice) + assert for_alice is not None and for_alice.description == "alice view" def test_signed_server_slot_splits_on_the_callers_key_not_only_the_user(self): """Two keys sharing a user_id get different signed JWTs, so they split; the same key