mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): opt every listing out of catalog recording unless it is served
The aggregate listing and _list_mcp_tools now default to record_listing=False, so a catalog fetched inside a tools/call no longer fills the caller's listed-tools slot. The /mcp/proxy meta-tools (call_tool, search_tools, get_tool_schema) and the tool-search virtual tool stop recording: /mcp/proxy serves only the meta-tools and the search serves only its hits, so a later pre_mcp_call hook was reading a description the caller never listed. The tools/list handler, the Responses MCP handler and the /v1/mcp/tools management listing opt in with record_listing=True, since each serves the catalog to the caller.
This commit is contained in:
parent
ad746dadb7
commit
0219521559
7 changed files with 171 additions and 3 deletions
|
|
@ -958,7 +958,7 @@ async def _get_tools_from_mcp_servers(
|
|||
client_ip: str | None = None,
|
||||
mcp_proxy_mode: bool = False,
|
||||
*,
|
||||
record_listing: bool = True,
|
||||
record_listing: bool = False,
|
||||
) -> AggregateToolListing:
|
||||
"""
|
||||
Helper method to fetch tools from MCP servers based on server filtering criteria.
|
||||
|
|
@ -1470,6 +1470,8 @@ async def _list_mcp_tools(
|
|||
list_tools_log_source: str | None = None,
|
||||
client_ip: str | None = None,
|
||||
mcp_proxy_mode: bool = False,
|
||||
*,
|
||||
record_listing: bool = False,
|
||||
) -> AggregateToolListing:
|
||||
"""
|
||||
List all available MCP tools.
|
||||
|
|
@ -1480,6 +1482,8 @@ async def _list_mcp_tools(
|
|||
mcp_servers: Optional list of server names/aliases to filter by
|
||||
mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value}
|
||||
client_ip: Client IP for IP-based server access control
|
||||
record_listing: Record each served catalog into the caller's listed-tools slot; only a
|
||||
listing actually served to the caller sets it
|
||||
|
||||
Returns:
|
||||
AggregateToolListing: Combined tools from all accessible servers plus each server's
|
||||
|
|
@ -1498,6 +1502,7 @@ async def _list_mcp_tools(
|
|||
list_tools_log_source=list_tools_log_source,
|
||||
client_ip=client_ip,
|
||||
mcp_proxy_mode=mcp_proxy_mode,
|
||||
record_listing=record_listing,
|
||||
)
|
||||
verbose_logger.debug("Successfully fetched %s tools from managed MCP servers", len(listing.tools))
|
||||
return listing
|
||||
|
|
@ -2757,6 +2762,7 @@ async def _execute_handle_list_tools(
|
|||
log_list_tools_to_spendlogs=log_list_tools_to_spendlogs,
|
||||
list_tools_log_source="mcp_protocol",
|
||||
client_ip=_client_ip,
|
||||
record_listing=True,
|
||||
)
|
||||
verbose_logger.info("MCP list_tools - Successfully returned %s tools", len(listing.tools))
|
||||
if not listing.outcomes:
|
||||
|
|
|
|||
|
|
@ -1010,6 +1010,7 @@ if MCP_AVAILABLE:
|
|||
mcp_auth_header=None,
|
||||
mcp_servers=None,
|
||||
mcp_server_auth_headers=None,
|
||||
record_listing=True,
|
||||
)
|
||||
tools: Final = listing.tools
|
||||
dumped_tools: Final = [tool.model_dump(by_alias=True) for tool in tools]
|
||||
|
|
|
|||
|
|
@ -328,6 +328,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
litellm_trace_id=litellm_trace_id,
|
||||
request_tags=request_tags,
|
||||
raw_headers=raw_headers,
|
||||
record_listing=True,
|
||||
)
|
||||
tools: Final = listing.tools
|
||||
|
||||
|
|
|
|||
|
|
@ -1,17 +1,29 @@
|
|||
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
|
||||
import asyncio
|
||||
import json
|
||||
from datetime import datetime
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from mcp.shared.exceptions import MCPError
|
||||
from mcp.types import CallToolResult, TextContent
|
||||
from mcp.types import Tool as MCPTool
|
||||
from pydantic import AnyUrl
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._experimental.mcp_server import server
|
||||
from litellm.proxy._experimental.mcp_server.mcp_context import _mcp_proxy_mode
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
|
||||
from litellm.proxy._experimental.mcp_server.tool_search import (
|
||||
handle_mcp_proxy_tool,
|
||||
mcp_proxy_tool_id,
|
||||
with_mcp_proxy_identity,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
AUTH = UserAPIKeyAuth(api_key="key")
|
||||
|
||||
|
|
@ -130,3 +142,46 @@ async def test_proxy_scope_exception_emits_failure_log(monkeypatch: pytest.Monke
|
|||
assert hook_payload["arguments"] == arguments
|
||||
assert "raw_headers" not in hook_payload
|
||||
assert "raw-scope-secret" not in recorder.events[1][1]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_call_tool_on_a_never_listed_tool_hands_the_pre_hook_no_listed_tool() -> None:
|
||||
"""/mcp/proxy tools/list serves only the meta-tools, so the catalog call_tool reads to resolve its
|
||||
tool_id was never served: it must not fill the caller's listed-tools slot, and the pre-call hook
|
||||
must see no listed tool for the call."""
|
||||
manager = mcp_operations.global_mcp_server_manager
|
||||
server = MCPServer(server_id="proxy-meta", name="proxy-meta", transport=MCPTransport.http, url="http://meta")
|
||||
auth = UserAPIKeyAuth(api_key="sk-proxy-meta", user_id="proxy-caller")
|
||||
upstream = [MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"})]
|
||||
served_as = with_mcp_proxy_identity(MCPTool(name="proxy-meta-echo", inputSchema={}), server.server_id)
|
||||
pre_call_tool_check = AsyncMock(return_value={})
|
||||
|
||||
async def call_regular_mcp_tool(*, tasks: list[asyncio.Task[object]], **_: object) -> CallToolResult:
|
||||
await asyncio.gather(*tasks)
|
||||
return CallToolResult(content=[TextContent(type="text", text="echoed")])
|
||||
|
||||
with (
|
||||
patch.dict(manager.registry, {server.server_id: server}),
|
||||
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
|
||||
patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
|
||||
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
|
||||
patch.object(manager, "pre_call_tool_check", pre_call_tool_check),
|
||||
patch.object(manager, "_call_regular_mcp_tool", call_regular_mcp_tool),
|
||||
):
|
||||
try:
|
||||
result = await handle_mcp_proxy_tool(
|
||||
name="call_tool",
|
||||
arguments={"tool_id": mcp_proxy_tool_id(served_as), "arguments": {}},
|
||||
user_api_key_dict=auth,
|
||||
)
|
||||
listed = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=auth))
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
|
||||
assert result.is_error is False
|
||||
assert result.content[0].text == "echoed"
|
||||
pre_call_tool_check.assert_awaited_once()
|
||||
assert pre_call_tool_check.await_args.kwargs["name"] == "echo"
|
||||
assert pre_call_tool_check.await_args.kwargs["tool"] is None
|
||||
assert listed is None
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ from mcp.types import Tool
|
|||
import litellm
|
||||
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
|
||||
from litellm.proxy._experimental.mcp_server.tool_search import (
|
||||
AGENT_SEARCH_TOOL_NAME,
|
||||
MCP_TOOL_CALL_TOOL_NAME,
|
||||
|
|
@ -31,12 +32,14 @@ from litellm.proxy._experimental.mcp_server.tool_search import (
|
|||
ToolSearchResult,
|
||||
coerce_top_k,
|
||||
get_virtual_tool_definitions,
|
||||
handle_mcp_tool_search,
|
||||
search_mcp_tools,
|
||||
search_tools,
|
||||
)
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.semantic_text_index import EmbeddingFailed, SemanticTextIndex, Vector
|
||||
from litellm.types.mcp import MCPToolSearchSettings
|
||||
from litellm.types.mcp import MCPToolSearchSettings, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
||||
def _make_tools(specs: list[tuple[str, str]]) -> tuple[Tool, ...]:
|
||||
|
|
@ -1268,3 +1271,33 @@ async def test_handle_mcp_tool_call_scoped_denial_names_the_binding_agent() -> N
|
|||
assert exc_info.value.status_code == 403
|
||||
assert "MCP server 'github'" in exc_info.value.detail["error"]
|
||||
assert "agent 'agent-123'" in exc_info.value.detail["error"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_tool_search_leaves_the_listed_tools_slot_empty(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""The search lists the whole catalog but serves only its hits, so the listing must not fill the
|
||||
caller's listed-tools slot: a later call to a tool the search never returned is not a listed tool."""
|
||||
monkeypatch.setattr(litellm, "mcp_tool_search", None)
|
||||
manager = mcp_operations.global_mcp_server_manager
|
||||
server = MCPServer(server_id="search-slot", name="search-slot", transport=MCPTransport.http, url="http://slot")
|
||||
user = UserAPIKeyAuth(api_key="sk-search-slot", user_id="searcher")
|
||||
upstream = [
|
||||
Tool(name="echo", description="Echo text back", inputSchema={"type": "object"}),
|
||||
Tool(name="delete_note", description="Delete a note", inputSchema={"type": "object"}),
|
||||
]
|
||||
with (
|
||||
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
|
||||
patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
|
||||
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
|
||||
):
|
||||
try:
|
||||
result = await handle_mcp_tool_search(query="echo", top_k=1, user_api_key_dict=user)
|
||||
caller = ListedToolsCaller(user_api_key_auth=user)
|
||||
listed = [manager.get_listed_tool(server, tool.name, caller) for tool in upstream]
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
|
||||
assert result.is_error is False
|
||||
assert [hit["name"] for hit in json.loads(result.content[0].text)] == ["search-slot-echo"]
|
||||
assert listed == [None, None]
|
||||
|
|
|
|||
|
|
@ -3,7 +3,10 @@ from unittest.mock import AsyncMock, patch
|
|||
|
||||
import pytest
|
||||
from mcp.types import GetPromptRequest, GetPromptRequestParams, GetPromptResult
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
|
||||
from litellm.proxy._experimental.mcp_server.operations import GatewayOperations, prepare_context
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
|
|
@ -665,3 +668,30 @@ async def test_tools_listing_preserves_explicit_spend_log_policy(log_enabled):
|
|||
)
|
||||
assert result.tools == []
|
||||
assert listing.await_args.kwargs["log_list_tools_to_spendlogs"] is log_enabled
|
||||
assert listing.await_args.kwargs["record_listing"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(("listing_kwargs", "recorded"), [({}, False), ({"record_listing": True}, True)])
|
||||
async def test_list_mcp_tools_records_the_catalog_only_when_asked(
|
||||
listing_kwargs: dict[str, bool], recorded: bool
|
||||
) -> None:
|
||||
"""The aggregate listing fills the caller's listed-tools slot only when asked: a listing an internal
|
||||
caller never serves must not hand a later tools/call a description the caller never saw."""
|
||||
manager = operations.global_mcp_server_manager
|
||||
server = MCPServer(server_id="listing-slot", name="listing-slot", transport=MCPTransport.http, url="http://slot")
|
||||
user = UserAPIKeyAuth(api_key="sk-listing-slot", user_id="lister")
|
||||
upstream = [MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"})]
|
||||
with (
|
||||
patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
|
||||
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
|
||||
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
|
||||
):
|
||||
try:
|
||||
listing = await operations._list_mcp_tools(user_api_key_auth=user, **listing_kwargs)
|
||||
listed = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=user))
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
assert [tool.name for tool in listing.tools] == ["listing-slot-echo"]
|
||||
assert (listed is not None) is recorded
|
||||
|
|
|
|||
|
|
@ -4,20 +4,25 @@ import sys
|
|||
import textwrap
|
||||
import types
|
||||
from typing import Any, Final, cast
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from mcp.types import CallToolResult, TextContent, Tool as MCPTool
|
||||
from openai.types.responses.tool_param import Mcp
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.responses import main as responses_main
|
||||
from litellm.responses.mcp import litellm_proxy_mcp_handler as mcp_handler_module
|
||||
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
|
||||
LiteLLM_Proxy_MCP_Handler,
|
||||
)
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
from litellm.types.responses.main import OutputFunctionToolCall
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
|
@ -1311,3 +1316,40 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py
|
|||
logged: Final = setup.call_args.kwargs["metadata"]["headers"]
|
||||
assert logged == {"x-app-id": "app-a", "x-nuid": "user-a", "x-user-id": "identity-a"}
|
||||
assert headers["x-mcp-deepwiki-authorization"] == "upstream-sentinel"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_mcp_tools_from_manager_records_the_served_catalog(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""The Responses bridge serves the listing to the model and its own tools/call reads the slot, so
|
||||
this listing records the caller's catalog."""
|
||||
manager: Final = mcp_operations.global_mcp_server_manager
|
||||
server: Final = MCPServer(server_id="responses-slot", name="responses-slot", transport=MCPTransport.http)
|
||||
user: Final = UserAPIKeyAuth(api_key="sk-responses-slot", user_id="responder")
|
||||
upstream: Final = [MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"})]
|
||||
fake_manager: Final = types.SimpleNamespace(
|
||||
get_registry=MagicMock(return_value={}),
|
||||
get_allowed_mcp_servers=AsyncMock(return_value=[]),
|
||||
get_mcp_servers_from_ids=MagicMock(return_value=[]),
|
||||
get_mcp_server_by_name=MagicMock(return_value=None),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
|
||||
fake_manager,
|
||||
)
|
||||
with (
|
||||
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
|
||||
patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
|
||||
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
|
||||
):
|
||||
try:
|
||||
tools, _server_names = await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager(
|
||||
user_api_key_auth=user,
|
||||
mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy/mcp/responses-slot"}],
|
||||
)
|
||||
listed: Final = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=user))
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
|
||||
assert [tool.name for tool in tools] == ["responses-slot-echo"]
|
||||
assert listed is not None and listed.description == "Echo text back"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue