diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index d0cefcb6086..61fef745991 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -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.""" diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index 9aab7f862ea..25caa948247 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -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 diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 8bc3b4a4dc9..93cba60ab21 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -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: diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index 49a4abc8247..15cf6c5fd87 100644 --- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -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 diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 2a652c98430..a3f33c2b07e 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -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