fix: tolerate empty MCP list results

This commit is contained in:
冯基魁 2026-06-08 14:45:06 +08:00
parent aaf1e2444b
commit b9b919d1a1
2 changed files with 81 additions and 5 deletions

View file

@ -34,14 +34,21 @@ except ImportError:
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
from mcp.types import CallToolResult as MCPCallToolResult
from mcp.types import (
ClientRequest,
GetPromptRequestParams,
GetPromptResult,
ListPromptsRequest,
ListPromptsResult,
ListResourceTemplatesRequest,
ListResourceTemplatesResult,
ListResourcesRequest,
ListResourcesResult,
Prompt,
ResourceTemplate,
TextContent,
)
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl
from pydantic import AnyUrl, Field
from litellm._logging import verbose_logger
from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_NPM_CACHE_DIR
from litellm.llms.custom_httpx.http_handler import get_ssl_configuration
@ -84,6 +91,18 @@ def _first_non_cancelled_cause(exc: BaseException) -> Optional[BaseException]:
TSessionResult = TypeVar("TSessionResult")
class _LenientListPromptsResult(ListPromptsResult):
prompts: List[Prompt] = Field(default_factory=list)
class _LenientListResourcesResult(ListResourcesResult):
resources: List[Resource] = Field(default_factory=list)
class _LenientListResourceTemplatesResult(ListResourceTemplatesResult):
resourceTemplates: List[ResourceTemplate] = Field(default_factory=list)
class MCPSigV4Auth(httpx.Auth):
"""
httpx Auth class that signs each request with AWS SigV4.
@ -628,7 +647,10 @@ class MCPClient:
)
async def _list_prompts_operation(session: ClientSession):
return await session.list_prompts()
return await session.send_request(
ClientRequest(ListPromptsRequest()),
_LenientListPromptsResult,
)
try:
result = await self.run_with_session(_list_prompts_operation)
@ -713,7 +735,10 @@ class MCPClient:
)
async def _list_resources_operation(session: ClientSession):
return await session.list_resources()
return await session.send_request(
ClientRequest(ListResourcesRequest()),
_LenientListResourcesResult,
)
try:
result = await self.run_with_session(_list_resources_operation)
@ -751,7 +776,10 @@ class MCPClient:
)
async def _list_resource_templates_operation(session: ClientSession):
return await session.list_resource_templates()
return await session.send_request(
ClientRequest(ListResourceTemplatesRequest()),
_LenientListResourceTemplatesResult,
)
try:
result = await self.run_with_session(_list_resource_templates_operation)

View file

@ -1,6 +1,5 @@
import asyncio
import os
import ssl
import sys
from unittest.mock import AsyncMock, MagicMock, patch
@ -15,9 +14,32 @@ from litellm.experimental_mcp_client.client import (
MCPClient,
_first_non_cancelled_cause,
)
from mcp.types import (
ListPromptsRequest,
ListResourceTemplatesRequest,
ListResourcesRequest,
)
from litellm.types.mcp import MCPAuth, MCPStdioConfig, MCPTransport
class _MissingListResultSession:
def __init__(self):
self.requests = []
async def send_request(self, request, result_type):
self.requests.append((request, result_type))
return result_type.model_validate({})
async def list_prompts(self):
raise AssertionError("strict list_prompts should not be used")
async def list_resources(self):
raise AssertionError("strict list_resources should not be used")
async def list_resource_templates(self):
raise AssertionError("strict list_resource_templates should not be used")
class _FakeExceptionGroup(Exception):
"""Duck-typed stand-in for an anyio/builtin ExceptionGroup.
@ -48,6 +70,32 @@ class TestMCPClient:
assert client.stdio_config.get("command") == "python"
assert client.stdio_config.get("args") == ["-m", "my_mcp_server"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
("client_method", "request_type"),
[
("list_prompts", ListPromptsRequest),
("list_resources", ListResourcesRequest),
("list_resource_templates", ListResourceTemplatesRequest),
],
)
async def test_list_methods_tolerate_missing_result_collections(
self, client_method, request_type
):
client = MCPClient(server_url="http://example.com/mcp", transport_type="http")
session = _MissingListResultSession()
async def _run_operation(operation):
return await operation(session)
client.run_with_session = AsyncMock(side_effect=_run_operation)
result = await getattr(client, client_method)()
assert result == []
assert len(session.requests) == 1
assert isinstance(session.requests[0][0].root, request_type)
@pytest.mark.asyncio
async def test_mcp_client_stdio_connect_error(self):
"""Test MCP client stdio connection error handling"""