diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py index 89a07234c6c..2f063341ba7 100644 --- a/litellm/proxy/db/spend_counter_reseed.py +++ b/litellm/proxy/db/spend_counter_reseed.py @@ -8,13 +8,15 @@ writes from other pods, so trusting it allows budget bypass in multi-pod deployments. This module reseeds from the authoritative DB instead. A per-counter singleflight lock collapses concurrent reseeds on the same pod -to one DB query per cold-cache window. The lock dict is bounded LRU to cap -memory in long-lived deployments. +to one DB query per cold-cache window. The lock registry retains active locks +and bounds idle entries with LRU eviction in long-lived deployments. """ import asyncio from collections import OrderedDict -from collections.abc import Mapping +from collections.abc import AsyncGenerator, Mapping +from contextlib import asynccontextmanager +from dataclasses import dataclass, replace from datetime import datetime, timezone from types import MappingProxyType from typing import TYPE_CHECKING, ClassVar, Final, Optional @@ -66,6 +68,12 @@ def _as_utc(value: datetime) -> datetime: return value if value.tzinfo is not None else value.replace(tzinfo=timezone.utc) +@dataclass(frozen=True, slots=True) +class _CounterLock: + lock: asyncio.Lock + users: int = 0 + + class SpendCounterReseed: """ Reseeds spend counters from the authoritative DB and warms the cache, @@ -87,29 +95,37 @@ class SpendCounterReseed: is the row the reset zeroed. """ - _locks: ClassVar["OrderedDict[str, asyncio.Lock]"] = OrderedDict() - _registry_lock: ClassVar[asyncio.Lock | None] = None + _locks: ClassVar["OrderedDict[str, _CounterLock]"] = OrderedDict() + _idle_locks: ClassVar["OrderedDict[str, None]"] = OrderedDict() @staticmethod - async def _get_lock(counter_key: str) -> asyncio.Lock: - if SpendCounterReseed._registry_lock is None: - SpendCounterReseed._registry_lock = asyncio.Lock() - async with SpendCounterReseed._registry_lock: - lock = SpendCounterReseed._locks.get(counter_key) - if lock is not None: - SpendCounterReseed._locks.move_to_end(counter_key) - return lock - lock = asyncio.Lock() - SpendCounterReseed._locks[counter_key] = lock - if len(SpendCounterReseed._locks) > SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE: - SpendCounterReseed._locks.popitem(last=False) - return lock + def _prune_idle_locks() -> None: + while len(SpendCounterReseed._locks) > SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE and SpendCounterReseed._idle_locks: + idle_key, _ = SpendCounterReseed._idle_locks.popitem(last=False) + SpendCounterReseed._locks.pop(idle_key) + + @staticmethod + @asynccontextmanager + async def _counter_lock(counter_key: str) -> AsyncGenerator[None]: + counter_lock: Final = SpendCounterReseed._locks.get(counter_key) or _CounterLock(lock=asyncio.Lock()) + SpendCounterReseed._locks[counter_key] = replace(counter_lock, users=counter_lock.users + 1) + SpendCounterReseed._idle_locks.pop(counter_key, None) + SpendCounterReseed._prune_idle_locks() + try: + async with counter_lock.lock: + yield + finally: + current: Final = SpendCounterReseed._locks[counter_key] + remaining: Final = replace(current, users=current.users - 1) + SpendCounterReseed._locks[counter_key] = remaining + if remaining.users == 0: + SpendCounterReseed._idle_locks[counter_key] = None + SpendCounterReseed._prune_idle_locks() @staticmethod async def increment_in_memory(spend_counter_cache: "DualCache", counter_key: str, increment: float) -> float | None: """Apply local deltas after an in-flight reseed establishes the spend balance.""" - lock: Final = await SpendCounterReseed._get_lock(counter_key) - async with lock: + async with SpendCounterReseed._counter_lock(counter_key): return await spend_counter_cache.async_increment_cache( key=counter_key, value=increment, local_only=True, refresh_ttl=True ) @@ -215,8 +231,7 @@ class SpendCounterReseed: Returns the spend value (including 0.0 from a fresh budget reset) when the DB read succeeds, or None when the DB is unavailable. """ - lock: Final = await SpendCounterReseed._get_lock(counter_key) - async with lock: + async with SpendCounterReseed._counter_lock(counter_key): batched: Final = await SpendCounterReseed._read_active_batch(counter_key) if batched is not None and batched[0] is not None: return batched[0] @@ -412,8 +427,7 @@ class SpendCounterReseed: window_duration: str | None, window_start: datetime, ) -> float | None: - lock: Final = await SpendCounterReseed._get_lock(counter_key) - async with lock: + async with SpendCounterReseed._counter_lock(counter_key): batched: Final = await SpendCounterReseed._read_active_batch(counter_key) if batched is not None and batched[0] is not None: return batched[0] diff --git a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py index bca6344b3f7..49b720ccb54 100644 --- a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py +++ b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py @@ -336,7 +336,9 @@ async def test_cold_reseed_preserves_concurrent_local_increment( monkeypatch: pytest.MonkeyPatch, window: bool, batch: bool, increment: float ) -> None: from litellm.proxy import proxy_server + from litellm.proxy.db import spend_counter_reseed + monkeypatch.setattr(spend_counter_reseed, "SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 2) cache: Final = DualCache(in_memory_cache=InMemoryCache()) counter_key: Final = ( f"spend:team:concurrent-{batch}-{increment}:window:1d" @@ -348,20 +350,59 @@ async def test_cold_reseed_preserves_concurrent_local_increment( reseed_task: Final = asyncio.create_task(_reseed_with_paused_table(table, cache, counter_key, window)) await asyncio.wait_for(table.read_started.wait(), timeout=5) - increment_task: Final = asyncio.create_task( - proxy_server._apply_spend_counter_increments( - pending=(proxy_server.PendingSpendIncrement(counter_key=counter_key, increment=increment),) - ) - if batch - else proxy_server._increment_spend_counter_cache(counter_key=counter_key, increment=increment) + pressure_tables: Final = tuple(_PausedSpendTable(20.0) for _ in range(4)) + pressure_tasks: Final = tuple( + asyncio.create_task(_reseed_with_paused_table(pressure_table, cache, f"spend:user:pressure-{index}", False)) + for index, pressure_table in enumerate(pressure_tables) ) - await asyncio.sleep(0) - table.resume_read.set() - await asyncio.wait_for(asyncio.gather(reseed_task, increment_task), timeout=5) + await asyncio.wait_for(asyncio.gather(*(item.read_started.wait() for item in pressure_tables)), timeout=5) + try: + increment_task: Final = asyncio.create_task( + proxy_server._apply_spend_counter_increments( + pending=(proxy_server.PendingSpendIncrement(counter_key=counter_key, increment=increment),) + ) + if batch + else proxy_server._increment_spend_counter_cache(counter_key=counter_key, increment=increment) + ) + await asyncio.sleep(0) + table.resume_read.set() + await asyncio.wait_for(asyncio.gather(reseed_task, increment_task), timeout=5) + finally: + for item in pressure_tables: + item.resume_read.set() + await asyncio.wait_for(asyncio.gather(*pressure_tasks), timeout=5) assert cache.in_memory_cache.get_cache(key=counter_key) == 100.0 + increment +@pytest.mark.asyncio +@pytest.mark.parametrize("cancel_holder", [False, True], ids=["waiter", "holder"]) +async def test_cancelled_counter_operation_releases_usage(monkeypatch: pytest.MonkeyPatch, cancel_holder: bool) -> None: + from litellm.proxy.db import spend_counter_reseed + + monkeypatch.setattr(spend_counter_reseed, "SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 1) + cache: Final = DualCache(in_memory_cache=InMemoryCache()) + counter_key: Final = f"spend:user:cancel-{cancel_holder}" + table: Final = _PausedSpendTable(100.0) + holder: Final = asyncio.create_task(_reseed_with_paused_table(table, cache, counter_key, False)) + await asyncio.wait_for(table.read_started.wait(), timeout=5) + waiter: Final = asyncio.create_task(SpendCounterReseed.increment_in_memory(cache, counter_key, 5.0)) + await asyncio.sleep(0) + assert not waiter.done() + + cancelled: Final = holder if cancel_holder else waiter + cancelled.cancel() + with pytest.raises(asyncio.CancelledError): + await cancelled + table.resume_read.set() + await asyncio.wait_for(waiter if cancel_holder else holder, timeout=5) + + assert await SpendCounterReseed.increment_in_memory(cache, counter_key, 2.0) == (7.0 if cancel_holder else 102.0) + async with SpendCounterReseed._counter_lock("spend:user:evict-cancelled"): + pass + assert counter_key not in SpendCounterReseed._locks + + @pytest.mark.asyncio async def test_end_user_from_db_reads_the_end_user_row_by_user_id(): prisma: Final = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="customer-42", spend=0.0)) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 03a24ec5e98..d136ebbddab 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -7915,7 +7915,7 @@ async def test_primary_spend_counter_redis_concurrent_seed_does_not_double_seed( 2 * db_spend. The per-counter asyncio.Lock is per-process, so it does NOT coordinate - across pods. We simulate two pods by patching _get_lock to return a + across pods. We simulate two pods by patching _counter_lock to return a fresh lock per call (each "pod" has its own lock registry in real life). """ from litellm.caching.dual_cache import DualCache @@ -7974,13 +7974,13 @@ async def test_primary_spend_counter_redis_concurrent_seed_does_not_double_seed( pod_b = DualCache() pod_b.redis_cache = fake_redis - # Each "pod" has its own per-process lock registry. Patch _get_lock to + # Each "pod" has its own per-process lock registry. Patch _counter_lock to # always return a fresh lock so the two coalesced calls do not serialize # via one in-process lock (which is what would happen across pods). - async def fresh_lock(_counter_key): + def fresh_lock(_counter_key): return asyncio.Lock() - with patch.object(SpendCounterReseed, "_get_lock", side_effect=fresh_lock): + with patch.object(SpendCounterReseed, "_counter_lock", side_effect=fresh_lock): results = await asyncio.gather( SpendCounterReseed.coalesced( prisma_client=fake_prisma, @@ -8857,7 +8857,9 @@ async def test_reseed_locks_dict_is_bounded(): from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed orig_locks = SpendCounterReseed._locks.copy() + orig_idle_locks = SpendCounterReseed._idle_locks.copy() SpendCounterReseed._locks.clear() + SpendCounterReseed._idle_locks.clear() orig_max = constants.SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE constants.SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE = 5 # The class reads the constant via module-level import, so patch the @@ -8868,7 +8870,8 @@ async def test_reseed_locks_dict_is_bounded(): scr.SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE = 5 try: for i in range(7): - await SpendCounterReseed._get_lock(f"spend:key:test-key-{i}") + async with SpendCounterReseed._counter_lock(f"spend:key:test-key-{i}"): + pass assert len(SpendCounterReseed._locks) == 5, f"got {len(SpendCounterReseed._locks)}" # Oldest two evicted assert "spend:key:test-key-0" not in SpendCounterReseed._locks @@ -8880,6 +8883,8 @@ async def test_reseed_locks_dict_is_bounded(): scr.SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE = orig_module_max SpendCounterReseed._locks.clear() SpendCounterReseed._locks.update(orig_locks) + SpendCounterReseed._idle_locks.clear() + SpendCounterReseed._idle_locks.update(orig_idle_locks) @pytest.mark.asyncio diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 7eadaa6c991..839aa52fa84 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -16781,7 +16781,6 @@ export interface paths { * - permissions: Optional[dict] - [Not Implemented Yet] User-specific permissions, eg. turning off pii masking. * - metadata: Optional[dict] - Metadata for user, store information for user. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" } * - max_parallel_requests: Optional[int] - Rate limit a user based on the number of parallel requests. Raises 429 error, if user's parallel requests > x. - * - soft_budget: Optional[float] - Get alerts when user crosses given budget, doesn't block requests. * - model_max_budget: Optional[dict] - Model-specific max budget for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-budgets-to-keys) * - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}. * - model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) @@ -16887,7 +16886,6 @@ export interface paths { * - permissions: Optional[dict] - [Not Implemented Yet] User-specific permissions, eg. turning off pii masking. * - metadata: Optional[dict] - Metadata for user, store information for user. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" } * - max_parallel_requests: Optional[int] - Rate limit a user based on the number of parallel requests. Raises 429 error, if user's parallel requests > x. - * - soft_budget: Optional[float] - Get alerts when user crosses given budget, doesn't block requests. * - model_max_budget: Optional[dict] - Model-specific max budget for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-budgets-to-keys) * - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}. * - model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)