mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
tighten budget counter cache recovery
This commit is contained in:
parent
5a619cf879
commit
926de696a1
4 changed files with 352 additions and 23 deletions
|
|
@ -14,6 +14,7 @@ memory in long-lived deployments.
|
|||
|
||||
import asyncio
|
||||
from collections import OrderedDict
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, ClassVar, Optional
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -118,6 +119,7 @@ class SpendCounterReseed:
|
|||
prisma_client: Optional["PrismaClient"],
|
||||
spend_counter_cache: "DualCache",
|
||||
counter_key: str,
|
||||
require_cache_warm: bool = False,
|
||||
) -> Optional[float]:
|
||||
"""
|
||||
Reseed a cold spend counter from the DB and warm the cache,
|
||||
|
|
@ -148,12 +150,136 @@ class SpendCounterReseed:
|
|||
return None
|
||||
# Warm even when 0 so subsequent reads hit cache, not DB.
|
||||
try:
|
||||
await spend_counter_cache.async_increment_cache(
|
||||
key=counter_key, value=db_spend
|
||||
)
|
||||
if require_cache_warm and spend_counter_cache.redis_cache is not None:
|
||||
current_value = (
|
||||
await spend_counter_cache.redis_cache.async_increment(
|
||||
key=counter_key,
|
||||
value=db_spend,
|
||||
)
|
||||
)
|
||||
spend_counter_cache.in_memory_cache.set_cache(
|
||||
key=counter_key,
|
||||
value=current_value,
|
||||
)
|
||||
else:
|
||||
await spend_counter_cache.async_increment_cache(
|
||||
key=counter_key, value=db_spend
|
||||
)
|
||||
except Exception:
|
||||
verbose_proxy_logger.exception(
|
||||
"SpendCounterReseed.coalesced: failed to warm counter %s",
|
||||
counter_key,
|
||||
)
|
||||
if require_cache_warm:
|
||||
raise
|
||||
return db_spend
|
||||
|
||||
@staticmethod
|
||||
async def window_from_spend_logs(
|
||||
prisma_client: Optional["PrismaClient"],
|
||||
entity_type: str,
|
||||
entity_id: str,
|
||||
window_start: datetime,
|
||||
) -> Optional[float]:
|
||||
if prisma_client is None:
|
||||
return None
|
||||
|
||||
if entity_type == "Key":
|
||||
group_field = "api_key"
|
||||
where = {
|
||||
"api_key": entity_id,
|
||||
"startTime": {"gte": window_start},
|
||||
}
|
||||
elif entity_type == "Team":
|
||||
group_field = "team_id"
|
||||
where = {
|
||||
"team_id": entity_id,
|
||||
"startTime": {"gte": window_start},
|
||||
}
|
||||
else:
|
||||
return None
|
||||
|
||||
try:
|
||||
response = await prisma_client.db.litellm_spendlogs.group_by(
|
||||
by=[group_field],
|
||||
where=where, # type: ignore[arg-type]
|
||||
sum={"spend": True},
|
||||
)
|
||||
except Exception:
|
||||
verbose_proxy_logger.exception(
|
||||
"SpendCounterReseed.window_from_spend_logs: failed for %s=%s",
|
||||
entity_type,
|
||||
entity_id,
|
||||
)
|
||||
return None
|
||||
|
||||
if not response:
|
||||
return 0.0
|
||||
first_row = response[0]
|
||||
sum_row = (
|
||||
first_row.get("_sum")
|
||||
if isinstance(first_row, dict)
|
||||
else getattr(first_row, "_sum", None)
|
||||
)
|
||||
spend = (
|
||||
sum_row.get("spend")
|
||||
if isinstance(sum_row, dict)
|
||||
else getattr(sum_row, "spend", None)
|
||||
)
|
||||
return float(spend or 0.0)
|
||||
|
||||
@staticmethod
|
||||
async def coalesced_window(
|
||||
prisma_client: Optional["PrismaClient"],
|
||||
spend_counter_cache: "DualCache",
|
||||
counter_key: str,
|
||||
entity_type: str,
|
||||
entity_id: str,
|
||||
window_start: datetime,
|
||||
) -> Optional[float]:
|
||||
lock = await SpendCounterReseed._get_lock(counter_key)
|
||||
async with lock:
|
||||
if spend_counter_cache.redis_cache is not None:
|
||||
try:
|
||||
val = await spend_counter_cache.redis_cache.async_get_cache(
|
||||
key=counter_key
|
||||
)
|
||||
if val is not None:
|
||||
return float(val)
|
||||
except Exception:
|
||||
pass
|
||||
val = spend_counter_cache.in_memory_cache.get_cache(key=counter_key)
|
||||
if val is not None:
|
||||
return float(val)
|
||||
|
||||
window_spend = await SpendCounterReseed.window_from_spend_logs(
|
||||
prisma_client=prisma_client,
|
||||
entity_type=entity_type,
|
||||
entity_id=entity_id,
|
||||
window_start=window_start,
|
||||
)
|
||||
if window_spend is None:
|
||||
return None
|
||||
try:
|
||||
if spend_counter_cache.redis_cache is not None:
|
||||
current_value = (
|
||||
await spend_counter_cache.redis_cache.async_increment(
|
||||
key=counter_key,
|
||||
value=window_spend,
|
||||
)
|
||||
)
|
||||
spend_counter_cache.in_memory_cache.set_cache(
|
||||
key=counter_key,
|
||||
value=current_value,
|
||||
)
|
||||
else:
|
||||
await spend_counter_cache.async_increment_cache(
|
||||
key=counter_key, value=window_spend
|
||||
)
|
||||
except Exception:
|
||||
verbose_proxy_logger.exception(
|
||||
"SpendCounterReseed.coalesced_window: failed to warm counter %s",
|
||||
counter_key,
|
||||
)
|
||||
raise
|
||||
return window_spend
|
||||
|
|
|
|||
|
|
@ -1897,9 +1897,16 @@ async def increment_spend_counters(
|
|||
)
|
||||
key_window_counter = f"spend:key:{hashed_token}:window:{duration}"
|
||||
if key_window_counter not in reserved_counter_keys:
|
||||
await spend_counter_cache.async_increment_cache(
|
||||
key=key_window_counter,
|
||||
value=response_cost,
|
||||
from litellm.proxy.spend_tracking.budget_reservation import (
|
||||
get_budget_window_start,
|
||||
)
|
||||
|
||||
await _init_and_increment_window_spend_counter(
|
||||
counter_key=key_window_counter,
|
||||
entity_type="Key",
|
||||
entity_id=hashed_token,
|
||||
window_start=get_budget_window_start(window),
|
||||
increment=response_cost,
|
||||
)
|
||||
|
||||
if team_id is not None:
|
||||
|
|
@ -1928,9 +1935,16 @@ async def increment_spend_counters(
|
|||
)
|
||||
team_window_counter = f"spend:team:{team_id}:window:{duration}"
|
||||
if team_window_counter not in reserved_counter_keys:
|
||||
await spend_counter_cache.async_increment_cache(
|
||||
key=team_window_counter,
|
||||
value=response_cost,
|
||||
from litellm.proxy.spend_tracking.budget_reservation import (
|
||||
get_budget_window_start,
|
||||
)
|
||||
|
||||
await _init_and_increment_window_spend_counter(
|
||||
counter_key=team_window_counter,
|
||||
entity_type="Team",
|
||||
entity_id=team_id,
|
||||
window_start=get_budget_window_start(window),
|
||||
increment=response_cost,
|
||||
)
|
||||
|
||||
if user_id is not None and team_id is not None:
|
||||
|
|
@ -1986,7 +2000,24 @@ async def _init_and_increment_spend_counter(
|
|||
counter_key=counter_key,
|
||||
source_cache_key=source_cache_key,
|
||||
)
|
||||
await spend_counter_cache.async_increment_cache(key=counter_key, value=increment)
|
||||
await _increment_spend_counter_cache(counter_key=counter_key, increment=increment)
|
||||
|
||||
|
||||
async def _init_and_increment_window_spend_counter(
|
||||
counter_key: str,
|
||||
entity_type: str,
|
||||
entity_id: str,
|
||||
window_start: Optional[datetime],
|
||||
increment: float,
|
||||
):
|
||||
if window_start is not None:
|
||||
await _ensure_window_spend_counter_initialized(
|
||||
counter_key=counter_key,
|
||||
entity_type=entity_type,
|
||||
entity_id=entity_id,
|
||||
window_start=window_start,
|
||||
)
|
||||
await _increment_spend_counter_cache(counter_key=counter_key, increment=increment)
|
||||
|
||||
|
||||
async def _ensure_spend_counter_initialized(
|
||||
|
|
@ -2000,6 +2031,7 @@ async def _ensure_spend_counter_initialized(
|
|||
prisma_client=prisma_client,
|
||||
spend_counter_cache=spend_counter_cache,
|
||||
counter_key=counter_key,
|
||||
require_cache_warm=True,
|
||||
)
|
||||
if db_spend is None:
|
||||
# DB unavailable - fall back to in-process cache (may be stale).
|
||||
|
|
@ -2011,11 +2043,66 @@ async def _ensure_spend_counter_initialized(
|
|||
else:
|
||||
base_spend = getattr(source, "spend", 0.0) or 0.0
|
||||
if base_spend > 0:
|
||||
await spend_counter_cache.async_increment_cache(
|
||||
key=counter_key, value=base_spend
|
||||
await _increment_spend_counter_cache(
|
||||
counter_key=counter_key, increment=base_spend
|
||||
)
|
||||
|
||||
|
||||
async def _ensure_window_spend_counter_initialized(
|
||||
counter_key: str,
|
||||
entity_type: str,
|
||||
entity_id: str,
|
||||
window_start: datetime,
|
||||
):
|
||||
current = await spend_counter_cache.async_get_cache(key=counter_key)
|
||||
if current is None:
|
||||
window_spend = await SpendCounterReseed.coalesced_window(
|
||||
prisma_client=prisma_client,
|
||||
spend_counter_cache=spend_counter_cache,
|
||||
counter_key=counter_key,
|
||||
entity_type=entity_type,
|
||||
entity_id=entity_id,
|
||||
window_start=window_start,
|
||||
)
|
||||
if window_spend is None:
|
||||
await _increment_spend_counter_cache(counter_key=counter_key, increment=0.0)
|
||||
|
||||
|
||||
async def _increment_spend_counter_cache(counter_key: str, increment: float):
|
||||
if spend_counter_cache.redis_cache is not None:
|
||||
try:
|
||||
current_value = await spend_counter_cache.redis_cache.async_increment(
|
||||
key=counter_key,
|
||||
value=increment,
|
||||
)
|
||||
except Exception:
|
||||
await _invalidate_spend_counter(counter_key=counter_key)
|
||||
raise
|
||||
spend_counter_cache.in_memory_cache.set_cache(
|
||||
key=counter_key,
|
||||
value=current_value,
|
||||
)
|
||||
return current_value
|
||||
|
||||
return await spend_counter_cache.async_increment_cache(
|
||||
key=counter_key,
|
||||
value=increment,
|
||||
)
|
||||
|
||||
|
||||
async def _invalidate_spend_counter(counter_key: str):
|
||||
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:
|
||||
verbose_proxy_logger.debug(
|
||||
"Unable to delete stale spend counter %s after increment failure",
|
||||
counter_key,
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
|
||||
async def update_cache( # noqa: PLR0915
|
||||
token: Optional[str],
|
||||
user_id: Optional[str],
|
||||
|
|
|
|||
|
|
@ -2,11 +2,13 @@ from __future__ import annotations
|
|||
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any, Dict, List, Optional, Sequence, cast
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching import DualCache
|
||||
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_TeamMembership,
|
||||
LiteLLM_TeamTable,
|
||||
|
|
@ -27,6 +29,8 @@ class _BudgetCounter:
|
|||
entity_type: str
|
||||
entity_id: str
|
||||
source_cache_key: Optional[str] = None
|
||||
spend_log_entity_id: Optional[str] = None
|
||||
window_start: Optional[datetime] = None
|
||||
|
||||
|
||||
def get_reserved_counter_keys(budget_reservation: Optional[dict]) -> set:
|
||||
|
|
@ -349,8 +353,7 @@ def _get_budget_limit_counters(
|
|||
max_budget = window_dict.get("max_budget")
|
||||
if not budget_duration or max_budget is None or max_budget <= 0:
|
||||
continue
|
||||
# Window counters intentionally have no source_cache_key: the DB stores
|
||||
# window definitions/reset times, but not accumulated per-window spend.
|
||||
window_start = get_budget_window_start(window_dict)
|
||||
counters.append(
|
||||
_BudgetCounter(
|
||||
counter_key=f"{entity_prefix}:window:{budget_duration}",
|
||||
|
|
@ -358,6 +361,8 @@ def _get_budget_limit_counters(
|
|||
fallback_spend=0.0,
|
||||
entity_type=entity_type,
|
||||
entity_id=f"{entity_id}:{budget_duration}",
|
||||
spend_log_entity_id=entity_id,
|
||||
window_start=window_start,
|
||||
)
|
||||
)
|
||||
return counters
|
||||
|
|
@ -409,7 +414,8 @@ async def _reserve_counter(
|
|||
) -> Optional[float]:
|
||||
from litellm.proxy.proxy_server import (
|
||||
_ensure_spend_counter_initialized,
|
||||
spend_counter_cache,
|
||||
_ensure_window_spend_counter_initialized,
|
||||
_increment_spend_counter_cache,
|
||||
)
|
||||
|
||||
if counter.source_cache_key is not None:
|
||||
|
|
@ -417,10 +423,17 @@ async def _reserve_counter(
|
|||
counter_key=counter.counter_key,
|
||||
source_cache_key=counter.source_cache_key,
|
||||
)
|
||||
elif counter.spend_log_entity_id is not None and counter.window_start is not None:
|
||||
await _ensure_window_spend_counter_initialized(
|
||||
counter_key=counter.counter_key,
|
||||
entity_type=counter.entity_type,
|
||||
entity_id=counter.spend_log_entity_id,
|
||||
window_start=counter.window_start,
|
||||
)
|
||||
|
||||
reserved_value = await spend_counter_cache.async_increment_cache(
|
||||
key=counter.counter_key,
|
||||
value=reservation_cost,
|
||||
reserved_value = await _increment_spend_counter_cache(
|
||||
counter_key=counter.counter_key,
|
||||
increment=reservation_cost,
|
||||
)
|
||||
return float(reserved_value) if reserved_value is not None else None
|
||||
|
||||
|
|
@ -438,7 +451,7 @@ async def _set_reserved_entries_adjustment(
|
|||
entries: List[dict],
|
||||
target_adjustment: float,
|
||||
) -> None:
|
||||
from litellm.proxy.proxy_server import spend_counter_cache
|
||||
from litellm.proxy.proxy_server import _increment_spend_counter_cache
|
||||
|
||||
for entry in entries:
|
||||
counter_key = entry.get("counter_key")
|
||||
|
|
@ -448,9 +461,9 @@ async def _set_reserved_entries_adjustment(
|
|||
adjustment = target_adjustment - applied_adjustment
|
||||
if adjustment == 0:
|
||||
continue
|
||||
await spend_counter_cache.async_increment_cache(
|
||||
key=counter_key,
|
||||
value=adjustment,
|
||||
await _increment_spend_counter_cache(
|
||||
counter_key=counter_key,
|
||||
increment=adjustment,
|
||||
)
|
||||
entry["applied_adjustment"] = target_adjustment
|
||||
|
||||
|
|
@ -464,6 +477,37 @@ def _counter_to_reservation_entry(counter: _BudgetCounter) -> Dict[str, Any]:
|
|||
}
|
||||
|
||||
|
||||
def get_budget_window_start(window: Any) -> Optional[datetime]:
|
||||
window_dict = _coerce_window(window)
|
||||
budget_duration = window_dict.get("budget_duration")
|
||||
if budget_duration is None:
|
||||
return None
|
||||
try:
|
||||
duration_seconds = duration_in_seconds(str(budget_duration))
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
reset_at = _coerce_datetime(window_dict.get("reset_at"))
|
||||
if reset_at is None:
|
||||
reset_at = datetime.now(timezone.utc) + timedelta(seconds=duration_seconds)
|
||||
if reset_at.tzinfo is None:
|
||||
reset_at = reset_at.replace(tzinfo=timezone.utc)
|
||||
return reset_at - timedelta(seconds=duration_seconds)
|
||||
|
||||
|
||||
def _coerce_datetime(value: Any) -> Optional[datetime]:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, datetime):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
return datetime.fromisoformat(value.replace("Z", "+00:00"))
|
||||
except ValueError:
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def estimate_request_max_cost(
|
||||
request_body: dict,
|
||||
route: str,
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ import os
|
|||
import socket
|
||||
import subprocess
|
||||
import sys
|
||||
from datetime import datetime, timezone
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from pathlib import Path
|
||||
from unittest import mock
|
||||
from unittest.mock import AsyncMock, MagicMock, mock_open, patch
|
||||
|
|
@ -5101,6 +5101,78 @@ async def test_reseed_spend_from_db_skips_window_variant_keys():
|
|||
fake_prisma.db.litellm_teamtable.find_unique.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss():
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
|
||||
|
||||
counter_cache = DualCache()
|
||||
window_start = datetime.now(timezone.utc) - timedelta(hours=1)
|
||||
fake_prisma = MagicMock()
|
||||
fake_prisma.db.litellm_spendlogs.group_by = AsyncMock(
|
||||
return_value=[{"api_key": "key-window", "_sum": {"spend": 2.25}}]
|
||||
)
|
||||
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client
|
||||
ps.spend_counter_cache = counter_cache
|
||||
ps.prisma_client = fake_prisma
|
||||
try:
|
||||
await _init_and_increment_window_spend_counter(
|
||||
counter_key="spend:key:key-window:window:1h",
|
||||
entity_type="Key",
|
||||
entity_id="key-window",
|
||||
window_start=window_start,
|
||||
increment=0.5,
|
||||
)
|
||||
|
||||
fake_prisma.db.litellm_spendlogs.group_by.assert_awaited_once_with(
|
||||
by=["api_key"],
|
||||
where={"api_key": "key-window", "startTime": {"gte": window_start}},
|
||||
sum={"spend": True},
|
||||
)
|
||||
assert counter_cache.in_memory_cache.get_cache(
|
||||
key="spend:key:key-window:window:1h"
|
||||
) == pytest.approx(2.75)
|
||||
finally:
|
||||
ps.spend_counter_cache = orig_counter
|
||||
ps.prisma_client = orig_prisma
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_increment_spend_counter_invalidates_stale_cache_on_redis_failure():
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.proxy_server import _increment_spend_counter_cache
|
||||
|
||||
counter_cache = DualCache()
|
||||
counter_cache.in_memory_cache.set_cache(key="spend:team:redis-fail", value=4.0)
|
||||
fake_redis = AsyncMock()
|
||||
fake_redis.async_increment = AsyncMock(side_effect=RuntimeError("redis down"))
|
||||
fake_redis.async_delete_cache = AsyncMock()
|
||||
counter_cache.redis_cache = fake_redis
|
||||
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
orig_counter = ps.spend_counter_cache
|
||||
ps.spend_counter_cache = counter_cache
|
||||
try:
|
||||
with pytest.raises(RuntimeError):
|
||||
await _increment_spend_counter_cache(
|
||||
counter_key="spend:team:redis-fail",
|
||||
increment=0.5,
|
||||
)
|
||||
|
||||
assert (
|
||||
counter_cache.in_memory_cache.get_cache(key="spend:team:redis-fail") is None
|
||||
)
|
||||
fake_redis.async_delete_cache.assert_awaited_once_with(
|
||||
key="spend:team:redis-fail"
|
||||
)
|
||||
finally:
|
||||
ps.spend_counter_cache = orig_counter
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_current_spend_reseeds_from_db_when_counter_missing():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue