fix(mcp): follow nextCursor on paginated tools/prompts/resources list operations

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-01 20:48:22 +00:00
parent 6d6c9af4ab
commit 4a6a387ca1
7 changed files with 258 additions and 36 deletions

View file

@ -136,6 +136,7 @@ MCP_CLIENT_TIMEOUT: Final = float(os.getenv("LITELLM_MCP_CLIENT_TIMEOUT", "60.0"
MCP_TOOL_LISTING_TIMEOUT: Final = float(os.getenv("LITELLM_MCP_TOOL_LISTING_TIMEOUT", "30.0"))
MCP_METADATA_TIMEOUT: Final = float(os.getenv("LITELLM_MCP_METADATA_TIMEOUT", "10.0"))
MCP_HEALTH_CHECK_TIMEOUT: Final = float(os.getenv("LITELLM_MCP_HEALTH_CHECK_TIMEOUT", "10.0"))
MCP_LIST_MAX_PAGES: Final = int(os.getenv("LITELLM_MCP_LIST_MAX_PAGES", "100"))
# Allowlist of commands permitted for MCP stdio transport.
# Prevents arbitrary command execution via /mcp-rest/test/* endpoints or server creation.

View file

@ -48,6 +48,12 @@ from pydantic import AnyUrl
from litellm._logging import verbose_logger
from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_NPM_CACHE_DIR
from litellm.experimental_mcp_client.pagination import (
list_all_prompts,
list_all_resource_templates,
list_all_resources,
list_all_tools,
)
from litellm.llms.custom_httpx.http_handler import get_ssl_configuration
from litellm.types.llms.custom_http import VerifyTypes
from litellm.types.mcp import (
@ -603,17 +609,17 @@ class MCPClient:
"""
verbose_logger.debug("MCP client listing tools from %s", self.server_url or "stdio")
async def _list_tools_operation(session: ClientSession):
return await session.list_tools()
async def _list_tools_operation(session: ClientSession) -> tuple[MCPTool, ...]:
return await list_all_tools(session, self.server_url or "stdio")
try:
result: Final = await self.run_with_session(_list_tools_operation, quiet_on_error=raise_on_error)
tool_count: Final = len(result.tools)
tool_names: Final = [tool.name for tool in result.tools]
tools: Final = await self.run_with_session(_list_tools_operation, quiet_on_error=raise_on_error)
tool_count: Final = len(tools)
tool_names: Final = [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
)
return result.tools
return list(tools)
except asyncio.CancelledError:
verbose_logger.warning("MCP client list_tools was cancelled")
raise
@ -734,17 +740,17 @@ class MCPClient:
"""List available prompts from the server."""
verbose_logger.debug("MCP client listing tools from %s", self.server_url or "stdio")
async def _list_prompts_operation(session: ClientSession):
return await session.list_prompts()
async def _list_prompts_operation(session: ClientSession) -> tuple[Prompt, ...]:
return await list_all_prompts(session, self.server_url or "stdio")
try:
result: Final = await self.run_with_session(_list_prompts_operation)
prompt_count: Final = len(result.prompts)
prompt_names: Final = [prompt.name for prompt in result.prompts]
prompts: Final = await self.run_with_session(_list_prompts_operation)
prompt_count: Final = len(prompts)
prompt_names: Final = [prompt.name for prompt in prompts]
verbose_logger.info(
"MCP client listed %s tools from %s: %s", prompt_count, self.server_url or "stdio", prompt_names
"MCP client listed %s prompts from %s: %s", prompt_count, self.server_url or "stdio", prompt_names
)
return result.prompts
return list(prompts)
except asyncio.CancelledError:
verbose_logger.warning("MCP client list_prompts was cancelled")
raise
@ -811,17 +817,17 @@ class MCPClient:
"""List available resources from the server."""
verbose_logger.debug("MCP client listing resources from %s", self.server_url or "stdio")
async def _list_resources_operation(session: ClientSession):
return await session.list_resources()
async def _list_resources_operation(session: ClientSession) -> tuple[Resource, ...]:
return await list_all_resources(session, self.server_url or "stdio")
try:
result: Final = await self.run_with_session(_list_resources_operation)
resource_count: Final = len(result.resources)
resource_names: Final = [resource.name for resource in result.resources]
resources: Final = await self.run_with_session(_list_resources_operation)
resource_count: Final = len(resources)
resource_names: Final = [resource.name for resource in resources]
verbose_logger.info(
"MCP client listed %s resources from %s: %s", resource_count, self.server_url or "stdio", resource_names
)
return result.resources
return list(resources)
except asyncio.CancelledError:
verbose_logger.warning("MCP client list_resources was cancelled")
raise
@ -847,20 +853,20 @@ class MCPClient:
"""List available resource templates from the server."""
verbose_logger.debug("MCP client listing resource templates from %s", self.server_url or "stdio")
async def _list_resource_templates_operation(session: ClientSession):
return await session.list_resource_templates()
async def _list_resource_templates_operation(session: ClientSession) -> tuple[ResourceTemplate, ...]:
return await list_all_resource_templates(session, self.server_url or "stdio")
try:
result: Final = await self.run_with_session(_list_resource_templates_operation)
resource_template_count: Final = len(result.resourceTemplates)
resource_template_names: Final = [resourceTemplate.name for resourceTemplate in result.resourceTemplates]
resource_templates: Final = await self.run_with_session(_list_resource_templates_operation)
resource_template_count: Final = len(resource_templates)
resource_template_names: Final = [resource_template.name for resource_template in resource_templates]
verbose_logger.info(
"MCP client listed %s resource templates from %s: %s",
resource_template_count,
self.server_url or "stdio",
resource_template_names,
)
return result.resourceTemplates
return list(resource_templates)
except asyncio.CancelledError:
verbose_logger.warning("MCP client list_resource_templates was cancelled")
raise

View file

@ -0,0 +1,92 @@
"""
Follows ``nextCursor`` on the paginated MCP list operations so a multi-page catalog is read in full.
"""
from collections.abc import Awaitable, Callable, Sequence
from typing import Final, TypeVar
from mcp import ClientSession, Resource
from mcp.types import PaginatedRequestParams, PaginatedResult, Prompt, ResourceTemplate
from mcp.types import Tool as MCPTool
from litellm._logging import verbose_logger
from litellm.constants import MCP_LIST_MAX_PAGES
TPage = TypeVar("TPage", bound=PaginatedResult)
TItem = TypeVar("TItem")
async def collect_pages(
fetch_page: Callable[[PaginatedRequestParams | None], Awaitable[TPage]],
items_of: Callable[[TPage], Sequence[TItem]],
*,
method: str,
server: str,
cursor: str | None = None,
seen_cursors: frozenset[str] = frozenset(),
) -> tuple[TItem, ...]:
page: Final = await fetch_page(None if cursor is None else PaginatedRequestParams(cursor=cursor))
items: Final = tuple(items_of(page))
next_cursor: Final = page.nextCursor
pages_read: Final = len(seen_cursors) + 1
if next_cursor is None:
return items
if next_cursor in seen_cursors:
verbose_logger.warning(
"MCP %s from %s repeated cursor %r; returning the %s page(s) read so far",
method,
server,
next_cursor,
pages_read,
)
return items
if pages_read >= MCP_LIST_MAX_PAGES:
verbose_logger.warning(
"MCP %s from %s still paginating after %s pages (LITELLM_MCP_LIST_MAX_PAGES); returning what was read",
method,
server,
pages_read,
)
return items
rest: Final = await collect_pages(
fetch_page,
items_of,
method=method,
server=server,
cursor=next_cursor,
seen_cursors=seen_cursors | frozenset((next_cursor,)),
)
return items + rest
async def list_all_tools(session: ClientSession, server: str) -> tuple[MCPTool, ...]:
return await collect_pages(
lambda params: session.list_tools(params=params), lambda page: page.tools, method="tools/list", server=server
)
async def list_all_prompts(session: ClientSession, server: str) -> tuple[Prompt, ...]:
return await collect_pages(
lambda params: session.list_prompts(params=params),
lambda page: page.prompts,
method="prompts/list",
server=server,
)
async def list_all_resources(session: ClientSession, server: str) -> tuple[Resource, ...]:
return await collect_pages(
lambda params: session.list_resources(params=params),
lambda page: page.resources,
method="resources/list",
server=server,
)
async def list_all_resource_templates(session: ClientSession, server: str) -> tuple[ResourceTemplate, ...]:
return await collect_pages(
lambda params: session.list_resource_templates(params=params),
lambda page: page.resourceTemplates,
method="resources/templates/list",
server=server,
)

View file

@ -9,6 +9,7 @@ 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.experimental_mcp_client.pagination import list_all_tools
from litellm.types.llms.anthropic import AnthropicMessagesTool
from litellm.types.utils import ChatCompletionMessageToolCall
@ -103,10 +104,10 @@ 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_all_tools(session, "upstream")
if format == "openai":
return [transform_mcp_tool_to_openai_tool(mcp_tool=tool) for tool in tools.tools]
return tools.tools
return [transform_mcp_tool_to_openai_tool(mcp_tool=tool) for tool in tools]
return list(tools)
########################################################

View file

@ -1402,11 +1402,7 @@ if MCP_AVAILABLE:
oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers)
async def _list_tools_operation(client):
async def _list_tools_session_operation(session):
return await session.list_tools()
list_tools_response: Final = await client.run_with_session(_list_tools_session_operation)
list_tools_result: Final[list[MCPTool]] = list_tools_response.tools
list_tools_result: Final[list[MCPTool]] = await client.list_tools(raise_on_error=True)
model_dumped_tools: Final[list[dict]] = [tool.model_dump() for tool in list_tools_result]
return {
"tools": model_dumped_tools,

View file

@ -20,8 +20,10 @@ from mcp.types import (
JSONRPCError,
JSONRPCMessage,
JSONRPCResponse,
ListToolsResult,
ServerCapabilities,
)
from mcp.types import Tool as MCPTool
# Add the parent directory to the path so we can import litellm
@ -740,8 +742,14 @@ class _ScriptedUpstream:
error, the shape an upstream application uses to report its own failure.
"""
def __init__(self, tools_list_error: ErrorData | None = None):
def __init__(
self,
tools_list_error: ErrorData | None = None,
tool_pages: tuple[tuple[MCPTool, ...], ...] = (),
):
self._tools_list_error = tools_list_error
self._tool_pages = tool_pages
self.tools_list_cursors: list[str | None] = []
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
@ -778,15 +786,37 @@ 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._tool_pages:
cursor = (request.params or {}).get("cursor")
self.tools_list_cursors.append(cursor)
page_index = int(cursor) if cursor else 0
has_more = page_index + 1 < len(self._tool_pages)
page = ListToolsResult(
tools=list(self._tool_pages[page_index]),
nextCursor=str(page_index + 1) if has_more else None,
)
await self._send(
JSONRPCResponse(
jsonrpc="2.0",
id=request.id,
result=page.model_dump(by_alias=True, mode="json", exclude_none=True),
)
)
class _ScriptedClient(MCPClient):
"""An MCPClient whose transport is a scripted in-memory upstream instead of a real connection,
so the real ``ClientSession`` and its real timeout machinery are what run."""
def __init__(self, *, timeout: float, tools_list_error: ErrorData | None = None):
def __init__(
self,
*,
timeout: float,
tools_list_error: ErrorData | None = None,
tool_pages: tuple[tuple[MCPTool, ...], ...] = (),
):
super().__init__(server_url="http://upstream.local/mcp", timeout=timeout)
self._upstream = _ScriptedUpstream(tools_list_error=tools_list_error)
self._upstream = _ScriptedUpstream(tools_list_error=tools_list_error, tool_pages=tool_pages)
def _create_transport_context(self):
return self._upstream, None
@ -821,6 +851,22 @@ async def test_list_tools_fails_on_its_own_timeout_when_the_upstream_never_answe
assert list_fault_http_status(fault) == 504
@pytest.mark.asyncio
async def test_list_tools_follows_tools_list_pagination_across_the_whole_catalog():
"""An upstream that pages tools/list (72 tools, 30 per page) must have every page read within the
one session, each request carrying the cursor the previous page returned. Reading only the first
page made 42 tools invisible to the proxy and every call to them fail as unknown."""
tools = tuple(
MCPTool(name=f"tool_{i:02d}", inputSchema={"type": "object", "properties": {}}) for i in range(72)
)
client = _ScriptedClient(timeout=30, tool_pages=(tools[:30], tools[30:60], tools[60:]))
listed = await asyncio.wait_for(client.list_tools(raise_on_error=True), timeout=10)
assert [tool.name for tool in listed] == [tool.name for tool in tools]
assert client._upstream.tools_list_cursors == [None, "1", "2"]
@pytest.mark.asyncio
async def test_upstream_json_rpc_error_408_is_not_reported_as_a_client_timeout():
"""The SDK reports its own elapsed read timeout and relays an upstream JSON-RPC error through

View file

@ -0,0 +1,80 @@
import logging
import pytest
from mcp.types import ListToolsResult, PaginatedRequestParams
from mcp.types import Tool as MCPTool
import litellm.experimental_mcp_client.pagination as pagination_module
from litellm.experimental_mcp_client.pagination import collect_pages
def _tool(index: int) -> MCPTool:
return MCPTool(name=f"tool_{index:02d}", inputSchema={"type": "object", "properties": {}})
class _PagedTools:
"""A tools/list upstream serving ``total`` tools ``page_size`` at a time, cursors being offsets."""
def __init__(self, total: int, page_size: int):
self._tools = tuple(_tool(i) for i in range(total))
self._page_size = page_size
self.cursors_seen: list[str | None] = []
async def fetch(self, params: PaginatedRequestParams | None) -> ListToolsResult:
cursor = params.cursor if params is not None else None
self.cursors_seen.append(cursor)
start = int(cursor) if cursor else 0
end = start + self._page_size
return ListToolsResult(
tools=list(self._tools[start:end]),
nextCursor=str(end) if end < len(self._tools) else None,
)
@pytest.mark.asyncio
async def test_collect_pages_follows_next_cursor_until_exhausted():
upstream = _PagedTools(total=72, page_size=30)
tools = await collect_pages(upstream.fetch, lambda page: page.tools, method="tools/list", server="s")
assert [t.name for t in tools] == [f"tool_{i:02d}" for i in range(72)]
assert upstream.cursors_seen == [None, "30", "60"], "each page must be requested with the cursor the previous one returned"
@pytest.mark.asyncio
async def test_collect_pages_single_page_makes_one_request():
upstream = _PagedTools(total=5, page_size=30)
tools = await collect_pages(upstream.fetch, lambda page: page.tools, method="tools/list", server="s")
assert len(tools) == 5
assert upstream.cursors_seen == [None]
@pytest.mark.asyncio
async def test_collect_pages_stops_on_a_repeated_cursor_and_keeps_what_it_read(caplog):
calls: list[str | None] = []
async def fetch(params: PaginatedRequestParams | None) -> ListToolsResult:
calls.append(params.cursor if params else None)
return ListToolsResult(tools=[_tool(len(calls))], nextCursor="same")
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
tools = await collect_pages(fetch, lambda page: page.tools, method="tools/list", server="s")
assert calls == [None, "same"], "the cursor must be followed once and refused the second time it comes back"
assert len(tools) == 2
assert any("repeated cursor" in record.getMessage() for record in caplog.records)
@pytest.mark.asyncio
async def test_collect_pages_honors_the_page_cap(monkeypatch, caplog):
monkeypatch.setattr(pagination_module, "MCP_LIST_MAX_PAGES", 3)
upstream = _PagedTools(total=1000, page_size=10)
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
tools = await collect_pages(upstream.fetch, lambda page: page.tools, method="tools/list", server="s")
assert len(upstream.cursors_seen) == 3
assert len(tools) == 30
assert any("LITELLM_MCP_LIST_MAX_PAGES" in record.getMessage() for record in caplog.records)