fix(proxy): settle rate-limit reservations at a failed stream's partial usage

This commit is contained in:
mateo-berri 2026-09-03 12:52:58 -07:00
parent f3021937c6
commit 5acb81888d
2 changed files with 204 additions and 21 deletions

View file

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

View file

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