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:
Aliaksei Venski 2026-08-14 14:33:55 +02:00
parent 423b791ee0
commit 3b7d9e74cb
3 changed files with 189 additions and 13 deletions

View file

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

View 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

View file

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