fix(proxy): make budget-reset and cancel-path invalidation race-safe against concurrent reservations

Addresses two concurrency findings from automated security review on
the reset/invalidation fallbacks added for #30460.

reset_budget_job.py: a reset that failed and retried replayed an
unconditional SET to new_spend on every attempt. A request's Redis
INCR landing in the delay between a failed attempt and its retry (a
legitimate reservation against the just-reset budget) would be
silently erased by the next attempt's write, since retrying and the
final delete both treat the counter as static rather than possibly
having moved. Replaced the retry's SET with
async_reset_preserving_delta, a single Lua GET/compute/SET that resets
to new_spend plus whatever was added on top of a snapshot taken once
before the first attempt and held fixed across retries, so a
concurrent increment survives no matter which attempt eventually
succeeds. If the snapshot read itself fails, there's no safe baseline
to preserve against, so that case now skips straight to the existing
delete fallback rather than attempting a reset that could guess wrong.

budget_reservation.py: release_budget_reservation_on_cancel's fallback
on a reconcile failure invalidates the shared key/user/team counters
outright. Unlike reconcile, which only ever adjusts this reservation's
own recorded contribution, that invalidation deletes state a
concurrent reservation or recorded spend also shares on the same
counter. A transient failure (e.g. a Redis timeout) during
cancellation would take the destructive path immediately. Now retries
the reconcile itself, which is safe and idempotent, a bounded number
of times first, and only falls back to invalidating the aggregate
counters once every retry hits the same failure.

Both fixes are covered by regression tests that simulate the race
directly (a concurrent increment landing during the reset retry delay,
and a reconcile that fails once then recovers on the cancel path) and
assert the concurrent write survives / the aggregate counter is not
touched.
This commit is contained in:
Abhyuday 2026-09-14 06:24:30 -04:00
parent fba32eea20
commit bdd409430e
5 changed files with 342 additions and 57 deletions

View file

@ -1403,6 +1403,49 @@ class RedisCache(BaseCache):
result = result.decode()
return float(result)
@_redis_circuit_breaker_guard
async def async_reset_preserving_delta(
self,
key: str,
new_base: float,
snapshot: float,
ttl: int | None = None,
) -> float:
"""Atomically reset ``key`` to ``new_base`` while preserving any amount
added since ``snapshot`` was read, so a reset racing a concurrent
``async_increment`` cannot erase spend reserved after the reset boundary.
``snapshot`` is the value read once before the first attempt and held
fixed across retries by the caller, not re-read each attempt: replaying
this call after a transient failure stays correct no matter how many
increments landed in between, because the delta is always measured
against that same original baseline. The GET/compute/SET runs in a
single Lua call, atomic across racing callers and pods, mirroring
``async_set_max``. Returns the resulting value.
"""
_redis_client: Final = self.init_async_client()
_used_ttl: Final = self.get_ttl(ttl=ttl)
key = self.check_and_fix_namespace(key=key)
lua: Final = (
"local cur = redis.call('GET', KEYS[1]) "
"local cur_num = cur and tonumber(cur) or tonumber(ARGV[2]) "
"local delta = cur_num - tonumber(ARGV[2]) "
"if delta < 0 then delta = 0 end "
"local result = tonumber(ARGV[1]) + delta "
"redis.call('SET', KEYS[1], result) "
"if tonumber(ARGV[3]) > 0 then redis.call('EXPIRE', KEYS[1], ARGV[3]) end "
"return tostring(result)"
)
result = cast(
"str | bytes",
await _redis_client.eval(
lua, 1, key, str(new_base), str(snapshot), str(int(_used_ttl or 0))
),
)
if isinstance(result, bytes):
result = result.decode()
return float(result)
@_redis_circuit_breaker_guard
async def async_increment_with_floor(self, key: str, value: int, ttl: int) -> int:
"""Async twin of ``increment_with_floor``, sharing its Lua script and its guarantees."""

View file

@ -545,11 +545,21 @@ class ResetBudgetJob:
commit opens a window where get_current_spend reads 0 from Redis
while the DB still holds the pre-reset value, allowing bypass.
A SET that keeps failing after retrying falls back to deleting the key rather
The reset is delta-preserving: a snapshot of the pre-reset value is read
once up front and held fixed across every retry, so a concurrent
async_increment landing during a retry (e.g. a request reserving spend
against the just-reset budget) is carried forward on top of new_spend
instead of being erased by a later attempt's write. See
async_reset_preserving_delta and _reset_redis_spend_counter.
A reset that keeps failing after retrying falls back to deleting the key rather
than leaving the pre-reset (possibly far higher) value authoritative in Redis
until its TTL expires: a missing counter reads as cold and reseeds from the
DB on the next request (_ensure_spend_counter_initialized), which is always
closer to the truth than the stale value a failed SET would otherwise leave behind.
closer to the truth than the stale value a failed reset would otherwise leave
behind. The same is true when the pre-reset snapshot itself cannot be read:
with no safe baseline to preserve increments against, deleting is the only
option that cannot silently erase a concurrent reservation.
"""
try:
from litellm.proxy.proxy_server import spend_counter_cache
@ -558,8 +568,12 @@ class ResetBudgetJob:
redis_cache = spend_counter_cache.redis_cache
if redis_cache is None:
return
if await ResetBudgetJob._reset_redis_spend_counter(
redis_cache=redis_cache, counter_key=counter_key, new_spend=new_spend
snapshot = await ResetBudgetJob._snapshot_spend_counter(
redis_cache=redis_cache, counter_key=counter_key
)
if snapshot is not None and await ResetBudgetJob._reset_redis_spend_counter(
redis_cache=redis_cache, counter_key=counter_key, new_spend=new_spend, snapshot=snapshot
):
return
verbose_proxy_logger.error(
@ -581,10 +595,31 @@ class ResetBudgetJob:
verbose_proxy_logger.warning("Failed to reset spend counter %s: %s", counter_key, e)
@staticmethod
async def _reset_redis_spend_counter(redis_cache: RedisCache, counter_key: str, new_spend: float) -> bool:
async def _snapshot_spend_counter(redis_cache: RedisCache, counter_key: str) -> float | None:
"""Read the pre-reset baseline the delta-preserving reset holds fixed across
retries. A failed read leaves no safe baseline to reconcile against, so the
caller skips the atomic reset and falls straight back to delete."""
try:
current = await redis_cache.async_get_cache(key=counter_key)
return float(current) if current is not None else 0.0
except Exception as redis_err:
verbose_proxy_logger.warning(
"Failed to read spend counter %s in Redis before reset; skipping the delta-preserving "
"reset and falling back to delete: %s",
counter_key,
redis_err,
)
return None
@staticmethod
async def _reset_redis_spend_counter(
redis_cache: RedisCache, counter_key: str, new_spend: float, snapshot: float
) -> bool:
for attempt in range(RESET_BUDGET_SPEND_COUNTER_RESET_MAX_ATTEMPTS):
try:
await redis_cache.async_set_cache(key=counter_key, value=new_spend, ttl=60)
await redis_cache.async_reset_preserving_delta(
key=counter_key, new_base=new_spend, snapshot=snapshot, ttl=60
)
return True
except Exception as redis_err:
is_last_attempt = attempt == RESET_BUDGET_SPEND_COUNTER_RESET_MAX_ATTEMPTS - 1

View file

@ -53,6 +53,15 @@ class _BudgetCounter:
window_start: datetime | None = None
_RELEASE_ON_CANCEL_RECONCILE_MAX_ATTEMPTS: Final = 2
"""Bounded retries of the per-reservation reconcile on the cancel path before
falling back to invalidating the shared aggregate counters. reconcile_budget_reservation
only ever adjusts this reservation's own recorded contribution to each counter (see
_set_reserved_entries_actual_cost), so it is safe and idempotent to retry; the counter
invalidation fallback below is not, since it deletes state every concurrent reservation
and recorded spend on that counter shares."""
_COUNTER_ENTITY_TYPES: Final[Mapping[str, str]] = {
"Key": Litellm_EntityType.KEY.value,
"Team": Litellm_EntityType.TEAM.value,
@ -372,7 +381,12 @@ async def release_budget_reservation_on_cancel(
when success/failure handling already reconciled, so calling it on every
cancellation path is safe.
A reconcile failure here (e.g. a Redis timeout) falls back to
A reconcile failure here (e.g. a Redis timeout) is retried a bounded number of
times before falling back to invalidate_budget_reservation_counters: unlike that
fallback, the reconcile only ever adjusts this reservation's own contribution to
each counter, so it cannot clobber a concurrent request's reservation or recorded
spend on the same key/user/team counter the way deleting the aggregate can. Only
once every retry hits the same failure does this fall back to
invalidate_budget_reservation_counters, mirroring release_or_invalidate_budget_reservation's
handling of the same failure on the non-cancel release path: dropping the reserved counters
forces the next read to reseed from the DB instead of leaving the pre-charge stuck in Redis
@ -389,8 +403,12 @@ async def release_budget_reservation_on_cancel(
pass
except Exception:
verbose_proxy_logger.exception(
"Failed to reconcile budget reservation on cancel; invalidating reserved counters"
"Failed to reconcile budget reservation on cancel; retrying before invalidating reserved counters"
)
if await asyncio.shield(
_retry_reconcile_reservation_on_cancel(budget_reservation=budget_reservation, incurred_cost=incurred_cost)
):
return
try:
await invalidate_budget_reservation_counters(budget_reservation=budget_reservation)
except Exception:
@ -401,6 +419,35 @@ async def release_budget_reservation_on_cancel(
budget_reservation["finalized"] = True
async def _retry_reconcile_reservation_on_cancel(
budget_reservation: dict, # mutable-ok: reconcile_budget_reservation stamps finalized on success
incurred_cost: float,
) -> bool:
"""Bounded retry of the reconcile that just failed once on the cancel path.
Unlike invalidate_budget_reservation_counters, reconcile_budget_reservation only
ever adjusts this reservation's own recorded contribution to each counter (see
_set_reserved_entries_actual_cost's applied_adjustment bookkeeping), so retrying it
is safe and idempotent, and correct on any attempt that gets through: it does not
touch concurrent reservations or spend the counter also aggregates. Returns whether
a retry succeeded, so the caller falls back to the destructive invalidation only
after every retry hits the same (presumably persistent) failure.
"""
for attempt in range(_RELEASE_ON_CANCEL_RECONCILE_MAX_ATTEMPTS):
try:
await reconcile_budget_reservation(budget_reservation=budget_reservation, actual_cost=incurred_cost)
return True
except Exception:
is_last_attempt = attempt == _RELEASE_ON_CANCEL_RECONCILE_MAX_ATTEMPTS - 1
verbose_proxy_logger.warning(
"Retry %d/%d to reconcile budget reservation on cancel failed",
attempt + 1,
_RELEASE_ON_CANCEL_RECONCILE_MAX_ATTEMPTS,
exc_info=not is_last_attempt,
)
return False
async def invalidate_budget_reservation_counters(
budget_reservation: dict | None,
) -> None:

View file

@ -1230,7 +1230,8 @@ def _make_counter_invalidation_job(monkeypatch):
spend_counter_cache = MagicMock()
spend_counter_cache.in_memory_cache.set_cache = MagicMock()
spend_counter_cache.redis_cache = MagicMock()
spend_counter_cache.redis_cache.async_set_cache = AsyncMock()
spend_counter_cache.redis_cache.async_get_cache = AsyncMock(return_value=0.0)
spend_counter_cache.redis_cache.async_reset_preserving_delta = AsyncMock()
user_api_key_cache = MagicMock()
user_api_key_cache.async_delete_cache = AsyncMock()
@ -1532,7 +1533,9 @@ def test_budget_table_reset_invalidates_counters_and_management_cache(
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
counter_cache.in_memory_cache.set_cache.assert_any_call(key=counter_key, value=0.0, ttl=60)
counter_cache.redis_cache.async_set_cache.assert_any_await(key=counter_key, value=0.0, ttl=60)
counter_cache.redis_cache.async_reset_preserving_delta.assert_any_await(
key=counter_key, new_base=0.0, snapshot=0.0, ttl=60
)
deleted = {call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list}
assert cache_keys <= deleted
@ -1571,7 +1574,9 @@ def test_budget_table_reset_invalidates_enduser_counter_and_cache(reset_budget_j
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:end_user:customer-42", value=0.0, ttl=60)
counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:end_user:customer-42", value=0.0, ttl=60)
counter_cache.redis_cache.async_reset_preserving_delta.assert_any_await(
key="spend:end_user:customer-42", new_base=0.0, snapshot=0.0, ttl=60
)
deleted: Final = {call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list}
assert "end_user_id:customer-42" in deleted
@ -3137,7 +3142,9 @@ def test_budget_cascade_carries_default_tier_enduser_counter_when_rollover_enabl
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:end_user:enduser-implicit", value=5.0, ttl=60)
counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:end_user:enduser-implicit", value=5.0, ttl=60)
counter_cache.redis_cache.async_reset_preserving_delta.assert_any_await(
key="spend:end_user:enduser-implicit", new_base=5.0, snapshot=0.0, ttl=60
)
deleted: Final = {call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list}
assert "end_user_id:enduser-implicit" in deleted
@ -3444,52 +3451,142 @@ async def test_reset_budget_for_keys_broadcasts_cache_invalidation_to_other_pods
def test_invalidate_spend_counter_retries_then_deletes_on_persistent_redis_failure(monkeypatch):
"""
Every reset-to-zero attempt fails: the fix must retry a bounded number of
Every reset attempt fails: the fix must retry a bounded number of
times, then fall back to deleting the counter (a missing counter reads as
cold and reseeds from the DB, see _ensure_spend_counter_initialized)
instead of leaving the old, inflated value authoritative until its TTL.
"""
from unittest.mock import patch
spend_counter_cache = MagicMock()
spend_counter_cache.in_memory_cache.set_cache = MagicMock()
spend_counter_cache.redis_cache = MagicMock()
spend_counter_cache.redis_cache.async_set_cache = AsyncMock(side_effect=RuntimeError("elasticache timeout"))
spend_counter_cache.redis_cache.async_delete_cache = AsyncMock()
fake_module = types.ModuleType("litellm.proxy.proxy_server")
fake_module.spend_counter_cache = spend_counter_cache
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_module)
with patch("litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.error") as mock_error, patch(
"asyncio.sleep", new=AsyncMock()
):
asyncio.run(ResetBudgetJob._invalidate_spend_counter("spend:key:sk-inflated", new_spend=0.0))
assert spend_counter_cache.redis_cache.async_set_cache.await_count == RESET_BUDGET_SPEND_COUNTER_RESET_MAX_ATTEMPTS
spend_counter_cache.redis_cache.async_delete_cache.assert_awaited_once_with(key="spend:key:sk-inflated")
assert mock_error.called
def test_invalidate_spend_counter_recovers_after_a_transient_redis_failure(monkeypatch):
"""A SET that fails once and then succeeds must not fall back to delete:
the counter ends up reset to the real value, not merely absent."""
from unittest.mock import patch
spend_counter_cache = MagicMock()
spend_counter_cache.in_memory_cache.set_cache = MagicMock()
spend_counter_cache.redis_cache = MagicMock()
spend_counter_cache.redis_cache.async_set_cache = AsyncMock(
side_effect=[RuntimeError("elasticache timeout"), None]
spend_counter_cache.redis_cache.async_get_cache = AsyncMock(return_value=80.0)
spend_counter_cache.redis_cache.async_reset_preserving_delta = AsyncMock(
side_effect=RuntimeError("elasticache timeout")
)
spend_counter_cache.redis_cache.async_delete_cache = AsyncMock()
fake_module = types.ModuleType("litellm.proxy.proxy_server")
fake_module.spend_counter_cache = spend_counter_cache
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_module)
monkeypatch.setattr(asyncio, "sleep", AsyncMock())
with patch("asyncio.sleep", new=AsyncMock()):
asyncio.run(ResetBudgetJob._invalidate_spend_counter("spend:key:sk-recovered", new_spend=0.0))
asyncio.run(ResetBudgetJob._invalidate_spend_counter("spend:key:sk-inflated", new_spend=0.0))
assert spend_counter_cache.redis_cache.async_set_cache.await_count == 2
assert (
spend_counter_cache.redis_cache.async_reset_preserving_delta.await_count
== RESET_BUDGET_SPEND_COUNTER_RESET_MAX_ATTEMPTS
)
spend_counter_cache.redis_cache.async_delete_cache.assert_awaited_once_with(key="spend:key:sk-inflated")
def test_invalidate_spend_counter_recovers_after_a_transient_redis_failure(monkeypatch):
"""A reset that fails once and then succeeds must not fall back to delete:
the counter ends up reset to the real value, not merely absent."""
spend_counter_cache = MagicMock()
spend_counter_cache.in_memory_cache.set_cache = MagicMock()
spend_counter_cache.redis_cache = MagicMock()
spend_counter_cache.redis_cache.async_get_cache = AsyncMock(return_value=80.0)
spend_counter_cache.redis_cache.async_reset_preserving_delta = AsyncMock(
side_effect=[RuntimeError("elasticache timeout"), 0.0]
)
spend_counter_cache.redis_cache.async_delete_cache = AsyncMock()
fake_module = types.ModuleType("litellm.proxy.proxy_server")
fake_module.spend_counter_cache = spend_counter_cache
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_module)
monkeypatch.setattr(asyncio, "sleep", AsyncMock())
asyncio.run(ResetBudgetJob._invalidate_spend_counter("spend:key:sk-recovered", new_spend=0.0))
assert spend_counter_cache.redis_cache.async_reset_preserving_delta.await_count == 2
spend_counter_cache.redis_cache.async_delete_cache.assert_not_awaited()
def test_invalidate_spend_counter_skips_straight_to_delete_when_snapshot_read_fails(monkeypatch):
"""No safe baseline to preserve a concurrent increment against: guessing with
a blind reset would reopen the exact race this fix closes, so a failed
pre-reset snapshot read must fall straight back to delete instead of
attempting the atomic reset at all."""
spend_counter_cache = MagicMock()
spend_counter_cache.in_memory_cache.set_cache = MagicMock()
spend_counter_cache.redis_cache = MagicMock()
spend_counter_cache.redis_cache.async_get_cache = AsyncMock(side_effect=RuntimeError("elasticache timeout"))
spend_counter_cache.redis_cache.async_reset_preserving_delta = AsyncMock()
spend_counter_cache.redis_cache.async_delete_cache = AsyncMock()
fake_module = types.ModuleType("litellm.proxy.proxy_server")
fake_module.spend_counter_cache = spend_counter_cache
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_module)
asyncio.run(ResetBudgetJob._invalidate_spend_counter("spend:key:sk-no-snapshot", new_spend=0.0))
spend_counter_cache.redis_cache.async_reset_preserving_delta.assert_not_awaited()
spend_counter_cache.redis_cache.async_delete_cache.assert_awaited_once_with(key="spend:key:sk-no-snapshot")
class _FakeRedisSpendCounter:
"""A minimal in-process stand-in for the Redis key under test, faithful to
async_reset_preserving_delta's GET/compute/SET contract (new_base + max(0,
current - snapshot)), so what's actually under test is the retry loop's own
behavior of holding one fixed snapshot across attempts rather than re-reading
it, which is what makes a concurrent increment during the retry delay survive."""
def __init__(self, initial: float):
self.value = initial
self.attempts = 0
self.fail_first_n = 0
async def async_get_cache(self, key):
return self.value
async def async_reset_preserving_delta(self, key, new_base, snapshot, ttl=None):
self.attempts += 1
if self.attempts <= self.fail_first_n:
raise RuntimeError("elasticache timeout")
delta = max(0.0, self.value - snapshot)
self.value = new_base + delta
return self.value
async def async_increment(self, key, value, refresh_ttl=False):
self.value += value
return self.value
async def async_delete_cache(self, key):
self.value = None
def test_invalidate_spend_counter_preserves_a_concurrent_increment_made_during_the_retry_delay(monkeypatch):
"""
Regression for the reset-retry race (veria-ai finding, reset_budget_job.py):
if a request's Redis INCR lands in the delay between a failed reset attempt
and its retry, the eventual successful reset must not erase it. Before the
fix, the retry replayed an unconditional SET to new_spend and would have
wiped the concurrent increment out; the delta-preserving reset must carry
it forward instead, since it represents real spend reserved after the
reset boundary.
"""
store = _FakeRedisSpendCounter(initial=80.0)
store.fail_first_n = 1
spend_counter_cache = MagicMock()
spend_counter_cache.in_memory_cache.set_cache = MagicMock()
spend_counter_cache.redis_cache = store
fake_module = types.ModuleType("litellm.proxy.proxy_server")
fake_module.spend_counter_cache = spend_counter_cache
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_module)
concurrent_increment_applied = False
async def sleep_and_race(_delay):
nonlocal concurrent_increment_applied
if not concurrent_increment_applied:
await store.async_increment(key="spend:key:sk-race", value=5.0)
concurrent_increment_applied = True
monkeypatch.setattr(asyncio, "sleep", sleep_and_race)
asyncio.run(ResetBudgetJob._invalidate_spend_counter("spend:key:sk-race", new_spend=0.0))
assert concurrent_increment_applied
assert store.value == 5.0

View file

@ -2230,6 +2230,23 @@ class _ExpiringRedisCache:
return None
class _FlakyPipelineRedisCache(_ExpiringRedisCache):
"""_ExpiringRedisCache, but the first ``fail_first_n`` calls to
async_increment_pipeline raise instead of applying, so a reconcile retry
recovering from a transient Redis failure is what's under test."""
def __init__(self, fail_first_n: int = 0) -> None:
super().__init__()
self.fail_first_n = fail_first_n
self.pipeline_calls = 0
async def async_increment_pipeline(self, increment_list, **kwargs):
self.pipeline_calls += 1
if self.pipeline_calls <= self.fail_first_n:
raise RuntimeError("redis down")
return await super().async_increment_pipeline(increment_list, **kwargs)
@pytest.mark.asyncio
async def test_reconcile_after_redis_counter_expiry_keeps_request_cost_enforced(
spend_counter_state,
@ -2772,20 +2789,22 @@ async def test_release_budget_reservation_on_cancel_swallows_release_errors():
@pytest.mark.asyncio
async def test_release_budget_reservation_on_cancel_invalidates_counter_when_reconcile_fails(
async def test_release_budget_reservation_on_cancel_invalidates_counter_when_reconcile_persistently_fails(
spend_counter_state,
):
"""
Regression for #30460 Path 1: if reconcile itself fails on the cancel path
(e.g. a Redis timeout), the pre-charge must not be left stuck in the
counter with nothing left to correct it. This mirrors what
release_or_invalidate_budget_reservation already does on the non-cancel
release path: fall back to invalidate_budget_reservation_counters so the
next read reseeds from the DB instead of enforcing the stale reservation
forever.
Regression for #30460 Path 1: if reconcile keeps failing on the cancel path
(e.g. a persistent Redis outage) across every retry, the pre-charge must
not be left stuck in the counter with nothing left to correct it. This
mirrors what release_or_invalidate_budget_reservation already does on the
non-cancel release path: fall back to invalidate_budget_reservation_counters
so the next read reseeds from the DB instead of enforcing the stale
reservation forever.
"""
counter_cache, _key_cache = spend_counter_state
counter_key = "spend:key:key-cancel-redis-down"
counter_cache.redis_cache = _FlakyPipelineRedisCache(fail_first_n=99)
counter_cache.redis_cache.store[counter_key] = 3.0
counter_cache.in_memory_cache.set_cache(key=counter_key, value=3.0)
reservation = {
@ -2795,13 +2814,57 @@ async def test_release_budget_reservation_on_cancel_invalidates_counter_when_rec
"entries": [{"counter_key": counter_key, "reserved_cost": 3.0, "applied_adjustment": 0.0}],
}
with patch(
"litellm.proxy.spend_tracking.budget_reservation.reconcile_budget_reservation",
new=AsyncMock(side_effect=RuntimeError("redis down")),
await release_budget_reservation_on_cancel(reservation)
assert counter_cache.in_memory_cache.get_cache(key=counter_key) is None
assert reservation["finalized"] is True
# Only the first retry reaches the pipeline: its own failure already
# invalidates the counter, so later retries take the (here also
# unavailable, with no real DB) reseed path instead -- both exhausted
# before falling back to invalidate_budget_reservation_counters.
assert counter_cache.redis_cache.pipeline_calls == 1
@pytest.mark.asyncio
async def test_release_budget_reservation_on_cancel_retries_before_invalidating(
spend_counter_state,
):
"""
Regression for the cancellation-fallback race (veria-ai finding on
budget_reservation.py): invalidate_budget_reservation_counters deletes the
whole aggregate counter unconditionally, which would also erase a
concurrent request's reservation or recorded spend sharing that same
key/user/team counter, not just this reservation's own contribution. A
reconcile failure that clears up on retry must settle through the
ordinary reconcile machinery instead of falling back to that unconditional
delete: here, the first attempt's own failure already invalidates the
counter (existing increment_spend_counters_pipeline cleanup), so the retry
settles it by reseeding the DB floor and adding this reservation's actual
cost, the same recovery path an expired counter already takes elsewhere in
this file, rather than being deleted a second time with nothing settled.
"""
import litellm.proxy.proxy_server as ps
counter_cache, _key_cache = spend_counter_state
counter_key = "spend:key:key-cancel-retry-recovers"
counter_cache.redis_cache = _FlakyPipelineRedisCache(fail_first_n=1)
counter_cache.redis_cache.store[counter_key] = 5.0
counter_cache.in_memory_cache.set_cache(key=counter_key, value=5.0)
reservation = {
"reserved_cost": 3.0,
"input_cost": 0.5,
"finalized": False,
"entries": [{"counter_key": counter_key, "reserved_cost": 3.0, "applied_adjustment": 0.0}],
}
with patch.object( # test-quality-ok: the reseed reads the DB floor through a Prisma client the test has no seam for
ps.SpendCounterReseed, "from_db", AsyncMock(return_value=0.2)
):
await release_budget_reservation_on_cancel(reservation)
assert counter_cache.in_memory_cache.get_cache(key=counter_key) is None
assert counter_cache.redis_cache.pipeline_calls == 1
assert counter_cache.in_memory_cache.get_cache(key=counter_key) == pytest.approx(0.7)
assert reservation["finalized"] is True