fix(mcp): preserve base TPM reservations for listed tool calls

This commit is contained in:
Devin AI 2026-10-02 21:43:56 +00:00
parent 7f7df081b9
commit 29baaf3f5e
2 changed files with 91 additions and 2 deletions

View file

@ -1028,7 +1028,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
def _estimate_tokens_for_request(
self,
data: dict,
data: dict[str, object],
model: str | None = None,
min_configured_tpm_limit: int | None = None,
call_type: str | None = None,
@ -1054,8 +1054,21 @@ 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
)
estimated_input_tokens, max_tokens_estimate = self._estimate_input_and_output_tokens(
data=data,
data=reservation_data,
min_configured_tpm_limit=min_configured_tpm_limit,
call_type=call_type,
configured_output_tokens=configured_output_tokens,

View file

@ -13,6 +13,7 @@ from typing import Any, Dict, Final, List, Optional
import pytest
from fastapi import HTTPException
from pydantic import TypeAdapter
import litellm
from litellm import Router
@ -21,6 +22,7 @@ from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
PARALLEL_REQUEST_SLOT_TTL_SECONDS,
ParallelSlotAcquisition,
@ -39,6 +41,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import (
from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.mcp import MCPPreCallRequestObject
from litellm.types.utils import (
EmbeddingResponse,
ModelResponse,
@ -108,6 +111,79 @@ def test_api_key_descriptor_applies_budget_throttle(
assert api_key_descriptor["rate_limit"]["tokens_per_unit"] == expected_tpm
@pytest.mark.asyncio
@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:
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"]
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")
stash: Final = get_request_stash()
assert stash is not None
assert stash.reserved_tokens == 25
assert (
await cache.async_get_cache(
key=handler.create_rate_limit_keys("api_key", caller.api_key, "tokens"), local_only=True
)
== 25
)
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]] = {
"messages": [{"role": "user", "content": "x" * 400}],
"max_tokens": 1,
"mcp_tool_name": "echo",
"mcp_arguments": {},
}
assert handler._estimate_tokens_for_request(data, call_type="acompletion") == 101
@pytest.mark.asyncio
async def test_unconverted_mcp_request_keeps_its_reservation() -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
caller: Final = UserAPIKeyAuth(api_key=hash_token("sk-raw-mcp-request"), tpm_limit=64)
data: Final[dict[str, object]] = {"name": "echo", "arguments": {"q": "hello"}, "server_id": "fixture"}
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 == 16
assert (
await cache.async_get_cache(
key=handler.create_rate_limit_keys("api_key", caller.api_key, "tokens"), local_only=True
)
== 16
)
@pytest.mark.flaky(reruns=3)
@pytest.mark.asyncio
async def test_sliding_window_rate_limit_v3(monkeypatch, time_controller):