diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index b60a6c99fd0..b0410e87103 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 3f31ec1a1dd..c8e12ce370a 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -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 diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 999f9ea7331..b33cea5742d 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -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 = ( diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index cf73fc0b99c..df7d859b15a 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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}" ) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_operations.py b/tests/unit/proxy/_experimental/mcp_server/test_operations.py index 47c146ae1bd..fc1e23229f9 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_operations.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_operations.py @@ -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 diff --git a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py index 5682c73503e..d3a723fe1d3 100644 --- a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py @@ -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'}", } ]