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.
This commit is contained in:
Yucheng He 2026-09-15 04:52:17 -07:00
parent a690b5a543
commit 2c9e153f25
3 changed files with 49 additions and 10 deletions

View file

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

View file

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

View file

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