mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix: tolerate empty MCP list results
This commit is contained in:
parent
aaf1e2444b
commit
b9b919d1a1
2 changed files with 81 additions and 5 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue