tighten budget counter cache recovery

This commit is contained in:
user 2026-04-29 20:51:07 -07:00
parent 5a619cf879
commit 926de696a1
4 changed files with 352 additions and 23 deletions

View file

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

View file

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

View file

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

View file

@ -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():
"""