mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(mcp): fix CI failures from the tolerant-session change
Three separate issues surfaced by CI:
ruff: Mapping/Sequence must come from collections.abc, not typing (UP035), and
the now-unused N815 noqa on _ResourcePayload.mimeType was flagged as dead
(RUF100) since that rule isn't enabled in this repo's config.
tests/mcp_tests/test_mcp_client_unit.py: three tests patch ClientSession via
patch.object(mcp_client_module, "ClientSession") to intercept construction, a
different idiom from the string-form patch("...client.ClientSession") already
updated in the other test file, and were missed in that earlier pass.
Construction now goes through TolerantClientSession, so update these three too.
tests/test_litellm/experimental_mcp_client/test_mcp_client.py: the two new
scripted-upstream tests timed out. TolerantClientSession.call_tool calls the
private _validate_tool_result, which fetches the tool list (once per tool
name) to check output-schema compliance. The scripted upstream only answers
tools/list when explicitly configured to, by design, for the pre-existing
timeout tests, so the new tests' tools/call response was received but the
follow-up tools/list request went unanswered forever. Give both new tests an
empty tools/list response so that internal call resolves immediately.
This commit is contained in:
parent
50a6c70a2f
commit
6ba50e388c
3 changed files with 38 additions and 9 deletions
|
|
@ -13,8 +13,9 @@ survives.
|
|||
|
||||
import base64
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import timedelta
|
||||
from typing import Any, Final, Literal, Mapping, Sequence, TypedDict
|
||||
from typing import Any, Final, Literal, TypedDict
|
||||
|
||||
from mcp import ClientSession, types
|
||||
from mcp.shared.session import ProgressFnT
|
||||
|
|
@ -38,7 +39,7 @@ class _ResourcePayload(BaseModel):
|
|||
|
||||
text: str | None = None
|
||||
blob: str | None = None
|
||||
mimeType: str | None = None # noqa: N815 # mirrors the MCP wire field name
|
||||
mimeType: str | None = None
|
||||
|
||||
|
||||
class _TextContentBlock(TypedDict):
|
||||
|
|
|
|||
|
|
@ -117,7 +117,7 @@ class TestMCPClientUnitTests:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@patch.object(mcp_client_module, "streamable_http_client")
|
||||
@patch.object(mcp_client_module, "ClientSession")
|
||||
@patch.object(mcp_client_module, "TolerantClientSession")
|
||||
async def test_run_with_session(self, mock_session_class, mock_transport):
|
||||
"""Test run_with_session establishes session with auth headers."""
|
||||
# Setup mocks
|
||||
|
|
@ -152,7 +152,7 @@ class TestMCPClientUnitTests:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@patch.object(mcp_client_module, "streamable_http_client")
|
||||
@patch.object(mcp_client_module, "ClientSession")
|
||||
@patch.object(mcp_client_module, "TolerantClientSession")
|
||||
async def test_list_tools(self, mock_session_class, mock_transport):
|
||||
"""Test listing tools from the server."""
|
||||
# Setup mocks
|
||||
|
|
@ -190,7 +190,7 @@ class TestMCPClientUnitTests:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@patch.object(mcp_client_module, "streamable_http_client")
|
||||
@patch.object(mcp_client_module, "ClientSession")
|
||||
@patch.object(mcp_client_module, "TolerantClientSession")
|
||||
async def test_call_tool(self, mock_session_class, mock_transport):
|
||||
"""Test calling a tool."""
|
||||
from mcp.types import CallToolRequestParams
|
||||
|
|
|
|||
|
|
@ -731,9 +731,15 @@ class _ScriptedUpstream:
|
|||
error, the shape an upstream application uses to report its own failure.
|
||||
"""
|
||||
|
||||
def __init__(self, tools_list_error: ErrorData | None = None, tools_call_result: dict | None = None):
|
||||
def __init__(
|
||||
self,
|
||||
tools_list_error: ErrorData | None = None,
|
||||
tools_call_result: dict | None = None,
|
||||
tools_list_result: dict | None = None,
|
||||
):
|
||||
self._tools_list_error = tools_list_error
|
||||
self._tools_call_result = tools_call_result
|
||||
self._tools_list_result = tools_list_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
|
||||
|
|
@ -770,6 +776,8 @@ 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/list" and self._tools_list_result is not None:
|
||||
await self._send(JSONRPCResponse(jsonrpc="2.0", id=request.id, result=self._tools_list_result))
|
||||
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.
|
||||
|
|
@ -786,9 +794,14 @@ class _ScriptedClient(MCPClient):
|
|||
timeout: float,
|
||||
tools_list_error: ErrorData | None = None,
|
||||
tools_call_result: dict | None = None,
|
||||
tools_list_result: dict | None = None,
|
||||
):
|
||||
super().__init__(server_url="http://upstream.local/mcp", timeout=timeout)
|
||||
self._upstream = _ScriptedUpstream(tools_list_error=tools_list_error, tools_call_result=tools_call_result)
|
||||
self._upstream = _ScriptedUpstream(
|
||||
tools_list_error=tools_list_error,
|
||||
tools_call_result=tools_call_result,
|
||||
tools_list_result=tools_list_result,
|
||||
)
|
||||
|
||||
def _create_transport_context(self):
|
||||
return self._upstream, None
|
||||
|
|
@ -813,6 +826,13 @@ _MALFORMED_TOOLS_CALL_RESULT = {
|
|||
}
|
||||
|
||||
|
||||
# call_tool's private _validate_tool_result fetches the tool list (once per tool name, cached
|
||||
# after) to check the result's structuredContent against the tool's output schema. An empty list
|
||||
# means the tool isn't found, which _validate_tool_result treats as "nothing to validate against"
|
||||
# rather than an error, so it doesn't affect the content-degradation behavior under test here.
|
||||
_EMPTY_TOOLS_LIST_RESULT = {"tools": []}
|
||||
|
||||
|
||||
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
|
||||
|
|
@ -824,7 +844,11 @@ class TestCallToolToleratesMalformedContentBlocks:
|
|||
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)
|
||||
client = _ScriptedClient(
|
||||
timeout=5,
|
||||
tools_call_result=_MALFORMED_TOOLS_CALL_RESULT,
|
||||
tools_list_result=_EMPTY_TOOLS_LIST_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)
|
||||
|
|
@ -839,7 +863,11 @@ class TestCallToolToleratesMalformedContentBlocks:
|
|||
from mcp.types import CallToolRequestParams
|
||||
|
||||
well_formed = {"content": [{"type": "text", "text": "ok"}], "isError": False}
|
||||
client = _ScriptedClient(timeout=5, tools_call_result=well_formed)
|
||||
client = _ScriptedClient(
|
||||
timeout=5,
|
||||
tools_call_result=well_formed,
|
||||
tools_list_result=_EMPTY_TOOLS_LIST_RESULT,
|
||||
)
|
||||
params = CallToolRequestParams(name="some_tool", arguments={})
|
||||
|
||||
result = await asyncio.wait_for(client.call_tool(params, raise_on_error=True), timeout=10)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue