From 2c9e153f25247253544da8586f27e27eb4d22cab Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Tue, 15 Sep 2026 04:52:17 -0700 Subject: [PATCH] fix(proxy): record a counter increment that lands after the request was cancelled A cancellation delivered while the Redis increment is in flight cannot tell whether it was applied. The increment now keeps running shielded and its entry joins the rollback list once it lands, so the cancelled request's reservation is released instead of pinning the counter until expiry. --- .../spend_tracking/budget_reservation.py | 39 ++++++++++++++++++- .../spend_tracking/test_budget_reservation.py | 15 ++++--- .../proxy/test_budget_reservation.py | 5 ++- 3 files changed, 49 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 2272221a5db..5bd858ab26c 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -326,9 +326,11 @@ async def _reserve_counters( reserved_cost=reservation_cost, ) try: - reserved_value = await _reserve_counter( + reserved_value = await _acquire_counter( counter=counter, reservation_cost=reservation_cost, + entry=entry, + applied_entries=applied_entries, ) except _CounterReservationUnavailable as exc: if exc.touched_counter and not exc.counter_invalidated: @@ -340,7 +342,6 @@ async def _reserve_counters( _raise_reservation_unavailable(counter_key=counter.counter_key) continue - applied_entries.append(entry) if reserved_value is not None: current_spend = reserved_value else: @@ -913,6 +914,40 @@ def _coerce_window(window: object) -> Mapping[str, object]: return dumped if isinstance(dumped, Mapping) else {} +async def _acquire_counter( + counter: _BudgetCounter, + reservation_cost: float, + entry: dict[str, float | str], + applied_entries: list[dict[str, float | str]], # mutable-ok: the caller's rollback list +) -> float | None: + """Increment the counter and record ``entry`` as applied once the increment went through. + + A cancellation delivered while the increment is in flight leaves the request unable to tell whether + Redis applied it, so the increment keeps running shielded and the entry is recorded if it lands, + letting the caller's rollback release exactly what was reserved. + """ + increment: Final = asyncio.ensure_future(_reserve_counter(counter=counter, reservation_cost=reservation_cost)) + try: + reserved_value: Final = await asyncio.shield(increment) + except asyncio.CancelledError: + if await _increment_landed(increment): + applied_entries.append(entry) # rebind-ok: the rollback list must see the landed increment + raise + applied_entries.append(entry) # rebind-ok: the caller settles and rolls back through this list + return reserved_value + + +async def _increment_landed(increment: asyncio.Future[float | None]) -> bool: + try: + await asyncio.shield(increment) + except _CounterReservationUnavailable: + return False + except asyncio.CancelledError: + # cancelled again while waiting: leave the counter to its TTL rather than refund what may not exist + return False + return True + + async def _reserve_counter( counter: _BudgetCounter, reservation_cost: float, 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 2d7cfc87091..a30914fbfd9 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py +++ b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py @@ -269,28 +269,30 @@ async def test_release_budget_reservation_on_cancel_settles_each_entry_to_its_ow class _ParkingIncrementCache(DualCache): - """A spend-counter cache whose increment of ``parked_key`` never returns, so a test can cancel mid-reservation.""" + """A spend-counter cache whose increment of ``parked_key`` waits for ``release``, like a Redis INCR whose + reply is still on the wire, so a test can cancel the request while that increment is in flight.""" def __init__(self, parked_key: str) -> None: super().__init__() self.parked_key: Final = parked_key self.parked: Final = asyncio.Event() + self.release: Final = asyncio.Event() async def async_increment_cache(self, key: str, value: float, **kwargs: object) -> float | None: if key == self.parked_key: self.parked.set() - await asyncio.Event().wait() + await self.release.wait() return await super().async_increment_cache(key=key, value=value, **kwargs) @pytest.mark.asyncio @pytest.mark.timeout(30) -async def test_reserve_budget_for_added_tags_releases_the_reserved_tag_when_cancelled_mid_acquisition( +async def test_reserve_budget_for_added_tags_releases_every_counter_it_took_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, and the - second tag, whose increment never happened, must not be refunded for it either.""" + """A client disconnect under SSE keepalives cancels the request while the second tag's increment is in + flight. Neither tag may stay charged for a request that never reached the provider: the first was + reserved before the cancel, the second lands after it.""" 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) @@ -303,6 +305,7 @@ async def test_reserve_budget_for_added_tags_releases_the_reserved_tag_when_canc await cache.parked.wait() assert cache.in_memory_cache.get_cache(key=f"spend:tag:{HOOK_TAG}") > 0 reserving.cancel() + cache.release.set() with pytest.raises(asyncio.CancelledError): await reserving diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 40ebc03781c..551ae64bff7 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -2764,11 +2764,12 @@ async def test_release_budget_reservation_on_cancel_swallows_release_errors(): "input_cost": 0.5, } with patch( - "litellm.proxy.spend_tracking.budget_reservation.reconcile_budget_reservation", + "litellm.proxy.spend_tracking.budget_reservation._set_reserved_entries_actual_cost", new=AsyncMock(side_effect=RuntimeError("redis down")), - ): + ) as release: # must return without raising await release_budget_reservation_on_cancel(reservation) + release.assert_awaited_once() @pytest.mark.asyncio