mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(mcp): degrade buggy pagination to partial results and bound the preview walk
A repeated nextCursor now returns the tools collected so far instead of discarding every page with a RuntimeError, an empty-string cursor is treated as terminal, load_mcp_tools shares the same pagination walk instead of returning only the first page, and the tools/list preview is bounded by the listing timeout instead of only the per-request timeout times the page cap
This commit is contained in:
parent
772dba1d11
commit
74a640d430
6 changed files with 264 additions and 49 deletions
|
|
@ -39,7 +39,6 @@ from mcp.types import CallToolResult as MCPCallToolResult
|
|||
from mcp.types import (
|
||||
GetPromptRequestParams,
|
||||
GetPromptResult,
|
||||
PaginatedRequestParams,
|
||||
Prompt,
|
||||
ResourceTemplate,
|
||||
TextContent,
|
||||
|
|
@ -48,7 +47,8 @@ from mcp.types import Tool as MCPTool
|
|||
from pydantic import AnyUrl
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_NPM_CACHE_DIR, MCP_TOOL_LISTING_MAX_PAGES
|
||||
from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_NPM_CACHE_DIR
|
||||
from litellm.experimental_mcp_client.tools import list_tools_with_pagination
|
||||
from litellm.llms.custom_httpx.http_handler import get_ssl_configuration
|
||||
from litellm.types.llms.custom_http import VerifyTypes
|
||||
from litellm.types.mcp import (
|
||||
|
|
@ -604,42 +604,10 @@ class MCPClient:
|
|||
"""
|
||||
verbose_logger.debug("MCP client listing tools from %s", self.server_url or "stdio")
|
||||
|
||||
async def _list_tools_operation(session: ClientSession) -> list[MCPTool]:
|
||||
tools: list[MCPTool] = []
|
||||
cursor: Optional[str] = None
|
||||
pages_fetched = 0
|
||||
seen_cursors: set[str] = set()
|
||||
|
||||
while True:
|
||||
result = (
|
||||
await session.list_tools()
|
||||
if cursor is None
|
||||
else await session.list_tools(params=PaginatedRequestParams(cursor=cursor))
|
||||
)
|
||||
pages_fetched += 1
|
||||
tools.extend(result.tools)
|
||||
|
||||
next_cursor = getattr(result, "nextCursor", None)
|
||||
if not isinstance(next_cursor, str):
|
||||
return tools
|
||||
if next_cursor in seen_cursors:
|
||||
raise RuntimeError(
|
||||
f"MCP server returned a repeated tools/list cursor while listing tools: {next_cursor}"
|
||||
)
|
||||
if pages_fetched >= MCP_TOOL_LISTING_MAX_PAGES:
|
||||
verbose_logger.warning(
|
||||
"MCP server tools/list pagination exceeded the maximum "
|
||||
f"of {MCP_TOOL_LISTING_MAX_PAGES} pages while listing tools; "
|
||||
f"returning {len(tools)} tools collected so far"
|
||||
)
|
||||
return tools
|
||||
seen_cursors.add(next_cursor)
|
||||
cursor = next_cursor
|
||||
|
||||
try:
|
||||
tools: Final = await self.run_with_session(_list_tools_operation, quiet_on_error=raise_on_error)
|
||||
tools: Final = await self.run_with_session(list_tools_with_pagination, quiet_on_error=raise_on_error)
|
||||
tool_count: Final = len(tools)
|
||||
tool_names: Final = [tool.name for tool in tools]
|
||||
tool_names: Final = tuple(tool.name for tool in tools)
|
||||
verbose_logger.info(
|
||||
"MCP client listed %s tools from %s: %s", tool_count, self.server_url or "stdio", tool_names
|
||||
)
|
||||
|
|
|
|||
|
|
@ -4,11 +4,14 @@ from typing import Final, Literal
|
|||
from mcp import ClientSession
|
||||
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
|
||||
from mcp.types import CallToolResult as MCPCallToolResult
|
||||
from mcp.types import PaginatedRequestParams
|
||||
from mcp.types import Tool as MCPTool
|
||||
from openai.types.chat import ChatCompletionToolParam
|
||||
from openai.types.responses.function_tool_param import FunctionToolParam
|
||||
from openai.types.shared_params.function_definition import FunctionDefinition
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import MCP_TOOL_LISTING_MAX_PAGES
|
||||
from litellm.types.llms.anthropic import AnthropicMessagesTool
|
||||
from litellm.types.utils import ChatCompletionMessageToolCall
|
||||
|
||||
|
|
@ -90,6 +93,45 @@ def transform_mcp_tool_to_anthropic_tool(mcp_tool: MCPTool) -> AnthropicMessages
|
|||
)
|
||||
|
||||
|
||||
async def list_tools_with_pagination(session: ClientSession) -> list[MCPTool]: # mutable-ok: list return contract
|
||||
"""Collect tools from every tools/list page by following nextCursor.
|
||||
|
||||
Stops and returns the tools collected so far when the upstream repeats a
|
||||
cursor or the page cap is reached, so a buggy upstream yields a partial
|
||||
catalog instead of an error.
|
||||
"""
|
||||
tools: Final[list[MCPTool]] = [] # mutable-ok: accumulates each page's tools
|
||||
seen_cursors: Final[set[str]] = set() # mutable-ok: guards against cursor loops
|
||||
cursor: str | None = None # rebind-ok: advances to each page's nextCursor
|
||||
|
||||
for _ in range(MCP_TOOL_LISTING_MAX_PAGES):
|
||||
result = (
|
||||
await session.list_tools()
|
||||
if cursor is None
|
||||
else await session.list_tools(params=PaginatedRequestParams(cursor=cursor))
|
||||
)
|
||||
tools.extend(result.tools)
|
||||
|
||||
next_cursor = getattr(result, "nextCursor", None)
|
||||
if not isinstance(next_cursor, str) or not next_cursor:
|
||||
return tools
|
||||
if next_cursor in seen_cursors:
|
||||
verbose_logger.warning(
|
||||
"MCP server repeated a tools/list cursor while listing tools; returning %s tools collected so far",
|
||||
len(tools),
|
||||
)
|
||||
return tools
|
||||
seen_cursors.add(next_cursor)
|
||||
cursor = next_cursor
|
||||
|
||||
verbose_logger.warning(
|
||||
"MCP server tools/list pagination exceeded the maximum of %s pages; returning %s tools collected so far",
|
||||
MCP_TOOL_LISTING_MAX_PAGES,
|
||||
len(tools),
|
||||
)
|
||||
return tools
|
||||
|
||||
|
||||
async def load_mcp_tools(
|
||||
session: ClientSession, format: Literal["mcp", "openai"] = "mcp"
|
||||
) -> list[MCPTool] | list[ChatCompletionToolParam]:
|
||||
|
|
@ -103,10 +145,12 @@ async def load_mcp_tools(
|
|||
|
||||
If format is set to "openai", the tools are converted to OpenAI API compatible tools.
|
||||
"""
|
||||
tools: Final = await session.list_tools()
|
||||
tools: Final = await list_tools_with_pagination(session)
|
||||
if format == "openai":
|
||||
return [transform_mcp_tool_to_openai_tool(mcp_tool=tool) for tool in tools.tools]
|
||||
return tools.tools
|
||||
return [ # mutable-ok: public API returns a list
|
||||
transform_mcp_tool_to_openai_tool(mcp_tool=tool) for tool in tools
|
||||
]
|
||||
return tools
|
||||
|
||||
|
||||
########################################################
|
||||
|
|
|
|||
|
|
@ -4,10 +4,12 @@ from collections.abc import Awaitable, Callable, Mapping
|
|||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
|
||||
import anyio
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_TOOL_LISTING_TIMEOUT
|
||||
from litellm.exceptions import (
|
||||
BlockedPiiEntityError,
|
||||
GuardrailRaisedException,
|
||||
|
|
@ -86,8 +88,6 @@ def _connection_error_message(exc: BaseException) -> str:
|
|||
|
||||
|
||||
if MCP_AVAILABLE:
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES,
|
||||
global_mcp_server_manager,
|
||||
|
|
@ -1402,7 +1402,24 @@ if MCP_AVAILABLE:
|
|||
oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers)
|
||||
|
||||
async def _list_tools_operation(client):
|
||||
list_tools_result: Final[list[MCPTool]] = await client.list_tools(raise_on_error=True)
|
||||
# Bound the whole pagination walk: without this the preview is limited only by the
|
||||
# per-request timeout times the page cap. max() keeps the pre-pagination guarantee
|
||||
# that a single slow page within the client timeout still succeeds.
|
||||
listing_deadline: Final = max(MCP_CLIENT_TIMEOUT, MCP_TOOL_LISTING_TIMEOUT)
|
||||
list_tools_result = None # rebind-ok: set inside the timeout scope below
|
||||
with anyio.move_on_after(listing_deadline):
|
||||
list_tools_result = await client.list_tools(raise_on_error=True)
|
||||
if list_tools_result is None:
|
||||
verbose_logger.warning(
|
||||
"MCP tools/list preview timed out after %s seconds while paginating upstream tools",
|
||||
listing_deadline,
|
||||
)
|
||||
return { # mutable-ok: error response payload
|
||||
"status": "error",
|
||||
"error": True,
|
||||
"message": f"Timed out listing tools after {listing_deadline} seconds. "
|
||||
"The MCP server may be responding slowly or paginating excessively.",
|
||||
}
|
||||
model_dumped_tools: Final[list[dict]] = [tool.model_dump() for tool in list_tools_result]
|
||||
return {
|
||||
"tools": model_dumped_tools,
|
||||
|
|
|
|||
|
|
@ -238,7 +238,7 @@ class TestMCPClientUnitTests:
|
|||
monkeypatch,
|
||||
):
|
||||
"""Test listing tools returns accumulated tools if an upstream keeps returning new cursors."""
|
||||
monkeypatch.setattr(mcp_client_module, "MCP_TOOL_LISTING_MAX_PAGES", 2, raising=False)
|
||||
monkeypatch.setattr("litellm.experimental_mcp_client.tools.MCP_TOOL_LISTING_MAX_PAGES", 2)
|
||||
|
||||
mock_transport_ctx = AsyncMock()
|
||||
mock_transport.return_value = mock_transport_ctx
|
||||
|
|
@ -273,12 +273,12 @@ class TestMCPClientUnitTests:
|
|||
@pytest.mark.asyncio
|
||||
@patch.object(mcp_client_module, "streamable_http_client")
|
||||
@patch.object(mcp_client_module, "ClientSession")
|
||||
async def test_list_tools_raises_on_repeated_next_cursor(
|
||||
async def test_list_tools_stops_on_repeated_next_cursor(
|
||||
self,
|
||||
mock_session_class,
|
||||
mock_transport,
|
||||
):
|
||||
"""Test listing tools fails if an upstream repeats a cursor."""
|
||||
"""Test listing tools returns collected tools when an upstream repeats a cursor."""
|
||||
mock_transport_ctx = AsyncMock()
|
||||
mock_transport.return_value = mock_transport_ctx
|
||||
mock_transport_instance = MagicMock()
|
||||
|
|
@ -290,14 +290,85 @@ class TestMCPClientUnitTests:
|
|||
mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance)
|
||||
|
||||
mock_session_instance.list_tools.side_effect = [
|
||||
ListToolsResult(tools=[], nextCursor="same-cursor"),
|
||||
ListToolsResult(tools=[], nextCursor="same-cursor"),
|
||||
ListToolsResult(
|
||||
tools=[MCPTool(name="tool_0", description="Tool 0", inputSchema={})],
|
||||
nextCursor="same-cursor",
|
||||
),
|
||||
ListToolsResult(
|
||||
tools=[MCPTool(name="tool_1", description="Tool 1", inputSchema={})],
|
||||
nextCursor="same-cursor",
|
||||
),
|
||||
]
|
||||
|
||||
client = MCPClient("http://example.com")
|
||||
with pytest.raises(RuntimeError, match="repeated tools/list cursor"):
|
||||
await client.list_tools(raise_on_error=True)
|
||||
result = await client.list_tools(raise_on_error=True)
|
||||
|
||||
assert [tool.name for tool in result] == ["tool_0", "tool_1"]
|
||||
assert mock_session_instance.list_tools.call_count == 2
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch.object(mcp_client_module, "streamable_http_client")
|
||||
@patch.object(mcp_client_module, "ClientSession")
|
||||
async def test_list_tools_treats_empty_cursor_as_terminal(
|
||||
self,
|
||||
mock_session_class,
|
||||
mock_transport,
|
||||
):
|
||||
"""Test listing tools stops when an upstream returns an empty-string cursor."""
|
||||
mock_transport_ctx = AsyncMock()
|
||||
mock_transport.return_value = mock_transport_ctx
|
||||
mock_transport_instance = MagicMock()
|
||||
mock_transport_ctx.__aenter__ = AsyncMock(return_value=mock_transport_instance)
|
||||
|
||||
mock_session_ctx = AsyncMock()
|
||||
mock_session_class.return_value = mock_session_ctx
|
||||
mock_session_instance = AsyncMock()
|
||||
mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance)
|
||||
|
||||
mock_session_instance.list_tools.side_effect = [
|
||||
ListToolsResult(
|
||||
tools=[MCPTool(name="tool_0", description="Tool 0", inputSchema={})],
|
||||
nextCursor="",
|
||||
),
|
||||
]
|
||||
|
||||
client = MCPClient("http://example.com")
|
||||
result = await client.list_tools()
|
||||
|
||||
assert [tool.name for tool in result] == ["tool_0"]
|
||||
mock_session_instance.list_tools.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch.object(mcp_client_module, "streamable_http_client")
|
||||
@patch.object(mcp_client_module, "ClientSession")
|
||||
async def test_list_tools_swallows_mid_walk_error_without_raise_on_error(
|
||||
self,
|
||||
mock_session_class,
|
||||
mock_transport,
|
||||
):
|
||||
"""Test a mid-walk failure returns [] when raise_on_error is False."""
|
||||
mock_transport_ctx = AsyncMock()
|
||||
mock_transport.return_value = mock_transport_ctx
|
||||
mock_transport_instance = MagicMock()
|
||||
mock_transport_ctx.__aenter__ = AsyncMock(return_value=mock_transport_instance)
|
||||
|
||||
mock_session_ctx = AsyncMock()
|
||||
mock_session_class.return_value = mock_session_ctx
|
||||
mock_session_instance = AsyncMock()
|
||||
mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance)
|
||||
|
||||
mock_session_instance.list_tools.side_effect = [
|
||||
ListToolsResult(
|
||||
tools=[MCPTool(name="tool_0", description="Tool 0", inputSchema={})],
|
||||
nextCursor="page-2",
|
||||
),
|
||||
RuntimeError("transient upstream failure"),
|
||||
]
|
||||
|
||||
client = MCPClient("http://example.com")
|
||||
result = await client.list_tools()
|
||||
|
||||
assert result == []
|
||||
assert mock_session_instance.list_tools.call_count == 2
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ from mcp.types import (
|
|||
CallToolRequestParams,
|
||||
CallToolResult,
|
||||
ListToolsResult,
|
||||
PaginatedRequestParams,
|
||||
TextContent,
|
||||
)
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
|
@ -106,6 +107,39 @@ async def test_load_mcp_tools_openai_format(mock_session, mock_list_tools_result
|
|||
mock_session.list_tools.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio()
|
||||
async def test_load_mcp_tools_follows_pagination(mock_session):
|
||||
mock_session.list_tools.side_effect = [
|
||||
ListToolsResult(
|
||||
tools=[
|
||||
MCPTool(name="tool_a", description="a", inputSchema={}),
|
||||
MCPTool(name="tool_b", description="b", inputSchema={}),
|
||||
],
|
||||
nextCursor="page-2",
|
||||
),
|
||||
ListToolsResult(tools=[MCPTool(name="tool_c", description="c", inputSchema={})]),
|
||||
]
|
||||
result = await load_mcp_tools(mock_session, format="mcp")
|
||||
assert [tool.name for tool in result] == ["tool_a", "tool_b", "tool_c"]
|
||||
assert mock_session.list_tools.call_count == 2
|
||||
second_call_params = mock_session.list_tools.call_args_list[1].kwargs["params"]
|
||||
assert isinstance(second_call_params, PaginatedRequestParams)
|
||||
assert second_call_params.cursor == "page-2"
|
||||
|
||||
|
||||
@pytest.mark.asyncio()
|
||||
async def test_load_mcp_tools_openai_format_spans_pages(mock_session):
|
||||
mock_session.list_tools.side_effect = [
|
||||
ListToolsResult(
|
||||
tools=[MCPTool(name="tool_a", description="a", inputSchema={})],
|
||||
nextCursor="page-2",
|
||||
),
|
||||
ListToolsResult(tools=[MCPTool(name="tool_b", description="b", inputSchema={})]),
|
||||
]
|
||||
result = await load_mcp_tools(mock_session, format="openai")
|
||||
assert [t["function"]["name"] for t in result] == ["tool_a", "tool_b"]
|
||||
|
||||
|
||||
def test_get_function_arguments():
|
||||
# Test with string arguments
|
||||
function = {"arguments": '{"test": "value"}'}
|
||||
|
|
|
|||
|
|
@ -524,6 +524,87 @@ class TestTestToolsList:
|
|||
assert captured["oauth2_headers"] is None
|
||||
assert oauth_call_counter["count"] == 0
|
||||
|
||||
async def test_preview_tools_list_times_out_on_slow_pagination(self, monkeypatch):
|
||||
"""A preview whose upstream paginates past the listing deadline returns a
|
||||
timeout error instead of holding the request open."""
|
||||
monkeypatch.setattr(rest_endpoints, "MCP_CLIENT_TIMEOUT", 0.05, raising=False)
|
||||
monkeypatch.setattr(rest_endpoints, "MCP_TOOL_LISTING_TIMEOUT", 0.05, raising=False)
|
||||
|
||||
class SlowClient:
|
||||
async def list_tools(self, raise_on_error=False):
|
||||
await asyncio.sleep(1)
|
||||
return []
|
||||
|
||||
async def fake_execute(
|
||||
request,
|
||||
operation,
|
||||
mcp_auth_header=None,
|
||||
oauth2_headers=None,
|
||||
raw_headers=None,
|
||||
):
|
||||
return await operation(SlowClient())
|
||||
|
||||
monkeypatch.setattr(rest_endpoints, "_execute_with_mcp_client", fake_execute, raising=False)
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
request = _build_request()
|
||||
payload = NewMCPServerRequest(
|
||||
server_name="example",
|
||||
url="https://example.com",
|
||||
auth_type=MCPAuth.api_key,
|
||||
credentials={"auth_value": "secret-key"},
|
||||
)
|
||||
|
||||
result = await rest_endpoints.test_tools_list(
|
||||
request,
|
||||
payload,
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
assert result["status"] == "error"
|
||||
assert result["error"] is True
|
||||
assert "Timed out listing tools" in result["message"]
|
||||
|
||||
async def test_preview_tools_list_succeeds_within_deadline(self, monkeypatch):
|
||||
"""The preview timeout scope passes a fast listing through untouched."""
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
class QuickClient:
|
||||
async def list_tools(self, raise_on_error=False):
|
||||
return [MCPTool(name="quick_tool", description="q", inputSchema={})]
|
||||
|
||||
async def fake_execute(
|
||||
request,
|
||||
operation,
|
||||
mcp_auth_header=None,
|
||||
oauth2_headers=None,
|
||||
raw_headers=None,
|
||||
):
|
||||
return await operation(QuickClient())
|
||||
|
||||
monkeypatch.setattr(rest_endpoints, "_execute_with_mcp_client", fake_execute, raising=False)
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
request = _build_request()
|
||||
payload = NewMCPServerRequest(
|
||||
server_name="example",
|
||||
url="https://example.com",
|
||||
auth_type=MCPAuth.api_key,
|
||||
credentials={"auth_value": "secret-key"},
|
||||
)
|
||||
|
||||
result = await rest_endpoints.test_tools_list(
|
||||
request,
|
||||
payload,
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
assert result["error"] is None
|
||||
assert result["message"] == "Successfully retrieved tools"
|
||||
assert [tool["name"] for tool in result["tools"]] == ["quick_tool"]
|
||||
|
||||
async def test_extracts_oauth2_headers(self, monkeypatch):
|
||||
"""Ensure oauth2 auth type pulls oauth headers and omits MCP auth header."""
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue