This commit is contained in:
Joscha Götzer 2026-09-15 15:53:56 -06:00 • committed by GitHub
commit 3b257f2a50
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 98 additions and 38 deletions

View file

@ -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]

View file

@ -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))

View file

@ -8002,7 +8002,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
@ -8061,13 +8061,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,
@ -8944,7 +8944,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
@ -8955,7 +8957,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
@ -8967,6 +8970,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