From 0ea166c160e64ca62f4781385b9f531d01675e08 Mon Sep 17 00:00:00 2001 From: joshua-berri Date: Sat, 3 Oct 2026 15:00:43 -0700 Subject: [PATCH] fix(mcp): preserve upstream tool schemas and parameter headers (#44425) * fix(mcp): preserve tool schemas and modern parameter headers * fix(mcp): allow bounded cold schema worker startup * fix(mcp): align catalog deadlines and bound schema traversal --------- Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> --- litellm/experimental_mcp_client/client.py | 6 + .../mcp_server/rest_endpoints.py | 6 +- .../_experimental/mcp_server/tool_search.py | 92 +++++- pyproject.toml | 1 + tests/mcp_tests/test_proxy_mcp_e2e.py | 177 ++++++++++- .../test_mcp_client.py | 293 +++++++++++++++++- .../mcp_server/test_mcp_server_manager.py | 40 +-- .../mcp_server/test_mcp_tool_search.py | 89 +++++- .../mcp_server/test_rest_endpoints.py | 92 +++--- uv.lock | 2 + 10 files changed, 697 insertions(+), 101 deletions(-) diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index bf6780c7812..c3e4881d689 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -927,6 +927,12 @@ class MCPClient: async def _call_tool_operation(session: ClientSession): verbose_logger.debug("MCP client sending tool call to session") + if self.protocol_version == "2026-07-28": + tools: Final = await list_tools_with_pagination( + session, listing_deadline=max(self.timeout, MCP_TOOL_LISTING_TIMEOUT) + ) + if not any(tool.name == call_tool_request_params.name for tool in tools): + raise MCPError(code=-32603, message="Tool schema is unavailable from the bounded upstream catalog") return await session.call_tool( name=call_tool_request_params.name, arguments=call_tool_request_params.arguments, diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index d17b849750b..7f936ec9269 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -555,7 +555,7 @@ if MCP_AVAILABLE: verbose_logger.warning("_prefetch_user_oauth_creds: failed to prefetch for user=%s: %s", user_id, e) return {} - def _create_tool_response_objects(tools, server: MCPServer): + def _create_tool_response_objects(tools: Sequence[MCPTool], server: MCPServer): """Helper function to create tool response objects. Enriches the server's ``mcp_info`` with ``server_id`` and ``alias`` so @@ -569,9 +569,7 @@ if MCP_AVAILABLE: } return [ ListMCPToolsRestAPIResponseObject( - name=tool.name, - description=tool.description, - inputSchema=tool.input_schema, + **tool.model_dump(by_alias=True, exclude={"mcp_info"}), mcp_info=enriched_mcp_info, ) for tool in tools diff --git a/litellm/proxy/_experimental/mcp_server/tool_search.py b/litellm/proxy/_experimental/mcp_server/tool_search.py index 63d7127ce98..217ffcd6b5f 100644 --- a/litellm/proxy/_experimental/mcp_server/tool_search.py +++ b/litellm/proxy/_experimental/mcp_server/tool_search.py @@ -2,13 +2,19 @@ from __future__ import annotations import hashlib import json -from collections.abc import Mapping, Sequence +import re +from collections import deque +from collections.abc import Iterator, Mapping, Sequence from dataclasses import dataclass from datetime import datetime +from itertools import chain, islice from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, TypedDict -from pydantic import ValidationError +import anyio +from anyio import to_process +from anyio.lowlevel import RunVar +from pydantic import JsonValue, ValidationError from typing_extensions import ReadOnly, Required, assert_never import litellm @@ -44,6 +50,77 @@ VIRTUAL_TOOL_NAMES: Final = frozenset( (MCP_TOOL_SEARCH_TOOL_NAME, MCP_TOOL_CALL_TOOL_NAME, AGENT_SEARCH_TOOL_NAME, SKILL_SEARCH_TOOL_NAME) ) +_SCHEMA_VALIDATION_LIMITER: Final = RunVar[anyio.CapacityLimiter]("mcp_schema_validation_limiter") +_MAX_VALIDATION_DEPTH: Final = 64 +_MAX_VALIDATION_NODES: Final = 10_000 +_MAX_VALIDATION_CHARACTERS: Final = 1_048_576 +_VALIDATION_TIMEOUT_SECONDS: Final = 30 + + +def _validation_nodes(value: JsonValue | Mapping[str, JsonValue]) -> Iterator[tuple[int, int]]: + pending: Final[deque[Iterator[JsonValue | Mapping[str, JsonValue]]]] = deque((iter((value,)),)) + while pending: + try: + node, depth = next(pending[-1]), len(pending) - 1 + except StopIteration: + pending.pop() + continue + yield depth, len(node) if isinstance(node, str) else 0 + if depth > _MAX_VALIDATION_DEPTH: + continue + if isinstance(node, Mapping): + pending.append(chain(node.keys(), node.values())) + elif isinstance(node, list): + pending.append(iter(node)) + + +def _validation_limit_error(schema: Mapping[str, JsonValue], arguments: Mapping[str, JsonValue]) -> str | None: + nodes: Final = tuple( + islice(chain(_validation_nodes(schema), _validation_nodes(arguments)), _MAX_VALIDATION_NODES + 1) + ) + if len(nodes) > _MAX_VALIDATION_NODES or any(depth > _MAX_VALIDATION_DEPTH for depth, _ in nodes): + return "Tool schema or arguments exceed validation size or depth limits" + if sum(size for _, size in nodes) > _MAX_VALIDATION_CHARACTERS: + return "Tool schema or arguments exceed validation size or depth limits" + return None + + +def _validate_tool_arguments(schema: Mapping[str, JsonValue], arguments: Mapping[str, JsonValue]) -> str | None: + from jsonschema import validate + from jsonschema.exceptions import SchemaError + from jsonschema.exceptions import ValidationError as JsonSchemaValidationError + from referencing import Registry + from referencing.exceptions import Unresolvable + + try: + validate(instance=arguments, schema=dict(schema), registry=Registry()) + except JsonSchemaValidationError as exc: + return f"Invalid arguments: {exc.message}" + except (SchemaError, Unresolvable, RecursionError, re.error): + return "Unable to validate tool arguments against the supplied schema" + return None + + +async def _tool_argument_validation_error( + schema: Mapping[str, JsonValue], arguments: Mapping[str, JsonValue] +) -> str | None: + limit_error: Final = _validation_limit_error(schema, arguments) + if limit_error is not None: + return limit_error + existing: Final = _SCHEMA_VALIDATION_LIMITER.get(None) + limiter: Final = existing if existing is not None else anyio.CapacityLimiter(4) + if existing is None: + _SCHEMA_VALIDATION_LIMITER.set(limiter) + try: + with anyio.fail_after(_VALIDATION_TIMEOUT_SECONDS): + return await to_process.run_sync( + _validate_tool_arguments, schema, arguments, cancellable=True, limiter=limiter + ) + except TimeoutError: + return "Tool argument validation exceeded its time limit" + except anyio.BrokenWorkerProcess: + return "Tool argument validation worker failed" + def coerce_top_k(value: Any, default: int = 5) -> int: try: @@ -502,7 +579,7 @@ async def handle_mcp_tool_search( async def handle_mcp_proxy_tool( name: str, - arguments: dict[str, object], # mutable-ok: MCP dispatcher passes mutable call arguments + arguments: dict[str, JsonValue], # mutable-ok: MCP dispatcher passes mutable JSON call arguments user_api_key_dict: UserAPIKeyAuth, client_ip: str | None = None, mcp_servers: list[str] | None = None, # mutable-ok: preserve MCP scope container for existing resolver @@ -513,8 +590,6 @@ async def handle_mcp_proxy_tool( litellm_logging_obj: LiteLLMLoggingObj | None = None, ) -> CallToolResult: from fastapi import HTTPException - from jsonschema import ValidationError as JsonSchemaValidationError - from jsonschema import validate from litellm.proxy import proxy_server from litellm.proxy._experimental.mcp_server.operations import ( @@ -571,10 +646,9 @@ async def handle_mcp_proxy_tool( tool_arguments: Final = arguments.get("arguments", {}) if not isinstance(tool_arguments, dict): return _text_tool_result("arguments must be an object", is_error=True) - try: - validate(instance=tool_arguments, schema=tool.input_schema) - except JsonSchemaValidationError as exc: - return _text_tool_result(f"Invalid arguments: {exc.message}", is_error=True) + validation_error: Final = await _tool_argument_validation_error(tool.input_schema, tool_arguments) + if validation_error is not None: + return _text_tool_result(validation_error, is_error=True) return await handle_mcp_tool_call( tool_name=_mcp_proxy_identity(tool)["tool_name"], diff --git a/pyproject.toml b/pyproject.toml index 229e71b9d79..d765bccb552 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -254,6 +254,7 @@ proxy-dev = [ "a2a-sdk==1.1.0", ] ci = [ + "psutil==7.2.2", # These are lazily imported at call sites; keep them out of core deps to # avoid bloating the base SDK install (google-generativeai pulls grpcio + # protobuf, Pillow is a compiled C extension). diff --git a/tests/mcp_tests/test_proxy_mcp_e2e.py b/tests/mcp_tests/test_proxy_mcp_e2e.py index 92f26e4ab7e..64e7483df6b 100644 --- a/tests/mcp_tests/test_proxy_mcp_e2e.py +++ b/tests/mcp_tests/test_proxy_mcp_e2e.py @@ -9,10 +9,12 @@ import tempfile import threading import time import typing +from collections.abc import Mapping from contextlib import asynccontextmanager, contextmanager from dataclasses import dataclass from datetime import datetime from pathlib import Path +from typing import Final import httpx import httpx2 @@ -21,9 +23,12 @@ import uvicorn import yaml from mcp import ClientSession from mcp.client.streamable_http import streamable_http_client +from mcp.shared._httpx_utils import create_mcp_http_client from mcp.types import CallToolResult from starlette.requests import Request +from tests.integration._support.wire import Reply, Request as WireRequest, Wire, wire_server + from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._experimental.mcp_server.tool_search import handle_mcp_proxy_tool from litellm.proxy._types import LiteLLM_ObjectPermissionTable, ProxyException, UserAPIKeyAuth @@ -45,6 +50,173 @@ PROXY_START_TIMEOUT = 30 PROXY_AUTHORIZATION_HEADER = "Bearer sk-1234" +@pytest.mark.asyncio +async def test_cold_concurrent_schema_validation_accepts_valid_arguments() -> None: + from litellm.proxy._experimental.mcp_server.tool_search import _tool_argument_validation_error + + schema: Final = { + "type": "object", + "$defs": {"amount": {"type": "number", "multipleOf": 0.25}}, + "properties": {"amount": {"$ref": "#/$defs/amount"}}, + } + original_affinity: Final = os.sched_getaffinity(0) if sys.platform == "linux" else None + try: + if original_affinity is not None: + os.sched_setaffinity(0, {min(original_affinity)}) + for _ in range(2): + results: Final = await asyncio.gather( + *(_tool_argument_validation_error(schema, {"amount": 0.75}) for _ in range(8)) + ) + assert results == [None] * 8, "Cold and warm workers must accept valid concurrent tool arguments" + finally: + if original_affinity is not None: + os.sched_setaffinity(0, original_affinity) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("cancel", (False, True)) +@pytest.mark.parametrize("concurrency", (1, 8)) +async def test_schema_validation_stops_expensive_work_and_recovers(cancel: bool, concurrency: int) -> None: + import psutil + + from litellm.proxy._experimental.mcp_server.tool_search import _tool_argument_validation_error + + existing_children: Final = frozenset(child.pid for child in psutil.Process().children()) + warm_count: Final = min(concurrency, 4) + assert await asyncio.gather( + *(_tool_argument_validation_error({"type": "object"}, {}) for _ in range(warm_count)) + ) == [None] * warm_count + started: Final = time.monotonic() + tasks: Final = tuple( + asyncio.create_task( + _tool_argument_validation_error( + {"type": "object", "properties": {"value": {"type": "string", "pattern": "^(a+)+$"}}}, + {"value": "a" * 80 + "!"}, + ) + ) + for _ in range(concurrency) + ) + group: Final = asyncio.gather(*tasks, return_exceptions=True) + try: + with pytest.raises(TimeoutError): + await asyncio.wait_for(asyncio.shield(group), timeout=0.2) + workers: Final = tuple( + child + for child in psutil.Process().children() + if child.pid not in existing_children and "anyio.to_process" in child.cmdline() + ) + assert 1 <= len(workers) <= 4 + if cancel: + for task in tasks: + task.cancel() + assert all(isinstance(result, asyncio.CancelledError) for result in await group) + else: + assert await group == ["Tool argument validation exceeded its time limit"] * concurrency + assert time.monotonic() - started < 35 + assert all(not worker.is_running() for worker in workers), "Cancelled validation must terminate worker CPU work" + assert await _tool_argument_validation_error({"type": "object"}, {}) is None + finally: + for task in tasks: + task.cancel() + await group + + +@pytest.fixture(scope="session") +def schema_peer() -> typing.Iterator[Wire]: + def respond(request: WireRequest) -> Reply: + if request.method == "GET": + return Reply(body=b'{"type":"number"}') + body: Final = json.loads(request.body) + if "id" not in body: + return Reply(status=202) + result: Final[Mapping[str, object]] + if body["method"] == "initialize": + result = { + "protocolVersion": body["params"]["protocolVersion"], + "capabilities": {"tools": {}}, + "serverInfo": {"name": "schema-peer", "version": "1"}, + } + elif body["method"] == "tools/list": + result = { + "tools": [ + { + "name": name, + "description": "Schema validation", + "inputSchema": { + "type": "object", + "$defs": {"amount": {"type": "number", "minimum": 0.25, "multipleOf": 0.25}}, + "properties": { + "value": {"$ref": peer.url + "/ref" if name == "external" else "#/$defs/amount"} + }, + }, + } + for name in ("external", "local") + ] + } + else: + assert body["method"] == "tools/call" + result = {"content": [{"type": "text", "text": "called"}], "isError": False} + return Reply(body=json.dumps({"jsonrpc": "2.0", "id": body["id"], "result": result}).encode()) + + with wire_server(respond) as peer: + yield peer + + +@pytest.mark.parametrize("external", (True, False)) +def test_proxy_schema_validation_resolves_only_local_references( + proxy_server_url: str, schema_peer: Wire, external: bool +) -> None: + with httpx.Client( + base_url=proxy_server_url, + headers={ + "Authorization": "Bearer sk-schema", + "Accept": "application/json, text/event-stream", + "x-mcp-servers": "schema", + }, + timeout=30, + ) as client: + search: Final = _rpc_result( + client.post( + "/mcp/proxy", + json={ + "jsonrpc": "2.0", + "id": 1, + "method": "tools/call", + "params": {"name": "search_tools", "arguments": {"query": "schema"}}, + }, + ) + ) + name: Final = "schema-external" if external else "schema-local" + tool_id: Final = next(hit["tool_id"] for hit in json.loads(search["content"][0]["text"]) if hit["name"] == name) + schema_peer.drain() + called: Final = _rpc_result( + client.post( + "/mcp/proxy", + json={ + "jsonrpc": "2.0", + "id": 2, + "method": "tools/call", + "params": { + "name": "call_tool", + "arguments": {"tool_id": tool_id, "arguments": {"value": 0.75}}, + }, + }, + ) + ) + observed: Final = schema_peer.drain() + assert not any(request.method == "GET" for request in observed), ( + "Schema validation must not retrieve external references" + ) + calls: Final = tuple( + request + for request in observed + if request.method == "POST" and json.loads(request.body).get("method") == "tools/call" + ) + assert len(calls) == (0 if external else 1), "Only valid local schemas may reach upstream execution" + assert called["isError"] is external + if not external: + assert called["content"][0]["text"] == "called" + @pytest.fixture(scope="session", autouse=True) def _clear_proxy_database_env() -> typing.Iterator[None]: """Ensure local proxy DB settings don't leak into tests.""" @@ -175,10 +347,12 @@ def _proxy_server( tmp_path_factory: pytest.TempPathFactory, math_streamable_http_server: str, math_restricted_server: str, + schema_peer: Wire, ): config_dir = tmp_path_factory.mktemp("mcp_e2e") config_path = config_dir / "config.yaml" config = yaml.safe_load(CONFIG_TEMPLATE_PATH.read_text()) + config["mcp_servers"]["schema"] = {"transport": "http", "url": schema_peer.url + "/mcp"} config["mcp_servers"]["math_stdio"]["command"] = MCP_PEER_PYTHON config["mcp_servers"]["math_streamable_http"]["url"] = f"{math_streamable_http_server}/mcp" config["mcp_servers"]["math_restricted"]["url"] = f"{math_restricted_server}/mcp" @@ -208,7 +382,7 @@ def proxy_server_url(_proxy_server: ProxyRig, setup_and_teardown: None) -> str: @asynccontextmanager async def _http_streams(url: str, headers: dict[str, str]): - async with httpx2.AsyncClient(headers=headers) as http_client: + async with create_mcp_http_client(headers=headers) as http_client: async with streamable_http_client(url, http_client=http_client) as streams: yield streams @@ -528,6 +702,7 @@ class TestProxyMcpSchemaDiscoveryMode: async def authorize_proxy_key(request: Request, api_key: str) -> UserAPIKeyAuth: permissions = { + "sk-schema": LiteLLM_ObjectPermissionTable(object_permission_id="schema", mcp_servers=["schema"]), "sk-1234": LiteLLM_ObjectPermissionTable(object_permission_id="open", mcp_servers=["math_stdio"]), "sk-restricted": LiteLLM_ObjectPermissionTable( object_permission_id="restricted", mcp_servers=["math_restricted"] diff --git a/tests/unit/experimental_mcp_client/test_mcp_client.py b/tests/unit/experimental_mcp_client/test_mcp_client.py index 015d12c3d5e..83b669e68f0 100644 --- a/tests/unit/experimental_mcp_client/test_mcp_client.py +++ b/tests/unit/experimental_mcp_client/test_mcp_client.py @@ -3114,8 +3114,8 @@ async def test_modern_upstream_requests_are_self_contained_without_initializatio assert client._last_initialize_instructions == "modern instructions" assert tuple(methods.get_nowait() for _ in range(methods.qsize())) == ( "server/discover", - "tools/call", "tools/list", + "tools/call", ) else: with pytest.raises((MCPError, RuntimeError), match="protocol version"): @@ -3123,6 +3123,297 @@ async def test_modern_upstream_requests_are_self_contained_without_initializatio assert tuple(methods.get_nowait() for _ in range(methods.qsize())) == ("server/discover",) +@pytest.mark.asyncio +@pytest.mark.parametrize("paginated", (False, True)) +@pytest.mark.parametrize("valid_annotation", (True, False)) +async def test_modern_call_emits_listed_argument_headers(paginated: bool, valid_annotation: bool) -> None: + from collections.abc import Mapping + from queue import SimpleQueue + + from mcp.types import DiscoverResult, ToolsCapability + + calls: Final[SimpleQueue[str]] = SimpleQueue() + + def response_result(request: httpx2.Request, payload: JSONRPCRequest) -> Mapping[str, object]: + assert payload.params is not None + caller: Final = request.headers["authorization"].removeprefix("Bearer ") + if payload.method == "server/discover": + return DiscoverResult( + supported_versions=["2026-07-28"], capabilities=ServerCapabilities(tools=ToolsCapability()) + ).model_dump(by_alias=True, exclude_none=True) + if payload.method == "tools/list": + if paginated and "cursor" not in payload.params: + return { + "resultType": "complete", + "cacheScope": "private", + "ttlMs": 0, + "tools": [], + "nextCursor": "second", + } + return { + "resultType": "complete", + "cacheScope": "private", + "ttlMs": 0, + "tools": [ + { + "name": "quote", + "inputSchema": { + "type": "object", + "properties": { + "workspace": { + "type": "string" if valid_annotation else "number", + "x-mcp-header": caller, + }, + "count": {"type": "integer", "x-mcp-header": "Count"}, + "preview": {"type": "boolean", "x-mcp-header": "Preview"}, + }, + }, + } + ], + } + assert payload.method == "tools/call" + calls.put(caller) + assert valid_annotation, "Invalid header definitions must prevent dispatch" + assert request.headers.get(f"mcp-param-{caller.lower()}") == caller + assert request.headers.get("mcp-param-count") == "3" + assert request.headers.get("mcp-param-preview") == "true" + assert set(key for key in request.headers if key.startswith("mcp-param-")) == { + f"mcp-param-{caller.lower()}", + "mcp-param-count", + "mcp-param-preview", + } + assert payload.params["arguments"] == {"workspace": caller, "count": 3, "preview": True} + return {"resultType": "complete", "content": [{"type": "text", "text": "quoted"}], "isError": False} + + def respond(request: httpx2.Request) -> httpx2.Response: + payload: Final = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content) + assert isinstance(payload, JSONRPCRequest) + return httpx2.Response( + 200, json={"jsonrpc": "2.0", "id": payload.id, "result": response_result(request, payload)} + ) + + async def call_as(caller: str) -> None: + client: Final = _MockTransportClient( + respond, + server_url="https://example.com/mcp", + protocol_version="2026-07-28", + auth_type=MCPAuth.bearer_token, + auth_value=caller, + ) + params: Final = CallToolRequestParams( + name="quote", arguments={"workspace": caller, "count": 3, "preview": True} + ) + if not valid_annotation: + with pytest.raises(MCPError, match="schema is unavailable"): + await client.call_tool(params, raise_on_error=True) + return + result: Final = await client.call_tool(params, raise_on_error=True) + assert result.is_error is False + assert result.content[0].text == "quoted" + + await asyncio.gather(call_as("Engineering"), call_as("Finance")) + assert calls.qsize() == (2 if valid_annotation else 0) + +@pytest.mark.asyncio +@pytest.mark.parametrize("stop", ("page_cap", "repeated_cursor")) +@pytest.mark.parametrize("schema_available", (False, True)) +async def test_modern_call_requires_schema_from_bounded_listing( + stop: str, schema_available: bool, monkeypatch: pytest.MonkeyPatch +) -> None: + from queue import SimpleQueue + + from mcp.types import DiscoverResult, ToolsCapability + + monkeypatch.setattr("litellm.experimental_mcp_client.tools.MCP_TOOL_LISTING_MAX_PAGES", 2) + methods: Final[SimpleQueue[str]] = SimpleQueue() + + def respond(request: httpx2.Request) -> httpx2.Response: + payload: Final = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content) + assert isinstance(payload, JSONRPCRequest) + methods.put(payload.method) + if payload.method == "server/discover": + return httpx2.Response( + 200, + json={ + "jsonrpc": "2.0", + "id": payload.id, + "result": DiscoverResult( + supported_versions=["2026-07-28"], capabilities=ServerCapabilities(tools=ToolsCapability()) + ).model_dump(by_alias=True, exclude_none=True), + }, + ) + if payload.method == "tools/list": + assert payload.params is not None + return httpx2.Response( + 200, + json={ + "jsonrpc": "2.0", + "id": payload.id, + "result": { + "resultType": "complete", + "cacheScope": "private", + "ttlMs": 0, + "nextCursor": "loop" + if stop == "repeated_cursor" + else str(int(payload.params.get("cursor", "0")) + 1), + "tools": [ + { + "name": "quote", + "inputSchema": { + "type": "object", + "properties": {"workspace": {"type": "string", "x-mcp-header": "Workspace"}}, + }, + } + ] + if schema_available + else [], + }, + }, + ) + assert payload.method == "tools/call" + assert schema_available, "An unavailable schema must not permit dispatch without required headers" + assert request.headers["mcp-param-workspace"] == "engineering" + return httpx2.Response( + 200, + json={ + "jsonrpc": "2.0", + "id": payload.id, + "result": {"resultType": "complete", "content": [{"type": "text", "text": "quoted"}], "isError": False}, + }, + ) + + client: Final = _MockTransportClient(respond, server_url="https://example.com/mcp", protocol_version="2026-07-28") + params: Final = CallToolRequestParams(name="quote", arguments={"workspace": "engineering"}) + if schema_available: + result: Final = await client.call_tool(params, raise_on_error=True) + assert not result.is_error + assert result.content[0].text == "quoted" + else: + with pytest.raises(MCPError, match="schema is unavailable") as error: + await client.call_tool(params, raise_on_error=True) + assert error.value.code == -32603 + observed: Final = tuple(methods.get_nowait() for _ in range(methods.qsize())) + assert observed.count("tools/list") == 2 + assert observed.count("tools/call") == int(schema_available) + + +def test_modern_call_uses_discovery_deadline_for_multiple_pages() -> None: + from mcp.types import DiscoverResult, ToolsCapability + + loop: Final = _AutojumpClockLoop() + + async def respond(request: httpx2.Request) -> httpx2.Response: + payload: Final = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content) + assert isinstance(payload, JSONRPCRequest) + if payload.method == "server/discover": + return httpx2.Response( + 200, + json={ + "jsonrpc": "2.0", + "id": payload.id, + "result": DiscoverResult( + supported_versions=["2026-07-28"], capabilities=ServerCapabilities(tools=ToolsCapability()) + ).model_dump(by_alias=True, exclude_none=True), + }, + ) + if payload.method == "tools/list": + await asyncio.sleep(0.06) + assert payload.params is not None + page: Final = ( + {"tools": [], "nextCursor": "second"} + if "cursor" not in payload.params + else { + "tools": [ + { + "name": "quote", + "inputSchema": { + "type": "object", + "properties": {"workspace": {"type": "string", "x-mcp-header": "Workspace"}}, + }, + } + ] + } + ) + return httpx2.Response( + 200, + json={ + "jsonrpc": "2.0", + "id": payload.id, + "result": {"resultType": "complete", "cacheScope": "private", "ttlMs": 0, **page}, + }, + ) + assert payload.method == "tools/call" + assert request.headers["mcp-param-workspace"] == "engineering" + return httpx2.Response( + 200, + json={ + "jsonrpc": "2.0", + "id": payload.id, + "result": {"resultType": "complete", "content": [{"type": "text", "text": "quoted"}], "isError": False}, + }, + ) + + async def run() -> None: + client: Final = _MockTransportClient( + respond, server_url="https://example.com/mcp", protocol_version="2026-07-28", timeout=0.1 + ) + assert [tool.name for tool in await client.list_tools(raise_on_error=True)] == ["quote"] + result: Final = await client.call_tool( + CallToolRequestParams(name="quote", arguments={"workspace": "engineering"}), raise_on_error=True + ) + assert result.is_error is False + assert result.content[0].text == "quoted" + + try: + loop.run_until_complete(run()) + finally: + loop.run_until_complete(loop.shutdown_asyncgens()) + loop.close() + + +@pytest.mark.asyncio +async def test_cancelled_modern_catalog_load_prevents_tool_execution() -> None: + from queue import SimpleQueue + + from mcp.types import DiscoverResult, ToolsCapability + + listing_started: Final = asyncio.Event() + hold_listing: Final = asyncio.Event() + methods: Final[SimpleQueue[str]] = SimpleQueue() + + async def respond(request: httpx2.Request) -> httpx2.Response: + payload: Final = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content) + assert isinstance(payload, JSONRPCRequest) + methods.put(payload.method) + if payload.method == "tools/list": + listing_started.set() + await hold_listing.wait() + assert payload.method == "server/discover", "A cancelled listing must not dispatch a tool call" + discovery: Final = DiscoverResult( + supported_versions=["2026-07-28"], capabilities=ServerCapabilities(tools=ToolsCapability()) + ) + return httpx2.Response( + 200, + json={"jsonrpc": "2.0", "id": payload.id, "result": discovery.model_dump(by_alias=True, exclude_none=True)}, + ) + + client: Final = _MockTransportClient(respond, server_url="https://example.com/mcp", protocol_version="2026-07-28") + task: Final = asyncio.create_task( + client.call_tool(CallToolRequestParams(name="quote", arguments={}), raise_on_error=True) + ) + observed_listing: Final = asyncio.create_task(listing_started.wait()) + try: + await asyncio.wait((task, observed_listing), return_when=asyncio.FIRST_COMPLETED) + assert listing_started.is_set() + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + assert tuple(methods.get_nowait() for _ in range(methods.qsize())) == ("server/discover", "tools/list") + finally: + task.cancel() + observed_listing.cancel() + await asyncio.gather(task, observed_listing, return_exceptions=True) + def test_modern_upstream_rejects_legacy_sse_transport() -> None: with pytest.raises(ValueError, match="transport"): MCPClient(protocol_version="2026-07-28", transport_type=MCPTransport.sse) 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 4b460fc67ae..cf6cf93f35b 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 @@ -6637,20 +6637,11 @@ class TestMCPServerManager: server.mcp_info = {"server_name": "test-server"} # Mock tools returned from manager (3 tools, but only 2 are allowed) - tool1 = MagicMock() - tool1.name = "allowed_tool_1" - tool1.description = "This tool is allowed" - tool1.input_schema = {} + tool1 = MCPTool(name="allowed_tool_1", description="This tool is allowed", inputSchema={}) - tool2 = MagicMock() - tool2.name = "blocked_tool" - tool2.description = "This tool is not allowed" - tool2.input_schema = {} + tool2 = MCPTool(name="blocked_tool", description="This tool is not allowed", inputSchema={}) - tool3 = MagicMock() - tool3.name = "allowed_tool_2" - tool3.description = "This tool is also allowed" - tool3.input_schema = {} + tool3 = MCPTool(name="allowed_tool_2", description="This tool is also allowed", inputSchema={}) # Mock the global_mcp_server_manager._get_tools_from_server from litellm.proxy._experimental.mcp_server import rest_endpoints @@ -6687,20 +6678,11 @@ class TestMCPServerManager: server.mcp_info = {"server_name": "test-server"} # Mock tools returned from manager - tool1 = MagicMock() - tool1.name = "tool_1" - tool1.description = "Tool 1" - tool1.input_schema = {} + tool1 = MCPTool(name="tool_1", description="Tool 1", inputSchema={}) - tool2 = MagicMock() - tool2.name = "tool_2" - tool2.description = "Tool 2" - tool2.input_schema = {} + tool2 = MCPTool(name="tool_2", description="Tool 2", inputSchema={}) - tool3 = MagicMock() - tool3.name = "tool_3" - tool3.description = "Tool 3" - tool3.input_schema = {} + tool3 = MCPTool(name="tool_3", description="Tool 3", inputSchema={}) # Mock the global_mcp_server_manager._get_tools_from_server from litellm.proxy._experimental.mcp_server import rest_endpoints @@ -6737,15 +6719,9 @@ class TestMCPServerManager: server.mcp_info = {"server_name": "test-server"} # Mock tools returned from manager - tool1 = MagicMock() - tool1.name = "tool_1" - tool1.description = "Tool 1" - tool1.input_schema = {} + tool1 = MCPTool(name="tool_1", description="Tool 1", inputSchema={}) - tool2 = MagicMock() - tool2.name = "tool_2" - tool2.description = "Tool 2" - tool2.input_schema = {} + tool2 = MCPTool(name="tool_2", description="Tool 2", inputSchema={}) # Mock the global_mcp_server_manager._get_tools_from_server from litellm.proxy._experimental.mcp_server import rest_endpoints diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py index 4575741aa8b..1cfe6198b6f 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py @@ -11,12 +11,13 @@ Covers: """ import json -from collections.abc import Sequence +from collections.abc import Mapping, Sequence from types import SimpleNamespace -from typing import Any +from typing import Any, Final from unittest.mock import AsyncMock, MagicMock, patch import pytest +from pydantic import JsonValue from mcp.types import Tool import litellm @@ -60,6 +61,90 @@ SAMPLE_TOOLS = _make_tools( ) +@pytest.mark.parametrize( + ("schema", "arguments", "error"), + ( + ( + { + "type": "object", + "$defs": {"amount": {"type": "number", "minimum": 0.25, "multipleOf": 0.25}}, + "properties": {"amount": {"$ref": "#/$defs/amount"}}, + "required": ["amount"], + }, + {"amount": 0.75}, + None, + ), + ({"type": "object", "anyOf": [{"required": ["amount"]}, {"required": ["trace"]}]}, {"trace": "a"}, None), + ( + {"type": "object", "properties": {"amount": {"type": "number", "minimum": 0.25}}}, + {"amount": 0.1}, + "Invalid arguments:", + ), + ({"$ref": "https://schemas.example.invalid/amount"}, {}, "Unable to validate"), + ({"$ref": "#/$defs/missing"}, {}, "Unable to validate"), + ({"$ref": "#"}, {}, "Unable to validate"), + ({"type": "not-a-type"}, {}, "Unable to validate"), + ({"properties": {"value": {"pattern": "["}}}, {"value": "a"}, "Unable to validate"), + ), +) +def test_offline_schema_validation_preserves_supported_constraints( + schema: Mapping[str, JsonValue], arguments: Mapping[str, JsonValue], error: str | None +) -> None: + from litellm.proxy._experimental.mcp_server.tool_search import _validate_tool_arguments + + result: Final = _validate_tool_arguments(schema, arguments) + if error is None: + assert result is None + else: + assert result is not None and result.startswith(error) + + +@pytest.mark.parametrize(("count", "allowed"), ((4_999, True), (5_000, False))) +def test_validation_node_limit(count: int, allowed: bool) -> None: + from litellm.proxy._experimental.mcp_server.tool_search import _validation_limit_error + + result: Final = _validation_limit_error({}, {str(index): None for index in range(count)}) + assert (result is None) is allowed + + +@pytest.mark.parametrize(("size", "allowed"), ((1_048_575, True), (1_048_576, False))) +def test_validation_text_limit(size: int, allowed: bool) -> None: + from litellm.proxy._experimental.mcp_server.tool_search import _validation_limit_error + + assert (_validation_limit_error({}, {"x": "a" * size}) is None) is allowed + + +@pytest.mark.parametrize(("depth", "allowed"), ((64, True), (65, False))) +def test_validation_depth_limit(depth: int, allowed: bool) -> None: + from functools import reduce + + from litellm.proxy._experimental.mcp_server.tool_search import _validation_limit_error + + nested: Final = reduce(lambda value, _: {"x": value}, range(depth), {}) + assert (_validation_limit_error({}, nested) is None) is allowed + + +@pytest.mark.asyncio +async def test_oversized_arguments_are_rejected_before_worker_submission() -> None: + from litellm.proxy._experimental.mcp_server.tool_search import _tool_argument_validation_error + + with patch("anyio.to_process.run_sync", new_callable=AsyncMock) as submit: + result: Final = await _tool_argument_validation_error({}, {"value": "x" * 1_048_576}) + assert result == "Tool schema or arguments exceed validation size or depth limits" + submit.assert_not_called() + + +@pytest.mark.asyncio +async def test_failed_validation_worker_returns_tool_error() -> None: + from anyio import BrokenWorkerProcess + + from litellm.proxy._experimental.mcp_server.tool_search import _tool_argument_validation_error + + with patch("anyio.to_process.run_sync", side_effect=BrokenWorkerProcess): + result: Final = await _tool_argument_validation_error({}, {}) + assert result == "Tool argument validation worker failed" + + FX_TOOL = Tool( name="treasury-get_rates", description="Get foreign exchange rates for a currency pair", diff --git a/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py index 680b84469d9..20a8ee05eb5 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -14,7 +14,7 @@ if sys.version_info < (3, 11): # BaseExceptionGroup is a builtin only from 3.11 import httpx import pytest from fastapi import HTTPException -from mcp.types import CallToolResult, TextContent +from mcp.types import CallToolResult, TextContent, Tool from starlette.requests import Request from litellm.constants import MCP_TOOL_LISTING_TIMEOUT @@ -3404,17 +3404,10 @@ class TestGetToolsForSingleServer: from litellm.proxy._types import LiteLLM_ObjectPermissionTable from litellm.types.mcp import MCPTransport - # Create mock tools - class MockTool: - def __init__(self, name, description): - self.name = name - self.description = description - self.input_schema = {} - mock_tools = [ - MockTool("tool1", "First tool"), - MockTool("tool2", "Second tool"), - MockTool("tool3", "Third tool"), + Tool(name="tool1", description="First tool", inputSchema={}), + Tool(name="tool2", description="Second tool", inputSchema={}), + Tool(name="tool3", description="Third tool", inputSchema={}), ] # Mock _get_tools_from_server to return all tools @@ -3466,15 +3459,9 @@ class TestGetToolsForSingleServer: from litellm.proxy._experimental.mcp_server.server import MCPServer from litellm.types.mcp import MCPTransport - class MockTool: - def __init__(self, name, description): - self.name = name - self.description = description - self.input_schema = {} - mock_tools = [ - MockTool("tool1", "First tool"), - MockTool("tool2", "Second tool"), + Tool(name="tool1", description="First tool", inputSchema={}), + Tool(name="tool2", description="Second tool", inputSchema={}), ] async def fake_get_tools_from_server(**kwargs): @@ -3514,15 +3501,9 @@ class TestGetToolsForSingleServer: from litellm.proxy._types import LiteLLM_ObjectPermissionTable from litellm.types.mcp import MCPTransport - class MockTool: - def __init__(self, name, description): - self.name = name - self.description = description - self.input_schema = {} - mock_tools = [ - MockTool("tool1", "First tool"), - MockTool("tool2", "Second tool"), + Tool(name="tool1", description="First tool", inputSchema={}), + Tool(name="tool2", description="Second tool", inputSchema={}), ] async def fake_get_tools_from_server(**kwargs): @@ -3567,15 +3548,9 @@ class TestGetToolsForSingleServer: from litellm.proxy._types import LiteLLM_ObjectPermissionTable from litellm.types.mcp import MCPTransport - class MockTool: - def __init__(self, name, description): - self.name = name - self.description = description - self.input_schema = {} - mock_tools = [ - MockTool("tool1", "First tool"), - MockTool("tool2", "Second tool"), + Tool(name="tool1", description="First tool", inputSchema={}), + Tool(name="tool2", description="Second tool", inputSchema={}), ] async def fake_get_tools_from_server(**kwargs): @@ -3620,17 +3595,11 @@ class TestGetToolsForSingleServer: from litellm.proxy._types import LiteLLM_ObjectPermissionTable from litellm.types.mcp import MCPTransport - class MockTool: - def __init__(self, name, description): - self.name = name - self.description = description - self.input_schema = {} - mock_tools = [ - MockTool("tool1", "First tool"), - MockTool("tool2", "Second tool"), - MockTool("tool3", "Third tool"), - MockTool("tool4", "Fourth tool"), + Tool(name="tool1", description="First tool", inputSchema={}), + Tool(name="tool2", description="Second tool", inputSchema={}), + Tool(name="tool3", description="Third tool", inputSchema={}), + Tool(name="tool4", description="Fourth tool", inputSchema={}), ] async def fake_get_tools_from_server(**kwargs): @@ -3682,13 +3651,7 @@ class TestGetToolsForSingleServer: from litellm.proxy._experimental.mcp_server.server import MCPServer from litellm.types.mcp import MCPTransport - class MockTool: - def __init__(self, name): - self.name = name - self.description = name - self.input_schema = {} - - mock_tools = [MockTool("tool1"), MockTool("tool2"), MockTool("tool3")] + mock_tools = [Tool(name="tool1", description="tool1", inputSchema={}), Tool(name="tool2", description="tool2", inputSchema={}), Tool(name="tool3", description="tool3", inputSchema={})] async def fake_get_tools_from_server(**kwargs): return mock_tools @@ -4401,6 +4364,31 @@ class TestToolResponseMcpInfoEnrichment: "alias": None, } + def test_preserves_complete_sdk_tool_definition(self) -> None: + from mcp.types import Tool + + tool: Final = Tool.model_validate( + { + "name": "quote", + "title": "Quote", + "description": "Return a quote", + "inputSchema": { + "type": "object", + "$defs": {"amount": {"type": "number", "minimum": 0.25}}, + "properties": {"amount": {"$ref": "#/$defs/amount"}}, + "anyOf": [{"required": ["amount"]}, {"maxProperties": 0}], + }, + "outputSchema": {"type": "object", "properties": {"price": {"type": "number", "multipleOf": 0.25}}}, + "annotations": {"readOnlyHint": True}, + "_meta": {"display": {"priority": 0.75}}, + "icons": [{"src": "https://example.com/icon.png"}], + } + ) + original: Final = tool.model_dump(by_alias=True) + server: Final = MCPServer(server_id="quotes", name="quotes", transport=MCPTransport.http) + response: Final = rest_endpoints._create_tool_response_objects([tool], server)[0] + assert response.model_dump(by_alias=True, exclude={"mcp_info"}) == original + assert tool.model_dump(by_alias=True) == original class TestRestListToolsetFiltering: @pytest.mark.asyncio diff --git a/uv.lock b/uv.lock index b227ba183b5..0d1b17e6a6d 100644 --- a/uv.lock +++ b/uv.lock @@ -4669,6 +4669,7 @@ ci = [ { name = "lunary", version = "1.4.36", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, { name = "lunary", version = "1.4.37", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11'" }, { name = "pillow" }, + { name = "psutil" }, { name = "psycopg2-binary" }, { name = "pyarrow" }, { name = "pygithub" }, @@ -4883,6 +4884,7 @@ ci = [ { name = "lunary", marker = "python_full_version == '3.10.*'", specifier = "==1.4.36" }, { name = "lunary", marker = "python_full_version >= '3.11'", specifier = "==1.4.37" }, { name = "pillow", specifier = "==12.3.0" }, + { name = "psutil", specifier = "==7.2.2" }, { name = "psycopg2-binary", specifier = "==2.9.11" }, { name = "pyarrow", specifier = "==23.0.1" }, { name = "pygithub", specifier = "==2.8.1" },