fix(model-budget): add DB fallback for cold cache in spend lookups

This commit is contained in:
Josh 2026-04-11 10:15:00 -04:00
parent 5c1f7d99bf
commit f8aa6591df

View file

@ -1,10 +1,12 @@
import json
from datetime import datetime, timedelta, timezone
from typing import List, Optional
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import Span
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.proxy._types import UserAPIKeyAuth
from litellm.router_strategy.budget_limiter import RouterBudgetLimiting
from litellm.types.llms.openai import AllMessageValues
@ -145,6 +147,16 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
model: str,
key_budget_config: BudgetConfig,
) -> Optional[float]:
"""
Get the current spend for an end-user for a model.
Lookup order:
1. Cache keyed by exact `model`
2. Cache keyed by model without custom_llm_provider prefix
3. DB fallback (cold-cache guard) — sums LiteLLM_SpendLogs within
the current budget window so budget limits are enforced even on
the first request after a cache flush or proxy restart.
"""
# 1. model: directly look up `model`
end_user_model_spend_cache_key = f"{END_USER_SPEND_CACHE_KEY_PREFIX}:{end_user_id}:{model}:{key_budget_config.budget_duration}"
_current_spend = await self.dual_cache.async_get_cache(
@ -157,6 +169,26 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
_current_spend = await self.dual_cache.async_get_cache(
key=end_user_model_spend_cache_key,
)
if _current_spend is None and key_budget_config.budget_duration:
# 3. Both cache tiers cold — query DB so we don't accidentally pass
# through a request that has already exceeded its budget.
_current_spend = await self._get_spend_from_db(
model=model,
budget_duration=key_budget_config.budget_duration,
budget_start_time_key=f"end_user_budget_start_time:{end_user_id}",
entity_filter={"end_user": end_user_id},
)
# Seed the canonical cache key so subsequent requests don't hit DB
# again until async_log_success_event writes the first increment.
if _current_spend is not None and key_budget_config.budget_duration:
ttl = duration_in_seconds(key_budget_config.budget_duration)
await self.dual_cache.async_set_cache(
key=end_user_model_spend_cache_key,
value=_current_spend,
ttl=ttl,
)
return _current_spend
async def _get_virtual_key_spend_for_model(
@ -166,11 +198,14 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
key_budget_config: BudgetConfig,
) -> Optional[float]:
"""
Get the current spend for a virtual key for a model
Get the current spend for a virtual key for a model.
Lookup model in this order:
1. model: directly look up `model`
2. If 1, does not exist, check if passed as {custom_llm_provider}/model
Lookup order:
1. Cache keyed by exact `model`
2. Cache keyed by model without custom_llm_provider prefix
3. DB fallback (cold-cache guard) — sums LiteLLM_SpendLogs within
the current budget window so budget limits are enforced even on
the first request after a cache flush or proxy restart.
"""
# 1. model: directly look up `model`
@ -186,6 +221,30 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
_current_spend = await self.dual_cache.async_get_cache(
key=virtual_key_model_spend_cache_key,
)
if (
_current_spend is None
and user_api_key_hash is not None
and key_budget_config.budget_duration
):
# 3. Both cache tiers cold — query DB so we don't accidentally pass
# through a request that has already exceeded its budget.
_current_spend = await self._get_spend_from_db(
model=model,
budget_duration=key_budget_config.budget_duration,
budget_start_time_key=f"virtual_key_budget_start_time:{user_api_key_hash}",
entity_filter={"api_key": user_api_key_hash},
)
# Seed the canonical cache key so subsequent requests don't hit DB
# again until async_log_success_event writes the first increment.
if _current_spend is not None and key_budget_config.budget_duration:
ttl = duration_in_seconds(key_budget_config.budget_duration)
await self.dual_cache.async_set_cache(
key=virtual_key_model_spend_cache_key,
value=_current_spend,
ttl=ttl,
)
return _current_spend
def _get_request_model_budget_config(
@ -208,6 +267,62 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
return model.split("/")[-1]
return model
async def _get_spend_from_db(
self,
model: str,
budget_duration: str,
budget_start_time_key: str,
entity_filter: dict,
) -> Optional[float]:
"""
DB fallback when cache is cold.
Queries LiteLLM_SpendLogs and sums spend for the given entity + model
within the current budget window. This prevents requests from slipping
through unenforced on the very first call after a cache flush or proxy
restart.
Returns None if prisma_client is unavailable (budget check is skipped
gracefully — same behaviour as before this fallback was added).
"""
from litellm.proxy.proxy_server import prisma_client # circular-import exception
if prisma_client is None:
return None
if not budget_duration:
return None
ttl_seconds = duration_in_seconds(budget_duration)
budget_start = await self.dual_cache.async_get_cache(budget_start_time_key)
if budget_start is not None:
window_start = datetime.fromtimestamp(float(budget_start), tz=timezone.utc)
else:
window_start = datetime.now(timezone.utc) - timedelta(seconds=ttl_seconds)
model_without_provider = self._get_model_without_custom_llm_provider(model)
# Match model_group (preferred — what spend logger writes) with a
# fallback to the raw model field for rows written before model_group
# was populated. Also strip provider prefix so a budget keyed on
# "gpt-4o" matches logs written as "openai/gpt-4o".
model_candidates = list(dict.fromkeys([model, model_without_provider]))
model_or_clauses = [{"model_group": m} for m in model_candidates] + [
{"model": m} for m in model_candidates
]
where: dict = {
**entity_filter,
"OR": model_or_clauses,
"startTime": {"gte": window_start},
}
rows = await prisma_client.db.litellm_spendlogs.group_by(
by=["model_group"],
where=where,
sum={"spend": True},
)
return sum((row.get("_sum") or {}).get("spend") or 0.0 for row in rows)
async def async_filter_deployments(
self,
model: str,