From a690b5a54355dde3d1bb08b1d89597a524a35f8b Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Tue, 15 Sep 2026 04:35:07 -0700 Subject: [PATCH] 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. --- litellm/proxy/common_request_processing.py | 4 +- .../spend_tracking/budget_reservation.py | 43 ++++++++++++++++--- .../proxy/auth/test_auth_checks.py | 2 + .../spend_tracking/test_budget_reservation.py | 37 +++++++++++++++- .../proxy/test_common_request_processing.py | 8 ++-- 5 files changed, 82 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 44163ee9e53..fed9dcd6d18 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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) diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 11bd770842c..2272221a5db 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -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: diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index d24b83a69e1..608d3d169aa 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -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( diff --git a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py index e0070318a96..2d7cfc87091 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py +++ b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py @@ -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 diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 77a03087f97..9e56f47fb41 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -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