fix(mcp): record served catalogs and preserve call message bytes

This commit is contained in:
Devin AI 2026-10-03 03:08:40 +00:00
parent acc02ce8d4
commit c669328ba1
6 changed files with 148 additions and 63 deletions

View file

@ -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

View file

@ -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

View file

@ -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 = (

View file

@ -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}"
)

View file

@ -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

View file

@ -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'}",
}
]