mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(proxy): preserve spend adjustments under lock registry pressure
This commit is contained in:
parent
c2c2a623c0
commit
d03c209a2e
4 changed files with 98 additions and 40 deletions
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
2
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
2
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue