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:
Devin AI 2026-07-24 21:13:46 +00:00
parent 35dc982692
commit 9b6a40e541
8 changed files with 423 additions and 136 deletions

View file

@ -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
#########################################################

View file

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

View file

@ -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.

View file

@ -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):
"""

View file

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

View file

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

View file

@ -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)

View file

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