diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index d979d56afa7..999f9ea7331 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -1026,6 +1026,26 @@ 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], @@ -1054,19 +1074,7 @@ 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 = ( - { - **data, - "messages": [ - { - "role": "user", - "content": f"Tool: {data.get('mcp_tool_name')}\nArguments: {data.get('mcp_arguments')}", - } - ], - } - if call_type == CallTypes.call_mcp_tool.value and "mcp_tool_name" in data and "mcp_arguments" in data - else data - ) + 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, min_configured_tpm_limit=min_configured_tpm_limit, @@ -3787,13 +3795,14 @@ 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=data, + data=reservation_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=data, model=requested_model, call_type=call_type + data=reservation_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/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py index 1d8731945b8..5682c73503e 100644 --- a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py @@ -153,6 +153,83 @@ async def test_mcp_description_does_not_change_admission_or_reserved_tokens(desc ] +@pytest.mark.asyncio +@pytest.mark.parametrize( + "description", [None, "Gateway metadata, not caller input. " * 100], ids=["unlisted", "listed"] +) +@pytest.mark.parametrize("itpm_limit,otpm_limit", [(64, 4096), (4096, 64), (4096, 4096)]) +async def test_mcp_description_preserves_project_input_and_output_reservations( + description: str | None, itpm_limit: int, otpm_limit: int +) -> None: + cache: Final = DualCache() + handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache)) + logger: Final = ProxyLogging(user_api_key_cache=UserApiKeyCache()) + schema: Final = {"type": "object", "properties": {"q": {"type": "string", "description": "Schema text " * 100}}} + request: Final = MCPPreCallRequestObject( + tool_name="echo", arguments={"q": "hello"}, tool_description=description, tool_input_schema=schema + ) + data: Final = TypeAdapter(dict[str, object]).validate_python(logger._convert_mcp_to_llm_format(request, {})) + messages: Final = data["messages"] + base_data: Final[dict[str, object]] = { + "messages": [{"role": "user", "content": "Tool: echo\nArguments: {'q': 'hello'}"}] + } + expected_input: Final = handler._estimate_precise_input_tokens(base_data, "mcp-tool-call", "call_mcp_tool") + expected_output: Final = handler.no_max_tokens_output_floor(otpm_limit) + expected_combined: Final = handler._estimate_tokens_for_request( + base_data, min_configured_tpm_limit=4096, call_type="call_mcp_tool" + ) + caller: Final = UserAPIKeyAuth( + api_key=hash_token("sk-mcp-project-reservation"), + tpm_limit=4096, + project_id="mcp-project-reservation", + project_metadata={ + "model_itpm_limit": {"mcp-tool-call": itpm_limit}, + "model_otpm_limit": {"mcp-tool-call": otpm_limit}, + }, + ) + + await handler.async_pre_call_hook(user_api_key_dict=caller, cache=cache, data=data, call_type="call_mcp_tool") + + stash: Final = get_request_stash() + assert stash is not None + assert (stash.reserved_tokens, stash.itpm_reserved_tokens, stash.otpm_reserved_tokens) == ( + expected_combined, + expected_input, + expected_output, + ) + assert ( + await cache.async_get_cache( + key=handler.create_rate_limit_keys( + "model_per_project_itpm", f"{caller.project_id}:mcp-tool-call", "tokens" + ), + local_only=True, + ) + == expected_input + ) + assert ( + await cache.async_get_cache( + key=handler.create_rate_limit_keys( + "model_per_project_otpm", f"{caller.project_id}:mcp-tool-call", "tokens" + ), + local_only=True, + ) + == expected_output + ) + assert data["messages"] is messages + assert data.get("mcp_tool_description") == description + assert data["mcp_input_schema"] == schema + assert messages == [ + { + "role": "user", + "content": ( + f"Tool: echo\nDescription: {description}\nArguments: {{'q': 'hello'}}" + if description + else "Tool: echo\nArguments: {'q': 'hello'}" + ), + } + ] + + def test_llm_tpm_estimation_still_counts_messages_with_mcp_metadata() -> None: handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) data: Final[dict[str, object]] = {