mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
fix(mcp): accept omitted arguments on protocol tool calls
This commit is contained in:
parent
00fb0ab673
commit
bcaf442f95
3 changed files with 41 additions and 5 deletions
|
|
@ -2785,7 +2785,7 @@ async def _execute_mcp_server_tool_call(
|
|||
return virtual_tool_result
|
||||
|
||||
# Create a body date for logging
|
||||
body_data: Final = {"name": params.name, "arguments": params.arguments} # mutable-ok: logging payload
|
||||
body_data: Final = {"name": params.name, "arguments": params.arguments or {}} # mutable-ok: logging payload
|
||||
# Set trace/session id from raw_headers so spend logs and logging_obj stay consistent (same as A2A)
|
||||
chain_id: Final = get_chain_id_from_headers(raw_headers)
|
||||
if chain_id:
|
||||
|
|
|
|||
|
|
@ -123,3 +123,39 @@ def test_configured_revision_blocks_unadvertised_handshake_and_keeps_allowed_con
|
|||
assert result.is_error is False and result.content[0].text == "7"
|
||||
|
||||
asyncio.run(exercise())
|
||||
|
||||
|
||||
def test_omitted_tool_arguments_reach_the_upstream(gateway: Gateway) -> None:
|
||||
import asyncio
|
||||
|
||||
import httpx
|
||||
from integration._support.mcp import mcp_peer, tool_calls
|
||||
from mcp import ClientSession
|
||||
from mcp.client.streamable_http import streamable_http_client
|
||||
from mcp.types import CallToolRequest, CallToolRequestParams, CallToolResult
|
||||
|
||||
async def invoke(url: str, name: str, key: str | None = None) -> CallToolResult:
|
||||
async with httpx.AsyncClient(headers={"Authorization": f"Bearer {key}"} if key else {}) as http:
|
||||
async with streamable_http_client(url, http_client=http) as streams:
|
||||
async with ClientSession(streams[0], streams[1]) as session:
|
||||
await session.initialize()
|
||||
return await session.send_request(
|
||||
CallToolRequest(params=CallToolRequestParams(name=name)), CallToolResult
|
||||
)
|
||||
|
||||
with mcp_peer() as reference, gateway.scenario() as scenario:
|
||||
alias: Final = "optional" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, reference, alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
direct: Final = asyncio.run(invoke(reference.url, "fail"))
|
||||
assert direct.is_error is True and "synthetic tool failure" in direct.content[0].text, direct
|
||||
reference.drain()
|
||||
endpoint: Final = str(gateway.client.base_url).rstrip("/") + f"/{alias}/mcp"
|
||||
forwarded: Final = asyncio.run(invoke(endpoint, f"{alias}-fail", key))
|
||||
assert forwarded.is_error is True and "synthetic tool failure" in forwarded.content[0].text, forwarded
|
||||
assert len(tool_calls(reference.drain())) == 1, "Omitted arguments never reached the upstream"
|
||||
denied_key: Final = scenario.key(object_permission={"mcp_servers": []})
|
||||
aggregate: Final = str(gateway.client.base_url).rstrip("/") + "/mcp"
|
||||
denied: Final = asyncio.run(invoke(aggregate, f"{alias}-fail", denied_key))
|
||||
assert denied.is_error is True and "not allowed" in denied.content[0].text.lower(), denied
|
||||
assert tool_calls(reference.drain()) == (), "Denied caller reached the upstream"
|
||||
|
|
|
|||
|
|
@ -1715,8 +1715,9 @@ async def test_mcp_server_tool_call_renders_denial_message_not_detail_dict(_mcp_
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_server_tool_call_body_with_none_arguments(_mcp_request_ctx):
|
||||
"""Test that proxy_server_request body handles None arguments correctly"""
|
||||
@pytest.mark.parametrize("tool_arguments,expected", ((None, {}), ({}, {}), ({"value": 7}, {"value": 7})))
|
||||
async def test_mcp_server_tool_call_body_with_optional_arguments(_mcp_request_ctx, tool_arguments, expected):
|
||||
"""Omitted MCP arguments are an empty object; explicit inputs remain intact."""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
mcp_server_tool_call,
|
||||
|
|
@ -1727,7 +1728,6 @@ async def test_mcp_server_tool_call_body_with_none_arguments(_mcp_request_ctx):
|
|||
|
||||
# Setup test data
|
||||
tool_name = "test_tool_no_args"
|
||||
tool_arguments = None
|
||||
|
||||
# Mock user auth
|
||||
user_api_key_auth = UserAPIKeyAuth(api_key="test_key", user_id="test_user")
|
||||
|
|
@ -1771,7 +1771,7 @@ async def test_mcp_server_tool_call_body_with_none_arguments(_mcp_request_ctx):
|
|||
|
||||
body = captured_data["proxy_server_request"]["body"]
|
||||
assert body["name"] == tool_name
|
||||
assert body["arguments"] == tool_arguments # Should be None
|
||||
assert body["arguments"] == expected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue