diff --git a/litellm/constants.py b/litellm/constants.py index b3f5b0471f4..3ee586a5f51 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1817,6 +1817,12 @@ RESET_BUDGET_JOB_NAME: Final = "reset_budget_job" # leader keeps the lease across its own run, and a crashed one strands the sweep for # at most a single tick. RESET_BUDGET_JOB_LOCK_TTL_SECONDS: Final[int] = 900 +RESET_BUDGET_SPEND_COUNTER_RESET_MAX_ATTEMPTS: Final = max( + 1, int(os.getenv("RESET_BUDGET_SPEND_COUNTER_RESET_MAX_ATTEMPTS", "3")) +) +RESET_BUDGET_SPEND_COUNTER_RESET_RETRY_DELAY_SECONDS: Final = float( + os.getenv("RESET_BUDGET_SPEND_COUNTER_RESET_RETRY_DELAY_SECONDS", "0.25") +) PROXY_BATCH_POLLING_INTERVAL: Final = int(os.getenv("PROXY_BATCH_POLLING_INTERVAL", 3600)) MAX_OBJECTS_PER_POLL_CYCLE: Final = max(1, int(os.getenv("MAX_OBJECTS_PER_POLL_CYCLE", 50))) MANAGED_OBJECT_STALENESS_CUTOFF_DAYS: Final = max(1, int(os.getenv("MANAGED_OBJECT_STALENESS_CUTOFF_DAYS", 7))) diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index b35b876b475..ca1247bed4e 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -7,7 +7,7 @@ from dataclasses import dataclass, field from datetime import datetime, timedelta, timezone from enum import Enum from types import MappingProxyType -from typing import Final, Generic, Literal, Protocol, TypeVar +from typing import TYPE_CHECKING, Final, Generic, Literal, Protocol, TypeVar from typing_extensions import assert_never @@ -68,6 +68,9 @@ from litellm.repositories.verification_token_repository import ( ) from litellm.types.services import ServiceTypes +if TYPE_CHECKING: + from litellm.caching.redis_cache import RedisCache + _RowT = TypeVar("_RowT") @@ -574,27 +577,48 @@ class ResetBudgetJob: """Drop a spend counter so the next read reseeds from the committed DB row, the only value that includes increments that raced the reset. - Call AFTER the DB write commits. Clearing Redis before the DB - commit opens a window where get_current_spend reads 0 from Redis - while the DB still holds the pre-reset value, allowing bypass. + The Redis delete is retried a bounded number of times: a failed one + leaves the inflated pre-reset value authoritative until its TTL. """ try: from litellm.proxy.proxy_server import spend_counter_cache spend_counter_cache.in_memory_cache.delete_cache(key=counter_key) if spend_counter_cache.redis_cache is not None: - try: - await spend_counter_cache.redis_cache.async_delete_cache(key=counter_key) - except Exception as redis_err: - verbose_proxy_logger.warning( - "Failed to reset spend counter %s in Redis: %s. " - "Budget may be over-enforced until counter expires.", - counter_key, - redis_err, - ) + await ResetBudgetJob._delete_redis_spend_counter( + redis_cache=spend_counter_cache.redis_cache, counter_key=counter_key + ) except Exception as e: verbose_proxy_logger.warning("Failed to reset spend counter %s: %s", counter_key, e) + @staticmethod + async def _delete_redis_spend_counter(redis_cache: "RedisCache", counter_key: str) -> None: + from litellm.constants import ( + RESET_BUDGET_SPEND_COUNTER_RESET_MAX_ATTEMPTS, + RESET_BUDGET_SPEND_COUNTER_RESET_RETRY_DELAY_SECONDS, + ) + + for attempt in range(RESET_BUDGET_SPEND_COUNTER_RESET_MAX_ATTEMPTS): + try: + await redis_cache.async_delete_cache(key=counter_key) + return + except Exception as redis_err: # noqa: BLE001 # any Redis failure here is worth a retry, not just specific ones + verbose_proxy_logger.warning( + "Attempt %d/%d to reset spend counter %s in Redis failed: %s", + attempt + 1, + RESET_BUDGET_SPEND_COUNTER_RESET_MAX_ATTEMPTS, + counter_key, + redis_err, + ) + if attempt < RESET_BUDGET_SPEND_COUNTER_RESET_MAX_ATTEMPTS - 1: + await asyncio.sleep(RESET_BUDGET_SPEND_COUNTER_RESET_RETRY_DELAY_SECONDS) + verbose_proxy_logger.error( + "Failed to reset spend counter %s in Redis after %d attempts; budget may be " + "over-enforced until the counter expires.", + counter_key, + RESET_BUDGET_SPEND_COUNTER_RESET_MAX_ATTEMPTS, + ) + @staticmethod async def _invalidate_global_proxy_spend_cache() -> None: """Drop the cached global-proxy spend accumulator after the proxy @@ -605,20 +629,16 @@ class ResetBudgetJob: @staticmethod async def _invalidate_user_api_key_cache_entry(cache_key: str) -> None: - """Drop a stale management-cache entry so the next read fetches from DB. - - Tags and end-users are not reseeded by ``SpendCounterReseed.from_db``; - for those, when the spend counter expires the budget check falls back - to ``cached_obj.spend``. Keys, orgs, and team memberships are reseeded - from the DB, but auth still may consult ``user_api_key_cache`` objects - whose ``.spend`` field can lag a cross-pod DB reset. Deleting the cache - entry forces the next auth-time fetch to reload the zeroed row from - Postgres. - """ + """Drop a stale management-cache entry on every pod (LIT-3803 cross-pod + eviction), so a pod other than the one that ran the reset doesn't keep + serving a cached object with the pre-reset spend.""" try: + from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import ( + evict_and_broadcast, + ) from litellm.proxy.proxy_server import user_api_key_cache - await user_api_key_cache.async_delete_cache(key=cache_key) + await evict_and_broadcast(cache_keys=(cache_key,), user_api_key_cache=user_api_key_cache) except Exception as e: verbose_proxy_logger.warning( "Failed to invalidate user_api_key_cache entry %s: %s", @@ -1089,6 +1109,7 @@ class ResetBudgetJob: token = getattr(k.row, "token", None) if token: await self._invalidate_spend_counter(f"spend:key:{token}") + await self._invalidate_user_api_key_cache_entry(token) end_time = time.time() outcome: Final = _ChunkOutcome( @@ -1200,6 +1221,7 @@ class ResetBudgetJob: user_id = getattr(u.row, "user_id", None) if user_id: await self._invalidate_spend_counter(f"spend:user:{user_id}") + await self._invalidate_user_api_key_cache_entry(user_id) if user_id == LITELLM_PROXY_BUDGET_NAME: await self._invalidate_global_proxy_spend_cache() @@ -1315,6 +1337,7 @@ class ResetBudgetJob: team_id = getattr(t.row, "team_id", None) if team_id: await self._invalidate_spend_counter(f"spend:team:{team_id}") + await self._invalidate_user_api_key_cache_entry(f"team_id:{team_id}") end_time = time.time() outcome: Final = _ChunkOutcome( diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index c094e91c6c0..ee50153873e 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -394,26 +394,19 @@ async def release_budget_reservation(budget_reservation: dict | None) -> None: async def release_budget_reservation_on_cancel( - budget_reservation: dict | None, + budget_reservation: dict | None, # mutable-ok: stamps finalized on the caller's shared reservation dict ) -> None: - """Reconcile a still-open reservation when the request is cancelled mid-flight. + """Reconcile a still-open reservation when the request is cancelled mid-flight, + since neither the success nor failure hook runs in that case. Reconciles to the + request's input-token cost, not zero, since that portion was already billed to + the provider. asyncio.shield keeps this running through the surrounding + cancellation. - A client disconnect or timeout cancels the request task, which surfaces as - CancelledError / GeneratorExit rather than a normal exception, so neither the - success cost callback nor the failure hook runs and the pre-call reservation - is never reconciled. Left alone it pins the spend counter above real spend - and 429s subsequent requests until the counter's TTL expires. - - Reconcile to the request's input-token cost rather than refunding to zero: - by the time a request is cancelled in-flight the provider call was already - dispatched, so the input tokens were billed even if no chunk reached the - client. Refunding to zero would let a caller abort pre-token to dodge that - charge; the worst-case output portion of the reservation is still released. - - asyncio.shield keeps the reconcile running to completion even though the - surrounding task is being cancelled. The `finalized` guard makes this a no-op - when success/failure handling already reconciled, so calling it on every - cancellation path is safe. + A failed reconcile is not retried: the reconcile is an additive INCRBYFLOAT, and + a failure (e.g. a timeout after Redis applied it) leaves it unknown whether the + refund landed, so re-applying it could refund twice. Instead the reserved + counters are dropped so the next read reseeds from the DB, the same fallback + release_or_invalidate_budget_reservation uses on the non-cancel path. """ if not budget_reservation or budget_reservation.get("finalized") is True: return @@ -422,8 +415,18 @@ async def release_budget_reservation_on_cancel( await asyncio.shield( reconcile_budget_reservation(budget_reservation=budget_reservation, actual_cost=incurred_cost) ) - except (asyncio.CancelledError, Exception): - pass + except asyncio.CancelledError: + pass # a second cancellation while shielded; the reconcile keeps running detached regardless + except Exception: # noqa: BLE001 # a reconcile failure must not pin the counter; drop it directly instead + verbose_proxy_logger.exception("Failed to reconcile budget reservation on cancel; invalidating counters") + try: + await invalidate_budget_reservation_counters(budget_reservation=budget_reservation) + except Exception: # noqa: BLE001 # nothing left to try; the finalized stamp below keeps it from being reprocessed + verbose_proxy_logger.exception( + "Failed to invalidate budget reservation counters after cancel-path reconcile failed" + ) + finally: + budget_reservation["finalized"] = True async def invalidate_budget_reservation_counters( 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 131db55ee01..894e060343f 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 @@ -10,18 +10,23 @@ from unittest.mock import AsyncMock, MagicMock import httpx import prisma import pytest +from redis.asyncio import Redis - -from litellm.proxy._types import LiteLLM_VerificationToken +from litellm.proxy._types import LiteLLM_VerificationToken, UserAPIKeyAuth from litellm.proxy.common_utils import reset_budget_job as reset_budget_job_module from litellm.constants import ( PROXY_BUDGET_RESCHEDULER_MIN_TIME, RESET_BUDGET_JOB_BATCH_SIZE, RESET_BUDGET_JOB_LOCK_TTL_SECONDS, RESET_BUDGET_JOB_NAME, + RESET_BUDGET_SPEND_COUNTER_RESET_MAX_ATTEMPTS, +) +from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import ( + AuthCacheInvalidationSubscriber, ) from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob, _RowReset from litellm.proxy.common_utils.timezone_utils import BudgetResetSettings +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache # Mock classes for testing @@ -3578,3 +3583,245 @@ def test_reset_deletes_spend_counter_instead_of_seeding(reset_budget_job, mock_p counter_cache.redis_cache.async_delete_cache.assert_any_await(key="spend:user:carol") counter_cache.in_memory_cache.set_cache.assert_not_called() counter_cache.redis_cache.async_set_cache.assert_not_awaited() + + +# --------------------------------------------------------------------------- +# Path 2 (#30460): resetting a key/user/team's own budget must also drop its +# user_api_key_cache entry, not just the Redis spend counter, and that +# invalidation must reach every pod, not only the one that ran the sweep. +# --------------------------------------------------------------------------- + + +def _direct_reset_fake_module(monkeypatch, user_api_key_cache, redis_usage_cache=None): + spend_counter_cache = MagicMock() + spend_counter_cache.in_memory_cache.set_cache = MagicMock() + spend_counter_cache.redis_cache = None + + fake_module = types.ModuleType("litellm.proxy.proxy_server") + fake_module.spend_counter_cache = spend_counter_cache + fake_module.user_api_key_cache = user_api_key_cache + fake_module.redis_usage_cache = redis_usage_cache + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_module) + return spend_counter_cache + + +def test_reset_budget_for_keys_invalidates_user_api_key_cache(reset_budget_job, mock_prisma_client, monkeypatch): + """ + A stale user_api_key_cache entry must not survive a key's own budget + reset. Before the fix, _get_source_cache_base_spend could read this + cached object's pre-reset .spend straight back and re-seed the Redis + counter from it on the next request whose DB read fails, reinflating a + just-reset budget with no corresponding spend log (#30460 Path 2). + """ + user_api_key_cache = UserApiKeyCache() + user_api_key_cache.set_cache( + "sk-abc", UserAPIKeyAuth(token="sk-abc", spend=100.0, max_budget=50.0), model_type=UserAPIKeyAuth + ) + assert user_api_key_cache.get_cache("sk-abc", model_type=UserAPIKeyAuth) is not None + _direct_reset_fake_module(monkeypatch, user_api_key_cache) + + now = datetime.now(timezone.utc) + mock_prisma_client.data["key"] = [ + type( + "Key", + (), + {"spend": 100.0, "budget_duration": "30d", "budget_reset_at": now, "id": "key-1", "token": "sk-abc"}, + ) + ] + + asyncio.run(reset_budget_job.reset_budget_for_litellm_keys()) + + assert user_api_key_cache.get_cache("sk-abc", model_type=UserAPIKeyAuth) is None + + +def test_reset_budget_for_users_invalidates_user_api_key_cache(reset_budget_job, mock_prisma_client, monkeypatch): + """Same Path 2 gap on the user-scoped reset chunk.""" + from litellm.proxy._types import LiteLLM_UserTable + + user_api_key_cache = UserApiKeyCache() + user_api_key_cache.set_cache( + "alice", LiteLLM_UserTable(user_id="alice", spend=50.0, max_budget=20.0), model_type=LiteLLM_UserTable + ) + assert user_api_key_cache.get_cache("alice", model_type=LiteLLM_UserTable) is not None + _direct_reset_fake_module(monkeypatch, user_api_key_cache) + + now = datetime.now(timezone.utc) + mock_prisma_client.data["user"] = [ + type( + "User", + (), + {"spend": 50.0, "budget_duration": "7d", "budget_reset_at": now, "id": "user-1", "user_id": "alice"}, + ) + ] + + asyncio.run(reset_budget_job.reset_budget_for_litellm_users()) + + assert user_api_key_cache.get_cache("alice", model_type=LiteLLM_UserTable) is None + + +def test_reset_budget_for_teams_invalidates_user_api_key_cache(reset_budget_job, mock_prisma_client, monkeypatch): + """Same Path 2 gap on the team-scoped reset chunk; teams cache under + ``team_id:{team_id}`` (see auth_checks.get_team_object), not the bare id.""" + from litellm.proxy._types import LiteLLM_TeamTable + + user_api_key_cache = UserApiKeyCache() + user_api_key_cache.set_cache( + "team_id:team-x", + LiteLLM_TeamTable(team_id="team-x", spend=200.0, max_budget=100.0), + model_type=LiteLLM_TeamTable, + ) + assert user_api_key_cache.get_cache("team_id:team-x", model_type=LiteLLM_TeamTable) is not None + _direct_reset_fake_module(monkeypatch, user_api_key_cache) + + now = datetime.now(timezone.utc) + mock_prisma_client.data["team"] = [ + type( + "Team", + (), + {"spend": 200.0, "budget_duration": "1mo", "budget_reset_at": now, "id": "team-1", "team_id": "team-x"}, + ) + ] + + asyncio.run(reset_budget_job.reset_budget_for_litellm_teams()) + + assert user_api_key_cache.get_cache("team_id:team-x", model_type=LiteLLM_TeamTable) is None + + +class _LoopbackPubSubRedisClient(Redis): + """One fake client good enough to both publish() and pubsub() against the + same in-process queue, so a test can prove a message one pod publishes is + actually delivered to another pod's subscriber.""" + + def __init__(self) -> None: + self._queue: asyncio.Queue = asyncio.Queue() + + async def publish(self, channel: str, message: str) -> int: + self._queue.put_nowait({"type": "message", "data": message.encode()}) + return 1 + + def pubsub(self): + return self + + async def subscribe(self, *channels: str) -> None: + pass + + async def get_message(self, *, ignore_subscribe_messages: bool, timeout: float): + try: + return await asyncio.wait_for(self._queue.get(), timeout) + except asyncio.TimeoutError: + return None + + async def aclose(self) -> None: + pass + + +class _FakeCoordinationRedisCache: + def __init__(self, client: object) -> None: + self._client = client + self.namespace = None + + def init_pubsub_client(self) -> object: + return self._client + + +async def test_reset_budget_for_keys_broadcasts_cache_invalidation_to_other_pods( + reset_budget_job, mock_prisma_client, monkeypatch +): + """ + Only one pod runs the reset job per tick (_acquire_lease elects one + sweeper), so a local-only cache delete would leave every other pod's + user_api_key_cache serving the pre-reset spend until its TTL. This proves + the invalidation is actually broadcast (LIT-3803) to a second pod that + never ran the sweep, closing the multi-pod half of #30460 Path 2. + """ + shared_redis_cache = _FakeCoordinationRedisCache(client=_LoopbackPubSubRedisClient()) + + leader_cache = UserApiKeyCache() + follower_cache = UserApiKeyCache() + for cache in (leader_cache, follower_cache): + cache.set_cache( + "sk-abc", UserAPIKeyAuth(token="sk-abc", spend=100.0, max_budget=50.0), model_type=UserAPIKeyAuth + ) + + _direct_reset_fake_module(monkeypatch, leader_cache, redis_usage_cache=shared_redis_cache) + + subscriber = AuthCacheInvalidationSubscriber(redis_cache=shared_redis_cache, user_api_key_cache=follower_cache) + subscriber.start() + + now = datetime.now(timezone.utc) + mock_prisma_client.data["key"] = [ + type( + "Key", + (), + {"spend": 100.0, "budget_duration": "30d", "budget_reset_at": now, "id": "key-1", "token": "sk-abc"}, + ) + ] + + try: + await reset_budget_job.reset_budget_for_litellm_keys() + + for _ in range(200): + if follower_cache.get_cache("sk-abc", model_type=UserAPIKeyAuth) is None: + break + await asyncio.sleep(0.01) + finally: + await subscriber.stop() + + assert leader_cache.get_cache("sk-abc", model_type=UserAPIKeyAuth) is None + assert follower_cache.get_cache("sk-abc", model_type=UserAPIKeyAuth) is None + + +# --------------------------------------------------------------------------- +# Path 3 (#30460): a failed Redis delete on budget reset must not silently +# leave the inflated pre-reset counter authoritative in Redis. +# --------------------------------------------------------------------------- + + +def _counter_cache_with_redis_delete(monkeypatch, delete_side_effect=None): + spend_counter_cache = MagicMock() + spend_counter_cache.redis_cache = MagicMock() + spend_counter_cache.redis_cache.async_delete_cache = AsyncMock(side_effect=delete_side_effect) + + 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()) + return spend_counter_cache + + +def test_invalidate_spend_counter_retries_the_redis_delete_a_bounded_number_of_times(monkeypatch): + """Every delete fails: retried up to the configured limit, then given up on + without raising out of the reset job.""" + spend_counter_cache = _counter_cache_with_redis_delete( + monkeypatch, delete_side_effect=RuntimeError("elasticache timeout") + ) + + asyncio.run(ResetBudgetJob._invalidate_spend_counter("spend:key:sk-inflated")) + + spend_counter_cache.in_memory_cache.delete_cache.assert_called_once_with(key="spend:key:sk-inflated") + assert ( + spend_counter_cache.redis_cache.async_delete_cache.await_count + == RESET_BUDGET_SPEND_COUNTER_RESET_MAX_ATTEMPTS + ) + assert asyncio.sleep.await_count == RESET_BUDGET_SPEND_COUNTER_RESET_MAX_ATTEMPTS - 1 + + +def test_invalidate_spend_counter_recovers_after_a_transient_redis_failure(monkeypatch): + """A delete that fails once and then succeeds stops retrying.""" + spend_counter_cache = _counter_cache_with_redis_delete( + monkeypatch, delete_side_effect=[RuntimeError("elasticache timeout"), None] + ) + + asyncio.run(ResetBudgetJob._invalidate_spend_counter("spend:key:sk-recovered")) + + assert spend_counter_cache.redis_cache.async_delete_cache.await_count == 2 + assert asyncio.sleep.await_count == 1 + + +def test_invalidate_spend_counter_deletes_once_when_redis_is_healthy(monkeypatch): + spend_counter_cache = _counter_cache_with_redis_delete(monkeypatch) + + asyncio.run(ResetBudgetJob._invalidate_spend_counter("spend:key:sk-healthy")) + + spend_counter_cache.redis_cache.async_delete_cache.assert_awaited_once_with(key="spend:key:sk-healthy") + asyncio.sleep.assert_not_awaited() diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index c8e4df1030f..529cda6c83a 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -2483,6 +2483,34 @@ class _TeamMembershipFloorDb: return SimpleNamespace(find_unique=AsyncMock(return_value=row)) +class _FlakyPipelineRedisCache(_ExpiringRedisCache): + """_ExpiringRedisCache, but the first ``fail_first_n`` calls to + async_increment_pipeline time out. With ``apply_before_failing`` the + increments land before the timeout surfaces, the ambiguous case where the + caller can't tell whether its write applied; with ``fail_deletes`` the + counter delete that would otherwise clean that up fails too.""" + + def __init__(self, fail_first_n: int = 0, apply_before_failing: bool = False, fail_deletes: bool = False) -> None: + super().__init__() + self.fail_first_n = fail_first_n + self.apply_before_failing = apply_before_failing + self.fail_deletes = fail_deletes + 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: + if self.apply_before_failing: + await super().async_increment_pipeline(increment_list, **kwargs) + raise TimeoutError("redis timeout") + return await super().async_increment_pipeline(increment_list, **kwargs) + + async def async_delete_cache(self, key: str, *args: object, **kwargs: object) -> None: + if self.fail_deletes: + raise ConnectionError("redis unreachable") + await super().async_delete_cache(key, *args, **kwargs) + + @pytest.mark.asyncio async def test_reconcile_after_redis_counter_expiry_keeps_request_cost_enforced( spend_counter_state, @@ -3143,6 +3171,120 @@ async def test_release_budget_reservation_on_cancel_swallows_release_errors(): await release_budget_reservation_on_cancel(reservation) +@pytest.mark.asyncio +async def test_release_budget_reservation_on_cancel_swallows_a_second_cancellation_while_shielded(): + # A second CancelledError arriving while the shielded reconcile is in flight must not + # propagate: the reconcile keeps running detached regardless, and there is nothing more + # for this call to do but return. + reservation = { + "reserved_cost": 3.0, + "entries": [{"counter_key": "spend:key:key-cancel-twice"}], + "finalized": False, + "input_cost": 0.5, + } + with patch( # test-quality-ok: reconcile_budget_reservation is a module function, no DI seam for this test + "litellm.proxy.spend_tracking.budget_reservation.reconcile_budget_reservation", + new=AsyncMock(side_effect=asyncio.CancelledError()), + ): + # must return without raising, and without taking the finalize-on-failure path either: + # a second cancellation isn't a failure, it's the reconcile still running in the background + await release_budget_reservation_on_cancel(reservation) + + assert reservation["finalized"] is False + + +@pytest.mark.asyncio +async def test_release_budget_reservation_on_cancel_swallows_invalidate_failure_after_reconcile_fails(): + # If both the reconcile and the invalidate fallback fail (e.g. a persistent + # outage), there is nothing left to try: the failure must be logged and swallowed, + # not propagated, and the reservation still ends up finalized so it is not reprocessed. + reservation = { + "reserved_cost": 3.0, + "entries": [{"counter_key": "spend:key:key-cancel-double-failure"}], + "finalized": False, + "input_cost": 0.5, + } + with patch( # test-quality-ok: reconcile_budget_reservation is a module function, no DI seam for this test + "litellm.proxy.spend_tracking.budget_reservation.reconcile_budget_reservation", + new=AsyncMock(side_effect=RuntimeError("redis down")), + ), patch( # test-quality-ok: invalidate_budget_reservation_counters is a module function, no DI seam for this test + "litellm.proxy.spend_tracking.budget_reservation.invalidate_budget_reservation_counters", + new=AsyncMock(side_effect=RuntimeError("redis still down")), + ): + # must return without raising + await release_budget_reservation_on_cancel(reservation) + + assert reservation["finalized"] is True + + +@pytest.mark.asyncio +async def test_release_budget_reservation_on_cancel_invalidates_counter_when_reconcile_fails( + spend_counter_state, +): + """ + Regression for #30460 Path 1: if reconcile fails on the cancel path (e.g. a + Redis outage), 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: drop the counter so the next read reseeds from the DB instead + of enforcing the stale reservation until the TTL. + """ + 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 = { + "reserved_cost": 3.0, + "input_cost": 0.5, + "finalized": False, + "entries": [{"counter_key": counter_key, "reserved_cost": 3.0, "applied_adjustment": 0.0}], + } + + await release_budget_reservation_on_cancel(reservation) + + assert counter_key not in counter_cache.redis_cache.store + assert counter_cache.in_memory_cache.get_cache(key=counter_key) is None + assert reservation["finalized"] is True + assert counter_cache.redis_cache.pipeline_calls == 1 + + +@pytest.mark.asyncio +async def test_release_budget_reservation_on_cancel_does_not_refund_twice_when_reconcile_times_out_after_applying( + spend_counter_state, +): + """ + Regression for the duplicate-refund race (veria-ai finding on + budget_reservation.py): the reconcile's INCRBYFLOAT lands in Redis but the + pipeline response times out, and the counter delete meant to clean that up + fails too, so the counter still holds the already-refunded value. Re-applying + the reconcile would subtract this reservation's refund a second time and eat + into a concurrent request's spend sharing the counter. The refund must land + exactly once. + """ + counter_cache, _key_cache = spend_counter_state + counter_key = "spend:key:key-cancel-timeout-after-apply" + counter_cache.redis_cache = _FlakyPipelineRedisCache(fail_first_n=99, apply_before_failing=True, fail_deletes=True) + # 2.0 of a concurrent request's spend plus this request's 3.0 reservation + counter_cache.redis_cache.store[counter_key] = 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}], + } + + await release_budget_reservation_on_cancel(reservation) + await release_budget_reservation_on_cancel(reservation) + + # 5.0 - (3.0 - 0.5), applied once; a second application would leave 0.0 + assert counter_cache.redis_cache.store[counter_key] == pytest.approx(2.5) + assert counter_cache.redis_cache.pipeline_calls == 1 + assert reservation["finalized"] is True + + @pytest.mark.asyncio async def test_streaming_cancel_in_slow_path_before_yield_refunds(spend_counter_state): counter_cache, key_cache = spend_counter_state