fix(mcp): preserve project token reservations for listed calls
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled

This commit is contained in:
Devin AI 2026-10-03 00:00:46 +00:00
parent 29baaf3f5e
commit acc02ce8d4
2 changed files with 101 additions and 15 deletions

View file

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

View file

@ -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]] = {