diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 7bd0a847ad8..3b1ae016fd1 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -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), diff --git a/litellm/experimental_mcp_client/tolerant_result.py b/litellm/experimental_mcp_client/tolerant_result.py new file mode 100644 index 00000000000..e25612f8a0f --- /dev/null +++ b/litellm/experimental_mcp_client/tolerant_result.py @@ -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 diff --git a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py index 7beb1c43a94..350be9abafa 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -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.