mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
98337c9334
commit
0ea166c160
10 changed files with 697 additions and 101 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
2
uv.lock
generated
|
|
@ -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" },
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue