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:
Yucheng Zhu 2026-09-01 12:00:46 -07:00
parent 772dba1d11
commit 74a640d430
6 changed files with 264 additions and 49 deletions

View file

@ -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
)

View file

@ -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
########################################################

View file

@ -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,

View file

@ -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

View file

@ -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"}'}

View file

@ -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."""