mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): keep bridge tool metadata request-local
This commit is contained in:
parent
b0e6f66190
commit
68eec0b3c5
7 changed files with 134 additions and 73 deletions
|
|
@ -6664,6 +6664,7 @@ class MCPServerManager:
|
|||
wire_compat: WireCompat = WireCompat.LEGACY,
|
||||
*,
|
||||
catalog_auth_header: str | None | EllipsisType = ...,
|
||||
listed_tool: MCPTool | None | EllipsisType = ...,
|
||||
) -> CallToolResult | InputRequiredResult:
|
||||
"""
|
||||
Call a tool with the given name and arguments
|
||||
|
|
@ -6717,7 +6718,7 @@ class MCPServerManager:
|
|||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
guardrail_context=guardrail_context,
|
||||
tool=self.get_listed_tool(mcp_server, name, listed_caller),
|
||||
tool=self.get_listed_tool(mcp_server, name, listed_caller) if listed_tool is ... else listed_tool,
|
||||
)
|
||||
if "arguments" in hook_result:
|
||||
arguments = hook_result["arguments"]
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import asyncio
|
|||
import traceback
|
||||
import types
|
||||
import uuid
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import Any, Final, NoReturn, TypeAlias, overload
|
||||
|
||||
|
|
@ -81,7 +81,6 @@ 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,
|
||||
|
|
@ -958,7 +957,6 @@ 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.
|
||||
|
|
@ -971,7 +969,6 @@ 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
|
||||
|
|
@ -1069,12 +1066,12 @@ async def _get_tools_from_mcp_servers(
|
|||
|
||||
async def _fetch_and_filter_server_tools(
|
||||
server: MCPServer,
|
||||
) -> tuple[list[MCPTool], ServerOutcome, Callable[[frozenset[int]], None] | None]:
|
||||
) -> "tuple[list[MCPTool], ServerOutcome]":
|
||||
"""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), None
|
||||
return [], ServerListOk(tool_count=0)
|
||||
|
||||
server_auth_header, extra_headers = _prepare_mcp_server_headers(
|
||||
server=server,
|
||||
|
|
@ -1126,18 +1123,6 @@ 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,
|
||||
|
|
@ -1149,7 +1134,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 and served_tool_selector is None,
|
||||
record_listing=record_listing,
|
||||
)
|
||||
filtered_tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
||||
|
|
@ -1159,14 +1144,6 @@ 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
|
||||
|
||||
|
|
@ -1174,27 +1151,13 @@ 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)),
|
||||
record_served_tools if defer_recording else None,
|
||||
)
|
||||
return filtered_tools, ServerListOk(tool_count=len(filtered_tools))
|
||||
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
|
||||
|
|
@ -1203,25 +1166,20 @@ 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), None
|
||||
return [], classify_list_exception(e)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Error getting tools from server %s: %s", server.name, e)
|
||||
return [], classify_list_exception(e), None
|
||||
return [], classify_list_exception(e)
|
||||
|
||||
# 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, _, _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)
|
||||
all_tools: Final[list[MCPTool]] = [tool for tools, _ in results for tool in tools]
|
||||
server_outcomes: Final[dict[str, ServerOutcome]] = {
|
||||
_aggregate_server_key(server): outcome
|
||||
for server, (_, outcome, _record) in zip(allowed_mcp_servers, results)
|
||||
for server, (_, outcome) in zip(allowed_mcp_servers, results)
|
||||
if server is not None
|
||||
}
|
||||
|
||||
|
|
@ -1229,7 +1187,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, _, _record) in zip(allowed_mcp_servers, results)
|
||||
for server, (server_tools, _) in zip(allowed_mcp_servers, results)
|
||||
if server is not None
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -286,6 +286,7 @@ async def aresponses_api_with_mcp(
|
|||
call_params=call_params,
|
||||
previous_response_id=previous_response_id,
|
||||
tool_server_map=tool_server_map,
|
||||
served_tools=original_mcp_tools,
|
||||
**kwargs,
|
||||
)
|
||||
await mcp_streaming_response._create_initial_response_iterator()
|
||||
|
|
@ -339,6 +340,7 @@ async def aresponses_api_with_mcp(
|
|||
|
||||
tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=tool_server_map,
|
||||
served_tools=original_mcp_tools,
|
||||
tool_calls=tool_calls,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
|
|
@ -395,6 +397,7 @@ async def aresponses_api_with_mcp(
|
|||
|
||||
final_response = MCPEnhancedStreamingIterator(
|
||||
tool_server_map=tool_server_map,
|
||||
served_tools=original_mcp_tools,
|
||||
base_iterator=final_response,
|
||||
mcp_events=tool_execution_events,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
|
|
|
|||
|
|
@ -435,6 +435,7 @@ async def acompletion_with_mcp(
|
|||
# Execute tool calls
|
||||
self.tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=self.tool_server_map,
|
||||
served_tools=deduplicated_mcp_tools,
|
||||
tool_calls=self.tool_calls,
|
||||
user_api_key_auth=self.user_api_key_auth,
|
||||
mcp_auth_header=self.mcp_auth_header,
|
||||
|
|
@ -609,6 +610,7 @@ async def acompletion_with_mcp(
|
|||
tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=tool_server_map,
|
||||
tool_calls=tool_calls,
|
||||
served_tools=deduplicated_mcp_tools,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
|
|
|
|||
|
|
@ -318,13 +318,6 @@ 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,
|
||||
|
|
@ -335,8 +328,6 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
litellm_trace_id=litellm_trace_id,
|
||||
request_tags=request_tags,
|
||||
raw_headers=raw_headers,
|
||||
record_listing=True,
|
||||
served_tool_selector=served_tools,
|
||||
)
|
||||
tools: Final = listing.tools
|
||||
|
||||
|
|
@ -705,6 +696,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
litellm_trace_id: str | None = None,
|
||||
request_tags: list[str] | None = None,
|
||||
guardrail_context: Mapping[str, object] | None = None,
|
||||
served_tools: Sequence[MCPTool] | None = None,
|
||||
) -> list[MCPToolResult]:
|
||||
"""Execute tool calls and return results."""
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -869,6 +861,11 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
guardrail_context=guardrail_context,
|
||||
listed_tool=(
|
||||
next((tool for tool in served_tools if tool.name == tool_name), None)
|
||||
if served_tools is not None
|
||||
else ...
|
||||
),
|
||||
)
|
||||
|
||||
if proxy_logging_obj:
|
||||
|
|
@ -1161,6 +1158,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
call_params: Mapping[str, object],
|
||||
previous_response_id: str | None,
|
||||
tool_server_map: dict[str, str],
|
||||
served_tools: Sequence[MCPTool] | None = None,
|
||||
**kwargs,
|
||||
) -> Any:
|
||||
"""
|
||||
|
|
@ -1190,6 +1188,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
base_iterator=None, # Will be created internally
|
||||
mcp_events=mcp_discovery_events, # Pre-generated MCP discovery events
|
||||
tool_server_map=tool_server_map,
|
||||
served_tools=served_tools,
|
||||
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy,
|
||||
user_api_key_auth=kwargs.get("user_api_key_auth")
|
||||
or kwargs.get("litellm_metadata", {}).get("user_api_key_auth"),
|
||||
|
|
|
|||
|
|
@ -281,6 +281,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]] | None = None,
|
||||
user_api_key_auth: "UserAPIKeyAuth | None" = None,
|
||||
original_request_params: dict[str, Any] | None = None,
|
||||
served_tools: Sequence[MCPTool] | None = None,
|
||||
):
|
||||
# MCP setup
|
||||
self.mcp_tools_with_litellm_proxy = mcp_tools_with_litellm_proxy or []
|
||||
|
|
@ -300,6 +301,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
self.mcp_discovery_generated = True # Events are already generated
|
||||
self.mcp_events = mcp_events # Store the initial MCP events for backward compatibility
|
||||
self.tool_server_map = tool_server_map
|
||||
self.served_tools = tuple(served_tools) if served_tools is not None else None
|
||||
|
||||
# Iterator references
|
||||
self.base_iterator: BaseResponsesAPIStreamingIterator | ResponsesAPIResponse | None = (
|
||||
|
|
@ -796,6 +798,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
# Execute the tools
|
||||
tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=self.tool_server_map,
|
||||
served_tools=self.served_tools,
|
||||
tool_calls=tool_calls,
|
||||
user_api_key_auth=self.user_api_key_auth,
|
||||
mcp_auth_header=self.mcp_auth_header,
|
||||
|
|
|
|||
|
|
@ -1,9 +1,10 @@
|
|||
import asyncio
|
||||
import importlib
|
||||
import subprocess
|
||||
import sys
|
||||
import textwrap
|
||||
import types
|
||||
from typing import Any, Final, cast
|
||||
from typing import Any, Final, Literal, cast
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -12,20 +13,26 @@ from mcp.types import CallToolResult, TextContent
|
|||
from mcp.types import Tool as MCPTool
|
||||
from openai.types.responses.tool_param import Mcp
|
||||
|
||||
import litellm
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
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._experimental.mcp_server.mcp_server_manager import ListedToolsCaller, MCPServerManager
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
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.guardrails import GuardrailEventHooks
|
||||
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
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs, ModelResponse
|
||||
|
||||
|
||||
class _DummyMCPResult:
|
||||
|
|
@ -1316,6 +1323,7 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("real_listing", [False, True])
|
||||
@pytest.mark.parametrize(
|
||||
("allowed_tools", "expected_names"),
|
||||
[
|
||||
|
|
@ -1325,11 +1333,9 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py
|
|||
(["absent"], []),
|
||||
],
|
||||
)
|
||||
async def test_get_mcp_tools_from_manager_records_the_served_catalog(
|
||||
monkeypatch: pytest.MonkeyPatch, allowed_tools: list[str], expected_names: list[str]
|
||||
async def test_bridge_listing_leaves_the_callers_catalog_unchanged(
|
||||
monkeypatch: pytest.MonkeyPatch, allowed_tools: list[str], expected_names: list[str], real_listing: bool
|
||||
) -> 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", alias="responses_slot", transport=MCPTransport.http
|
||||
|
|
@ -1357,6 +1363,15 @@ async def test_get_mcp_tools_from_manager_records_the_served_catalog(
|
|||
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
|
||||
):
|
||||
try:
|
||||
if real_listing:
|
||||
await manager._get_tools_from_server(server, user_api_key_auth=user, record_listing=True)
|
||||
caller: Final = ListedToolsCaller(user_api_key_auth=user)
|
||||
before: 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 bool(before) is real_listing
|
||||
tools, _server_names = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform(
|
||||
user_api_key_auth=user,
|
||||
mcp_tools_with_litellm_proxy=[
|
||||
|
|
@ -1367,15 +1382,12 @@ async def test_get_mcp_tools_from_manager_records_the_served_catalog(
|
|||
}
|
||||
],
|
||||
)
|
||||
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 recorded == before
|
||||
assert (
|
||||
manager.get_listed_tool(
|
||||
server, "echo", ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-other-caller"))
|
||||
|
|
@ -1388,6 +1400,89 @@ async def test_get_mcp_tools_from_manager_records_the_served_catalog(
|
|||
assert [tool.name for tool in tools] == expected_names
|
||||
|
||||
|
||||
class _BridgeMetadataGuardrail(CustomGuardrail):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(guardrail_name="bridge-metadata", event_hook=GuardrailEventHooks.pre_mcp_call, default_on=True)
|
||||
self.calls: tuple[tuple[object, object], ...] = ()
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict[str, object],
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Logging | None = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
if request_data.get("mcp_arguments") == {"probe": "bridge"}:
|
||||
self.calls += ((request_data.get("mcp_tool_description"), request_data.get("mcp_input_schema")),)
|
||||
return inputs
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_bridge_calls_use_their_own_served_metadata(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
manager: Final = MCPServerManager()
|
||||
server: Final = MCPServer(server_id="bridge", name="bridge", transport=MCPTransport.http, url="http://upstream")
|
||||
manager.registry = {server.server_id: server}
|
||||
user: Final = UserAPIKeyAuth(api_key="sk-bridge", user_id="bridge-user")
|
||||
upstream: Final = [
|
||||
MCPTool(
|
||||
name="echo",
|
||||
description="Echo text",
|
||||
inputSchema={"type": "object", "properties": {"text": {"type": "string"}}},
|
||||
),
|
||||
MCPTool(name="status", description="Read status", inputSchema={"type": "object"}),
|
||||
]
|
||||
client: Final = AsyncMock()
|
||||
client.call_tool.return_value = CallToolResult(content=[TextContent(type="text", text="ok")])
|
||||
manager._create_mcp_client = AsyncMock(return_value=client)
|
||||
manager._fetch_tools_with_timeout = AsyncMock(return_value=upstream)
|
||||
guardrail: Final = _BridgeMetadataGuardrail()
|
||||
logger: Final = ProxyLogging(user_api_key_cache=DualCache())
|
||||
monkeypatch.setattr(litellm, "callbacks", [guardrail])
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", logger)
|
||||
monkeypatch.setattr(mcp_operations, "global_mcp_server_manager", manager)
|
||||
monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server]))
|
||||
monkeypatch.setattr("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", manager)
|
||||
first_listed: Final = asyncio.Event()
|
||||
second_listed: Final = asyncio.Event()
|
||||
|
||||
async def bridge(name: str, first: bool) -> None:
|
||||
if not first:
|
||||
await first_listed.wait()
|
||||
tools, server_map = 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/bridge", "allowed_tools": [name]}
|
||||
],
|
||||
)
|
||||
(first_listed if first else second_listed).set()
|
||||
await second_listed.wait()
|
||||
result: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=server_map,
|
||||
tool_calls=[
|
||||
{"type": "function_call", "name": f"bridge-{name}", "arguments": '{"probe":"bridge"}', "call_id": name}
|
||||
],
|
||||
user_api_key_auth=user,
|
||||
served_tools=tools,
|
||||
)
|
||||
assert [entry["result"] for entry in result] == ["ok"]
|
||||
|
||||
try:
|
||||
await asyncio.gather(bridge("echo", True), bridge("status", False))
|
||||
assert sorted(guardrail.calls, key=str) == sorted(
|
||||
((tool.description, tool.input_schema) for tool in upstream), key=str
|
||||
)
|
||||
await manager.call_tool("bridge", "echo", {"probe": "bridge"}, user_api_key_auth=user, proxy_logging_obj=logger)
|
||||
assert guardrail.calls[-1] == (None, None)
|
||||
await manager._get_tools_from_server(
|
||||
server, user_api_key_auth=user, proxy_logging_obj=logger, record_listing=True
|
||||
)
|
||||
await manager.call_tool("bridge", "echo", {"probe": "bridge"}, user_api_key_auth=user, proxy_logging_obj=logger)
|
||||
assert guardrail.calls[-1] == (upstream[0].description, upstream[0].input_schema)
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
|
||||
|
||||
def _toolset_gateway_manager(toolset_id: str, server_id: str) -> types.SimpleNamespace:
|
||||
return types.SimpleNamespace(
|
||||
get_registry=MagicMock(return_value={}),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue