mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): hand tools/call hooks the exact catalog entry tools/list served
get_listed_tool re-applied the admin description override on top of the cached listing, so a guardrail-masked description was restored to its original wording at call time, and the OpenAPI / local-registry call path built its metadata from the registry instead of the guarded caller catalog. Both paths now return the cached entry as served, falling back to the registry only when no listing was recorded Adds tests/integration/mcp/test_mcp_listed_tool_metadata.py (red on the prior head for the two regressions, red on the merge base for the feature, green on this head) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
f8734fd815
commit
2e1c6bfc70
6 changed files with 380 additions and 71 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
169
tests/integration/mcp/test_mcp_listed_tool_metadata.py
Normal file
169
tests/integration/mcp/test_mcp_listed_tool_metadata.py
Normal file
|
|
@ -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"
|
||||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue