mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(mcp): degrade invalid tools/call content blocks instead of failing the whole result
Some upstream MCP servers emit content blocks that fail SDK-side validation, for example an EmbeddedResource with a relative uri instead of a URI with a scheme, or text/* content shipped as a base64 blob instead of text. ClientSession.call_tool lets a single such block fail the entire tools/call result with a ValidationError, which litellm's own MCPClient.call_tool then degrades into an opaque isError text block containing the pydantic traceback, discarding content the caller could otherwise use. Add TolerantCallToolResult, a CallToolResult whose model_validator degrades only the blocks that fail validation to text (decoding text/* blobs where possible) instead of raising, and TolerantClientSession, a ClientSession that parses tools/call results with it. Wire it in at the single ClientSession construction site in MCPClient._execute_session_operation. Verified against a real payload captured from an upstream Azure DevOps MCP server's repo_file tool, driven through the real ClientSession machinery over real anyio streams via the existing _ScriptedUpstream/_ScriptedClient test harness rather than a hand-built model. Since construction now uses TolerantClientSession, updated the nine existing tests that patched ClientSession directly to intercept construction.
This commit is contained in:
parent
423b791ee0
commit
3b7d9e74cb
3 changed files with 189 additions and 13 deletions
|
|
@ -35,6 +35,7 @@ from pydantic import AnyUrl
|
|||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_NPM_CACHE_DIR
|
||||
from litellm.experimental_mcp_client.tolerant_result import TolerantClientSession
|
||||
from litellm.llms.custom_httpx.http_handler import get_ssl_configuration
|
||||
from litellm.types.llms.custom_http import VerifyTypes
|
||||
from litellm.types.mcp import (
|
||||
|
|
@ -373,7 +374,7 @@ class MCPClient:
|
|||
session_kwargs["logging_callback"] = self._logging_callback
|
||||
# The SDK drops a response stream that ends without a JSON-RPC reply, so nothing else
|
||||
# ever fails the request.
|
||||
session_ctx: Final = ClientSession(
|
||||
session_ctx: Final = TolerantClientSession(
|
||||
read_stream,
|
||||
write_stream,
|
||||
read_timeout_seconds=timedelta(seconds=self.timeout),
|
||||
|
|
|
|||
108
litellm/experimental_mcp_client/tolerant_result.py
Normal file
108
litellm/experimental_mcp_client/tolerant_result.py
Normal file
|
|
@ -0,0 +1,108 @@
|
|||
"""
|
||||
Tolerant parsing of ``tools/call`` results from non-spec-compliant upstream MCP servers.
|
||||
|
||||
Some upstream MCP servers emit content blocks that fail SDK-side validation, e.g. an
|
||||
``EmbeddedResource`` whose ``uri`` is a relative path rather than a URI, or ``text/*``
|
||||
content shipped as a base64 ``blob`` instead of ``text``. The stock ``ClientSession``
|
||||
lets a single such block fail the entire ``tools/call`` result with a ``ValidationError``,
|
||||
discarding content the caller could otherwise use (litellm's own ``MCPClient.call_tool``
|
||||
degrades that into an opaque ``isError`` text block containing the pydantic traceback).
|
||||
Blocks that fail validation are degraded to text instead, so the rest of the result
|
||||
survives.
|
||||
"""
|
||||
|
||||
import base64
|
||||
import json
|
||||
from datetime import timedelta
|
||||
from typing import Any, Final, Mapping, Sequence
|
||||
|
||||
from mcp import ClientSession, types
|
||||
from mcp.shared.session import ProgressFnT
|
||||
from mcp.types import CallToolResult as MCPCallToolResult
|
||||
from pydantic import ValidationError, model_validator
|
||||
|
||||
|
||||
def _block_is_valid(block: object) -> bool:
|
||||
try:
|
||||
MCPCallToolResult.model_validate({"content": [block]})
|
||||
except ValidationError:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _resource_text(resource: Mapping[str, Any]) -> str | None:
|
||||
text: Final = resource.get("text")
|
||||
if isinstance(text, str):
|
||||
return text
|
||||
blob: Final = resource.get("blob")
|
||||
mime_type: Final = resource.get("mimeType")
|
||||
if not isinstance(blob, str) or not isinstance(mime_type, str) or not mime_type.startswith("text/"):
|
||||
return None
|
||||
try:
|
||||
return base64.b64decode(blob, validate=True).decode("utf-8")
|
||||
except (ValueError, UnicodeDecodeError):
|
||||
return None
|
||||
|
||||
|
||||
def _as_text_block(block: object) -> dict[str, Any]:
|
||||
if isinstance(block, Mapping):
|
||||
resource: Final = block.get("resource")
|
||||
if isinstance(resource, Mapping):
|
||||
text: Final = _resource_text(resource)
|
||||
if text is not None:
|
||||
return {"type": "text", "text": text}
|
||||
return {"type": "text", "text": json.dumps(block, default=str)}
|
||||
|
||||
|
||||
class TolerantCallToolResult(MCPCallToolResult):
|
||||
"""``CallToolResult`` that degrades invalid content blocks to text instead of raising."""
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def _degrade_invalid_content_blocks(cls, data: object) -> object:
|
||||
if not isinstance(data, Mapping):
|
||||
return data
|
||||
content: Final = data.get("content")
|
||||
if not isinstance(content, Sequence) or isinstance(content, (str, bytes)):
|
||||
return data
|
||||
if all(_block_is_valid(block) for block in content):
|
||||
return data
|
||||
return {
|
||||
**data,
|
||||
"content": [block if _block_is_valid(block) else _as_text_block(block) for block in content],
|
||||
}
|
||||
|
||||
|
||||
class TolerantClientSession(ClientSession):
|
||||
"""``ClientSession`` that parses ``tools/call`` results with ``TolerantCallToolResult``.
|
||||
|
||||
``ClientSession.call_tool`` hardcodes ``types.CallToolResult`` as the response type passed to
|
||||
``send_request`` with no seam to override just that argument, so this mirrors the mcp==1.28.1
|
||||
method body verbatim (including the private ``_validate_tool_result`` call) rather than calling
|
||||
``super().call_tool()``. A future mcp release changing that method's signature or behavior would
|
||||
not automatically propagate here.
|
||||
"""
|
||||
|
||||
async def call_tool(
|
||||
self,
|
||||
name: str,
|
||||
arguments: dict[str, Any] | None = None,
|
||||
read_timeout_seconds: timedelta | None = None,
|
||||
progress_callback: ProgressFnT | None = None,
|
||||
*,
|
||||
meta: dict[str, Any] | None = None,
|
||||
) -> types.CallToolResult:
|
||||
request_meta: Final = types.RequestParams.Meta(**meta) if meta is not None else None
|
||||
result: Final = await self.send_request(
|
||||
types.ClientRequest(
|
||||
types.CallToolRequest(
|
||||
params=types.CallToolRequestParams(name=name, arguments=arguments, _meta=request_meta),
|
||||
)
|
||||
),
|
||||
TolerantCallToolResult,
|
||||
request_read_timeout_seconds=read_timeout_seconds,
|
||||
progress_callback=progress_callback,
|
||||
)
|
||||
if not result.isError:
|
||||
await self._validate_tool_result(name, result)
|
||||
return result
|
||||
|
|
@ -78,7 +78,7 @@ class TestMCPClient:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.experimental_mcp_client.client.stdio_client")
|
||||
@patch("litellm.experimental_mcp_client.client.ClientSession")
|
||||
@patch("litellm.experimental_mcp_client.client.TolerantClientSession")
|
||||
async def test_mcp_client_stdio_connect_success(self, mock_session, mock_stdio_client):
|
||||
"""Test successful stdio connection"""
|
||||
# Setup mocks - create proper async context manager
|
||||
|
|
@ -130,7 +130,7 @@ class TestMCPClient:
|
|||
mock_streamable_http_client.return_value = mock_http_ctx
|
||||
|
||||
# Mock the session
|
||||
with patch("litellm.experimental_mcp_client.client.ClientSession") as mock_session:
|
||||
with patch("litellm.experimental_mcp_client.client.TolerantClientSession") as mock_session:
|
||||
mock_session_instance = AsyncMock()
|
||||
mock_session_instance.initialize = AsyncMock()
|
||||
mock_session_ctx = AsyncMock()
|
||||
|
|
@ -176,7 +176,7 @@ class TestMCPClient:
|
|||
mock_sse_client.return_value = mock_sse_ctx
|
||||
|
||||
# Mock the session
|
||||
with patch("litellm.experimental_mcp_client.client.ClientSession") as mock_session:
|
||||
with patch("litellm.experimental_mcp_client.client.TolerantClientSession") as mock_session:
|
||||
mock_session_instance = AsyncMock()
|
||||
mock_session_instance.initialize = AsyncMock()
|
||||
mock_session_ctx = AsyncMock()
|
||||
|
|
@ -228,7 +228,7 @@ class TestMCPClient:
|
|||
mock_streamable_http_client.return_value = mock_http_ctx
|
||||
|
||||
# Mock the session
|
||||
with patch("litellm.experimental_mcp_client.client.ClientSession") as mock_session:
|
||||
with patch("litellm.experimental_mcp_client.client.TolerantClientSession") as mock_session:
|
||||
mock_session_instance = AsyncMock()
|
||||
mock_session_instance.initialize = AsyncMock()
|
||||
mock_session_ctx = AsyncMock()
|
||||
|
|
@ -368,7 +368,7 @@ class TestMCPClientInstructionsCapture:
|
|||
assert client._last_initialize_instructions is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.experimental_mcp_client.client.ClientSession")
|
||||
@patch("litellm.experimental_mcp_client.client.TolerantClientSession")
|
||||
async def test_captures_instructions_from_initialize(self, mock_session_cls):
|
||||
"""Instructions from upstream initialize() are captured and stripped."""
|
||||
client = MCPClient(
|
||||
|
|
@ -397,7 +397,7 @@ class TestMCPClientInstructionsCapture:
|
|||
assert client._last_initialize_instructions == "upstream says hello"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.experimental_mcp_client.client.ClientSession")
|
||||
@patch("litellm.experimental_mcp_client.client.TolerantClientSession")
|
||||
async def test_none_instructions_stays_none(self, mock_session_cls):
|
||||
"""When upstream returns no instructions the field stays None."""
|
||||
client = MCPClient(
|
||||
|
|
@ -487,7 +487,7 @@ class TestExecuteSessionOperationSurfacesTransportError:
|
|||
return transport_ctx
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.experimental_mcp_client.client.ClientSession")
|
||||
@patch("litellm.experimental_mcp_client.client.TolerantClientSession")
|
||||
async def test_surfaces_connect_error_over_cancelled(self, mock_session_cls):
|
||||
client = MCPClient(server_url="http://example.com/mcp", transport_type="http")
|
||||
self._make_session(
|
||||
|
|
@ -504,7 +504,7 @@ class TestExecuteSessionOperationSurfacesTransportError:
|
|||
await client._execute_session_operation(transport_ctx, _op)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.experimental_mcp_client.client.ClientSession")
|
||||
@patch("litellm.experimental_mcp_client.client.TolerantClientSession")
|
||||
async def test_genuine_cancellation_is_not_replaced(self, mock_session_cls):
|
||||
client = MCPClient(server_url="http://example.com/mcp", transport_type="http")
|
||||
self._make_session(mock_session_cls, AsyncMock(side_effect=asyncio.CancelledError()))
|
||||
|
|
@ -517,7 +517,7 @@ class TestExecuteSessionOperationSurfacesTransportError:
|
|||
await client._execute_session_operation(transport_ctx, _op)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.experimental_mcp_client.client.ClientSession")
|
||||
@patch("litellm.experimental_mcp_client.client.TolerantClientSession")
|
||||
async def test_cleanup_error_after_success_is_swallowed(self, mock_session_cls):
|
||||
client = MCPClient(server_url="http://example.com/mcp", transport_type="http")
|
||||
init_result = MagicMock()
|
||||
|
|
@ -731,8 +731,9 @@ class _ScriptedUpstream:
|
|||
error, the shape an upstream application uses to report its own failure.
|
||||
"""
|
||||
|
||||
def __init__(self, tools_list_error: ErrorData | None = None):
|
||||
def __init__(self, tools_list_error: ErrorData | None = None, tools_call_result: dict | None = None):
|
||||
self._tools_list_error = tools_list_error
|
||||
self._tools_call_result = tools_call_result
|
||||
self._to_client_tx, self._to_client_rx = anyio.create_memory_object_stream(10)
|
||||
self._from_client_tx, self._from_client_rx = anyio.create_memory_object_stream(10)
|
||||
self._task_group = None
|
||||
|
|
@ -769,20 +770,86 @@ class _ScriptedUpstream:
|
|||
)
|
||||
elif method == "tools/list" and self._tools_list_error is not None:
|
||||
await self._send(JSONRPCError(jsonrpc="2.0", id=request.id, error=self._tools_list_error))
|
||||
elif method == "tools/call" and self._tools_call_result is not None:
|
||||
# Sent as a raw dict, not built from CallToolResult, so a non-compliant upstream's
|
||||
# actual wire bytes reach the client exactly as they would over a real connection.
|
||||
await self._send(JSONRPCResponse(jsonrpc="2.0", id=request.id, result=self._tools_call_result))
|
||||
|
||||
|
||||
class _ScriptedClient(MCPClient):
|
||||
"""An MCPClient whose transport is a scripted in-memory upstream instead of a real connection,
|
||||
so the real ``ClientSession`` and its real timeout machinery are what run."""
|
||||
|
||||
def __init__(self, *, timeout: float, tools_list_error: ErrorData | None = None):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
timeout: float,
|
||||
tools_list_error: ErrorData | None = None,
|
||||
tools_call_result: dict | None = None,
|
||||
):
|
||||
super().__init__(server_url="http://upstream.local/mcp", timeout=timeout)
|
||||
self._upstream = _ScriptedUpstream(tools_list_error=tools_list_error)
|
||||
self._upstream = _ScriptedUpstream(tools_list_error=tools_list_error, tools_call_result=tools_call_result)
|
||||
|
||||
def _create_transport_context(self):
|
||||
return self._upstream, None
|
||||
|
||||
|
||||
# A real payload captured from an upstream Azure DevOps MCP server's ``repo_file`` tool: an
|
||||
# ``EmbeddedResource`` with a repo-relative ``uri`` (the MCP spec requires a URI with a scheme) and
|
||||
# ``text/plain`` content shipped as a base64 ``blob`` instead of ``text``. mcp==1.28.1's
|
||||
# ``CallToolResult.model_validate`` raises a 14-error ``ValidationError`` on this exact payload.
|
||||
_MALFORMED_TOOLS_CALL_RESULT = {
|
||||
"content": [
|
||||
{
|
||||
"type": "resource",
|
||||
"resource": {
|
||||
"uri": "/templates/ev.job/template/docker-compose.yml",
|
||||
"mimeType": "text/plain",
|
||||
"blob": "dmVyc2lvbjogJzIuNCcK",
|
||||
},
|
||||
}
|
||||
],
|
||||
"isError": False,
|
||||
}
|
||||
|
||||
|
||||
class TestCallToolToleratesMalformedContentBlocks:
|
||||
"""A single content block that fails MCP-SDK validation must not fail the whole ``tools/call``
|
||||
result. Driven through ``_ScriptedClient`` over real anyio streams and the real session class
|
||||
the client actually wires in, so a change to which class gets constructed (not just the
|
||||
validator logic in isolation) fails this test too.
|
||||
"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_relative_uri_resource_degrades_to_text_instead_of_raising(self):
|
||||
from mcp.types import CallToolRequestParams
|
||||
|
||||
client = _ScriptedClient(timeout=5, tools_call_result=_MALFORMED_TOOLS_CALL_RESULT)
|
||||
params = CallToolRequestParams(name="repo_file", arguments={"action": "get_content"})
|
||||
|
||||
result = await asyncio.wait_for(client.call_tool(params, raise_on_error=True), timeout=10)
|
||||
|
||||
assert result.isError is False
|
||||
assert len(result.content) == 1
|
||||
assert result.content[0].type == "text"
|
||||
assert result.content[0].text == "version: '2.4'\n"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_well_formed_result_is_unaffected(self):
|
||||
from mcp.types import CallToolRequestParams
|
||||
|
||||
well_formed = {"content": [{"type": "text", "text": "ok"}], "isError": False}
|
||||
client = _ScriptedClient(timeout=5, tools_call_result=well_formed)
|
||||
params = CallToolRequestParams(name="some_tool", arguments={})
|
||||
|
||||
result = await asyncio.wait_for(client.call_tool(params, raise_on_error=True), timeout=10)
|
||||
|
||||
assert result.isError is False
|
||||
assert len(result.content) == 1
|
||||
assert result.content[0].type == "text"
|
||||
assert result.content[0].text == "ok"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tools_fails_on_its_own_timeout_when_the_upstream_never_answers():
|
||||
"""An upstream that accepts the request and never answers must fail the client's own timeout.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue