fix(proxy): settle folded tag reservations to their own input cost and roll back only acquired counters

A hook-added tag is priced after the hook rewrote the request, so its entries keep that input cost when they join the auth-time reservation and a cancellation settles each entry to the cost it was priced with. An entry counts as applied only once its counter increment completed, so a cancellation during the increment no longer refunds spend that was never reserved.
This commit is contained in:
Yucheng He 2026-09-15 04:35:07 -07:00
parent f85cb1dac5
commit a690b5a543
5 changed files with 82 additions and 12 deletions

View file

@ -77,7 +77,7 @@ from litellm.proxy.common_utils.sse_keepalive import (
from litellm.proxy.dd_span_tagger import DDSpanTagger
from litellm.proxy.guardrails.auto_router_compression import arm_pre_call as _arm_auto_router_compression
from litellm.proxy.route_llm_request import route_request
from litellm.proxy.spend_tracking.budget_reservation import reserve_budget_for_added_tags
from litellm.proxy.spend_tracking.budget_reservation import merge_budget_reservation, reserve_budget_for_added_tags
from litellm.proxy.utils import ProxyLogging, _check_and_merge_model_level_guardrails
from litellm.router import Router
from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict
@ -2156,7 +2156,7 @@ class ProxyBaseLLMRequestProcessing:
return
existing: Final = user_api_key_dict.budget_reservation
if existing is not None:
existing["entries"].extend(reservation["entries"])
merge_budget_reservation(existing=existing, added=reservation)
return
user_api_key_dict.budget_reservation = reservation # rebind-ok: the failure and cancel paths read it here
_, metadata_bucket = get_or_create_metadata_bucket(self.data)

View file

@ -325,7 +325,6 @@ async def _reserve_counters(
counter=counter,
reserved_cost=reservation_cost,
)
applied_entries.append(entry)
try:
reserved_value = await _reserve_counter(
counter=counter,
@ -337,11 +336,11 @@ async def _reserve_counters(
entries=[entry],
default_reserved_cost=reservation_cost,
)
applied_entries.remove(entry)
if fail_closed_budget_enforcement:
_raise_reservation_unavailable(counter_key=counter.counter_key)
continue
applied_entries.append(entry)
if reserved_value is not None:
current_spend = reserved_value
else:
@ -434,15 +433,47 @@ async def release_budget_reservation_on_cancel(
"""
if not budget_reservation or budget_reservation.get("finalized") is True:
return
incurred_cost: Final = float(budget_reservation.get("input_cost") or 0.0)
try:
await asyncio.shield(
reconcile_budget_reservation(budget_reservation=budget_reservation, actual_cost=incurred_cost)
)
await asyncio.shield(_reconcile_entries_to_their_input_cost(budget_reservation=budget_reservation))
except (asyncio.CancelledError, Exception):
pass
async def _reconcile_entries_to_their_input_cost(
budget_reservation: dict[str, object], # mutable-ok: the reservation is stamped finalized in place
) -> None:
"""Entries folded in by ``merge_budget_reservation`` carry the input cost of the request that
went upstream; the rest were priced at auth and settle to the reservation's own input cost."""
shared_input_cost: Final = float(cast(SupportsFloat, budget_reservation.get("input_cost") or 0.0))
reserved_cost: Final = float(cast(SupportsFloat, budget_reservation.get("reserved_cost") or 0.0))
entries: Final = cast(list[dict[str, float | str]], budget_reservation.get("entries") or [])
for input_cost in dict.fromkeys(_entry_input_cost(entry, shared_input_cost) for entry in entries):
await _set_reserved_entries_actual_cost(
entries=[entry for entry in entries if _entry_input_cost(entry, shared_input_cost) == input_cost],
actual_cost=input_cost,
default_reserved_cost=reserved_cost,
)
budget_reservation["finalized"] = True # rebind-ok: the settlement paths share this one dict
def _entry_input_cost(entry: Mapping[str, float | str], shared_input_cost: float) -> float:
return float(entry.get("input_cost", shared_input_cost))
def merge_budget_reservation(
existing: dict[str, object], # mutable-ok: the request's reservation, extended in place
added: Mapping[str, object],
) -> None:
"""Fold a later reservation's entries into the request's reservation so every settlement path sees one.
Each added entry keeps the input cost it was priced with, since a pre-call hook may have rewritten the
request between the two estimates and a cancellation settles each entry to that cost."""
added_entries: Final = cast(list[dict[str, float | str]], added.get("entries") or [])
added_input_cost: Final = float(cast(SupportsFloat, added.get("input_cost") or 0.0))
for entry in added_entries:
entry["input_cost"] = added_input_cost # rebind-ok: the entry settles to the cost it was priced with
cast(list[dict[str, float | str]], existing["entries"]).extend(added_entries) # rebind-ok: one dict for all paths
async def invalidate_budget_reservation_counters(
budget_reservation: dict | None,
) -> None:

View file

@ -2517,6 +2517,8 @@ def test_route_skips_budget_checks_matches_auth_scope(route, expected):
False,
False,
),
("/bria", {"pass_through_endpoints": "not-a-list"}, "sk-master", False, False),
("/bria", {"pass_through_endpoints": [{"path": "/bria", "auth": False}, "junk"]}, "sk-master", False, False),
],
)
def test_auth_skips_common_checks_names_the_requests_that_never_run_them(

View file

@ -18,6 +18,8 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.spend_tracking.budget_reservation import (
count_request_input_tokens,
estimate_request_max_cost,
merge_budget_reservation,
release_budget_reservation_on_cancel,
reserve_budget_for_added_tags,
reserve_budget_for_request,
)
@ -237,6 +239,35 @@ async def test_reserve_budget_for_added_tags_skips_routes_auth_never_reserves(sp
assert spend_counter_cache.in_memory_cache.get_cache(key=f"spend:tag:{HOOK_TAG}") is None
@pytest.mark.asyncio
async def test_release_budget_reservation_on_cancel_settles_each_entry_to_its_own_input_cost(
spend_counter_cache: DualCache,
):
"""Auth priced the key before the hook rewrote the prompt; the hook tag was priced after. A cancellation
charges each counter the input cost its own estimate saw instead of the auth-time one for both."""
spend_counter_cache.in_memory_cache.set_cache(key="spend:key:hashed-hook-tag-key", value=0.5)
auth_reservation: Final[dict[str, object]] = {
"reserved_cost": 0.5,
"entries": [{"counter_key": "spend:key:hashed-hook-tag-key", "entity_type": "Key", "reserved_cost": 0.5}],
"finalized": False,
"input_cost": 0.2,
"input_tokens": 40,
}
hook_reservation: Final = await _reserve_added_tags(
"/v1/chat/completions", _budgeted_tag_prisma((HOOK_TAG,), max_budget=1.0)
)
assert hook_reservation is not None
hook_input_cost: Final = hook_reservation["input_cost"]
assert isinstance(hook_input_cost, float) and 0 < hook_input_cost < 0.2
merge_budget_reservation(existing=auth_reservation, added=hook_reservation)
await release_budget_reservation_on_cancel(auth_reservation)
assert spend_counter_cache.in_memory_cache.get_cache(key="spend:key:hashed-hook-tag-key") == pytest.approx(0.2)
assert spend_counter_cache.in_memory_cache.get_cache(key=f"spend:tag:{HOOK_TAG}") == pytest.approx(hook_input_cost)
assert auth_reservation["finalized"] is True
class _ParkingIncrementCache(DualCache):
"""A spend-counter cache whose increment of ``parked_key`` never returns, so a test can cancel mid-reservation."""
@ -253,12 +284,15 @@ class _ParkingIncrementCache(DualCache):
@pytest.mark.asyncio
@pytest.mark.timeout(30)
async def test_reserve_budget_for_added_tags_releases_the_reserved_tag_when_cancelled_mid_acquisition(
monkeypatch: pytest.MonkeyPatch,
):
"""A client disconnect under SSE keepalives cancels the request while the second tag is being reserved;
the first tag's counter must not stay charged for a request that never reached the provider."""
the first tag's counter must not stay charged for a request that never reached the provider, and the
second tag, whose increment never happened, must not be refunded for it either."""
cache: Final = _ParkingIncrementCache(parked_key=f"spend:tag:{SECOND_HOOK_TAG}")
cache.in_memory_cache.set_cache(key=f"spend:tag:{SECOND_HOOK_TAG}", value=0.3)
monkeypatch.setattr(proxy_server, "spend_counter_cache", cache)
monkeypatch.setattr(proxy_server, "prisma_client", None)
prisma: Final = _budgeted_tag_prisma((HOOK_TAG, SECOND_HOOK_TAG), max_budget=1.0)
@ -273,6 +307,7 @@ async def test_reserve_budget_for_added_tags_releases_the_reserved_tag_when_canc
await reserving
assert cache.in_memory_cache.get_cache(key=f"spend:tag:{HOOK_TAG}") == pytest.approx(0.0)
assert cache.in_memory_cache.get_cache(key=f"spend:tag:{SECOND_HOOK_TAG}") == pytest.approx(0.3)
@pytest.mark.asyncio

View file

@ -656,14 +656,14 @@ class TestProxyBaseLLMRequestProcessing:
assert tag_check.await_count == (1 if checked else 0)
@staticmethod
def _reservation(counter_key: str) -> dict:
def _reservation(counter_key: str, input_cost: float = 0.1) -> dict:
return {
"reserved_cost": 0.5,
"entries": [
{"counter_key": counter_key, "entity_type": "Tag", "entity_id": counter_key, "reserved_cost": 0.5}
],
"finalized": False,
"input_cost": 0.1,
"input_cost": input_cost,
"input_tokens": 3,
}
@ -677,7 +677,7 @@ class TestProxyBaseLLMRequestProcessing:
mock_request, mock_proxy_logging_obj, _ = self._tag_budget_rig(
monkeypatch, request_data={"model": "live-mini", "metadata": {"tags": []}}, pre_call_hook=mock_pre_call_hook
)
reserve = AsyncMock(return_value=self._reservation("spend:tag:guardrail-tag"))
reserve = AsyncMock(return_value=self._reservation("spend:tag:guardrail-tag", input_cost=0.02))
monkeypatch.setattr(litellm.proxy.common_request_processing, "reserve_budget_for_added_tags", reserve)
await processing_obj.common_processing_pre_call_logic(
request=mock_request,
@ -715,6 +715,8 @@ class TestProxyBaseLLMRequestProcessing:
"spend:key:test-token",
"spend:tag:guardrail-tag",
]
assert [entry.get("input_cost") for entry in auth_reservation["entries"]] == [None, 0.02]
assert auth_reservation["input_cost"] == 0.1
assert "user_api_key_budget_reservation" not in processing_obj.data["metadata"]
@pytest.mark.asyncio