mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
fba32eea20
commit
bdd409430e
5 changed files with 342 additions and 57 deletions
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue