mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
parent
a690b5a543
commit
2c9e153f25
3 changed files with 49 additions and 10 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue