diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 63129602082..31437af7770 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -4518,12 +4518,25 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): statuses=statuses, ) + def _recovered_partial_usage_tokens(self, source: Mapping[str, object]) -> tuple[int, int, int]: + usage: Final = source.get("combined_usage_object") + if not isinstance(usage, Usage) or (usage.completion_tokens or 0) <= 0: + return 0, 0, 0 + billable_input, completion_tokens, _ = self._resolve_io_token_reconcile_usage(usage) + return ( + self._get_total_tokens_from_usage(usage=usage, rate_limit_type=self.get_rate_limit_type()), + billable_input, + completion_tokens, + ) + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): """ On failure: decrement max_parallel_requests and refund the upfront TPM reservation only against the scopes the reservation actually charged. Unreserved scopes were never incremented at pre-call, so - refunding them would drive their counter negative. + refunding them would drive their counter negative. A failed stream + whose partial usage was recovered settles the reservation at that + usage instead of refunding it. """ from litellm.litellm_core_utils.core_helpers import ( _get_parent_otel_span_from_kwargs, @@ -4552,31 +4565,31 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if stash is None or stash.reservation_released else (stash.reserved_tokens, stash.itpm_reserved_tokens, stash.otpm_reserved_tokens) ) + tpm_actual, itpm_actual, otpm_actual = self._recovered_partial_usage_tokens(kwargs) if stash is not None and reserved_tokens > 0: - verbose_proxy_logger.debug("Releasing reserved TPM tokens on failure: %s", reserved_tokens) - # Refund only against the scopes the reservation actually - # charged. _build_reservation_aware_tpm_ops with - # actual_tokens=0 emits -reserved on reserved scopes and 0 - # on unreserved (skipped), so unreserved scopes can't drift - # negative. + verbose_proxy_logger.debug( + "Settling reserved TPM tokens on failure: reserved=%s actual=%s", reserved_tokens, tpm_actual + ) + # Settle only against the scopes the reservation actually + # charged: unreserved scopes were never incremented, so a + # refund there would drive their counter negative. pipeline_operations.extend( self._build_reservation_aware_tpm_ops( targets=list(stash.reserved_scopes), reserved_scopes=stash.reserved_scopes, - actual_tokens=0, + actual_tokens=tpm_actual, reserved_tokens=reserved_tokens, ) ) - # Refund project ITPM/OTPM reservations the same way -- full - # refund, since a failed call has no billable usage to reconcile - # against. + # Settle project ITPM/OTPM reservations the same way: at the + # recovered partial usage, or a full refund when there is none. itpm_operations: Final = ( self._build_project_reservation_ops( targets=tuple(stash.itpm_reserved_scopes), reserved_scopes=stash.itpm_reserved_scopes, - actual_tokens=0, + actual_tokens=itpm_actual, reserved_tokens=itpm_reserved, reservation_window_identities=stash.itpm_reserved_window_identities, ) @@ -4584,7 +4597,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): else self._build_reservation_aware_tpm_ops( targets=tuple(stash.itpm_reserved_scopes), reserved_scopes=stash.itpm_reserved_scopes, - actual_tokens=0, + actual_tokens=itpm_actual, reserved_tokens=itpm_reserved, ) if stash is not None and itpm_reserved > 0 @@ -4595,7 +4608,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self._build_project_reservation_ops( targets=tuple(stash.otpm_reserved_scopes), reserved_scopes=stash.otpm_reserved_scopes, - actual_tokens=0, + actual_tokens=otpm_actual, reserved_tokens=otpm_reserved, reservation_window_identities=stash.otpm_reserved_window_identities, ) @@ -4603,7 +4616,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): else self._build_reservation_aware_tpm_ops( targets=tuple(stash.otpm_reserved_scopes), reserved_scopes=stash.otpm_reserved_scopes, - actual_tokens=0, + actual_tokens=otpm_actual, reserved_tokens=otpm_reserved, ) if stash is not None and otpm_reserved > 0 @@ -4742,7 +4755,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): removal is a no-op ZREM on a second run), and the TPM/ITPM/OTPM refund is guarded by the stash's ``reservation_released`` flag — if both this hook and async_log_failure_event end up running in the same - flow, only the first release/refund applies. + flow, only the first release/refund applies. A mid-stream failure + relayed here with recovered partial usage settles the reservation at + that usage instead of refunding it. """ try: stash: Final = get_request_stash() @@ -4769,12 +4784,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): otpm_reserved: Final = stash.otpm_reserved_tokens if reserved_tokens <= 0 and itpm_reserved <= 0 and otpm_reserved <= 0: return + tpm_actual, itpm_actual, otpm_actual = self._recovered_partial_usage_tokens(request_data) combined_ops: Final = ( self._build_reservation_aware_tpm_ops( targets=tuple(stash.reserved_scopes), reserved_scopes=stash.reserved_scopes, - actual_tokens=0, + actual_tokens=tpm_actual, reserved_tokens=reserved_tokens, ) if reserved_tokens > 0 @@ -4784,7 +4800,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self._build_project_reservation_ops( targets=tuple(stash.itpm_reserved_scopes), reserved_scopes=stash.itpm_reserved_scopes, - actual_tokens=0, + actual_tokens=itpm_actual, reserved_tokens=itpm_reserved, reservation_window_identities=stash.itpm_reserved_window_identities, ) @@ -4792,7 +4808,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): else self._build_reservation_aware_tpm_ops( targets=tuple(stash.itpm_reserved_scopes), reserved_scopes=stash.itpm_reserved_scopes, - actual_tokens=0, + actual_tokens=itpm_actual, reserved_tokens=itpm_reserved, ) if itpm_reserved > 0 @@ -4802,7 +4818,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self._build_project_reservation_ops( targets=tuple(stash.otpm_reserved_scopes), reserved_scopes=stash.otpm_reserved_scopes, - actual_tokens=0, + actual_tokens=otpm_actual, reserved_tokens=otpm_reserved, reservation_window_identities=stash.otpm_reserved_window_identities, ) @@ -4810,7 +4826,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): else self._build_reservation_aware_tpm_ops( targets=tuple(stash.otpm_reserved_scopes), reserved_scopes=stash.otpm_reserved_scopes, - actual_tokens=0, + actual_tokens=otpm_actual, reserved_tokens=otpm_reserved, ) if otpm_reserved > 0 diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index fc0088b28d7..4003286d887 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -3647,6 +3647,173 @@ async def test_stash_applies_when_owner_or_callback_call_id_missing(): assert claimed.reservation_released is True +async def _reserve_tpm_for_owner_call(handler, local_cache, api_key: str, call_id: str) -> int: + await handler.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key=api_key, tpm_limit=10_000), + cache=local_cache, + data={ + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 50, + "litellm_call_id": call_id, + }, + call_type="completion", + ) + stash = get_request_stash() + assert stash is not None and stash.reserved_tokens > 0 + return stash.reserved_tokens + + +@pytest.mark.asyncio +async def test_failure_event_settles_tpm_reservation_at_recovered_partial_usage_v3(): + """ + A stream that fails mid-way after the model already produced tokens is + logged as a failure carrying the recovered partial usage. Those tokens + were consumed, so the TPM window must settle at them instead of refunding + the whole reservation (which would let repeated timeouts burn output + tokens for free). + """ + _api_key = hash_token("sk-partial-stream-failure") + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) + tokens_key = handler.create_rate_limit_keys(key="api_key", value=_api_key, rate_limit_type="tokens") + await _reserve_tpm_for_owner_call(handler, local_cache, _api_key, "partial-call") + + await handler.async_log_failure_event( + kwargs={ + "litellm_call_id": "partial-call", + "standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}}, + "combined_usage_object": Usage(prompt_tokens=20, completion_tokens=7, total_tokens=27), + }, + response_obj=None, + start_time=None, + end_time=None, + ) + + assert int(await local_cache.async_get_cache(key=tokens_key) or 0) == 27 + stash = get_request_stash() + assert stash is not None and stash.reservation_released is True + + +@pytest.mark.asyncio +async def test_failure_event_refunds_reservation_for_input_only_estimate_v3(): + """ + A failure with no recovered output carries only the input-token estimate + the proxy lifts onto every failure; that is not consumed usage, so the + reservation is still refunded in full. + """ + _api_key = hash_token("sk-estimated-failure") + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) + tokens_key = handler.create_rate_limit_keys(key="api_key", value=_api_key, rate_limit_type="tokens") + await _reserve_tpm_for_owner_call(handler, local_cache, _api_key, "estimate-call") + + await handler.async_log_failure_event( + kwargs={ + "litellm_call_id": "estimate-call", + "standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}}, + "combined_usage_object": Usage(prompt_tokens=20, completion_tokens=0, total_tokens=20), + }, + response_obj=None, + start_time=None, + end_time=None, + ) + + assert int(await local_cache.async_get_cache(key=tokens_key) or 0) == 0 + + +@pytest.mark.asyncio +async def test_post_call_failure_hook_settles_reservation_at_recovered_partial_usage_v3(): + """ + Pass-through streams report a mid-stream failure through the proxy-level + failure hook first, with the recovered usage lifted onto request_data. + That hook must settle at the partial usage too, and the later failure + callback must not double-apply it. + """ + _api_key = hash_token("sk-partial-post-call") + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) + user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, tpm_limit=10_000) + tokens_key = handler.create_rate_limit_keys(key="api_key", value=_api_key, rate_limit_type="tokens") + await _reserve_tpm_for_owner_call(handler, local_cache, _api_key, "post-call") + + await handler.async_post_call_failure_hook( + request_data={ + "model": "gpt-4o-mini", + "litellm_call_id": "post-call", + "combined_usage_object": Usage(prompt_tokens=20, completion_tokens=7, total_tokens=27), + }, + original_exception=Exception("upstream dropped the stream"), + user_api_key_dict=user_api_key_dict, + ) + assert int(await local_cache.async_get_cache(key=tokens_key) or 0) == 27 + + await handler.async_log_failure_event( + kwargs={ + "litellm_call_id": "post-call", + "standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}}, + "combined_usage_object": Usage(prompt_tokens=20, completion_tokens=7, total_tokens=27), + }, + response_obj=None, + start_time=None, + end_time=None, + ) + assert int(await local_cache.async_get_cache(key=tokens_key) or 0) == 27 + + +@pytest.mark.asyncio +async def test_failure_event_settles_project_itpm_otpm_at_recovered_partial_usage_v3(): + """ + Project ITPM/OTPM reservations settle the same way: input at the billable + prompt tokens and output at the completion tokens the failed stream + actually produced. + """ + _api_key = hash_token("sk-partial-project-io") + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) + user_api_key_dict = UserAPIKeyAuth( + api_key=_api_key, + project_id="proj-partial", + project_metadata={ + "model_itpm_limit": {"gpt-4o-mini": 10_000}, + "model_otpm_limit": {"gpt-4o-mini": 10_000}, + }, + ) + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={ + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 50, + "litellm_call_id": "project-call", + }, + call_type="completion", + ) + stash = get_request_stash() + assert stash is not None and stash.itpm_reserved_tokens > 0 and stash.otpm_reserved_tokens > 0 + itpm_key = handler.create_rate_limit_keys( + key="model_per_project_itpm", value="proj-partial:gpt-4o-mini", rate_limit_type="tokens" + ) + otpm_key = handler.create_rate_limit_keys( + key="model_per_project_otpm", value="proj-partial:gpt-4o-mini", rate_limit_type="tokens" + ) + + await handler.async_log_failure_event( + kwargs={ + "litellm_call_id": "project-call", + "standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}}, + "combined_usage_object": Usage(prompt_tokens=20, completion_tokens=7, total_tokens=27), + }, + response_obj=None, + start_time=None, + end_time=None, + ) + + assert int(await local_cache.async_get_cache(key=itpm_key) or 0) == 20 + assert int(await local_cache.async_get_cache(key=otpm_key) or 0) == 7 + + # ----------------------- Per-MCP-server rate limiting (v3) -----------------------