mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): preserve base TPM reservations for listed tool calls
This commit is contained in:
parent
7f7df081b9
commit
29baaf3f5e
2 changed files with 91 additions and 2 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue