diff --git a/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py b/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py index d99896af513..aadbd3f900f 100644 --- a/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py +++ b/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py @@ -577,21 +577,25 @@ _INDEX_TTL_SECONDS: Final = 5.0 # branch's own completion checks the identical instance's # `has_logged_{event_type}` flag (`Logging.should_run_logging`) before # dispatching to any registered `CustomLogger`, and whichever branch gets -# there first flips it for all the others. Reserving one unit per matching -# branch and relying on that many independent releases was the wrong model: -# with only one release ever happening, every other branch's own -# reservation would leak until `_CONCURRENCY_MIN_SAFETY_TTL_SECONDS` on -# every ordinary multi-model batch call, not just a race. Admission -# (`async_filter_deployments`, via `_pending_concurrency_keys`) instead -# reserves at most one unit per key for the whole dispatch regardless of -# how many branches match it -- every branch after the first rides along on -# that one reservation for free -- so the single release that does happen -# always exactly balances what was reserved. Each entry also carries an -# admission-scoped token (see `_current_admission_token`), used only by +# there first flips it for all the others. +# +# That single-terminal-event fact is independent of how many branches are +# genuinely, concurrently in flight at the provider level -- also confirmed +# live (wall-clock: three branches with a 1s mock delay each, including two +# dialing the identical deployment, complete in ~1s total, not serially). +# A concurrency limit exists to cap that real simultaneous load, so +# admission must still reserve one unit per branch that actually admits, +# never deduped by key just because a sibling already holds one -- two +# branches racing the same deployment are two real concurrent calls, and +# collapsing them to one reservation would undercount exactly the load the +# limit exists to cap. What changes for the single terminal event is only +# the release side: since nothing else will ever come along afterward to +# release anything else, that one event releases everything still pending +# for the call in one shot, not "this branch's own key" -- there is no +# second, later event to race against by doing so. Each entry also carries +# an admission-scoped token (see `_current_admission_token`), used only by # `_release_stale_hop_reservations`'s *own* admission-time cleanup of a -# prior hop of the identical serial fallback chain, never by the terminal -# release hooks, which release everything still pending unconditionally -- -# safe now that at most one reservation per key ever exists at a time. +# prior hop of the identical serial fallback chain. _PENDING_CONCURRENCY_KEYS_FIELD: Final[str] = "_model_based_tag_rate_limits_pending_concurrency_keys" # Identifies which admission call queued a given reservation, scoped by @@ -1011,42 +1015,6 @@ def _queue_pending_reservations( ) # mutable-ok: see comment above -def _pending_concurrency_keys(request_kwargs: Mapping[str, object]) -> frozenset[str]: - """Every concurrency key already queued for this call, by any branch of - an `abatch_completion` dispatch -- see `_PENDING_CONCURRENCY_KEYS_FIELD`'s - docstring for why a dispatch reserves at most one unit per key regardless - of how many branches match it.""" - logging_obj: Final = request_kwargs.get("litellm_logging_obj") - model_call_details: Final = getattr(logging_obj, "model_call_details", None) - if not isinstance(model_call_details, dict): - return frozenset() - pending: Final = model_call_details.get(_PENDING_CONCURRENCY_KEYS_FIELD) - if not isinstance(pending, list): - return frozenset() - return frozenset(entry[0] for entry in pending) - - -def _discard_pending_concurrency_keys(request_kwargs: Mapping[str, object], keys: Iterable[str]) -> None: - """Rolls back this hop's own just-staked claim(s) for `keys` -- used - when the atomic batch they were staked ahead of ends up rejected, since - `_atomic_check_and_increment`'s all-or-nothing contract means nothing - was actually incremented for them; leaving the claim in place would - both corrupt release-time bookkeeping (nothing to release against) and - make a genuinely later sibling wrongly skip its own real reservation, - believing this one already covers it.""" - logging_obj: Final = request_kwargs.get("litellm_logging_obj") - model_call_details: Final = getattr(logging_obj, "model_call_details", None) - if not isinstance(model_call_details, dict): - return - pending: Final = model_call_details.get(_PENDING_CONCURRENCY_KEYS_FIELD) - if not isinstance(pending, list): - return - for key in keys: - entry = next((candidate for candidate in pending if candidate[0] == key), None) - if entry is not None: - pending.remove(entry) # mutable-ok: shared, request-scoped accumulator; see field's own docstring - - def _record_admission_time(request_kwargs: Mapping[str, object], model_group: str, now: float) -> None: """Stash this hop's admission timestamp under its own model_group -- see `_ADMISSION_TIME_FIELD`'s docstring for why keyed, not scalar. Silently a @@ -1366,7 +1334,7 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass] key_alias: Final = _extract_key_alias(resolved_request_kwargs, metadata_variable_name) now: Final = self._time_provider().timestamp() _record_admission_time(resolved_request_kwargs, model, now) - raw_classified: Final = tuple( + classified: Final = tuple( check for configured_limit in configured if ( @@ -1383,34 +1351,6 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass] ) is not None ) - # A dispatch reserves at most one concurrency unit per key regardless - # of how many abatch_completion branches match it -- see - # _PENDING_CONCURRENCY_KEYS_FIELD's own docstring for why. Staked - # synchronously below, before this function's next `await`: asyncio - # only switches tasks at an `await`, so a sibling branch admitting - # concurrently can never observe this hop mid-decision, only either - # fully before or fully after it -- closing the race that would - # otherwise let two branches both see "not yet reserved" and each - # stake their own. - already_pending_keys: Final = _pending_concurrency_keys(resolved_request_kwargs) - own_new_concurrency_claims: Final = [] # mutable-ok: staked synchronously below before any further await; see comment above - deduped_classified: Final = [] # mutable-ok: see comment above - for check in raw_classified: - # not Final: rebound each loop iteration - already_claimed = check.configured_limit.unit == "concurrency" and ( - check.key in already_pending_keys or any(key == check.key for key, _ in own_new_concurrency_claims) - ) - if already_claimed: - continue - if check.configured_limit.unit == "concurrency": - own_new_concurrency_claims.append((check.key, _partition_key(check.configured_limit.entry))) - deduped_classified.append(check) - classified: Final = tuple(deduped_classified) - if own_new_concurrency_claims: - _queue_pending_reservations( - resolved_request_kwargs, _PENDING_CONCURRENCY_KEYS_FIELD, own_new_concurrency_claims - ) - read_only_checks: Final = tuple((c.configured_limit, c.tag_value, c.key) for c in classified if not c.is_atomic) atomic_checks: Final = tuple((c.configured_limit, c.tag_value, c.key) for c in classified if c.is_atomic) @@ -1446,23 +1386,22 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass] ) ) if failing_index is not None: - if own_new_concurrency_claims: - # Nothing was actually incremented for these (the whole - # batch is all-or-nothing), so the claim staked above - # must be rolled back -- otherwise it would both corrupt - # release-time bookkeeping and make a genuinely later - # sibling wrongly skip its own real reservation. - _discard_pending_concurrency_keys( - resolved_request_kwargs, (key for key, _partition_key in own_new_concurrency_claims) - ) failing_limit, failing_tag_value, _ = atomic_checks[failing_index] self._raise_over_limit(failing_limit, failing_tag_value, model, current=values[0]) - if own_new_concurrency_claims: + concurrency_reservations: Final = tuple( + (key, _partition_key(configured_limit.entry)) + for configured_limit, _tag_value, key in atomic_checks + if configured_limit.unit == "concurrency" + ) + if concurrency_reservations: + _queue_pending_reservations( + resolved_request_kwargs, _PENDING_CONCURRENCY_KEYS_FIELD, concurrency_reservations + ) await self._mirror_pending_reservations( resolved_request_kwargs.get("litellm_call_id"), _extract_key_hash(resolved_request_kwargs, metadata_variable_name), - tuple(own_new_concurrency_claims), + concurrency_reservations, ) # Only genuinely new keys, never one already in stale_request_keys: @@ -1747,10 +1686,11 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # docstring for why that distinction, not just presence, decides # what's actually stale. The terminal release hooks # (`async_log_success_event`/`async_log_failure_event`) never pass - # it: with admission now reserving at most one unit per key for the - # whole dispatch (see `_PENDING_CONCURRENCY_KEYS_FIELD`'s docstring), - # everything still pending when the one terminal event for this - # dispatch fires is safe to release unconditionally. + # it: only one terminal event ever fires per call (see + # `_PENDING_CONCURRENCY_KEYS_FIELD`'s docstring), so nothing else + # will ever come along afterward to release anything else -- + # everything still pending when it fires is safe to release + # unconditionally, however many branches actually admitted. pending: Final = kwargs.get(_PENDING_CONCURRENCY_KEYS_FIELD) if not isinstance(pending, list) or not pending: return () @@ -2022,9 +1962,9 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # # Unconditional, not filtered to this hop's own lineage: only one # terminal event ever fires per abatch_completion dispatch (see - # _PENDING_CONCURRENCY_KEYS_FIELD's docstring), and admission now - # reserves at most one unit per key for the whole dispatch, so - # whatever's still pending here is exactly that one reservation. + # _PENDING_CONCURRENCY_KEYS_FIELD's docstring), so nothing else will + # ever come along afterward to release whatever every admitted + # branch of this dispatch actually reserved. release_keys: Final = await self._pop_pending_concurrency_keys(kwargs) if release_keys: await self._release_keys(release_keys) diff --git a/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py b/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py index 306e1ed03e6..b659ff2eace 100644 --- a/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py +++ b/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py @@ -26,10 +26,8 @@ from litellm.proxy.hooks.model_based_tag_rate_limits_hook import ( _build_limits_index, _ConfiguredLimit, _current_admission_token, - _discard_pending_concurrency_keys, _extract_team_id, _inflight_key, - _pending_concurrency_keys, _pending_reservations_cache_key, _PROXY_ModelBasedTagRateLimitsHook, _record_admission_time, @@ -3168,22 +3166,22 @@ async def test_next_hops_admission_releases_a_prior_hops_leaked_reservation(time @pytest.mark.asyncio -async def test_concurrent_batch_siblings_dedupe_to_exactly_one_concurrency_reservation(time_controller): +async def test_concurrent_batch_siblings_each_reserve_their_own_unit_and_a_tight_cap_still_rejects(time_controller): """ - Router.abatch_completion's comma-separated multi-model dispatch runs - each model concurrently as its own asyncio.Task, but every branch is - handed the identical litellm_logging_obj (the proxy attaches one to the - request before the comma-split), so two genuinely concurrent branches - share one model_call_details -- and, confirmed live, that shared - Logging instance means only ONE terminal success/failure event will - ever fire for the whole dispatch (litellm's own has_logged_{event_type} - dedup), never one per branch. Reserving one unit per branch and relying - on that many independent releases was the wrong model: with only one - release ever happening, every other branch's own reservation would leak - until its safety TTL on every ordinary multi-model batch call. Admission - instead reserves at most one unit per key for the whole dispatch: both - branches here admit successfully under a limit of 1, and only one real - reservation exists regardless. + Veria AI finding, independently verified by wall-clock timing: three + abatch_completion branches with a 1s mock delay each complete in ~1s + total, not serially, even when two dial the identical deployment -- + concurrency and terminal-event-count are two separate things. Only ONE + terminal event ever fires for the whole dispatch (litellm's own + has_logged_{event_type} dedup on the shared Logging instance every + branch reuses), but that does not mean the dispatch only ever needs + one concurrency unit: a concurrency limit exists to cap real + simultaneous load, so admission must reserve one unit per branch that + actually admits, never deduped by key just because a sibling of the + same dispatch already holds one. Two branches racing the identical + deployment under a limit of 1 are two real concurrent calls, so the + second is correctly rejected, exactly like two genuinely separate + concurrent requests would be. """ limiter = _make_limiter(time_controller) router = _concurrency_router(limit=1) @@ -3200,19 +3198,22 @@ async def test_concurrent_batch_siblings_dedupe_to_exactly_one_concurrency_reser asyncio.create_task(_admit()), asyncio.create_task(_admit()), return_exceptions=True ) rejections = [result for result in results if isinstance(result, ProxyRateLimitError)] - assert rejections == [] - assert len(_pending_concurrency_keys(request_kwargs)) == 1 + assert len(rejections) == 1 @pytest.mark.asyncio -async def test_three_batch_siblings_still_dedupe_to_exactly_one_concurrency_reservation(time_controller): - """A dispatch's own reservation count does not scale with how many - models are in its comma-separated list.""" +async def test_three_concurrent_batch_siblings_each_reserve_their_own_unit(time_controller): + """ + A dispatch's own reservation count scales with how many branches + genuinely admit, not with dispatch width collapsed to one, and not + with key uniqueness: three branches targeting the identical deployment + under a limit of 3 each reserve their own real unit. + """ limiter = _make_limiter(time_controller) - router = _concurrency_router(limit=1) + router = _concurrency_router(limit=3) limiter.update_variables(llm_router=router) healthy = router.model_list - request_kwargs, _kwargs = _call_context(["end_user_id:u1"]) + request_kwargs, kwargs = _call_context(["end_user_id:u1"]) async def _admit() -> None: await limiter.async_filter_deployments( @@ -3224,68 +3225,22 @@ async def test_three_batch_siblings_still_dedupe_to_exactly_one_concurrency_rese ) rejections = [result for result in results if isinstance(result, ProxyRateLimitError)] assert rejections == [] - assert len(_pending_concurrency_keys(request_kwargs)) == 1 + pending = kwargs.get(_PENDING_CONCURRENCY_KEYS_FIELD) + assert pending is not None and len(pending) == 3 @pytest.mark.asyncio -async def test_rejected_atomic_batch_rolls_back_its_own_newly_staked_concurrency_claim(time_controller): - """ - A concurrency claim is staked synchronously before the atomic - check-and-increment it's part of even runs, to close the race two - concurrent siblings could otherwise hit (see - _PENDING_CONCURRENCY_KEYS_FIELD's docstring). If that batch then - rejects because of a *different* check in it, nothing was actually - incremented for the concurrency claim either -- the whole batch is - all-or-nothing -- so the staked claim must be rolled back. Left in - place, it would falsely tell a later admission this key is already - covered by a real reservation that never actually happened. - """ - limiter = _make_limiter(time_controller) - router = litellm.Router( - model_list=[ - _deployment( - "grp", - "dep-1", - { - "concurrency_limits": { - "limits": [{"name": "inflight", "tag_id": "shared_pool", "limit": 2, "period_seconds": 300}] - }, - "request_limits": { - "limits": [{"name": "per_period", "tag_id": "end_user_id", "limit": 1, "period_seconds": 300}] - }, - }, - ) - ] - ) - limiter.update_variables(llm_router=router) - healthy = router.model_list - - first_request_kwargs, _first_kwargs = _call_context(["end_user_id:u1", "shared_pool:pool-a"]) - await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs=first_request_kwargs - ) - - second_request_kwargs, _second_kwargs = _call_context(["end_user_id:u1", "shared_pool:pool-a"]) - with pytest.raises(ProxyRateLimitError): - await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs=second_request_kwargs - ) - - assert _pending_concurrency_keys(second_request_kwargs) == frozenset() - - -@pytest.mark.asyncio -async def test_the_dispatchs_one_surviving_success_event_fully_releases_its_single_reservation(time_controller): +async def test_the_dispatchs_one_surviving_success_event_releases_every_branchs_own_reservation(time_controller): """ Only one terminal event -- success or failure -- ever fires for the - whole abatch_completion dispatch. Since admission reserves exactly one - unit for the dispatch regardless of how many branches matched it, that - one event's own unconditional release always exactly balances it: - whichever branch's data happens to be reflected when it fires, the - concurrency slot is fully freed for the request as a whole. + whole abatch_completion dispatch. Since admission reserves one real + unit per branch that actually admitted (not deduped to one for the + whole dispatch), that one event's own unconditional release must free + every one of them together, not just one: there is no second, later + event that could ever release the rest. """ limiter = _make_limiter(time_controller) - router = _concurrency_router(limit=1) + router = _concurrency_router(limit=2) limiter.update_variables(llm_router=router) healthy = router.model_list request_kwargs, kwargs = _call_context(["end_user_id:u1"]) @@ -3297,7 +3252,7 @@ async def test_the_dispatchs_one_surviving_success_event_fully_releases_its_sing await asyncio.create_task(_admit()) await asyncio.create_task(_admit()) - assert len(_pending_concurrency_keys(request_kwargs)) == 1 + assert len(kwargs.get(_PENDING_CONCURRENCY_KEYS_FIELD)) == 2 kwargs["standard_logging_object"] = { "model_group": "grp", @@ -3308,7 +3263,15 @@ async def test_the_dispatchs_one_surviving_success_event_fully_releases_its_sing await limiter.async_log_success_event(kwargs=kwargs, response_obj=None, start_time=0, end_time=0) await asyncio.sleep(0) - assert _pending_concurrency_keys(request_kwargs) == frozenset() + assert not kwargs.get(_PENDING_CONCURRENCY_KEYS_FIELD) + # Both units are free again: two fresh requests both fit under the + # limit=2 cap, proving neither branch's own unit leaked. + await limiter.async_filter_deployments( + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, + ) await limiter.async_filter_deployments( model="grp", healthy_deployments=healthy, @@ -3318,10 +3281,10 @@ async def test_the_dispatchs_one_surviving_success_event_fully_releases_its_sing @pytest.mark.asyncio -async def test_the_dispatchs_one_surviving_failure_event_fully_releases_its_single_reservation(time_controller): +async def test_the_dispatchs_one_surviving_failure_event_releases_every_branchs_own_reservation(time_controller): """Same as the success-event version above, but for the failure path.""" limiter = _make_limiter(time_controller) - router = _concurrency_router(limit=1) + router = _concurrency_router(limit=2) limiter.update_variables(llm_router=router) healthy = router.model_list request_kwargs, kwargs = _call_context(["end_user_id:u1"]) @@ -3333,7 +3296,7 @@ async def test_the_dispatchs_one_surviving_failure_event_fully_releases_its_sing await asyncio.create_task(_admit()) await asyncio.create_task(_admit()) - assert len(_pending_concurrency_keys(request_kwargs)) == 1 + assert len(kwargs.get(_PENDING_CONCURRENCY_KEYS_FIELD)) == 2 kwargs["standard_logging_object"] = { "model_group": "grp", @@ -3343,7 +3306,13 @@ async def test_the_dispatchs_one_surviving_failure_event_fully_releases_its_sing } await limiter.async_log_failure_event(kwargs=kwargs, response_obj=None, start_time=0, end_time=0) - assert _pending_concurrency_keys(request_kwargs) == frozenset() + assert not kwargs.get(_PENDING_CONCURRENCY_KEYS_FIELD) + await limiter.async_filter_deployments( + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, + ) await limiter.async_filter_deployments( model="grp", healthy_deployments=healthy, @@ -3415,7 +3384,7 @@ async def test_a_hop_matching_two_concurrency_scoped_entries_releases_both_reser @pytest.mark.asyncio -async def test_real_pipeline_abatch_completion_reserves_and_releases_exactly_one_concurrency_unit( +async def test_real_pipeline_abatch_completion_reserves_one_unit_per_branch_and_releases_them_all( time_controller, monkeypatch ): """ @@ -3428,17 +3397,21 @@ async def test_real_pipeline_abatch_completion_reserves_and_releases_exactly_one ONE terminal success/failure event ever fires for the whole dispatch, never one per branch, since every branch checks the identical instance's has_logged_{event_type} flag before dispatching to any - CustomLogger. This drives the real hook (registered in litellm.callbacks, + CustomLogger -- but also confirmed live (wall-clock timing) that the + branches themselves genuinely run concurrently at the provider level + regardless. This drives the real hook (registered in litellm.callbacks, invoked by Router.async_callback_filter_deployments and litellm's own logging worker, not called directly) through a real 3-branch dispatch - of the same model against a concurrency limit of 1: without dedup, the - second and third branches would each be rejected outright (confirmed - empirically without the admission-side fix); with it, all three succeed, - and the single reservation is fully released once the one surviving - terminal event fires -- checked after an explicit flush of litellm's - own async logging worker, not a fixed sleep, since that worker (unlike - this hook's own success/failure hooks in the other tests here) runs the - completion callback on its own background task, not inline. + of the same model against a concurrency limit of 2: if admission + deduped by key instead of reserving one real unit per branch, all + three would wrongly admit (one reservation is well within a limit of + 2); reserving genuinely means exactly one of the three is rejected, + and the two real reservations that did admit are released together by + the single surviving terminal event -- checked after an explicit flush + of litellm's own async logging worker, not a fixed sleep, since that + worker (unlike this hook's own success/failure hooks in the other + tests here) runs the completion callback on its own background task, + not inline. """ from unittest.mock import AsyncMock, MagicMock @@ -3450,7 +3423,32 @@ async def test_real_pipeline_abatch_completion_reserves_and_releases_exactly_one from litellm.proxy.utils import ProxyLogging limiter = _make_limiter(time_controller) - router = _concurrency_router(limit=1) + # mock_delay keeps every branch genuinely in flight at once, the same + # way the wall-clock verification did: without it, a branch's own + # near-instant mock completion can release its reservation before a + # later branch even reaches its own admission, masking the bug this + # test exists to catch. num_retries=0 keeps the rejected branch's own + # 429 as the final result instead of Router transparently retrying it + # (confirmed live: a retried rejection can still eventually succeed + # once an earlier branch's own release frees a slot, which would mask + # the same bug from a different angle). + router = litellm.Router( + model_list=[ + { + "model_name": "grp", + "litellm_params": {"model": "gpt-4o", "mock_response": "ok", "mock_delay": 0.2}, + "model_info": { + "id": "dep-1", + "tag_rate_limits": { + "concurrency_limits": { + "limits": [{"name": "inflight", "tag_id": "end_user_id", "limit": 2, "period_seconds": 300}] + } + }, + }, + } + ], + num_retries=0, + ) limiter.update_variables(llm_router=router) monkeypatch.setattr(litellm, "callbacks", [limiter]) @@ -3495,12 +3493,15 @@ async def test_real_pipeline_abatch_completion_reserves_and_releases_exactly_one if asyncio.iscoroutine(resolved): resolved = await resolved assert len(resolved) == 3 - assert not any(isinstance(branch_result, Exception) for branch_result in resolved) + rejections = [branch_result for branch_result in resolved if isinstance(branch_result, Exception)] + assert len(rejections) == 1 await GLOBAL_LOGGING_WORKER.flush() await asyncio.sleep(0.2) - assert _pending_concurrency_keys(data) == frozenset() + logging_obj = data.get("litellm_logging_obj") + model_call_details = getattr(logging_obj, "model_call_details", None) + assert not model_call_details.get(_PENDING_CONCURRENCY_KEYS_FIELD) @pytest.mark.asyncio