mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): record only tools served by the bridge
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
8ffd7c107a
commit
b0e6f66190
3 changed files with 104 additions and 18 deletions
|
|
@ -4,7 +4,7 @@ import asyncio
|
|||
import traceback
|
||||
import types
|
||||
import uuid
|
||||
from collections.abc import Mapping, Sequence
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import Any, Final, NoReturn, TypeAlias, overload
|
||||
|
||||
|
|
@ -81,6 +81,7 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
|
|||
outcome_wire_value,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
ListedToolsCaller,
|
||||
MCPServerManager,
|
||||
_caller_authorization_fans_out,
|
||||
_client_forwarded_authorization_headers,
|
||||
|
|
@ -957,6 +958,7 @@ async def _get_tools_from_mcp_servers(
|
|||
mcp_proxy_mode: bool = False,
|
||||
*,
|
||||
record_listing: bool = False,
|
||||
served_tool_selector: Callable[[list[MCPTool]], list[MCPTool]] | None = None,
|
||||
) -> AggregateToolListing:
|
||||
"""
|
||||
Helper method to fetch tools from MCP servers based on server filtering criteria.
|
||||
|
|
@ -969,6 +971,7 @@ async def _get_tools_from_mcp_servers(
|
|||
oauth2_headers: Optional dict of oauth2 headers
|
||||
record_listing: Record each served catalog into the caller's listed-tools slot; only a
|
||||
listing actually served to the caller sets it
|
||||
served_tool_selector: Select existing tool objects from the aggregate for deferred recording
|
||||
|
||||
Returns:
|
||||
AggregateToolListing: Combined tools from filtered servers plus each server's
|
||||
|
|
@ -1066,12 +1069,12 @@ async def _get_tools_from_mcp_servers(
|
|||
|
||||
async def _fetch_and_filter_server_tools(
|
||||
server: MCPServer,
|
||||
) -> "tuple[list[MCPTool], ServerOutcome]":
|
||||
) -> tuple[list[MCPTool], ServerOutcome, Callable[[frozenset[int]], None] | None]:
|
||||
"""Fetch and filter tools from a single server, classifying any failure into that
|
||||
server's outcome so the aggregate can keep serving the healthy subset without a
|
||||
broken server masquerading as an empty one."""
|
||||
if server is None:
|
||||
return [], ServerListOk(tool_count=0)
|
||||
return [], ServerListOk(tool_count=0), None
|
||||
|
||||
server_auth_header, extra_headers = _prepare_mcp_server_headers(
|
||||
server=server,
|
||||
|
|
@ -1123,6 +1126,18 @@ async def _get_tools_from_mcp_servers(
|
|||
try:
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
defer_recording: Final = record_listing and served_tool_selector is not None
|
||||
listed_generation: Final = (
|
||||
global_mcp_server_manager._listed_tools_generations.get(server.server_id, 0)
|
||||
if defer_recording
|
||||
else None
|
||||
)
|
||||
listed_caller: Final = ListedToolsCaller(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=catalog_auth_header,
|
||||
raw_headers=raw_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
)
|
||||
tools: Final = await global_mcp_server_manager._get_tools_from_server(
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
|
|
@ -1134,7 +1149,7 @@ async def _get_tools_from_mcp_servers(
|
|||
oauth2_headers=oauth2_headers,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
catalog_auth_header=catalog_auth_header,
|
||||
record_listing=record_listing,
|
||||
record_listing=record_listing and served_tool_selector is None,
|
||||
)
|
||||
filtered_tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
||||
|
|
@ -1144,6 +1159,14 @@ async def _get_tools_from_mcp_servers(
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
unprefixed_tools: Final = (
|
||||
tuple(
|
||||
tool.model_copy(update={"name": strip_known_server_prefix(tool.name, server) or tool.name})
|
||||
for tool in filtered_tools
|
||||
)
|
||||
if defer_recording
|
||||
else ()
|
||||
)
|
||||
if mcp_proxy_mode:
|
||||
from litellm.proxy._experimental.mcp_server.tool_search import with_mcp_proxy_identity
|
||||
|
||||
|
|
@ -1151,13 +1174,27 @@ async def _get_tools_from_mcp_servers(
|
|||
else:
|
||||
filtered_tools = apply_display_name_overrides(filtered_tools, server)
|
||||
|
||||
catalog_entries: Final = tuple(zip(filtered_tools, unprefixed_tools))
|
||||
|
||||
def record_served_tools(served_ids: frozenset[int]) -> None:
|
||||
global_mcp_server_manager._record_listed_tools(
|
||||
server,
|
||||
[original for exposed, original in catalog_entries if id(exposed) in served_ids],
|
||||
listed_caller,
|
||||
listed_generation,
|
||||
)
|
||||
|
||||
verbose_logger.debug(
|
||||
"Successfully fetched %s tools from server %s, %s after filtering",
|
||||
len(tools),
|
||||
server.name,
|
||||
len(filtered_tools),
|
||||
)
|
||||
return filtered_tools, ServerListOk(tool_count=len(filtered_tools))
|
||||
return (
|
||||
filtered_tools,
|
||||
ServerListOk(tool_count=len(filtered_tools)),
|
||||
record_served_tools if defer_recording else None,
|
||||
)
|
||||
except MCPUpstreamAuthError as e:
|
||||
# Absorb so one unauthenticated server does not empty every other server's
|
||||
# tools. Surfacing the upstream 401 to the client as a re-auth challenge is
|
||||
|
|
@ -1166,20 +1203,25 @@ async def _get_tools_from_mcp_servers(
|
|||
# error). Single-server routes surface it via the request-scope preemptive
|
||||
# check in _raise_preemptive_401_for_unauthenticated_servers instead.
|
||||
verbose_logger.debug("MCP list_tools: omitting %s; it needs upstream auth", server.name)
|
||||
return [], classify_list_exception(e)
|
||||
return [], classify_list_exception(e), None
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Error getting tools from server %s: %s", server.name, e)
|
||||
return [], classify_list_exception(e)
|
||||
return [], classify_list_exception(e), None
|
||||
|
||||
# Fetch tools from all servers in parallel
|
||||
tasks: Final = [_fetch_and_filter_server_tools(server) for server in allowed_mcp_servers]
|
||||
results: Final = await asyncio.gather(*tasks)
|
||||
|
||||
# Flatten results into single list
|
||||
all_tools: Final[list[MCPTool]] = [tool for tools, _ in results for tool in tools]
|
||||
all_tools: Final[list[MCPTool]] = [tool for tools, _, _record in results for tool in tools]
|
||||
if record_listing and served_tool_selector is not None:
|
||||
served_ids: Final = frozenset(id(tool) for tool in served_tool_selector(all_tools))
|
||||
for _tools, _outcome, record in results:
|
||||
if record is not None:
|
||||
record(served_ids)
|
||||
server_outcomes: Final[dict[str, ServerOutcome]] = {
|
||||
_aggregate_server_key(server): outcome
|
||||
for server, (_, outcome) in zip(allowed_mcp_servers, results)
|
||||
for server, (_, outcome, _record) in zip(allowed_mcp_servers, results)
|
||||
if server is not None
|
||||
}
|
||||
|
||||
|
|
@ -1187,7 +1229,7 @@ async def _get_tools_from_mcp_servers(
|
|||
if litellm_logging_obj:
|
||||
per_server_tool_counts: Final[dict[str, int]] = {
|
||||
_aggregate_server_key(server): len(server_tools)
|
||||
for server, (server_tools, _) in zip(allowed_mcp_servers, results)
|
||||
for server, (server_tools, _, _record) in zip(allowed_mcp_servers, results)
|
||||
if server is not None
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -318,6 +318,13 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
# names), so use None and let the auth object's mcp_servers do the filtering.
|
||||
effective_server_filter: Final = None if resolved_toolset_ids else (resolved_mcp_servers or None)
|
||||
|
||||
def served_tools(tools: list[MCPTool]) -> list[MCPTool]:
|
||||
filtered: Final = LiteLLM_Proxy_MCP_Handler._filter_mcp_tools_by_allowed_tools(
|
||||
tools, mcp_tools_with_litellm_proxy
|
||||
)
|
||||
deduplicated, _server_map = LiteLLM_Proxy_MCP_Handler._deduplicate_mcp_tools(filtered, [])
|
||||
return deduplicated
|
||||
|
||||
listing: Final = await _get_tools_from_mcp_servers(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
|
|
@ -329,6 +336,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
request_tags=request_tags,
|
||||
raw_headers=raw_headers,
|
||||
record_listing=True,
|
||||
served_tool_selector=served_tools,
|
||||
)
|
||||
tools: Final = listing.tools
|
||||
|
||||
|
|
|
|||
|
|
@ -1316,13 +1316,30 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_mcp_tools_from_manager_records_the_served_catalog(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
@pytest.mark.parametrize(
|
||||
("allowed_tools", "expected_names"),
|
||||
[
|
||||
([], ["responses_slot-echo", "responses_slot-status"]),
|
||||
(["echo"], ["responses_slot-echo"]),
|
||||
(["responses_slot-echo"], ["responses_slot-echo"]),
|
||||
(["absent"], []),
|
||||
],
|
||||
)
|
||||
async def test_get_mcp_tools_from_manager_records_the_served_catalog(
|
||||
monkeypatch: pytest.MonkeyPatch, allowed_tools: list[str], expected_names: list[str]
|
||||
) -> 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)
|
||||
server: Final = MCPServer(
|
||||
server_id="responses-slot", name="responses_slot", alias="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"})]
|
||||
upstream: Final = [
|
||||
MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"}),
|
||||
MCPTool(name="status", description="Report status", inputSchema={"type": "object"}),
|
||||
MCPTool(name="echo", description="Duplicate echo", inputSchema={"type": "object", "properties": {}}),
|
||||
]
|
||||
fake_manager: Final = types.SimpleNamespace(
|
||||
get_registry=MagicMock(return_value={}),
|
||||
get_allowed_mcp_servers=AsyncMock(return_value=[]),
|
||||
|
|
@ -1340,16 +1357,35 @@ async def test_get_mcp_tools_from_manager_records_the_served_catalog(monkeypatch
|
|||
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(
|
||||
tools, _server_names = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform(
|
||||
user_api_key_auth=user,
|
||||
mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy/mcp/responses-slot"}],
|
||||
mcp_tools_with_litellm_proxy=[
|
||||
{
|
||||
"type": "mcp",
|
||||
"server_url": "litellm_proxy/mcp/responses-slot",
|
||||
"allowed_tools": allowed_tools,
|
||||
}
|
||||
],
|
||||
)
|
||||
caller: Final = ListedToolsCaller(user_api_key_auth=user)
|
||||
recorded: Final = {
|
||||
tool.name: (listed.description, listed.input_schema)
|
||||
for tool in upstream
|
||||
if (listed := manager.get_listed_tool(server, tool.name, caller)) is not None
|
||||
}
|
||||
assert recorded == {
|
||||
tool.name.removeprefix("responses_slot-"): (tool.description, tool.input_schema) for tool in tools
|
||||
}
|
||||
assert (
|
||||
manager.get_listed_tool(
|
||||
server, "echo", ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-other-caller"))
|
||||
)
|
||||
is None
|
||||
)
|
||||
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"
|
||||
assert [tool.name for tool in tools] == expected_names
|
||||
|
||||
|
||||
def _toolset_gateway_manager(toolset_id: str, server_id: str) -> types.SimpleNamespace:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue