mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): record served catalogs and preserve call message bytes
This commit is contained in:
parent
acc02ce8d4
commit
c669328ba1
6 changed files with 148 additions and 63 deletions
|
|
@ -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,
|
||||
|
|
@ -1123,6 +1124,7 @@ async def _get_tools_from_mcp_servers(
|
|||
try:
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
listed_generation: Final = global_mcp_server_manager._listed_tools_generations.get(server.server_id, 0)
|
||||
tools: Final = await global_mcp_server_manager._get_tools_from_server(
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
|
|
@ -1134,7 +1136,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=False,
|
||||
)
|
||||
filtered_tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
||||
|
|
@ -1143,6 +1145,21 @@ async def _get_tools_from_mcp_servers(
|
|||
server_id=server.server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
global_mcp_server_manager._record_listed_tools(
|
||||
server,
|
||||
[
|
||||
tool.model_copy(update={"name": strip_known_server_prefix(tool.name, server)})
|
||||
for tool in filtered_tools
|
||||
],
|
||||
ListedToolsCaller(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=catalog_auth_header,
|
||||
raw_headers=raw_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
),
|
||||
listed_generation,
|
||||
record_listing=record_listing,
|
||||
)
|
||||
|
||||
if mcp_proxy_mode:
|
||||
from litellm.proxy._experimental.mcp_server.tool_search import with_mcp_proxy_identity
|
||||
|
|
|
|||
|
|
@ -223,6 +223,7 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES,
|
||||
ListedToolsCaller,
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
|
|
@ -719,8 +720,8 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
|
||||
async def _get_tools_for_single_server(
|
||||
server,
|
||||
server_auth_header,
|
||||
server: MCPServer,
|
||||
server_auth_header: dict[str, str] | str | None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
|
|
@ -736,7 +737,8 @@ if MCP_AVAILABLE:
|
|||
"""
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
tools = await _list_server_tools(
|
||||
listed_generation: Final = global_mcp_server_manager._listed_tools_generations.get(server.server_id, 0)
|
||||
tools: Final = await _list_server_tools(
|
||||
server,
|
||||
server_auth_header,
|
||||
raw_headers,
|
||||
|
|
@ -744,30 +746,31 @@ if MCP_AVAILABLE:
|
|||
extra_headers,
|
||||
client_ip,
|
||||
proxy_logging_obj,
|
||||
record_listing=True,
|
||||
record_listing=False,
|
||||
)
|
||||
|
||||
if not apply_tool_filters:
|
||||
return _create_tool_response_objects(tools, server)
|
||||
|
||||
# Always apply allowed_tools/disallowed_tools so the blacklist is
|
||||
# enforced even when no allowlist is set (matches the SSE/HTTP path).
|
||||
tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
||||
# Filter by the key's effective tool permissions through the same
|
||||
# function the MCP protocol path uses (direct grants, toolset grants,
|
||||
# and team/agent/org ceilings), so REST listing cannot drift from it.
|
||||
# Entries here are tool names on one server, written bare by every
|
||||
# writer, and dispatch compares them bare; matching a wider set of
|
||||
# spellings would advertise a tool that tools/call then refuses
|
||||
if user_api_key_auth:
|
||||
tools = await filter_tools_by_key_team_permissions(
|
||||
tools=tools,
|
||||
server_filtered: Final = filter_tools_by_allowed_tools(tools, server) if apply_tool_filters else tools
|
||||
served_tools: Final = (
|
||||
await filter_tools_by_key_team_permissions(
|
||||
tools=server_filtered,
|
||||
server_id=server.server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
if apply_tool_filters and user_api_key_auth
|
||||
else server_filtered
|
||||
)
|
||||
global_mcp_server_manager._record_listed_tools(
|
||||
server,
|
||||
served_tools,
|
||||
ListedToolsCaller(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=server_auth_header,
|
||||
raw_headers=raw_headers,
|
||||
),
|
||||
listed_generation,
|
||||
)
|
||||
|
||||
return _create_tool_response_objects(tools, server)
|
||||
return _create_tool_response_objects(served_tools, server)
|
||||
|
||||
async def fetch_pinnable_tool_catalog(
|
||||
server: MCPServer, request: Request, user_api_key_dict: UserAPIKeyAuth
|
||||
|
|
|
|||
|
|
@ -1026,29 +1026,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
if existing_cap is None or effective_cap < existing_cap:
|
||||
data[cap_field] = effective_cap # rebind-ok: downstream routing requires the bounded output cap
|
||||
|
||||
@staticmethod
|
||||
def _mcp_token_reservation_data(data: object, call_type: str | None) -> object:
|
||||
if (
|
||||
call_type != CallTypes.call_mcp_tool.value
|
||||
or not isinstance(data, dict)
|
||||
or "mcp_tool_name" not in data
|
||||
or "mcp_arguments" not in data
|
||||
):
|
||||
return data
|
||||
mcp_data: Final = TypeAdapter(dict[str, object]).validate_python(data)
|
||||
return {
|
||||
**mcp_data,
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": f"Tool: {mcp_data['mcp_tool_name']}\nArguments: {mcp_data['mcp_arguments']}",
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
def _estimate_tokens_for_request(
|
||||
self,
|
||||
data: dict[str, object],
|
||||
data: dict,
|
||||
model: str | None = None,
|
||||
min_configured_tpm_limit: int | None = None,
|
||||
call_type: str | None = None,
|
||||
|
|
@ -1074,9 +1054,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
floor entirely, so the reservation reflects what this tenant's model
|
||||
actually emits rather than one constant shared by every tenant.
|
||||
"""
|
||||
reservation_data: Final = self._mcp_token_reservation_data(data, call_type)
|
||||
estimated_input_tokens, max_tokens_estimate = self._estimate_input_and_output_tokens(
|
||||
data=reservation_data,
|
||||
data=data,
|
||||
min_configured_tpm_limit=min_configured_tpm_limit,
|
||||
call_type=call_type,
|
||||
configured_output_tokens=configured_output_tokens,
|
||||
|
|
@ -3795,14 +3774,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
if v is not None
|
||||
]
|
||||
min_configured_otpm_limit: Final = min(configured_otpm_limits) if configured_otpm_limits else None
|
||||
reservation_data: Final = self._mcp_token_reservation_data(data, call_type)
|
||||
_, raw_estimated_output_tokens = self._estimate_input_and_output_tokens(
|
||||
data=reservation_data,
|
||||
data=data,
|
||||
min_configured_tpm_limit=min_configured_otpm_limit,
|
||||
call_type=call_type,
|
||||
)
|
||||
raw_estimated_input_tokens: Final = await offload_token_count(self._estimate_precise_input_tokens)(
|
||||
data=reservation_data, model=requested_model, call_type=call_type
|
||||
data=data, model=requested_model, call_type=call_type
|
||||
)
|
||||
estimated_input_tokens: Final = max(raw_estimated_input_tokens, 1)
|
||||
estimated_output_tokens: Final = (
|
||||
|
|
|
|||
|
|
@ -1479,7 +1479,8 @@ class ProxyLogging:
|
|||
if request_obj.tool_input_schema is not None
|
||||
else kwargs.get("mcp_input_schema")
|
||||
)
|
||||
description_line: Final = f"\nDescription: {mcp_tool_description}" if mcp_tool_description else ""
|
||||
listing_description: Final = kwargs.get("mcp_tool_description")
|
||||
description_line: Final = f"\nDescription: {listing_description}" if listing_description else ""
|
||||
tool_call_content: Final = (
|
||||
f"Tool: {request_obj.tool_name}{description_line}\nArguments: {request_obj.arguments}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,18 +1,101 @@
|
|||
import asyncio
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from mcp.types import GetPromptRequest, GetPromptRequestParams, GetPromptResult
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
import litellm
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
from litellm.proxy._experimental.mcp_server import rest_endpoints
|
||||
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._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
||||
class _CatalogHookCapture(CustomLogger):
|
||||
data: dict[str, object] | None = None
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache, data: dict[str, object], call_type: str
|
||||
) -> None:
|
||||
if call_type == "call_mcp_tool":
|
||||
self.data = data.copy()
|
||||
|
||||
|
||||
async def _served_catalog_tool() -> str:
|
||||
return "ok"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("surface", ["mcp", "rest"])
|
||||
@pytest.mark.parametrize("restriction", ["key", "server"])
|
||||
async def test_listing_records_only_tools_the_caller_received(
|
||||
monkeypatch: pytest.MonkeyPatch, surface: str, restriction: str
|
||||
) -> None:
|
||||
manager: Final = operations.global_mcp_server_manager
|
||||
server: Final = MCPServer(
|
||||
server_id="served-catalog", name="served-catalog", transport=MCPTransport.http,
|
||||
spec_path="/catalog.yaml", allow_all_keys=True,
|
||||
allowed_tools=["echo"] if restriction == "server" else None,
|
||||
)
|
||||
auth: Final = UserAPIKeyAuth(
|
||||
api_key="sk-served-catalog", user_id="lister",
|
||||
object_permission={
|
||||
"object_permission_id": "served-permission",
|
||||
"mcp_servers": [server.server_id],
|
||||
"mcp_tool_permissions": {server.server_id: ["echo"]} if restriction == "key" else None,
|
||||
},
|
||||
)
|
||||
monkeypatch.setitem(manager.registry, server.server_id, server)
|
||||
monkeypatch.setitem(manager.tool_name_to_mcp_server_name_mapping, "status", server.server_id)
|
||||
monkeypatch.setitem(manager.tool_name_to_mcp_server_name_mapping, "served-catalog-status", server.server_id)
|
||||
capture: Final = _CatalogHookCapture()
|
||||
monkeypatch.setattr(litellm, "callbacks", [capture])
|
||||
for name in ("echo", "status"):
|
||||
global_mcp_tool_registry.register_tool(
|
||||
name=f"served-catalog-{name}", description=f"{name} description",
|
||||
input_schema={"type": "object"}, handler=_served_catalog_tool,
|
||||
)
|
||||
try:
|
||||
if surface == "mcp":
|
||||
listing: Final = await operations._list_mcp_tools(
|
||||
user_api_key_auth=auth, mcp_servers=[server.server_id], record_listing=True,
|
||||
)
|
||||
assert [tool.name for tool in listing.tools] == ["served-catalog-echo"]
|
||||
else:
|
||||
rest_listing: Final = await rest_endpoints._get_tools_for_single_server(
|
||||
server, None, user_api_key_auth=auth,
|
||||
)
|
||||
assert [tool.name for tool in rest_listing] == ["echo"]
|
||||
granted: Final = auth.model_copy(update={"object_permission": None})
|
||||
caller: Final = ListedToolsCaller(user_api_key_auth=granted)
|
||||
assert manager.get_listed_tool(server, "status", caller) is None
|
||||
served: Final = manager.get_listed_tool(server, "echo", caller)
|
||||
assert served is not None
|
||||
assert (served.description, served.input_schema) == ("echo description", {"type": "object"})
|
||||
server.allowed_tools = None
|
||||
result: Final = await manager.call_tool(
|
||||
server_name=server.server_id, name="status", arguments={}, user_api_key_auth=granted,
|
||||
proxy_logging_obj=ProxyLogging(user_api_key_cache=UserApiKeyCache()),
|
||||
)
|
||||
assert result.is_error is False
|
||||
assert capture.data is not None
|
||||
assert capture.data["messages"] == [{"role": "user", "content": "Tool: status\nArguments: {}"}]
|
||||
assert (capture.data.get("mcp_tool_description"), capture.data.get("mcp_input_schema")) == (None, None)
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
global_mcp_tool_registry.unregister_tools_with_prefix("served-catalog-")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_oauth_prefetch_failure_does_not_log_caller_or_exception_text(caplog):
|
||||
from litellm.proxy._experimental.mcp_server.operations import _prefetch_oauth_creds_for_user
|
||||
|
|
|
|||
|
|
@ -115,7 +115,10 @@ def test_api_key_descriptor_applies_budget_throttle(
|
|||
@pytest.mark.parametrize(
|
||||
"description", [None, "Gateway metadata, not caller input. " * 100], ids=["unlisted", "listed"]
|
||||
)
|
||||
async def test_mcp_description_does_not_change_admission_or_reserved_tokens(description: str | None) -> None:
|
||||
@pytest.mark.parametrize("arguments_rewritten", [False, True])
|
||||
async def test_mcp_description_does_not_change_admission_or_reserved_tokens(
|
||||
description: str | None, arguments_rewritten: bool, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
|
||||
logger: Final = ProxyLogging(user_api_key_cache=UserApiKeyCache())
|
||||
|
|
@ -127,7 +130,10 @@ async def test_mcp_description_does_not_change_admission_or_reserved_tokens(desc
|
|||
messages: Final = data["messages"]
|
||||
caller: Final = UserAPIKeyAuth(api_key=hash_token("sk-mcp-description-reservation"), tpm_limit=64)
|
||||
|
||||
await handler.async_pre_call_hook(user_api_key_dict=caller, cache=cache, data=data, call_type="call_mcp_tool")
|
||||
if arguments_rewritten:
|
||||
data["mcp_arguments"] = {"q": "Transformed arguments " * 100}
|
||||
monkeypatch.setattr(litellm, "callbacks", [handler])
|
||||
await logger.pre_call_hook(user_api_key_dict=caller, data=data, call_type="call_mcp_tool")
|
||||
|
||||
stash: Final = get_request_stash()
|
||||
assert stash is not None
|
||||
|
|
@ -144,11 +150,7 @@ async def test_mcp_description_does_not_change_admission_or_reserved_tokens(desc
|
|||
assert messages == [
|
||||
{
|
||||
"role": "user",
|
||||
"content": (
|
||||
f"Tool: echo\nDescription: {description}\nArguments: {{'q': 'hello'}}"
|
||||
if description
|
||||
else "Tool: echo\nArguments: {'q': 'hello'}"
|
||||
),
|
||||
"content": "Tool: echo\nArguments: {'q': 'hello'}",
|
||||
}
|
||||
]
|
||||
|
||||
|
|
@ -158,8 +160,10 @@ async def test_mcp_description_does_not_change_admission_or_reserved_tokens(desc
|
|||
"description", [None, "Gateway metadata, not caller input. " * 100], ids=["unlisted", "listed"]
|
||||
)
|
||||
@pytest.mark.parametrize("itpm_limit,otpm_limit", [(64, 4096), (4096, 64), (4096, 4096)])
|
||||
@pytest.mark.parametrize("arguments_rewritten", [False, True])
|
||||
async def test_mcp_description_preserves_project_input_and_output_reservations(
|
||||
description: str | None, itpm_limit: int, otpm_limit: int
|
||||
description: str | None, itpm_limit: int, otpm_limit: int,
|
||||
arguments_rewritten: bool, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
|
||||
|
|
@ -188,7 +192,10 @@ async def test_mcp_description_preserves_project_input_and_output_reservations(
|
|||
},
|
||||
)
|
||||
|
||||
await handler.async_pre_call_hook(user_api_key_dict=caller, cache=cache, data=data, call_type="call_mcp_tool")
|
||||
if arguments_rewritten:
|
||||
data["mcp_arguments"] = {"q": "Transformed arguments " * 100}
|
||||
monkeypatch.setattr(litellm, "callbacks", [handler])
|
||||
await logger.pre_call_hook(user_api_key_dict=caller, data=data, call_type="call_mcp_tool")
|
||||
|
||||
stash: Final = get_request_stash()
|
||||
assert stash is not None
|
||||
|
|
@ -221,11 +228,7 @@ async def test_mcp_description_preserves_project_input_and_output_reservations(
|
|||
assert messages == [
|
||||
{
|
||||
"role": "user",
|
||||
"content": (
|
||||
f"Tool: echo\nDescription: {description}\nArguments: {{'q': 'hello'}}"
|
||||
if description
|
||||
else "Tool: echo\nArguments: {'q': 'hello'}"
|
||||
),
|
||||
"content": "Tool: echo\nArguments: {'q': 'hello'}",
|
||||
}
|
||||
]
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue