mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(proxy): release max_parallel_requests slot when a stream is cancelled mid-flight (#27955)
This commit is contained in:
parent
51ba6e39cd
commit
b2cf030291
3 changed files with 222 additions and 4 deletions
|
|
@ -2933,6 +2933,42 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
f"Error in rate limit failure event: {str(e)}"
|
||||
)
|
||||
|
||||
async def async_release_max_parallel_requests_on_disconnect(
|
||||
self, user_api_key_dict: UserAPIKeyAuth
|
||||
) -> None:
|
||||
"""
|
||||
Release the api-key ``max_parallel_requests`` slot that
|
||||
``async_pre_call_hook`` reserved, for a request that ended without
|
||||
either logging callback firing.
|
||||
|
||||
The +1 is normally undone by ``async_log_success_event`` (natural
|
||||
stream completion) or ``async_log_failure_event`` (LLM error). When a
|
||||
client cancels a stream mid-flight, the cancellation surfaces as
|
||||
``asyncio.CancelledError`` / ``GeneratorExit`` and neither callback
|
||||
runs, so without this the counter leaks one slot per cancelled stream
|
||||
until the key wedges at its limit.
|
||||
"""
|
||||
if (
|
||||
not user_api_key_dict.api_key
|
||||
or user_api_key_dict.max_parallel_requests is None
|
||||
):
|
||||
return
|
||||
|
||||
await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline(
|
||||
increment_list=[
|
||||
RedisPipelineIncrementOperation(
|
||||
key=self.create_rate_limit_keys(
|
||||
key="api_key",
|
||||
value=user_api_key_dict.api_key,
|
||||
rate_limit_type="max_parallel_requests",
|
||||
),
|
||||
increment_value=-1,
|
||||
ttl=self.window_size,
|
||||
)
|
||||
],
|
||||
litellm_parent_otel_span=None,
|
||||
)
|
||||
|
||||
async def async_post_call_success_hook(
|
||||
self, data: dict, user_api_key_dict: UserAPIKeyAuth, response
|
||||
):
|
||||
|
|
|
|||
|
|
@ -137,6 +137,9 @@ from litellm.proxy.hooks.max_budget_limiter import _PROXY_MaxBudgetLimiter
|
|||
from litellm.proxy.hooks.parallel_request_limiter import (
|
||||
_PROXY_MaxParallelRequestsHandler,
|
||||
)
|
||||
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
||||
_PROXY_MaxParallelRequestsHandler_v3,
|
||||
)
|
||||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
|
||||
from litellm.proxy.policy_engine.pipeline_executor import PipelineExecutor
|
||||
from litellm.repositories.budget_repository import BudgetRepository
|
||||
|
|
@ -2613,8 +2616,12 @@ class ProxyLogging:
|
|||
# through each of them adds N pass-through trampolines per chunk for
|
||||
# zero behavior change. Skip the chain entirely and stream through.
|
||||
if not caps.iterator_overrides:
|
||||
async for chunk in response:
|
||||
yield chunk
|
||||
try:
|
||||
async for chunk in response:
|
||||
yield chunk
|
||||
except (asyncio.CancelledError, GeneratorExit):
|
||||
self._release_max_parallel_requests_on_disconnect(user_api_key_dict)
|
||||
raise
|
||||
ProxyLogging._fire_deferred_stream_logging(request_data)
|
||||
return
|
||||
|
||||
|
|
@ -2658,8 +2665,12 @@ class ProxyLogging:
|
|||
)
|
||||
|
||||
# Actually iterate through the chained async generator and yield chunks
|
||||
async for chunk in current_response:
|
||||
yield chunk
|
||||
try:
|
||||
async for chunk in current_response:
|
||||
yield chunk
|
||||
except (asyncio.CancelledError, GeneratorExit):
|
||||
self._release_max_parallel_requests_on_disconnect(user_api_key_dict)
|
||||
raise
|
||||
|
||||
# Fire deferred logging AFTER all guardrail end-of-stream blocks
|
||||
# completed. unified_guardrail writes guardrail_information during
|
||||
|
|
@ -2687,6 +2698,29 @@ class ProxyLogging:
|
|||
logging_obj._deferred_stream_complete_args = None
|
||||
asyncio.create_task(_deferred_cb(*_args))
|
||||
|
||||
def _release_max_parallel_requests_on_disconnect(
|
||||
self, user_api_key_dict: UserAPIKeyAuth
|
||||
) -> None:
|
||||
"""
|
||||
Release the api-key max_parallel_requests slot when a streaming
|
||||
response is cancelled mid-flight (client disconnect). Neither the
|
||||
success nor failure logging callback fires on the resulting
|
||||
CancelledError / GeneratorExit, so the pre-call +1 would otherwise
|
||||
leak. Scheduled fire-and-forget (no await) because awaiting is not
|
||||
permitted while unwinding a GeneratorExit.
|
||||
"""
|
||||
limiter = self.get_proxy_hook("parallel_request_limiter")
|
||||
if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3):
|
||||
return
|
||||
try:
|
||||
asyncio.create_task(
|
||||
limiter.async_release_max_parallel_requests_on_disconnect(
|
||||
user_api_key_dict
|
||||
)
|
||||
)
|
||||
except RuntimeError:
|
||||
pass
|
||||
|
||||
def _init_response_taking_too_long_task(self, data: Optional[dict] = None):
|
||||
"""
|
||||
Initialize the response taking too long task if user is using slack alerting
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
|||
_PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler,
|
||||
)
|
||||
from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
from litellm.types.utils import ModelResponse, Usage
|
||||
|
||||
|
||||
|
|
@ -3120,3 +3121,150 @@ def test_get_key_mcp_rpm_limit_precedence():
|
|||
none_set = UserAPIKeyAuth(api_key=hash_token("sk-mcp-key"))
|
||||
assert get_key_mcp_rpm_limit(none_set) is None
|
||||
assert get_team_mcp_rpm_limit(none_set) is None
|
||||
|
||||
|
||||
async def _seed_max_parallel_requests_counter(
|
||||
dual_cache: DualCache, counter_key: str, window_size: int
|
||||
) -> None:
|
||||
await dual_cache.async_increment_cache_pipeline(
|
||||
increment_list=[
|
||||
RedisPipelineIncrementOperation(
|
||||
key=counter_key, increment_value=1, ttl=window_size
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_release_max_parallel_requests_on_disconnect_v3():
|
||||
"""
|
||||
Regression for issue #27955: a stream cancelled mid-flight must release the
|
||||
pre-call +1 reservation. The success/failure logging callbacks never fire
|
||||
on cancellation, so without an explicit release the api-key counter climbs
|
||||
by one per cancelled request until the key wedges at its limit. The release
|
||||
must decrement the api-key max_parallel_requests counter by exactly one.
|
||||
"""
|
||||
_api_key = hash_token("sk-12345")
|
||||
local_cache = DualCache()
|
||||
handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(local_cache)
|
||||
)
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=2)
|
||||
counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests"
|
||||
|
||||
await _seed_max_parallel_requests_counter(
|
||||
local_cache, counter_key, handler.window_size
|
||||
)
|
||||
assert await local_cache.async_get_cache(key=counter_key) == 1
|
||||
|
||||
await handler.async_release_max_parallel_requests_on_disconnect(user_api_key_dict)
|
||||
|
||||
assert await local_cache.async_get_cache(key=counter_key) == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_release_max_parallel_requests_on_disconnect_noop_v3():
|
||||
"""
|
||||
The release must be a no-op when the key never reserved a parallel slot
|
||||
(no api_key, or max_parallel_requests unset). Otherwise a cancelled
|
||||
no-limit request would drive an unrelated counter negative.
|
||||
"""
|
||||
_api_key = hash_token("sk-12345")
|
||||
local_cache = DualCache()
|
||||
handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(local_cache)
|
||||
)
|
||||
counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests"
|
||||
|
||||
await handler.async_release_max_parallel_requests_on_disconnect(
|
||||
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(
|
||||
UserAPIKeyAuth(api_key=None, max_parallel_requests=5)
|
||||
)
|
||||
assert await local_cache.async_get_cache(key=counter_key) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_iterator_hook_releases_counter_on_cancel_v3():
|
||||
"""
|
||||
End-to-end regression for issue #27955. When the client cancels a stream
|
||||
mid-flight, async_post_call_streaming_iterator_hook must release the
|
||||
api-key max_parallel_requests reservation and re-raise the CancelledError.
|
||||
Before the fix the counter stayed at +1 because the cancellation slipped
|
||||
past the hook's post-loop cleanup.
|
||||
"""
|
||||
_api_key = hash_token("sk-12345")
|
||||
limiter_cache = DualCache()
|
||||
limiter = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(limiter_cache)
|
||||
)
|
||||
counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests"
|
||||
await _seed_max_parallel_requests_counter(
|
||||
limiter_cache, counter_key, limiter.window_size
|
||||
)
|
||||
assert await limiter_cache.async_get_cache(key=counter_key) == 1
|
||||
|
||||
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
|
||||
proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = limiter
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=2)
|
||||
|
||||
async def cancelling_stream():
|
||||
yield ModelResponse()
|
||||
raise asyncio.CancelledError()
|
||||
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
async for _ in proxy_logging_obj.async_post_call_streaming_iterator_hook(
|
||||
response=cancelling_stream(),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data={},
|
||||
):
|
||||
pass
|
||||
|
||||
# The release is scheduled fire-and-forget via create_task; let it run.
|
||||
for _ in range(5):
|
||||
await asyncio.sleep(0)
|
||||
assert await limiter_cache.async_get_cache(key=counter_key) == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_iterator_hook_releases_counter_on_aclose_v3():
|
||||
"""
|
||||
Companion to the cancel test for issue #27955. When the client disconnects,
|
||||
the nested streaming generators are closed, which throws GeneratorExit (not
|
||||
CancelledError) into the iterator hook. That path must also release the
|
||||
api-key max_parallel_requests reservation.
|
||||
"""
|
||||
_api_key = hash_token("sk-12345")
|
||||
limiter_cache = DualCache()
|
||||
limiter = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(limiter_cache)
|
||||
)
|
||||
counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests"
|
||||
await _seed_max_parallel_requests_counter(
|
||||
limiter_cache, counter_key, limiter.window_size
|
||||
)
|
||||
|
||||
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
|
||||
proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = limiter
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=2)
|
||||
|
||||
async def open_stream():
|
||||
while True:
|
||||
yield ModelResponse()
|
||||
|
||||
hook_iter = proxy_logging_obj.async_post_call_streaming_iterator_hook(
|
||||
response=open_stream(),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data={},
|
||||
)
|
||||
await hook_iter.__anext__()
|
||||
await hook_iter.aclose()
|
||||
|
||||
for _ in range(5):
|
||||
await asyncio.sleep(0)
|
||||
assert await limiter_cache.async_get_cache(key=counter_key) == 0
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue