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>
This commit is contained in:
joshua-berri 2026-10-03 15:00:43 -07:00 • committed by GitHub
parent 98337c9334
commit 0ea166c160
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 697 additions and 101 deletions

View file

@ -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,

View file

@ -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

View file

@ -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"],

View file

@ -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).

View file

@ -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"]

View file

@ -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)

View file

@ -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

View file

@ -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",

View file

@ -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

2
uv.lock generated
View file

@ -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" },