diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py index bf60a087c65..b4dc4a4d0a2 100644 --- a/litellm/proxy/db/spend_counter_reseed.py +++ b/litellm/proxy/db/spend_counter_reseed.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 811bdfbc303..f6655ba9285 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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], diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index cccf1745e80..130a3e20275 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -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, diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 1f4f82a64ef..8f3297780e6 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -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(): """