mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(mcp): release max_parallel_requests slot after MCP tool calls
MCP tool calls run the proxy pre-call hooks, which acquire a parallel-request slot, but they never fire the litellm success/failure logging callbacks that release it for LLM requests. Every tool call leaked one slot, so a key with max_parallel_requests: N started 429ing after N strictly sequential calls and stayed wedged until the slot TTL expired or the proxy restarted. The tool call path now owns the release: pre_call_tool_check hands the synthetic pre-call data back to its callers, which release the slot once the call is done, including on upstream failures and on a guardrail rejecting the call after the limiter admitted it.
This commit is contained in:
parent
35dc982692
commit
9b6a40e541
8 changed files with 423 additions and 136 deletions
|
|
@ -136,6 +136,9 @@ from litellm.proxy._types import (
|
|||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
|
||||
from litellm.proxy.common_utils.user_api_key_cache import get_management_object_ttl
|
||||
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
||||
_PROXY_MaxParallelRequestsHandler_v3,
|
||||
)
|
||||
from litellm.proxy.utils import ProxyLogging, get_server_root_path
|
||||
from litellm.repositories.table_repositories import MCPServerRepository
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
|
@ -621,6 +624,30 @@ def _openapi_forwarded_extra_headers(
|
|||
return forwarded or None
|
||||
|
||||
|
||||
MCP_PRE_CALL_DATA_KEY = "pre_call_data"
|
||||
|
||||
|
||||
async def release_parallel_request_slot(
|
||||
proxy_logging_obj: ProxyLogging | None,
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
pre_call_data: dict[str, Any] | None,
|
||||
) -> None:
|
||||
"""
|
||||
Release the ``max_parallel_requests`` slot the MCP pre-call hooks acquired.
|
||||
|
||||
The pre-call hooks run against a synthetic request dict that never reaches
|
||||
the litellm success/failure logging callbacks, which is where LLM requests
|
||||
release their slot, so the MCP tool call path owns the release itself. A
|
||||
no-op when no slot was acquired.
|
||||
"""
|
||||
if proxy_logging_obj is None or user_api_key_auth is None or pre_call_data is None:
|
||||
return
|
||||
limiter = proxy_logging_obj.get_proxy_hook("parallel_request_limiter")
|
||||
if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3):
|
||||
return
|
||||
await limiter.async_release_max_parallel_requests_slot(user_api_key_auth, pre_call_data)
|
||||
|
||||
|
||||
async def _resolve_byok_mcp_auth_header(
|
||||
mcp_server: MCPServer,
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth],
|
||||
|
|
@ -4518,8 +4545,14 @@ class MCPServerManager:
|
|||
) as e:
|
||||
# Re-raise guardrail exceptions to properly fail the MCP call
|
||||
verbose_logger.error(f"Guardrail blocked MCP tool call pre call: {str(e)}")
|
||||
await release_parallel_request_slot(
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
pre_call_data=synthetic_llm_data,
|
||||
)
|
||||
raise e
|
||||
|
||||
hook_result[MCP_PRE_CALL_DATA_KEY] = synthetic_llm_data
|
||||
return hook_result
|
||||
|
||||
def _create_during_hook_task(
|
||||
|
|
@ -5122,6 +5155,7 @@ class MCPServerManager:
|
|||
# Using standard pre_call_hook
|
||||
#########################################################
|
||||
hook_result: dict[str, Any] = {}
|
||||
pre_call_data: dict[str, Any] | None = None
|
||||
if proxy_logging_obj:
|
||||
hook_result = await self.pre_call_tool_check(
|
||||
name=name,
|
||||
|
|
@ -5132,86 +5166,96 @@ class MCPServerManager:
|
|||
server=mcp_server,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
pre_call_data = hook_result.pop(MCP_PRE_CALL_DATA_KEY, None)
|
||||
if "arguments" in hook_result:
|
||||
arguments = hook_result["arguments"]
|
||||
|
||||
# Prepare tasks for during hooks
|
||||
tasks = []
|
||||
if proxy_logging_obj:
|
||||
during_hook_task = self._create_during_hook_task(
|
||||
name=name,
|
||||
arguments=arguments,
|
||||
server_name_from_prefix=server_name,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
start_time=start_time,
|
||||
try:
|
||||
# Prepare tasks for during hooks
|
||||
tasks = []
|
||||
if proxy_logging_obj:
|
||||
during_hook_task = self._create_during_hook_task(
|
||||
name=name,
|
||||
arguments=arguments,
|
||||
server_name_from_prefix=server_name,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
start_time=start_time,
|
||||
)
|
||||
tasks.append(during_hook_task)
|
||||
|
||||
caller_oauth2_headers = oauth2_headers
|
||||
oauth2_headers = await self._resolve_oauth2_headers_for_tool_call(
|
||||
mcp_server, oauth2_headers, user_api_key_auth
|
||||
)
|
||||
tasks.append(during_hook_task)
|
||||
|
||||
caller_oauth2_headers = oauth2_headers
|
||||
oauth2_headers = await self._resolve_oauth2_headers_for_tool_call(mcp_server, oauth2_headers, user_api_key_auth)
|
||||
# For OpenAPI servers, call the tool handler directly instead of via MCP client
|
||||
if mcp_server.spec_path:
|
||||
verbose_logger.debug("Calling OpenAPI tool %s directly via HTTP handler", name)
|
||||
if hook_result.get("extra_headers"):
|
||||
verbose_logger.warning(
|
||||
"pre_mcp_call hook returned extra_headers for OpenAPI-backed "
|
||||
"MCP server '%s' — header injection is not supported for "
|
||||
"OpenAPI servers; headers will be ignored. Use SSE/HTTP "
|
||||
"transport to enable hook header injection.",
|
||||
server_name,
|
||||
)
|
||||
|
||||
# For OpenAPI servers, call the tool handler directly instead of via MCP client
|
||||
if mcp_server.spec_path:
|
||||
verbose_logger.debug("Calling OpenAPI tool %s directly via HTTP handler", name)
|
||||
if hook_result.get("extra_headers"):
|
||||
verbose_logger.warning(
|
||||
"pre_mcp_call hook returned extra_headers for OpenAPI-backed "
|
||||
"MCP server '%s' — header injection is not supported for "
|
||||
"OpenAPI servers; headers will be ignored. Use SSE/HTTP "
|
||||
"transport to enable hook header injection.",
|
||||
server_name,
|
||||
auth_header_value = (
|
||||
_format_byok_openapi_auth_header(mcp_server, mcp_auth_header) if mcp_auth_header else None
|
||||
)
|
||||
resolved_auth_headers, forwarded_headers = await self.resolve_openapi_upstream_auth(
|
||||
mcp_server=mcp_server,
|
||||
oauth2_headers=caller_oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
forwarded_headers=_openapi_forwarded_extra_headers(mcp_server, raw_headers, user_api_key_auth),
|
||||
)
|
||||
|
||||
auth_header_value = (
|
||||
_format_byok_openapi_auth_header(mcp_server, mcp_auth_header) if mcp_auth_header else None
|
||||
)
|
||||
resolved_auth_headers, forwarded_headers = await self.resolve_openapi_upstream_auth(
|
||||
mcp_server=mcp_server,
|
||||
oauth2_headers=caller_oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
forwarded_headers=_openapi_forwarded_extra_headers(mcp_server, raw_headers, user_api_key_auth),
|
||||
)
|
||||
async def _call_openapi_via_handler():
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
_request_auth_header,
|
||||
_request_extra_headers,
|
||||
_request_resolved_auth_headers,
|
||||
)
|
||||
|
||||
async def _call_openapi_via_handler():
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
_request_auth_header,
|
||||
_request_extra_headers,
|
||||
_request_resolved_auth_headers,
|
||||
auth_token = _request_auth_header.set(auth_header_value)
|
||||
extra_token = _request_extra_headers.set(forwarded_headers)
|
||||
resolved_token = _request_resolved_auth_headers.set(resolved_auth_headers)
|
||||
try:
|
||||
async with self._limit_outbound_concurrency(mcp_server):
|
||||
return await self._call_openapi_tool_handler(mcp_server, name, arguments)
|
||||
finally:
|
||||
_request_auth_header.reset(auth_token)
|
||||
_request_extra_headers.reset(extra_token)
|
||||
_request_resolved_auth_headers.reset(resolved_token)
|
||||
|
||||
tasks.append(asyncio.create_task(_call_openapi_via_handler()))
|
||||
else:
|
||||
return await self._call_regular_mcp_tool(
|
||||
mcp_server=mcp_server,
|
||||
original_tool_name=name,
|
||||
arguments=arguments,
|
||||
tasks=tasks,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
host_progress_callback=host_progress_callback,
|
||||
hook_extra_headers=hook_result.get("extra_headers"),
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
auth_token = _request_auth_header.set(auth_header_value)
|
||||
extra_token = _request_extra_headers.set(forwarded_headers)
|
||||
resolved_token = _request_resolved_auth_headers.set(resolved_auth_headers)
|
||||
try:
|
||||
async with self._limit_outbound_concurrency(mcp_server):
|
||||
return await self._call_openapi_tool_handler(mcp_server, name, arguments)
|
||||
finally:
|
||||
_request_auth_header.reset(auth_token)
|
||||
_request_extra_headers.reset(extra_token)
|
||||
_request_resolved_auth_headers.reset(resolved_token)
|
||||
|
||||
tasks.append(asyncio.create_task(_call_openapi_via_handler()))
|
||||
else:
|
||||
return await self._call_regular_mcp_tool(
|
||||
mcp_server=mcp_server,
|
||||
original_tool_name=name,
|
||||
arguments=arguments,
|
||||
tasks=tasks,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
return await self._gather_openapi_tool_tasks(tasks, proxy_logging_obj)
|
||||
finally:
|
||||
await release_parallel_request_slot(
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
host_progress_callback=host_progress_callback,
|
||||
hook_extra_headers=hook_result.get("extra_headers"),
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
pre_call_data=pre_call_data,
|
||||
)
|
||||
|
||||
return await self._gather_openapi_tool_tasks(tasks, proxy_logging_obj)
|
||||
|
||||
#########################################################
|
||||
# End of Methods that call the upstream MCP servers
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -347,12 +347,14 @@ if MCP_AVAILABLE:
|
|||
outcome_wire_value,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
MCP_PRE_CALL_DATA_KEY,
|
||||
MCPServerManager,
|
||||
_caller_authorization_fans_out,
|
||||
_client_forwarded_authorization_headers,
|
||||
_should_strip_caller_authorization,
|
||||
_without_authorization,
|
||||
global_mcp_server_manager,
|
||||
release_parallel_request_slot,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
_request_auth_header,
|
||||
|
|
@ -2734,73 +2736,81 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
# `pre_call_tool_check` may return guardrail-modified
|
||||
# arguments; honor them on the local path too.
|
||||
if isinstance(hook_result, dict) and "arguments" in hook_result:
|
||||
pre_call_data = hook_result.pop(MCP_PRE_CALL_DATA_KEY, None)
|
||||
if "arguments" in hook_result:
|
||||
arguments = hook_result["arguments"]
|
||||
|
||||
verbose_logger.debug(f"Executing local registry tool: {name}")
|
||||
# For BYOK servers the credential must be injected via a ContextVar
|
||||
# because the tool function has headers baked into its closure.
|
||||
# Pre-format the full Authorization header value using the server's
|
||||
# configured auth_type so the generator doesn't need to know the prefix.
|
||||
auth_header_value: Optional[str] = None
|
||||
if mcp_auth_header:
|
||||
server_auth_type = getattr(mcp_server, "auth_type", None) if mcp_server else None
|
||||
if server_auth_type == MCPAuth.api_key:
|
||||
auth_header_value = f"ApiKey {mcp_auth_header}"
|
||||
elif server_auth_type == MCPAuth.basic:
|
||||
auth_header_value = f"Basic {mcp_auth_header}"
|
||||
else:
|
||||
auth_header_value = f"Bearer {mcp_auth_header}"
|
||||
|
||||
# Forward named client headers to OpenAPI tool upstream requests.
|
||||
# MCPServer.extra_headers lists header names to copy from raw_headers.
|
||||
# The strip decision is centralized in _should_strip_caller_authorization so this
|
||||
# OpenAPI/local path agrees with the managed paths: M2M and the resolver-owned modes
|
||||
# (token_exchange's raw subject token, authorization_code's stored token) must never
|
||||
# have the caller's Authorization forwarded verbatim upstream.
|
||||
forwarded_headers: Optional[Dict[str, str]] = None
|
||||
if mcp_server and mcp_server.extra_headers and raw_headers:
|
||||
normalized_raw = {str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)}
|
||||
skip_caller_authorization = _should_strip_caller_authorization(
|
||||
mcp_server=mcp_server,
|
||||
raw_headers=raw_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
for header_name in mcp_server.extra_headers:
|
||||
if not isinstance(header_name, str):
|
||||
continue
|
||||
if skip_caller_authorization and header_name.lower() == "authorization":
|
||||
continue
|
||||
value = normalized_raw.get(header_name.lower())
|
||||
if value is not None:
|
||||
if forwarded_headers is None:
|
||||
forwarded_headers = {}
|
||||
forwarded_headers[header_name] = value
|
||||
|
||||
resolved_auth_headers: dict[str, str] | None = None
|
||||
if mcp_server:
|
||||
(
|
||||
resolved_auth_headers,
|
||||
forwarded_headers,
|
||||
) = await global_mcp_server_manager.resolve_openapi_upstream_auth(
|
||||
mcp_server=mcp_server,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
forwarded_headers=forwarded_headers,
|
||||
)
|
||||
|
||||
_auth_token = _request_auth_header.set(auth_header_value)
|
||||
_extra_token = _request_extra_headers.set(forwarded_headers)
|
||||
_resolved_token = _request_resolved_auth_headers.set(resolved_auth_headers)
|
||||
try:
|
||||
local_content = await _handle_local_mcp_tool(name, arguments)
|
||||
verbose_logger.debug(f"Executing local registry tool: {name}")
|
||||
# For BYOK servers the credential must be injected via a ContextVar
|
||||
# because the tool function has headers baked into its closure.
|
||||
# Pre-format the full Authorization header value using the server's
|
||||
# configured auth_type so the generator doesn't need to know the prefix.
|
||||
auth_header_value: str | None = None
|
||||
if mcp_auth_header:
|
||||
server_auth_type = getattr(mcp_server, "auth_type", None) if mcp_server else None
|
||||
if server_auth_type == MCPAuth.api_key:
|
||||
auth_header_value = f"ApiKey {mcp_auth_header}"
|
||||
elif server_auth_type == MCPAuth.basic:
|
||||
auth_header_value = f"Basic {mcp_auth_header}"
|
||||
else:
|
||||
auth_header_value = f"Bearer {mcp_auth_header}"
|
||||
|
||||
# Forward named client headers to OpenAPI tool upstream requests.
|
||||
# MCPServer.extra_headers lists header names to copy from raw_headers.
|
||||
# The strip decision is centralized in _should_strip_caller_authorization so this
|
||||
# OpenAPI/local path agrees with the managed paths: M2M and the resolver-owned modes
|
||||
# (token_exchange's raw subject token, authorization_code's stored token) must never
|
||||
# have the caller's Authorization forwarded verbatim upstream.
|
||||
forwarded_headers: dict[str, str] | None = None
|
||||
if mcp_server and mcp_server.extra_headers and raw_headers:
|
||||
normalized_raw = {str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)}
|
||||
skip_caller_authorization = _should_strip_caller_authorization(
|
||||
mcp_server=mcp_server,
|
||||
raw_headers=raw_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
for header_name in mcp_server.extra_headers:
|
||||
if not isinstance(header_name, str):
|
||||
continue
|
||||
if skip_caller_authorization and header_name.lower() == "authorization":
|
||||
continue
|
||||
value = normalized_raw.get(header_name.lower())
|
||||
if value is not None:
|
||||
if forwarded_headers is None:
|
||||
forwarded_headers = {}
|
||||
forwarded_headers[header_name] = value
|
||||
|
||||
resolved_auth_headers: dict[str, str] | None = None
|
||||
if mcp_server:
|
||||
(
|
||||
resolved_auth_headers,
|
||||
forwarded_headers,
|
||||
) = await global_mcp_server_manager.resolve_openapi_upstream_auth(
|
||||
mcp_server=mcp_server,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
forwarded_headers=forwarded_headers,
|
||||
)
|
||||
|
||||
_auth_token = _request_auth_header.set(auth_header_value)
|
||||
_extra_token = _request_extra_headers.set(forwarded_headers)
|
||||
_resolved_token = _request_resolved_auth_headers.set(resolved_auth_headers)
|
||||
try:
|
||||
local_content = await _handle_local_mcp_tool(name, arguments)
|
||||
finally:
|
||||
_request_auth_header.reset(_auth_token)
|
||||
_request_extra_headers.reset(_extra_token)
|
||||
_request_resolved_auth_headers.reset(_resolved_token)
|
||||
response = CallToolResult(content=cast(Any, local_content), isError=False)
|
||||
finally:
|
||||
_request_auth_header.reset(_auth_token)
|
||||
_request_extra_headers.reset(_extra_token)
|
||||
_request_resolved_auth_headers.reset(_resolved_token)
|
||||
response = CallToolResult(content=cast(Any, local_content), isError=False)
|
||||
await release_parallel_request_slot(
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
pre_call_data=pre_call_data,
|
||||
)
|
||||
|
||||
# Try managed MCP server tool (pass the full prefixed name)
|
||||
# Primary and recommended way to use external MCP servers
|
||||
|
|
|
|||
|
|
@ -3378,7 +3378,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error in rate limit failure event: {str(e)}")
|
||||
|
||||
async def async_release_max_parallel_requests_on_disconnect(
|
||||
async def async_release_max_parallel_requests_slot(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
request_data: dict | None = None,
|
||||
|
|
@ -3393,7 +3393,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
client cancels a stream mid-flight, the cancellation surfaces as
|
||||
``asyncio.CancelledError`` / ``GeneratorExit`` and neither callback
|
||||
runs, so without this the slot leaks per cancelled stream until its
|
||||
TTL prunes it. ``request_data`` carries the stashed acquisition;
|
||||
TTL prunes it. MCP tool calls are the other caller: they run the
|
||||
pre-call hooks against a synthetic request dict that no logging
|
||||
callback ever sees, so the tool call path releases its own slot.
|
||||
``request_data`` carries the stashed acquisition;
|
||||
its presence (not the key object's current max_parallel_requests
|
||||
configuration, which can change mid-request) decides whether there
|
||||
is anything to release.
|
||||
|
|
|
|||
|
|
@ -2609,7 +2609,7 @@ class ProxyLogging:
|
|||
limiter = self.get_proxy_hook("parallel_request_limiter")
|
||||
if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3):
|
||||
return
|
||||
await limiter.async_release_max_parallel_requests_on_disconnect(user_api_key_dict, request_data)
|
||||
await limiter.async_release_max_parallel_requests_slot(user_api_key_dict, request_data)
|
||||
|
||||
def _init_response_taking_too_long_task(self, data: Optional[dict] = None):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -17,7 +17,10 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
MCP_PRE_CALL_DATA_KEY,
|
||||
MCPServerManager,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
|
|
@ -118,6 +121,7 @@ class TestPreCallToolCheckReturnsHeaders:
|
|||
server=server,
|
||||
)
|
||||
|
||||
assert result.pop(MCP_PRE_CALL_DATA_KEY) == {"model": "fake"}
|
||||
assert result == {}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -179,6 +183,7 @@ class TestPreCallToolCheckReturnsHeaders:
|
|||
server=server,
|
||||
)
|
||||
|
||||
assert result.pop(MCP_PRE_CALL_DATA_KEY) == {"model": "fake"}
|
||||
assert result == {}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -7726,3 +7726,67 @@ class TestPreemptive401ModeAware:
|
|||
await self._run(delegate, None, has_stored_token=False)
|
||||
assert exc.value.status_code == 401
|
||||
await self._run(delegate, self.LITELLM_KEY_HEADERS, has_stored_token=False)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_registry_tool_calls_release_max_parallel_requests_slot(monkeypatch):
|
||||
"""
|
||||
The local (OpenAPI-backed) dispatch path runs the same pre-call hooks as the
|
||||
managed path, so it acquires a ``max_parallel_requests`` slot too. Nothing
|
||||
downstream releases it, so ``execute_mcp_tool`` must; otherwise a key with
|
||||
``max_parallel_requests: N`` wedges after N sequential tool calls.
|
||||
"""
|
||||
import litellm
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_module
|
||||
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
||||
_PROXY_MaxParallelRequestsHandler_v3,
|
||||
)
|
||||
from litellm.proxy.utils import InternalUsageCache, ProxyLogging
|
||||
|
||||
cache = DualCache()
|
||||
limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(cache))
|
||||
proxy_logging_obj = ProxyLogging(user_api_key_cache=cache)
|
||||
proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = limiter
|
||||
monkeypatch.setattr(litellm, "callbacks", [limiter])
|
||||
|
||||
openapi_server = MCPServer(
|
||||
server_id="openapi-server-id",
|
||||
name="openapi_server",
|
||||
server_name="openapi_server",
|
||||
url="https://example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
spec_path="/fake/openapi.json",
|
||||
)
|
||||
user_api_key_auth = UserAPIKeyAuth(api_key="hashed-key", max_parallel_requests=1)
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj),
|
||||
patch.object(mcp_module.global_mcp_tool_registry, "get_tool", return_value=MagicMock()),
|
||||
patch.object(
|
||||
mcp_module.global_mcp_server_manager,
|
||||
"_get_mcp_server_from_tool_name",
|
||||
return_value=openapi_server,
|
||||
),
|
||||
patch.object(
|
||||
mcp_module.global_mcp_server_manager,
|
||||
"resolve_openapi_upstream_auth",
|
||||
new=AsyncMock(return_value=(None, None)),
|
||||
),
|
||||
patch.object(mcp_module.MCPRequestHandler, "is_tool_allowed", return_value=True),
|
||||
patch.object(
|
||||
mcp_module,
|
||||
"_handle_local_mcp_tool",
|
||||
new=AsyncMock(return_value=[TextContent(type="text", text="ok")]),
|
||||
),
|
||||
):
|
||||
for _ in range(3):
|
||||
result = await mcp_module.execute_mcp_tool(
|
||||
name="openapi_server-tool",
|
||||
arguments={},
|
||||
allowed_mcp_servers=[openapi_server],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
assert result.isError is False
|
||||
|
|
|
|||
|
|
@ -9240,3 +9240,164 @@ class TestDiscoveryFailureLogging:
|
|||
assert "typo_row" in caplog.text
|
||||
assert "authorization_url, token_url" in caplog.text
|
||||
assert "unresolved" in caplog.text
|
||||
|
||||
|
||||
class TestMaxParallelRequestsSlotRelease:
|
||||
"""
|
||||
Every MCP tool call runs the proxy pre-call hooks, which acquire a
|
||||
``max_parallel_requests`` slot. Nothing in the MCP path fires the litellm
|
||||
success/failure logging callbacks that release the slot for LLM requests,
|
||||
so the tool call itself must release it; otherwise a key with
|
||||
``max_parallel_requests: N`` wedges after N strictly sequential tool calls
|
||||
and only a restart (or the hour-long slot TTL) frees it.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _proxy_logging_with_limiter():
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
||||
_PROXY_MaxParallelRequestsHandler_v3,
|
||||
)
|
||||
from litellm.proxy.utils import InternalUsageCache, ProxyLogging
|
||||
|
||||
cache = DualCache()
|
||||
limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(cache))
|
||||
proxy_logging_obj = ProxyLogging(user_api_key_cache=cache)
|
||||
proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = limiter
|
||||
return proxy_logging_obj, limiter
|
||||
|
||||
@staticmethod
|
||||
def _server() -> MCPServer:
|
||||
return MCPServer(
|
||||
server_id="slot-server",
|
||||
name="slot_server",
|
||||
server_name="slot_server",
|
||||
url="https://example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _user_api_key_auth(max_parallel_requests: int):
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
return UserAPIKeyAuth(api_key="hashed-key", max_parallel_requests=max_parallel_requests)
|
||||
|
||||
@staticmethod
|
||||
def _patched_manager(manager: MCPServerManager, server: MCPServer, tool_call):
|
||||
async def fake_create_mcp_client(server, **kwargs):
|
||||
class _Client:
|
||||
async def call_tool(self, params, host_progress_callback=None):
|
||||
return await tool_call()
|
||||
|
||||
return _Client()
|
||||
|
||||
return (
|
||||
patch.object(manager, "_resolve_mcp_server_for_tool_call", return_value=server),
|
||||
patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client),
|
||||
)
|
||||
|
||||
async def _call(self, manager, proxy_logging_obj, user_api_key_auth):
|
||||
return await manager.call_tool(
|
||||
server_name="slot_server",
|
||||
name="tool",
|
||||
arguments={},
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sequential_tool_calls_never_exhaust_the_limit(self, monkeypatch):
|
||||
proxy_logging_obj, limiter = self._proxy_logging_with_limiter()
|
||||
monkeypatch.setattr("litellm.callbacks", [limiter])
|
||||
manager = MCPServerManager()
|
||||
server = self._server()
|
||||
user_api_key_auth = self._user_api_key_auth(max_parallel_requests=1)
|
||||
|
||||
async def tool_call():
|
||||
return "ok"
|
||||
|
||||
resolve_patch, client_patch = self._patched_manager(manager, server, tool_call)
|
||||
with resolve_patch, client_patch:
|
||||
for _ in range(5):
|
||||
await self._call(manager, proxy_logging_obj, user_api_key_auth)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failed_tool_call_releases_its_slot(self, monkeypatch):
|
||||
proxy_logging_obj, limiter = self._proxy_logging_with_limiter()
|
||||
monkeypatch.setattr("litellm.callbacks", [limiter])
|
||||
manager = MCPServerManager()
|
||||
server = self._server()
|
||||
user_api_key_auth = self._user_api_key_auth(max_parallel_requests=1)
|
||||
|
||||
async def failing_tool_call():
|
||||
raise ValueError("upstream exploded")
|
||||
|
||||
resolve_patch, client_patch = self._patched_manager(manager, server, failing_tool_call)
|
||||
with resolve_patch, client_patch:
|
||||
with pytest.raises(ValueError):
|
||||
await self._call(manager, proxy_logging_obj, user_api_key_auth)
|
||||
|
||||
async def tool_call():
|
||||
return "ok"
|
||||
|
||||
resolve_patch, client_patch = self._patched_manager(manager, server, tool_call)
|
||||
with resolve_patch, client_patch:
|
||||
await self._call(manager, proxy_logging_obj, user_api_key_auth)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_guardrail_blocked_tool_call_releases_its_slot(self, monkeypatch):
|
||||
"""A guardrail that rejects the call after the limiter admitted it must
|
||||
not strand the acquired slot."""
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
class _BlockingGuardrail(CustomLogger):
|
||||
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
|
||||
raise HTTPException(status_code=400, detail="blocked by guardrail")
|
||||
|
||||
proxy_logging_obj, limiter = self._proxy_logging_with_limiter()
|
||||
monkeypatch.setattr("litellm.callbacks", [limiter, _BlockingGuardrail()])
|
||||
manager = MCPServerManager()
|
||||
server = self._server()
|
||||
user_api_key_auth = self._user_api_key_auth(max_parallel_requests=1)
|
||||
|
||||
async def tool_call():
|
||||
return "ok"
|
||||
|
||||
resolve_patch, client_patch = self._patched_manager(manager, server, tool_call)
|
||||
with resolve_patch, client_patch:
|
||||
with pytest.raises(HTTPException):
|
||||
await self._call(manager, proxy_logging_obj, user_api_key_auth)
|
||||
|
||||
monkeypatch.setattr("litellm.callbacks", [limiter])
|
||||
resolve_patch, client_patch = self._patched_manager(manager, server, tool_call)
|
||||
with resolve_patch, client_patch:
|
||||
await self._call(manager, proxy_logging_obj, user_api_key_auth)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_tool_calls_still_hit_the_limit(self, monkeypatch):
|
||||
"""Releasing after the call must not turn the limiter off: two tool
|
||||
calls in flight against a limit of 1 still reject the second."""
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
|
||||
proxy_logging_obj, limiter = self._proxy_logging_with_limiter()
|
||||
monkeypatch.setattr("litellm.callbacks", [limiter])
|
||||
manager = MCPServerManager()
|
||||
server = self._server()
|
||||
user_api_key_auth = self._user_api_key_auth(max_parallel_requests=1)
|
||||
|
||||
async def slow_tool_call():
|
||||
await asyncio.sleep(0.2)
|
||||
return "ok"
|
||||
|
||||
resolve_patch, client_patch = self._patched_manager(manager, server, slow_tool_call)
|
||||
with resolve_patch, client_patch:
|
||||
results = await asyncio.gather(
|
||||
self._call(manager, proxy_logging_obj, user_api_key_auth),
|
||||
self._call(manager, proxy_logging_obj, user_api_key_auth),
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
rejections = [r for r in results if isinstance(r, ProxyRateLimitError)]
|
||||
assert len(rejections) == 1
|
||||
assert "max_parallel_requests" in str(rejections[0].detail)
|
||||
|
|
|
|||
|
|
@ -3511,7 +3511,7 @@ async def test_release_max_parallel_requests_on_disconnect_v3():
|
|||
await local_cache.async_get_cache(key=counter_key)
|
||||
) == 1
|
||||
|
||||
await handler.async_release_max_parallel_requests_on_disconnect(
|
||||
await handler.async_release_max_parallel_requests_slot(
|
||||
user_api_key_dict,
|
||||
request_data={
|
||||
"metadata": {
|
||||
|
|
@ -3544,7 +3544,7 @@ async def test_release_on_disconnect_works_when_key_config_changed_v3():
|
|||
counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests"
|
||||
await _seed_max_parallel_requests_slots(local_cache, counter_key, [_TEST_SLOT_ID])
|
||||
|
||||
await handler.async_release_max_parallel_requests_on_disconnect(
|
||||
await handler.async_release_max_parallel_requests_slot(
|
||||
UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=None),
|
||||
request_data={
|
||||
"metadata": {
|
||||
|
|
@ -3886,12 +3886,12 @@ async def test_release_max_parallel_requests_on_disconnect_noop_v3():
|
|||
)
|
||||
counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests"
|
||||
|
||||
await handler.async_release_max_parallel_requests_on_disconnect(
|
||||
await handler.async_release_max_parallel_requests_slot(
|
||||
UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=None)
|
||||
)
|
||||
assert await local_cache.async_get_cache(key=counter_key) is None
|
||||
|
||||
await handler.async_release_max_parallel_requests_on_disconnect(
|
||||
await handler.async_release_max_parallel_requests_slot(
|
||||
UserAPIKeyAuth(api_key=None, max_parallel_requests=5)
|
||||
)
|
||||
assert await local_cache.async_get_cache(key=counter_key) is None
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue