From ad746dadb7d4019966dae970a05a6de9da944a91 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Thu, 1 Oct 2026 14:22:05 -0700 Subject: [PATCH] fix(mcp): record a listed-tool catalog only for a listing the caller is served _get_tools_from_server now records the catalog into the caller's listed-tools slot only when asked (record_listing=True), which the served listings pass: the /mcp and Responses API tools/list handlers via _get_tools_from_mcp_servers, MCPServerManager.list_tools, and the REST listing via _list_server_tools. Four internal listings stop recording, so a later tools/call hands pre_mcp_call hooks name and arguments only, as on main: - _list_tools_before_first_call, the implicit listing inside tools/call when this worker does not yet expose the tool - fetch_pinnable_tool_catalog, the admin pin snapshot listed without the catalog guard and without description overrides - _initialize_tool_name_to_mcp_server_name_mapping, the startup fill - get_tools_for_server, used by the semantic tool filter _create_prefixed_tools returns to its tool-name mapping job only; the record follows it in _get_tools_from_server. --- .../mcp_server/mcp_server_manager.py | 14 +- .../_experimental/mcp_server/operations.py | 6 + .../mcp_server/rest_endpoints.py | 13 +- .../mcp_server/test_mcp_server.py | 2 + .../mcp_server/test_mcp_server_manager.py | 166 +++++++++++++----- .../test_mcp_server_tool_calls_and_headers.py | 81 ++++++++- 6 files changed, 227 insertions(+), 55 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index f37916cad3f..2d13dc0dcf3 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -3870,6 +3870,7 @@ class MCPServerManager: server=server, mcp_auth_header=server_auth_header, user_api_key_auth=user_api_key_auth, + record_listing=True, ) return tools except Exception as e: @@ -4482,6 +4483,7 @@ class MCPServerManager: proxy_logging_obj: ProxyLogging | None = None, *, catalog_auth_header: str | dict[str, str] | None | EllipsisType = ..., + record_listing: bool = False, ) -> Sequence[MCPTool]: """ Helper method to get tools from a single MCP server with prefixed names. @@ -4491,6 +4493,8 @@ class MCPServerManager: mcp_auth_header: Optional auth header for MCP server catalog_auth_header: The header the client supplied, keying the caller's catalog slot; defaults to ``mcp_auth_header`` + record_listing: Record the served catalog into the caller's listed-tools slot; only a + listing actually served to the caller sets it Returns: List[MCPTool]: List of tools available on the server with prefixed names @@ -4608,7 +4612,8 @@ class MCPServerManager: # through _create_prefixed_tools — that would add the prefix a second # time producing "test_petstore-test_petstore-getinventory". unprefixed_tools: Final = guarded_openapi - self._record_listed_tools(server, unprefixed_tools, listed_caller, listed_generation) + if record_listing: + self._record_listed_tools(server, unprefixed_tools, listed_caller, listed_generation) if not add_prefix: return unprefixed_tools return [t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi] @@ -4624,8 +4629,10 @@ class MCPServerManager: raw_headers=raw_headers, ) prefixed_or_original_tools: Final = self._create_prefixed_tools( - guarded_tools, server, add_prefix=add_prefix, caller=listed_caller, generation=listed_generation + guarded_tools, server, add_prefix=add_prefix ) + if record_listing: + self._record_listed_tools(server, guarded_tools, listed_caller, listed_generation) return prefixed_or_original_tools @@ -5655,8 +5662,6 @@ class MCPServerManager: tools: Sequence[MCPTool], server: MCPServer, add_prefix: bool = True, - caller: ListedToolsCaller | None = None, - generation: int | None = None, ) -> list[MCPTool]: """ Create prefixed tools and update tool mapping. @@ -5682,7 +5687,6 @@ class MCPServerManager: for spelling in iter_known_tool_name_spellings(original_name, server): self.tool_name_to_mcp_server_name_mapping[spelling] = prefix - self._record_listed_tools(server, tools, caller, generation) verbose_logger.info("Successfully fetched %s tools from server %s", len(prefixed_tools), server.name) return prefixed_tools diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 151e3a87939..d8886e26da7 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -957,6 +957,8 @@ async def _get_tools_from_mcp_servers( request_tags: list[str] | None = None, client_ip: str | None = None, mcp_proxy_mode: bool = False, + *, + record_listing: bool = True, ) -> AggregateToolListing: """ Helper method to fetch tools from MCP servers based on server filtering criteria. @@ -967,6 +969,8 @@ async def _get_tools_from_mcp_servers( mcp_servers: Optional list of server names/aliases to filter by mcp_server_auth_headers: Optional dict of server-specific auth headers oauth2_headers: Optional dict of oauth2 headers + record_listing: Record each served catalog into the caller's listed-tools slot; only a + listing actually served to the caller sets it Returns: AggregateToolListing: Combined tools from filtered servers plus each server's @@ -1132,6 +1136,7 @@ async def _get_tools_from_mcp_servers( oauth2_headers=oauth2_headers, proxy_logging_obj=proxy_logging_obj, catalog_auth_header=catalog_auth_header, + record_listing=record_listing, ) filtered_tools = filter_tools_by_allowed_tools(tools, server) @@ -1805,6 +1810,7 @@ async def _list_tools_before_first_call( oauth2_headers=oauth2_headers, raw_headers=raw_headers, client_ip=client_ip, + record_listing=False, ) except Exception as e: # noqa: BLE001 # best effort: resolution below answers as it did before verbose_logger.debug("MCP tools/call: listing %s before its first call failed: %s", server.name, e) diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 12cdab59e0f..b9eb17b1758 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -703,6 +703,8 @@ if MCP_AVAILABLE: extra_headers: dict[str, str] | None, client_ip: str | None, proxy_logging_obj: "ProxyLogging | None", + *, + record_listing: bool, ) -> list[MCPTool]: return await global_mcp_server_manager._get_tools_from_server( server=server, @@ -713,6 +715,7 @@ if MCP_AVAILABLE: client_ip=client_ip, user_api_key_auth=user_api_key_auth, proxy_logging_obj=proxy_logging_obj, + record_listing=record_listing, ) async def _get_tools_for_single_server( @@ -734,7 +737,14 @@ if MCP_AVAILABLE: from litellm.proxy.proxy_server import proxy_logging_obj tools = await _list_server_tools( - server, server_auth_header, raw_headers, user_api_key_auth, extra_headers, client_ip, proxy_logging_obj + server, + server_auth_header, + raw_headers, + user_api_key_auth, + extra_headers, + client_ip, + proxy_logging_obj, + record_listing=True, ) if not apply_tool_filters: @@ -776,6 +786,7 @@ if MCP_AVAILABLE: await _get_user_oauth_extra_headers(server, user_api_key_dict), IPAddressUtils.get_mcp_client_ip(request), None, + record_listing=False, ) scan: Final = await scan_tool_descriptions( apply_description_overrides(upstream, server), server, proxy_logging_obj, user_api_key_dict, raw_headers 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 00ea7c0c78c..441cc57c292 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py @@ -964,6 +964,7 @@ async def test_get_tools_from_mcp_servers(): oauth2_headers=None, proxy_logging_obj=None, catalog_auth_header=None, + record_listing=True, ): if server.server_id == "server1_id": return [mock_tool_1] @@ -1999,6 +2000,7 @@ async def test_get_tools_for_single_server(): client_ip=None, user_api_key_auth=None, proxy_logging_obj=ANY, + record_listing=True, ) # Verify the result diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 2b4b025f97b..b0b40982ade 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -7050,7 +7050,8 @@ 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, caller=caller) + manager._create_prefixed_tools(listed_tools, server) + manager._record_listed_tools(server, listed_tools, caller) mock_client = AsyncMock() mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) @@ -7134,8 +7135,8 @@ class TestMCPServerManager: def test_get_listed_tool_resolves_the_bare_name_from_the_latest_listing(self): manager = MCPServerManager() server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") - manager._create_prefixed_tools([MCPTool(name="echo", description="v1", inputSchema={})], server) - manager._create_prefixed_tools([MCPTool(name="echo", description="v2", inputSchema={})], server) + manager._record_listed_tools(server, [MCPTool(name="echo", description="v1", inputSchema={})], None) + manager._record_listed_tools(server, [MCPTool(name="echo", description="v2", inputSchema={})], None) latest = manager.get_listed_tool(server, "echo") assert latest is not None and latest.description == "v2" @@ -7147,13 +7148,13 @@ class TestMCPServerManager: manager = MCPServerManager() server = MCPServer(server_id="srv-id", name="srv", alias="srv", transport=MCPTransport.http, url="http://srv") caller = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-user", user_id="alice")) - manager._create_prefixed_tools( + manager._record_listed_tools( + server, [ MCPTool(name="foo", description="Fetches foo records", inputSchema={"type": "object"}), MCPTool(name="bar", description="Fetches bar records", inputSchema={"type": "object"}), ], - server, - caller=caller, + caller, ) assert manager.get_listed_tool(server, "srv-foo", caller) is None @@ -7174,7 +7175,7 @@ class TestMCPServerManager: url="http://srv", tool_name_to_description={"echo": "Admin wording"}, ) - await manager._get_tools_from_server(server, add_prefix=True) + await manager._get_tools_from_server(server, add_prefix=True, record_listing=True) overridden = manager.get_listed_tool(server, "echo") assert overridden is not None @@ -7194,7 +7195,9 @@ class TestMCPServerManager: transport=MCPTransport.http, tool_name_to_description={"read_note": "Read a SECRET note"}, ) - served = await manager._get_tools_from_server(server, add_prefix=True, proxy_logging_obj=proxy_logging_obj) + served = await manager._get_tools_from_server( + server, add_prefix=True, proxy_logging_obj=proxy_logging_obj, record_listing=True + ) assert [tool.description for tool in served] == ["Read a [MASKED] note"] listed = manager.get_listed_tool(server, "read_note") @@ -7206,8 +7209,8 @@ class TestMCPServerManager: manager = MCPServerManager() server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") other = MCPServer(server_id="other", name="other", transport=MCPTransport.http, url="http://other") - manager._create_prefixed_tools([MCPTool(name="echo", description="old", inputSchema={})], server) - manager._create_prefixed_tools([MCPTool(name="ping", description="kept", inputSchema={})], other) + manager._record_listed_tools(server, [MCPTool(name="echo", description="old", inputSchema={})], None) + manager._record_listed_tools(other, [MCPTool(name="ping", description="kept", inputSchema={})], None) manager._invalidate_server_definition_caches(server.server_id) @@ -7236,7 +7239,7 @@ class TestMCPServerManager: caller = ListedToolsCaller(user_api_key_auth=user) async def list_tools() -> None: - await manager._get_tools_from_server(server=server, user_api_key_auth=user) + await manager._get_tools_from_server(server=server, user_api_key_auth=user, record_listing=True) listing = asyncio.create_task(list_tools()) await fetch_started.wait() @@ -7249,7 +7252,7 @@ class TestMCPServerManager: manager._fetch_tools_with_timeout = AsyncMock( return_value=[MCPTool(name="turn", description="after save", inputSchema={})] ) - await manager._get_tools_from_server(server=server, user_api_key_auth=user) + await manager._get_tools_from_server(server=server, user_api_key_auth=user, record_listing=True) listed = manager.get_listed_tool(server, "turn", caller) assert listed is not None and listed.description == "after save" @@ -7348,7 +7351,7 @@ class TestMCPServerManager: async def test_user_oauth_refresh_keeps_listed_tools(self): manager = MCPServerManager() server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") - manager._create_prefixed_tools([MCPTool(name="echo", description="shared", inputSchema={})], server) + manager._record_listed_tools(server, [MCPTool(name="echo", description="shared", inputSchema={})], None) await manager.invalidate_user_oauth_token_cache("alice", server.server_id) @@ -7368,15 +7371,15 @@ class TestMCPServerManager: bob = UserAPIKeyAuth(user_id="bob", token="hashed-bob") alice_schema = {"type": "object", "properties": {"path": {"type": "string"}}} bob_schema = {"type": "object", "properties": {"path": {"type": "string"}, "site": {"type": "string"}}} - manager._create_prefixed_tools( + manager._record_listed_tools( + server, [MCPTool(name="read", description="alice view", inputSchema=alice_schema)], - server, - caller=ListedToolsCaller(user_api_key_auth=alice), + ListedToolsCaller(user_api_key_auth=alice), ) - manager._create_prefixed_tools( - [MCPTool(name="read", description="bob view", inputSchema=bob_schema)], + manager._record_listed_tools( server, - caller=ListedToolsCaller(user_api_key_auth=bob), + [MCPTool(name="read", description="bob view", inputSchema=bob_schema)], + ListedToolsCaller(user_api_key_auth=bob), ) alice_tool = manager.get_listed_tool(server, "read", ListedToolsCaller(user_api_key_auth=alice)) @@ -7390,10 +7393,10 @@ class TestMCPServerManager: assert manager.get_listed_tool(server, "read", carol) is None shared = MCPServer(server_id="shared", name="shared", transport=MCPTransport.http, url="http://shared") - manager._create_prefixed_tools( - [MCPTool(name="echo", description="everyone", inputSchema={})], + manager._record_listed_tools( shared, - caller=ListedToolsCaller(user_api_key_auth=alice), + [MCPTool(name="echo", description="everyone", inputSchema={})], + ListedToolsCaller(user_api_key_auth=alice), ) for_bob = manager.get_listed_tool(shared, "echo", ListedToolsCaller(user_api_key_auth=bob)) assert for_bob is None, "keyed callers get their own slot even on servers without upstream per-user auth" @@ -7446,12 +7449,8 @@ class TestMCPServerManager: server = MCPServer( **{"server_id": "srv", "name": "srv", "transport": MCPTransport.http, "url": "http://srv", **server_kwargs} ) - manager._create_prefixed_tools( - [MCPTool(name="turn", description="Catalog A", inputSchema={})], server, caller=caller_a - ) - manager._create_prefixed_tools( - [MCPTool(name="turn", description="Catalog B", inputSchema={})], server, caller=caller_b - ) + manager._record_listed_tools(server, [MCPTool(name="turn", description="Catalog A", inputSchema={})], caller_a) + manager._record_listed_tools(server, [MCPTool(name="turn", description="Catalog B", inputSchema={})], caller_b) for_a = manager.get_listed_tool(server, "turn", caller_a) for_b = manager.get_listed_tool(server, "turn", caller_b) @@ -7462,10 +7461,10 @@ class TestMCPServerManager: def test_shared_server_ignores_headers_it_never_forwards(self): manager = MCPServerManager() server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") - manager._create_prefixed_tools( - [MCPTool(name="turn", description="everyone", inputSchema={})], + manager._record_listed_tools( server, - caller=ListedToolsCaller(raw_headers={"authorization": "Bearer sk-litellm", "x-workspace": "A"}), + [MCPTool(name="turn", description="everyone", inputSchema={})], + ListedToolsCaller(raw_headers={"authorization": "Bearer sk-litellm", "x-workspace": "A"}), ) other = ListedToolsCaller(raw_headers={"authorization": "Bearer sk-other", "x-workspace": "B"}) @@ -7499,7 +7498,7 @@ class TestMCPServerManager: AsyncMock(side_effect=RuntimeError("DB DOWN")), ), ): - await manager._get_tools_from_server(server=server, user_api_key_auth=user) + await manager._get_tools_from_server(server=server, user_api_key_auth=user, record_listing=True) listed = manager.get_listed_tool(server, "turn", listed_tools_caller_for(server, user, None, None, None, None)) assert listed is not None and listed.description == "listed while db down" @@ -7550,7 +7549,9 @@ class TestMCPServerManager: proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) cache_byok_credential("byok-user", "byok-catalog", "stored-secret") try: - await manager._get_tools_from_server(server=server, mcp_auth_header=list_header, user_api_key_auth=user) + await manager._get_tools_from_server( + server=server, mcp_auth_header=list_header, user_api_key_auth=user, record_listing=True + ) listed = manager.get_listed_tool( server, "turn", listed_tools_caller_for(server, user, list_header, None, None, None) ) @@ -7590,6 +7591,7 @@ class TestMCPServerManager: server=server, mcp_auth_header="Bearer hdr", user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm"), + record_listing=True, ) caller: Final = ListedToolsCaller( @@ -7659,7 +7661,7 @@ class TestMCPServerManager: signer_headers, ), ): - await manager._get_tools_from_server(server=server, user_api_key_auth=alice) + await manager._get_tools_from_server(server=server, user_api_key_auth=alice, record_listing=True) finally: byok_credential_cache.delete_cache(byok_credential_cache_key("alice", "cc1")) @@ -7692,8 +7694,8 @@ class TestMCPServerManager: "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", return_value=signer, ): - manager._create_prefixed_tools( - [MCPTool(name="turn", description="alice view", inputSchema={})], server, caller=alice + manager._record_listed_tools( + server, [MCPTool(name="turn", description="alice view", inputSchema={})], alice ) assert manager.get_listed_tool(server, "turn", bob) is None @@ -7712,9 +7714,7 @@ class TestMCPServerManager: "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", return_value=MagicMock(), ): - manager._create_prefixed_tools( - [MCPTool(name="turn", description="slot a", inputSchema={})], server, caller=alice - ) + manager._record_listed_tools(server, [MCPTool(name="turn", description="slot a", inputSchema={})], alice) assert manager.get_listed_tool(server, "turn", bob) is None same_key = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="same-user", api_key="sk-alpha")) @@ -7756,6 +7756,7 @@ class TestMCPServerManager: extra_headers={"X-Workspace": workspace}, raw_headers={"x-workspace": workspace, "authorization": "Bearer sk-litellm"}, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm", user_id="shared-key"), + record_listing=True, ) proxy_logging_obj = MagicMock() @@ -7789,20 +7790,18 @@ class TestMCPServerManager: url="http://srv", auth_type=MCPAuth.oauth2_token_exchange, ) - manager._create_prefixed_tools([MCPTool(name="read", description="shared", inputSchema={})], server) + manager._record_listed_tools(server, [MCPTool(name="read", description="shared", inputSchema={})], None) callers = [ ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id=f"u{i}", api_key=f"k{i}")) for i in range(_LISTED_TOOLS_CALLERS_PER_SERVER + 1) ] for caller in callers: - manager._create_prefixed_tools( - [MCPTool(name="read", description=caller.user_api_key_auth.user_id, inputSchema={})], + manager._record_listed_tools( server, - caller=caller, + [MCPTool(name="read", description=caller.user_api_key_auth.user_id, inputSchema={})], + caller, ) - manager._create_prefixed_tools( - [MCPTool(name="read", description="u1 again", inputSchema={})], server, caller=callers[1] - ) + manager._record_listed_tools(server, [MCPTool(name="read", description="u1 again", inputSchema={})], callers[1]) assert manager.get_listed_tool(server, "read", callers[0]) is None second = manager.get_listed_tool(server, "read", callers[1]) @@ -7840,7 +7839,7 @@ class TestMCPServerManager: handler=_handler, ) try: - listed = await manager._get_tools_from_server(server=server, add_prefix=add_prefix) + listed = await manager._get_tools_from_server(server=server, add_prefix=add_prefix, record_listing=True) finally: global_mcp_tool_registry.unregister_tools_with_prefix("petstore-") @@ -7882,7 +7881,7 @@ class TestMCPServerManager: handler=_handler, ) try: - listed = await manager._get_tools_from_server(server=server, add_prefix=True) + listed = await manager._get_tools_from_server(server=server, add_prefix=True, record_listing=True) finally: for prefix in ("pet-", "petstore-"): global_mcp_tool_registry.unregister_tools_with_prefix(prefix) @@ -7892,6 +7891,77 @@ class TestMCPServerManager: assert tool is not None and tool.description == "Local pet tool" assert tool.input_schema["properties"] == {"limit": {"type": "integer"}} + @pytest.mark.asyncio + @pytest.mark.parametrize("openapi", [False, True], ids=["remote", "openapi"]) + async def test_get_tools_from_server_records_the_catalog_only_when_asked_to(self, openapi): + """The startup fill, the implicit pre-call listing and the pin snapshot reuse this fetch without + serving its result, so only a listing that asks to be recorded sets what tools/call hooks see.""" + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + if openapi: + server = MCPServer( + server_id="srv", name="srv", alias="srv", transport=MCPTransport.http, url=None, spec_path="/spec.yaml" + ) + manager = MCPServerManager() + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + global_mcp_tool_registry.unregister_tools_with_prefix("srv-") + global_mcp_tool_registry.register_tool( + name="srv-echo", description="Echoes", input_schema={"type": "object"}, handler=lambda **kwargs: None + ) + else: + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager = _catalog_manager(MCPTool(name="echo", description="Echoes", inputSchema={"type": "object"})) + user = UserAPIKeyAuth(api_key="sk-litellm", user_id="lister") + + try: + listed = await manager._get_tools_from_server(server=server, user_api_key_auth=user) + assert [t.name for t in listed] == ["srv-echo"] + assert server.server_id not in manager._listed_tools_by_server_id + + await manager._get_tools_from_server(server=server, user_api_key_auth=user, record_listing=True) + finally: + global_mcp_tool_registry.unregister_tools_with_prefix("srv-") + + recorded = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=user)) + assert recorded is not None and recorded.description == "Echoes" + + @pytest.mark.asyncio + async def test_list_tools_records_the_served_catalog(self): + manager = _catalog_manager(MCPTool(name="echo", description="Echoes", inputSchema={"type": "object"})) + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager.registry = {"srv": server} + manager.get_allowed_mcp_servers = AsyncMock(return_value=["srv"]) + user = UserAPIKeyAuth(api_key="sk-litellm", user_id="lister") + + listed = await manager.list_tools(user_api_key_auth=user) + + assert [t.name for t in listed] == ["srv-echo"] + recorded = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=user)) + assert recorded is not None and recorded.description == "Echoes" + + @pytest.mark.asyncio + async def test_startup_tool_name_mapping_records_no_listed_catalog(self): + manager = _catalog_manager(MCPTool(name="echo", description="Echoes", inputSchema={"type": "object"})) + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager.registry = {"srv": server} + + await manager._initialize_tool_name_to_mcp_server_name_mapping() + + assert manager.server_exposes_tool(server, "echo") is True + assert server.server_id not in manager._listed_tools_by_server_id + assert manager.get_listed_tool(server, "echo") is None + + @pytest.mark.asyncio + async def test_get_tools_for_server_records_no_listed_catalog(self): + manager = _catalog_manager(MCPTool(name="echo", description="Echoes", inputSchema={"type": "object"})) + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager.registry = {"srv": server} + + listed = await manager.get_tools_for_server("srv") + + assert [t.name for t in listed] == ["srv-echo"] + assert server.server_id not in manager._listed_tools_by_server_id + @pytest.mark.asyncio async def test_get_allowed_mcp_servers_with_user_api_key_auth(self): """ 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 5aec9a9186f..545b2757ffd 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 @@ -37,7 +37,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.types.mcp import MCPAuth -from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer +from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer, PinnedMCPTool def test_mcp_available_on_sdk2(): @@ -8332,6 +8332,85 @@ async def test_execute_mcp_tool_hands_hooks_nothing_for_a_never_listed_operation assert result.content[0].text == "long" +@pytest.mark.asyncio +async def test_execute_mcp_tool_implicit_listing_before_the_first_call_hands_hooks_no_description(): + """The listing tools/call runs on its own when this worker does not yet expose the tool is never served + to the caller, so it leaves the caller's listed slot empty and the pre-call hooks still get name and + arguments only, as on main.""" + manager = mcp_operations.global_mcp_server_manager + server = _never_listed_passthrough_server() + manager.registry[server.server_id] = server + manager._listed_tools_by_server_id.pop(server.server_id, None) + upstream = AsyncMock() + upstream.call_tool.return_value = CallToolResult(content=[TextContent(type="text", text="ok")], isError=False) + proxy_logging = _mock_mcp_proxy_logging() + proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging.pre_call_hook = AsyncMock(return_value={}) + proxy_logging.during_call_hook = AsyncMock(return_value=None) + fetch_tools = AsyncMock( + return_value=[MCPTool(name="add", description="Adds. FLAGWORD", inputSchema={"type": "object"})] + ) + + with ( + patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=upstream)), + patch.object(manager, "_fetch_tools_with_timeout", new=fetch_tools), + patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging), + ): + result = await mcp_operations.execute_mcp_tool( + name="lazy_map-add", + arguments={"a": 1, "b": 2}, + allowed_mcp_servers=[server], + start_time=datetime.now(), + mcp_auth_header="Bearer caller-token", + raw_headers={"authorization": "Bearer caller-token"}, + ) + + assert fetch_tools.await_count == 1 + assert upstream.call_tool.await_count == 1 + assert result.content[0].text == "ok" + hook_kwargs = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0] + assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == (None, None) + assert server.server_id not in manager._listed_tools_by_server_id + + +@pytest.mark.asyncio +async def test_fetch_pinnable_tool_catalog_records_no_listed_catalog_for_the_admin(): + """The pin snapshot lists the raw upstream catalog, without the catalog guard or the admin's description + overrides, so it must not become what the admin's own later tools/call is evaluated against.""" + from litellm.caching.caching import DualCache + from litellm.proxy._experimental.mcp_server.rest_endpoints import fetch_pinnable_tool_catalog + from litellm.proxy.utils import ProxyLogging + + manager = mcp_operations.global_mcp_server_manager + server = MCPServer( + server_id="pin-srv", + name="pin_srv", + transport=MCPTransport.http, + url="https://up.example.com/mcp", + tool_name_to_description={"add": "Admin wording"}, + ) + manager._listed_tools_by_server_id.pop(server.server_id, None) + admin = UserAPIKeyAuth(api_key="sk-admin", user_id="admin") + request = MagicMock() + request.client.host = "10.1.2.3" + request.headers = {"x-litellm-api-key": "sk-admin"} + fetch_tools = AsyncMock( + return_value=[MCPTool(name="add", description="Upstream wording", inputSchema={"type": "object"})] + ) + + with ( + patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=MagicMock())), + patch.object(manager, "_fetch_tools_with_timeout", new=fetch_tools), + patch("litellm.proxy.proxy_server.proxy_logging_obj", ProxyLogging(user_api_key_cache=DualCache())), + ): + snapshot = await fetch_pinnable_tool_catalog(server, request, admin) + + assert snapshot == {"add": PinnedMCPTool(description="Upstream wording", input_schema={"type": "object"})} + assert server.server_id not in manager._listed_tools_by_server_id + assert manager.get_listed_tool(server, "add", ListedToolsCaller(user_api_key_auth=admin)) is None + + @pytest.mark.asyncio async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requested_server(): """A prefixed REST name that resolves to no tool must still dispatch to the server_id.