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