mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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
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:
parent
29baaf3f5e
commit
acc02ce8d4
2 changed files with 101 additions and 15 deletions
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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]] = {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue