diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 02d7255700d..a247ce64bf5 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -5647,11 +5647,7 @@ class MCPServerManager: listed: Final = self._listed_tools_by_server_id.get(server.server_id, MappingProxyType({})).get(identity) if not listed: return None - tool: Final = listed.get(name) or listed.get(strip_known_server_prefix(name, server)) - if tool is None: - return None - description: Final = server.tool_name_to_description.get(tool.name) if server.tool_name_to_description else None - return tool if description is None else tool.model_copy(update={"description": description}) + return listed.get(name) or listed.get(strip_known_server_prefix(name, server)) def _create_prefixed_prompts( self, prompts: Sequence[Prompt], server: MCPServer, add_prefix: bool = True diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index cccdf4c59c0..5a63fedce73 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -1608,6 +1608,11 @@ async def _list_mcp_resource_templates( def _registered_tool_metadata(name: str, registered: RegisteredTool, server: MCPServer) -> MCPTool: + """The tool as ``tools/list`` served it (pinned, overridden, guardrail-masked) when a listing was + recorded for ``server``, else the registry entry with the admin description override applied.""" + listed: Final = global_mcp_server_manager.get_listed_tool(server, name) + if listed is not None: + return listed overrides: Final = server.tool_name_to_description description: Final = overrides.get(name, registered.description) if overrides else registered.description return MCPTool(name=name, description=description, input_schema=registered.input_schema) diff --git a/tests/integration/_support/mcp.py b/tests/integration/_support/mcp.py index a3693433de4..d39d7e4029e 100644 --- a/tests/integration/_support/mcp.py +++ b/tests/integration/_support/mcp.py @@ -216,6 +216,16 @@ JsonRpc = Mapping[str, object] class ScriptedTool: name: str respond: Callable[[JsonRpc], Reply | JsonRpc] + description: str | Callable[[Mapping[str, str]], str] | None = None + input_schema: JsonRpc = field(default_factory=lambda: {"type": "object"}) + + def listing(self, headers: Mapping[str, str]) -> JsonRpc: + described: Final = self.description(headers) if callable(self.description) else self.description + return { + "name": self.name, + "inputSchema": self.input_schema, + **({} if described is None else {"description": described}), + } def jsonrpc_reply(identity: object, result: JsonRpc) -> Reply: @@ -253,9 +263,7 @@ def scripted_peer(*tools: ScriptedTool) -> Iterator[McpPeer]: }, ) if method == "tools/list": - return jsonrpc_reply( - identity, {"tools": [{"name": name, "inputSchema": {"type": "object"}} for name in by_name]} - ) + return jsonrpc_reply(identity, {"tools": [tool.listing(request.headers) for tool in by_name.values()]}) if method != "tools/call": return jsonrpc_error(identity, -32601, f"unsupported method {method}") tool: Final = by_name.get(body["params"]["name"]) diff --git a/tests/integration/mcp/test_mcp_listed_tool_metadata.py b/tests/integration/mcp/test_mcp_listed_tool_metadata.py new file mode 100644 index 00000000000..0381e8abb00 --- /dev/null +++ b/tests/integration/mcp/test_mcp_listed_tool_metadata.py @@ -0,0 +1,169 @@ +"""pre_mcp_call guardrails are handed the tool entry ``tools/list`` served to the caller. + +One owned proxy carries a default-on ``custom_code`` pre_mcp_call guardrail. At listing time it masks +``SECRET`` out of every scanned text. At call time, when an argument carries the probe marker, it +blocks and echoes the description and parameters it was handed, which is the only way to observe from outside +what metadata the gateway attached to the hook +""" + +import json +import uuid +from collections.abc import Iterator, Mapping +from pathlib import Path +from typing import Final + +import pytest +import yaml +from integration._support.client import Gateway, gateway_from_environment +from integration._support.mcp import ( + EntryPoint, + McpCaller, + ScriptedTool, + listed_tools, + openapi_peer, + register_mcp, + scripted_peer, + text_result, +) +from integration._support.process import owned_proxy + +_ECHO: Final = "catalog-echo:" +_PROBE: Final = "catalog-probe" +_GUARDRAIL_CODE: Final = ( + "def apply_guardrail(inputs, request_data, input_type):\n" + ' texts = list(inputs.get("texts") or [])\n' + ' function = inputs.get("tools", [{}])[0].get("function", {})\n' + f' if "{_PROBE}" in texts:\n' + f' return block("{_ECHO}" + json_stringify(' + '{"description": function.get("description"), "parameters": function.get("parameters")}))\n' + ' masked = [text.replace("SECRET", "[MASKED]") for text in texts]\n' + " if masked != texts:\n" + " return modify(texts=masked)\n" + " return allow()\n" +) + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]: + directory: Final = tmp_path_factory.mktemp("listed-tool-metadata") + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": "catalog-echo-" + uuid.uuid4().hex[:8], + "litellm_params": { + "guardrail": "custom_code", + "mode": "pre_mcp_call", + "default_on": True, + "custom_code": _GUARDRAIL_CODE, + }, + } + ] + path: Final = directory / "config.yaml" + path.write_text(yaml.safe_dump(config)) + with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path) as candidate: + yield candidate + + +def _strings(value: object) -> Iterator[str]: + if isinstance(value, str): + yield value + return + children: Final = value.values() if isinstance(value, Mapping) else value if isinstance(value, list) else () + for child in children: + yield from _strings(child) + + +def _decoded(raw: str) -> object: + data: Final = tuple(line[5:].strip() for line in raw.splitlines() if line.startswith("data:")) + return json.loads(data[-1] if data else raw) + + +def _echoed(raw: str) -> tuple[str | None, Mapping[str, object] | None]: + """The (description, parameters) the guardrail was handed, recovered from its block reason.""" + carrier: Final = next((text for text in _strings(_decoded(raw)) if _ECHO in text), None) + assert carrier is not None, raw + echoed, _ = json.JSONDecoder().raw_decode(carrier.split(_ECHO, 1)[1]) + assert isinstance(echoed, dict), carrier + return echoed.get("description"), echoed.get("parameters") + + +def _probe(caller: McpCaller, name: str, server_id: str) -> tuple[str | None, Mapping[str, object] | None]: + outcome: Final = caller.call(name, {"probe": _PROBE}, server_id=server_id) + assert outcome.error is not None, outcome.raw + return _echoed(outcome.raw) + + +@pytest.mark.parametrize("entry", ["rest", "mcp"]) +def test_pre_call_hook_receives_the_description_and_input_schema_the_caller_was_listed( + rig: Gateway, entry: EntryPoint +) -> None: + schema: Final = {"type": "object", "properties": {"probe": {"type": "string", "description": "a probe marker"}}} + tool: Final = ScriptedTool( + "lookup", lambda _: text_result("found"), description="Look up one record", input_schema=schema + ) + with scripted_peer(tool) as peer, rig.scenario() as scenario: + alias: Final = "meta" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + caller: Final = McpCaller(rig, key, entry, headers={"x-mcp-servers": alias}) + assert caller.initialize().ok + listed: Final = caller.list_tools(server_id=identity) + assert listed.ok, listed.raw + name: Final = next(full for full in listed.tools if full.endswith("lookup")) + description, parameters = _probe(caller, name, identity) + assert description == "Look up one record", (description, parameters) + assert parameters is not None and parameters.get("properties") == schema["properties"], parameters + + +def test_each_caller_is_evaluated_against_the_catalog_its_own_forwarded_headers_produced(rig: Gateway) -> None: + tool: Final = ScriptedTool( + "report", + lambda _: text_result("ok"), + description=lambda headers: f"Report for tenant {headers.get('x-tenant', 'nobody')}", + ) + with scripted_peer(tool) as peer, rig.scenario() as scenario: + alias: Final = "tenant" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias, extra_headers=["x-tenant"]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + acme: Final = McpCaller(rig, key, "server_mcp", alias, headers={"x-tenant": "acme"}) + globex: Final = McpCaller(rig, key, "server_mcp", alias, headers={"x-tenant": "globex"}) + acme_listing: Final = acme.list_tools() + globex_listing: Final = globex.list_tools() + assert acme_listing.ok and globex_listing.ok, (acme_listing.raw, globex_listing.raw) + name: Final = next(full for full in acme_listing.tools if full.endswith("report")) + acme_seen, _ = _probe(acme, name, identity) + globex_seen, _ = _probe(globex, name, identity) + assert (acme_seen, globex_seen) == ("Report for tenant acme", "Report for tenant globex"), ( + "each caller's tools/call must be evaluated against the catalog its own headers listed" + ) + + +def test_call_is_evaluated_against_the_masked_description_the_listing_served(rig: Gateway) -> None: + tool: Final = ScriptedTool("read_note", lambda _: text_result("note"), description="Read a note") + with scripted_peer(tool) as peer, rig.scenario() as scenario: + alias: Final = "note" + uuid.uuid4().hex[:8] + identity: Final = register_mcp( + scenario, peer, alias, tool_name_to_description={"read_note": "Read a SECRET note"} + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + served: Final = listed_tools(rig, key, identity) + name: Final = next(full for full in served if full.endswith("read_note")) + assert served[name]["description"] == "Read a [MASKED] note", served[name] + seen, _ = _probe(McpCaller(rig, key, "rest"), name, identity) + assert seen == "Read a [MASKED] note", "the admin override must not restore wording the listing masked" + + +def test_openapi_call_is_evaluated_against_the_masked_override_the_listing_served(rig: Gateway) -> None: + with openapi_peer() as peer, rig.scenario() as scenario: + alias: Final = "pets" + uuid.uuid4().hex[:8] + identity: Final = register_mcp( + scenario, peer, alias, tool_name_to_description={"getpet": "Fetch one SECRET pet"} + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + served: Final = listed_tools(rig, key, identity) + name: Final = next(full for full in served if full.endswith("getpet")) + assert served[name]["description"] == "Fetch one [MASKED] pet", served[name] + seen, parameters = _probe(McpCaller(rig, key, "rest"), name, identity) + assert seen == "Fetch one [MASKED] pet", "the OpenAPI call path must hand hooks the entry the listing served" + assert parameters is not None and "petId" in parameters.get("properties", {}), parameters + assert not [call for call in peer.drain() if call["path"].startswith("/pets")], "blocked before upstream" 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 62a7830a85d..57268902614 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 @@ -24,6 +24,7 @@ from mcp.types import ( TextContent, TextResourceContents, ) +from mcp.types import Tool as MCPTool from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION, MODERN_PROTOCOL_VERSIONS from pydantic import TypeAdapter from starlette.types import Message, Receive, Scope, Send @@ -86,9 +87,6 @@ def cleanup_mcp_global_state(): yield - - - def _call_tool_params(name, arguments=None): from mcp.types import CallToolRequestParams @@ -100,6 +98,7 @@ def _paged_params(): return PaginatedRequestParams() + @pytest.mark.asyncio async def test_mcp_server_tool_call_body_contains_request_data(_mcp_request_ctx): """Test that proxy_server_request body contains name and arguments""" @@ -296,7 +295,9 @@ async def test_mcp_server_tool_call_relays_upstream_auth_error_as_iserror(_mcp_r ): with patch("litellm.proxy.proxy_server.proxy_config", MagicMock()): with patch("litellm.proxy._experimental.mcp_server.operations.verbose_logger", mock_logger): - result = await mcp_server_tool_call(_mcp_request_ctx(), _call_tool_params("test_tool", {"param": "value"})) + result = await mcp_server_tool_call( + _mcp_request_ctx(), _call_tool_params("test_tool", {"param": "value"}) + ) assert result.is_error is True # The dedicated MCPUpstreamAuthError branch (not the generic Exception fallthrough) produces this @@ -1168,20 +1169,32 @@ async def test_read_resource_preserves_content_metadata(_mcp_request_ctx, kind, else BlobResourceContents(uri=uri, blob="aGVsbG8=", mimeType="image/png", meta=metadata) ) with ( - patch.object(server, "get_or_extract_auth_context", AsyncMock(return_value=(caller, None, ["catalog"], None, None, None, None))), + patch.object( + server, + "get_or_extract_auth_context", + AsyncMock(return_value=(caller, None, ["catalog"], None, None, None, None)), + ), patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream_server])), - patch.object(operations.global_mcp_server_manager, "read_resource_from_server", AsyncMock(return_value=ReadResourceResult(contents=[content]))), + patch.object( + operations.global_mcp_server_manager, + "read_resource_from_server", + AsyncMock(return_value=ReadResourceResult(contents=[content])), + ), ): result: Final = await server.read_resource(_mcp_request_ctx(), ReadResourceRequestParams(uri=uri)) assert result.model_dump(mode="json", by_alias=True, exclude_none=True) == { - "cacheScope": "private", "resultType": "complete", "ttlMs": 0, - "contents": [{ - "uri": uri, - "mimeType": "text/plain" if kind == "text" else "image/png", - "text" if kind == "text" else "blob": "hello world" if kind == "text" else "aGVsbG8=", - **({"_meta": metadata} if metadata is not None else {}), - }], + "cacheScope": "private", + "resultType": "complete", + "ttlMs": 0, + "contents": [ + { + "uri": uri, + "mimeType": "text/plain" if kind == "text" else "image/png", + "text" if kind == "text" else "blob": "hello world" if kind == "text" else "aGVsbG8=", + **({"_meta": metadata} if metadata is not None else {}), + } + ], } @@ -1675,7 +1688,9 @@ async def test_handle_list_tools_converts_permission_httpexception_to_mcp_error( with ( patch( # test-quality-ok: the protocol handler reads auth from module context; no injection seam "litellm.proxy._experimental.mcp_server.server.get_or_extract_auth_context", - new=AsyncMock(return_value=(None, None, None, None, None, None, None), side_effect=denial if denial_at_auth else None), + new=AsyncMock( + return_value=(None, None, None, None, None, None, None), side_effect=denial if denial_at_auth else None + ), ), patch( # test-quality-ok: the listing helper is the handler's only collaborator; the suite's seam "litellm.proxy._experimental.mcp_server.operations._list_mcp_tools", @@ -1914,8 +1929,8 @@ async def test_streamable_http_session_manager_is_stateless(): ("DELETE", b"", False), ), ) -async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(_mcp_request_ctx, - debug: bool, method: str, request_body: bytes, stateful: bool +async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless( + _mcp_request_ctx, debug: bool, method: str, request_body: bytes, stateful: bool ) -> None: from starlette.requests import Request from starlette.types import Message, Receive, Scope, Send @@ -4057,7 +4072,8 @@ async def test_truncated_jsonrpc_response_with_nested_method_skips_lock( # parsed, with a nested "method" key in the first bytes to trip a flat # substring heuristic. response_prefix: Final = ( - '{"jsonrpc":"2.0","id":99,"' + response_field + '{"jsonrpc":"2.0","id":99,"' + + response_field + '":{"code":-32000,"message":"test","data":{"method":"GET","payload":"' ).encode() response_body: Final = ( @@ -6565,8 +6581,12 @@ class TestGatewayCreateInitializationOptions: yield (None, None) async def record_request( - serving_server: object, read_stream: object, write_stream: object, - *, lifespan_state: object, init_options: InitializationOptions, + serving_server: object, + read_stream: object, + write_stream: object, + *, + lifespan_state: object, + init_options: InitializationOptions, ) -> None: captured["server_name"] = init_options.server_name @@ -7412,7 +7432,8 @@ async def test_execute_mcp_tool_rest_server_id_authoritative_for_unprefixed_tool return_value=oauth_server, ), patch.object( - mcp_operations, "_handle_managed_mcp_tool", + mcp_operations, + "_handle_managed_mcp_tool", new=fake_handle_managed_mcp_tool, ), patch.object( @@ -7658,7 +7679,8 @@ async def test_execute_mcp_tool_strips_a_prefix_that_contains_the_separator(): return_value=alias_less_server, ), patch.object( - mcp_operations, "_handle_managed_mcp_tool", + mcp_operations, + "_handle_managed_mcp_tool", new=fake_handle_managed_mcp_tool, ), patch.object( @@ -7933,7 +7955,8 @@ async def test_execute_mcp_tool_rest_hyphenated_upstream_tool_name_routes_to_req return_value=None, ), patch.object( - mcp_operations, "_handle_managed_mcp_tool", + mcp_operations, + "_handle_managed_mcp_tool", new=fake_handle_managed_mcp_tool, ), patch.object( @@ -8132,6 +8155,57 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_admin_description_client assert (handed_tool.description, handed_tool.input_schema) == ("ADMIN DESC", schema) +@pytest.mark.asyncio +async def test_execute_mcp_tool_hands_openapi_hooks_the_guarded_catalog_entry_clients_saw(): + """When tools/list pinned the schema and masked the description of an OpenAPI tool, the local-registry + call path must hand the pre-call hooks that served entry, not the raw registry one.""" + from litellm.proxy._experimental.mcp_server import operations as mcp_module + + petstore = MCPServer( + server_id="petstore-id", + name="petstore", + server_name="petstore", + transport=MCPTransport.http, + url=None, + spec_path="https://example.com/petstore.yaml", + tool_name_to_description={"getpetbyid": "Find a SECRET pet"}, + ) + registry_schema = {"type": "object", "properties": {"petId": {"type": "integer"}, "dump_all": {"type": "boolean"}}} + pinned_schema = {"type": "object", "properties": {"petId": {"type": "integer"}}} + mcp_module.global_mcp_tool_registry.register_tool( + name="petstore-getpetbyid", + description="Find pet by ID", + input_schema=registry_schema, + handler=lambda petId: "ok", + ) + manager = mcp_module.global_mcp_server_manager + manager._record_listed_tools( + petstore, [MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=pinned_schema)], None + ) + pre_call_tool_check = AsyncMock(return_value={}) + + try: + with ( + patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore), + patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check), + ): + await mcp_module.execute_mcp_tool( + name="petstore-getpetbyid", + arguments={"petId": 1}, + allowed_mcp_servers=[petstore], + start_time=datetime.now(), + user_api_key_auth=UserAPIKeyAuth(api_key="sk-user", user_id="alice"), + ) + finally: + mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-") + manager._listed_tools_by_server_id.pop(petstore.server_id, None) + + handed_tool = pre_call_tool_check.call_args.kwargs["tool"] + assert (handed_tool.description, handed_tool.input_schema) == ("Find a [MASKED] pet", pinned_schema), ( + "the pre-call policy must evaluate the entry tools/list served, not the raw registry entry" + ) + + @pytest.mark.asyncio async def test_execute_mcp_tool_hands_hooks_the_metadata_of_the_operation_it_runs_when_names_collide(): """An OpenAPI operation whose name starts with its own server prefix must not be reported to the @@ -8236,7 +8310,8 @@ async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requeste return_value=None, ), patch.object( - mcp_operations, "_handle_managed_mcp_tool", + mcp_operations, + "_handle_managed_mcp_tool", new=fake_handle_managed_mcp_tool, ), patch.object( @@ -8735,7 +8810,9 @@ class TestMCPMetaTraceCarrier: assert _mcp_meta_trace_carrier(None) is None assert _mcp_meta_trace_carrier(SimpleNamespace(meta=None)) is None - only_progress = CallToolRequestParams.model_validate({"name": "t", "_meta": {"progressToken": "p1"}}, by_name=False).meta + only_progress = CallToolRequestParams.model_validate( + {"name": "t", "_meta": {"progressToken": "p1"}}, by_name=False + ).meta assert _mcp_meta_trace_carrier(SimpleNamespace(meta=only_progress)) is None @@ -10532,7 +10609,9 @@ async def test_mcp_origin_admission_precedes_authentication( patch("litellm.proxy.proxy_server.origins", allowed_origins), patch.object(server, "extract_mcp_auth_context", authenticate), ): - async with httpx.AsyncClient(transport=httpx.ASGITransport(app=server.app), base_url="http://gateway") as client: + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=server.app), base_url="http://gateway" + ) as client: response: Final = await client.request(method, path, headers=(*session_headers, *origin_headers)) assert response.status_code == expected_status @@ -10615,12 +10694,15 @@ async def test_streamable_http_rejects_modern_protocol_version( @pytest.mark.asyncio -@pytest.mark.parametrize("handler_name,field", [ - ("handle_list_tools", "tools"), - ("list_prompts", "prompts"), - ("list_resources", "resources"), - ("list_resource_templates", "resource_templates"), -]) +@pytest.mark.parametrize( + "handler_name,field", + [ + ("handle_list_tools", "tools"), + ("list_prompts", "prompts"), + ("list_resources", "resources"), + ("list_resource_templates", "resource_templates"), + ], +) async def test_native_listing_preserves_empty_result_on_auth_failure(_mcp_request_ctx, handler_name, field): from litellm.proxy._experimental.mcp_server import server @@ -10638,7 +10720,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai auth = UserAPIKeyAuth(user_id="denied-caller") denial = HTTPException(status_code=403, detail="scope denied") logger = MagicMock() - logger.post_call_failure_hook = AsyncMock(side_effect=RuntimeError("log unavailable") if failure_hook_raises else None) + logger.post_call_failure_hook = AsyncMock( + side_effect=RuntimeError("log unavailable") if failure_hook_raises else None + ) upstream = AsyncMock() with ( patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(side_effect=denial)), @@ -10647,7 +10731,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai patch.object(operations.global_mcp_server_manager, "_get_tools_from_server", upstream), ): with pytest.raises(HTTPException) as rejected: - await operations._get_tools_from_mcp_servers(user_api_key_auth=auth, mcp_auth_header=None, mcp_servers=["catalog"], log_list_tools_to_spendlogs=True) + await operations._get_tools_from_mcp_servers( + user_api_key_auth=auth, mcp_auth_header=None, mcp_servers=["catalog"], log_list_tools_to_spendlogs=True + ) assert rejected.value is denial upstream.assert_not_awaited() logger.post_call_failure_hook.assert_awaited_once() @@ -10659,7 +10745,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai @pytest.mark.parametrize("prefix,suffix", (("", ""), ("/gateway", "/"))) @pytest.mark.parametrize("opening_protocol", (None, *MODERN_PROTOCOL_VERSIONS)) async def test_legacy_sse_mount_emits_message_endpoint( - prefix: str, suffix: str, opening_protocol: str | None, + prefix: str, + suffix: str, + opening_protocol: str | None, ) -> None: from starlette.applications import Starlette from starlette.routing import Mount @@ -10724,16 +10812,20 @@ async def test_legacy_sse_mount_emits_message_endpoint( return (await messages.get())["status"] if opening_protocol is not None: - discover: Final = json.dumps({ - "jsonrpc": "2.0", - "id": 0, - "method": "server/discover", - "params": {"_meta": { - "io.modelcontextprotocol/protocolVersion": opening_protocol, - "io.modelcontextprotocol/clientInfo": {"name": "modern-client", "version": "1"}, - "io.modelcontextprotocol/clientCapabilities": {}, - }}, - }).encode() + discover: Final = json.dumps( + { + "jsonrpc": "2.0", + "id": 0, + "method": "server/discover", + "params": { + "_meta": { + "io.modelcontextprotocol/protocolVersion": opening_protocol, + "io.modelcontextprotocol/clientInfo": {"name": "modern-client", "version": "1"}, + "io.modelcontextprotocol/clientCapabilities": {}, + } + }, + } + ).encode() assert await post(discover) == 202 discovered_frame: Final = (await asyncio.wait_for(outgoing.get(), 2))["body"].decode() discovered: Final = json.loads(discovered_frame.split("data: ", 1)[1].splitlines()[0]) @@ -10766,7 +10858,16 @@ async def test_legacy_sse_mount_emits_message_endpoint( patch.object( mcp_server, "extract_mcp_auth_context", - AsyncMock(return_value=(post_auth, None, [marker], {marker: {"Authorization": marker}}, {"Authorization": marker}, {"x-request-marker": marker})), + AsyncMock( + return_value=( + post_auth, + None, + [marker], + {marker: {"Authorization": marker}}, + {"Authorization": marker}, + {"x-request-marker": marker}, + ) + ), ), patch.object(mcp_server.operations, "_get_tools_from_mcp_servers", listing), ): @@ -10815,7 +10916,11 @@ async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ct dispatched = AsyncMock(return_value=expected) auth = UserAPIKeyAuth(user_id="discover-caller") with ( - patch.object(server, "get_or_extract_auth_context", AsyncMock(return_value=(auth, None, ["allowed"], None, None, None, None))), + patch.object( + server, + "get_or_extract_auth_context", + AsyncMock(return_value=(auth, None, ["allowed"], None, None, None, None)), + ), patch.object(server.operations.GatewayOperations, "execute", dispatched), ): result = await server.discover(_mcp_request_ctx(), RequestParams()) 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 4cfc9d36714..7000497769f 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 @@ -7137,8 +7137,13 @@ class TestMCPServerManager: assert by_prefixed_name is not None and by_prefixed_name.description == "v2" assert manager.get_listed_tool(server, "missing") is None - def test_get_listed_tool_uses_admin_description_override_clients_saw(self): - manager = MCPServerManager() + @pytest.mark.asyncio + async def test_get_listed_tool_uses_admin_description_override_clients_saw(self): + schema = {"type": "object", "properties": {"text": {"type": "string"}}} + manager = _catalog_manager( + MCPTool(name="echo", description="Upstream wording", inputSchema=schema), + MCPTool(name="ping", description="Untouched", inputSchema={}), + ) server = MCPServer( server_id="srv", name="srv", @@ -7146,14 +7151,7 @@ class TestMCPServerManager: url="http://srv", tool_name_to_description={"echo": "Admin wording"}, ) - schema = {"type": "object", "properties": {"text": {"type": "string"}}} - manager._create_prefixed_tools( - [ - MCPTool(name="echo", description="Upstream wording", inputSchema=schema), - MCPTool(name="ping", description="Untouched", inputSchema={}), - ], - server, - ) + await manager._get_tools_from_server(server, add_prefix=True) overridden = manager.get_listed_tool(server, "srv-echo") assert overridden is not None @@ -7161,6 +7159,26 @@ class TestMCPServerManager: untouched = manager.get_listed_tool(server, "ping") assert untouched is not None and untouched.description == "Untouched" + @pytest.mark.asyncio + async def test_get_listed_tool_keeps_the_masked_description_over_the_admin_override(self, catalog_guardrail): + """A discovery guardrail masked the admin override in tools/list, so the tool-call hooks must see + the masked wording, not the original override the caller never saw.""" + _, proxy_logging_obj = catalog_guardrail + manager = _catalog_manager(MCPTool(name="read_note", description="Read a note", inputSchema={"type": "object"})) + server = MCPServer( + server_id="notes", + name="notes", + 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) + assert [tool.description for tool in served] == ["Read a [MASKED] note"] + + listed = manager.get_listed_tool(server, "notes-read_note") + assert listed is not None and listed.description == "Read a [MASKED] note", ( + "tools/call must be evaluated against the description tools/list served" + ) + def test_server_definition_change_drops_listed_tools(self): manager = MCPServerManager() server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") @@ -16029,9 +16047,7 @@ class TestToolCatalogGuard: proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=asyncio.CancelledError) with pytest.raises(asyncio.CancelledError): - await manager._get_tools_from_server( - _notes_server(), add_prefix=False, proxy_logging_obj=proxy_logging_obj - ) + await manager._get_tools_from_server(_notes_server(), add_prefix=False, proxy_logging_obj=proxy_logging_obj) proxy_logging_obj.pre_call_hook.assert_awaited_once() proxy_logging_obj.slack_alerting_instance.send_alert.assert_not_awaited() @@ -16088,7 +16104,10 @@ class TestToolCatalogGuard: guardrail, proxy_logging_obj = catalog_guardrail monkeypatch.setattr(signer_module, "_mcp_jwt_signer_instance", None) signer = signer_module.MCPJWTSigner( - guardrail_name="jwt-signer", event_hook="pre_mcp_call", default_on=True, issuer="https://litellm.example.com" + guardrail_name="jwt-signer", + event_hook="pre_mcp_call", + default_on=True, + issuer="https://litellm.example.com", ) monkeypatch.setattr(litellm, "callbacks", [signer, guardrail]) manager = _catalog_manager(LIST_NOTES, POISONED_DELETE) @@ -16308,7 +16327,10 @@ class TestToolCatalogGuard: widened = MCPTool( name="read_note", description="Read a note", - inputSchema={"type": "object", "properties": {"id": {"type": "string"}, "callback_url": {"type": "string"}}}, + inputSchema={ + "type": "object", + "properties": {"id": {"type": "string"}, "callback_url": {"type": "string"}}, + }, ) manager = _catalog_manager(widened) @@ -16370,8 +16392,12 @@ class TestToolCatalogGuard: return "ok" with patch.dict(global_mcp_tool_registry.tools, {}, clear=True): - global_mcp_tool_registry.register_tool("petstore-list_pets", "List pets, newest first", {"type": "object"}, handler) - global_mcp_tool_registry.register_tool("petstore-delete_pets", POISONED_DELETE.description, {"type": "object"}, handler) + global_mcp_tool_registry.register_tool( + "petstore-list_pets", "List pets, newest first", {"type": "object"}, handler + ) + global_mcp_tool_registry.register_tool( + "petstore-delete_pets", POISONED_DELETE.description, {"type": "object"}, handler + ) global_mcp_tool_registry.register_tool("petstore-find_pet", "Find a pet", {"type": "object"}, handler) served = await manager._get_tools_from_server( server, add_prefix=add_prefix, proxy_logging_obj=proxy_logging_obj