feat: add Postgres fallback to per-model budget read path

The per-model budget limiter previously relied entirely on in-memory
cache for spend lookups.  This meant spend data was lost on pod restart
and was not shared across replicas, so budgets were not enforced
reliably in multi-pod deployments.

This commit adds a Postgres fallback to the read path.  On a cache
miss, the limiter now queries the LiteLLM_DailyUserSpend /
LiteLLM_DailyEndUserSpend tables (already populated by the existing
db_spend_update_writer pipeline) and caches the result for 60 seconds.

The write path is deliberately unchanged.  RouterBudgetLimiting remains
the parent class and _increment_spend_for_key still seeds and
increments the in-memory cache on every request.  This preserves
single-instance and no-Postgres deployments where the cache is the
only source of truth.

Performance impact:
- No budgets configured: zero, none of the new code runs.
- In-memory only (no Postgres): one extra function call on cache miss
  that checks prisma_client is None and returns immediately.
- Postgres enabled: one find_many query per model/entity combo when
  the cache is cold or after 60s TTL expiry.  Cache hits are unchanged.

Includes 36 new tests covering the DB fallback, cache population,
fail-open behavior, two-key cache lookup, helper functions, and
write-path preservation.
This commit is contained in:
Jason Kidd 2026-04-08 08:02:39 -07:00
parent 2dac54b732
commit 1342c1c25c
No known key found for this signature in database
GPG key ID: 72BF942827539044
2 changed files with 909 additions and 31 deletions

View file

@ -1,10 +1,12 @@
import json
from typing import List, Optional
from datetime import datetime, timedelta, timezone
from typing import Callable, 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
@ -17,17 +19,53 @@ from litellm.types.utils import (
VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX = "virtual_key_spend"
END_USER_SPEND_CACHE_KEY_PREFIX = "end_user_model_spend"
# In-memory cache TTL for per-model spend values (seconds).
# Controls how long a spend value lives in the fast in-memory cache before
# the next read re-fetches from Postgres. 60 s matches the TTL used by the
# overall end-user budget check (UserAPIKeyCacheTTLEnum).
_MODEL_SPEND_CACHE_TTL_SECONDS = 60
def _budget_window_start_date(budget_duration: str) -> str:
"""Return the earliest ISO date string (YYYY-MM-DD) that falls within the
current budget window for the given *budget_duration* (e.g. ``"30d"``).
The daily-spend tables store one row per calendar day, so the boundary is
always rounded to a full day.
"""
seconds = duration_in_seconds(budget_duration)
days = max(seconds // 86400, 1) # at least 1 day
now = datetime.now(timezone.utc)
start = now - timedelta(days=days)
return start.strftime("%Y-%m-%d")
def _strip_provider_prefix(model: str) -> str:
"""``"openai/gpt-4"`` -> ``"gpt-4"``, ``"gpt-4"`` -> ``"gpt-4"``."""
if "/" in model:
return model.split("/", 1)[-1]
return model
class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
"""
Handles budgets for model + virtual key
Handles budgets for model + virtual key.
Example: key=sk-1234567890, model=gpt-4o, max_budget=100, time_period=1d
The read path checks the in-memory cache first; on a miss it falls back
to querying the ``LiteLLM_DailyUserSpend`` / ``LiteLLM_DailyEndUserSpend``
Postgres tables that the proxy already populates on every request. This
makes per-model budgets survive pod restarts and work across replicas.
The write path is inherited from ``RouterBudgetLimiting`` and seeds /
increments the in-memory cache so that enforcement within the same pod
is near-instant without waiting for the DB round-trip.
"""
def __init__(self, dual_cache: DualCache):
self.dual_cache = dual_cache
self.redis_increment_operation_queue = []
self.redis_increment_operation_queue: list = []
async def is_key_within_model_budget(
self,
@ -139,25 +177,44 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
return True
# spend lookup (read path) overrides parent with DB fallback
async def _get_end_user_spend_for_model(
self,
end_user_id: str,
model: str,
key_budget_config: BudgetConfig,
) -> Optional[float]:
# 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(
key=end_user_model_spend_cache_key,
)
"""Return current spend for an end-user on a specific model within the
budget window.
if _current_spend is None:
# 2. If 1, does not exist, check if passed as {custom_llm_provider}/model
end_user_model_spend_cache_key = f"{END_USER_SPEND_CACHE_KEY_PREFIX}:{end_user_id}:{self._get_model_without_custom_llm_provider(model)}:{key_budget_config.budget_duration}"
_current_spend = await self.dual_cache.async_get_cache(
key=end_user_model_spend_cache_key,
)
return _current_spend
Checks the in-memory cache first (trying both ``model`` as-is and
with the provider prefix stripped, matching the upstream two-key
lookup). On a complete cache miss, queries
``LiteLLM_DailyEndUserSpend`` in Postgres and populates the cache.
"""
duration = key_budget_config.budget_duration
# Try cache with model as-is, then with provider prefix stripped
primary_key = f"{END_USER_SPEND_CACHE_KEY_PREFIX}:{end_user_id}:{model}:{duration}"
cached = await self.dual_cache.async_get_cache(key=primary_key)
if cached is not None:
return float(cached)
stripped = self._get_model_without_custom_llm_provider(model)
if stripped != model:
alt_key = f"{END_USER_SPEND_CACHE_KEY_PREFIX}:{end_user_id}:{stripped}:{duration}"
cached = await self.dual_cache.async_get_cache(key=alt_key)
if cached is not None:
return float(cached)
# Cache miss — fall back to Postgres
return await self._query_db_and_cache(
cache_key=primary_key,
db_lookup=self._query_end_user_model_spend,
entity_id=end_user_id,
model=model,
budget_duration=duration or "",
)
async def _get_virtual_key_spend_for_model(
self,
@ -165,28 +222,132 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
model: str,
key_budget_config: BudgetConfig,
) -> Optional[float]:
"""
Get the current spend for a virtual key for a model
"""Return current spend for a virtual key on a specific model within
the budget window.
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
Checks the in-memory cache first (trying both ``model`` as-is and
with the provider prefix stripped, matching the upstream two-key
lookup). On a complete cache miss, queries
``LiteLLM_DailyUserSpend`` in Postgres and populates the cache.
"""
duration = key_budget_config.budget_duration
# Try cache with model as-is, then with provider prefix stripped
primary_key = f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:{user_api_key_hash}:{model}:{duration}"
cached = await self.dual_cache.async_get_cache(key=primary_key)
if cached is not None:
return float(cached)
# 1. model: directly look up `model`
virtual_key_model_spend_cache_key = f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:{user_api_key_hash}:{model}:{key_budget_config.budget_duration}"
_current_spend = await self.dual_cache.async_get_cache(
key=virtual_key_model_spend_cache_key,
stripped = self._get_model_without_custom_llm_provider(model)
if stripped != model:
alt_key = f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:{user_api_key_hash}:{stripped}:{duration}"
cached = await self.dual_cache.async_get_cache(key=alt_key)
if cached is not None:
return float(cached)
# Cache miss — fall back to Postgres
return await self._query_db_and_cache(
cache_key=primary_key,
db_lookup=self._query_virtual_key_model_spend,
entity_id=user_api_key_hash or "",
model=model,
budget_duration=duration or "",
)
if _current_spend is None:
# 2. If 1, does not exist, check if passed as {custom_llm_provider}/model
# if "/" in model, remove first part before "/" - eg. openai/o1-preview -> o1-preview
virtual_key_model_spend_cache_key = f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:{user_api_key_hash}:{self._get_model_without_custom_llm_provider(model)}:{key_budget_config.budget_duration}"
_current_spend = await self.dual_cache.async_get_cache(
key=virtual_key_model_spend_cache_key,
async def _query_db_and_cache(
self,
cache_key: str,
db_lookup: Callable,
entity_id: str,
model: str,
budget_duration: str,
) -> Optional[float]:
"""Query Postgres for spend and cache the result on success."""
try:
db_spend = await db_lookup(
entity_id=entity_id,
model=model,
budget_duration=budget_duration,
)
return _current_spend
except Exception:
verbose_proxy_logger.debug(
"model_max_budget_limiter: failed to query DB for spend, "
"cache_key=%s — returning None (will not block request)",
cache_key,
exc_info=True,
)
return None
if db_spend is not None:
await self.dual_cache.async_set_cache(
key=cache_key,
value=db_spend,
ttl=_MODEL_SPEND_CACHE_TTL_SECONDS,
)
return db_spend
@staticmethod
async def _query_end_user_model_spend(
entity_id: str,
model: str,
budget_duration: str,
) -> Optional[float]:
"""Sum spend from ``LiteLLM_DailyEndUserSpend`` for *entity_id* +
*model* within the budget window."""
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
return None
start_date = _budget_window_start_date(budget_duration)
model_without_provider = _strip_provider_prefix(model)
rows = await prisma_client.db.litellm_dailyenduserspend.find_many(
where={
"end_user_id": entity_id,
"date": {"gte": start_date},
"OR": [
{"model_group": model},
{"model_group": model_without_provider},
],
},
)
if not rows:
return 0.0
return sum(r.spend for r in rows)
@staticmethod
async def _query_virtual_key_model_spend(
entity_id: str,
model: str,
budget_duration: str,
) -> Optional[float]:
"""Sum spend from ``LiteLLM_DailyUserSpend`` for *entity_id* (an
api_key hash) + *model* within the budget window.
Note: ``LiteLLM_DailyUserSpend`` is keyed by ``api_key`` which stores
the hashed token — the same value passed as *entity_id*.
"""
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
return None
start_date = _budget_window_start_date(budget_duration)
model_without_provider = _strip_provider_prefix(model)
rows = await prisma_client.db.litellm_dailyuserspend.find_many(
where={
"api_key": entity_id,
"date": {"gte": start_date},
"OR": [
{"model_group": model},
{"model_group": model_without_provider},
],
},
)
if not rows:
return 0.0
return sum(r.spend for r in rows)
def _get_request_model_budget_config(
self, model: str, internal_model_max_budget: GenericBudgetConfigType

View file

@ -0,0 +1,717 @@
"""
Tests for the Postgres fallback added to the per-model budget read path.
The read methods ``_get_virtual_key_spend_for_model`` and
``_get_end_user_spend_for_model`` now check the in-memory cache first and,
on a miss, query the daily-spend Postgres tables before returning.
These tests verify:
- Cache hits short-circuit without DB queries
- Cache misses query Postgres and populate the cache
- Postgres returning no rows yields 0.0
- Postgres errors fail open (return None, don't block)
- No prisma_client returns None
- Helper functions (_budget_window_start_date, _strip_provider_prefix)
- End-to-end: budget enforcement with DB fallback
"""
import os
import sys
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
sys.path.insert(
0, os.path.abspath("../../..")
)
import pytest
import litellm
from litellm.caching.caching import DualCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.hooks.model_max_budget_limiter import (
END_USER_SPEND_CACHE_KEY_PREFIX,
VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX,
_MODEL_SPEND_CACHE_TTL_SECONDS,
_PROXY_VirtualKeyModelMaxBudgetLimiter,
_budget_window_start_date,
_strip_provider_prefix,
)
from litellm.types.utils import BudgetConfig
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
@pytest.fixture
def budget_limiter():
dual_cache = DualCache()
return _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache)
@pytest.fixture
def budget_config_1d():
return BudgetConfig(budget_limit=100.0, time_period="1d")
# ---------------------------------------------------------------------------
# Unit tests: _strip_provider_prefix
# ---------------------------------------------------------------------------
class TestStripProviderPrefix:
def test_should_strip_openai_prefix(self):
assert _strip_provider_prefix("openai/gpt-4") == "gpt-4"
def test_should_strip_azure_prefix(self):
assert _strip_provider_prefix("azure/gpt-4") == "gpt-4"
def test_should_leave_bare_model_unchanged(self):
assert _strip_provider_prefix("gpt-4") == "gpt-4"
def test_should_handle_multiple_slashes(self):
assert _strip_provider_prefix("azure/openai/gpt-4") == "openai/gpt-4"
def test_should_handle_empty_string(self):
assert _strip_provider_prefix("") == ""
# ---------------------------------------------------------------------------
# Unit tests: _budget_window_start_date
# ---------------------------------------------------------------------------
class TestBudgetWindowStartDate:
def test_should_return_date_string_for_1d(self):
result = _budget_window_start_date("1d")
# Should be yesterday or today depending on timing
parsed = datetime.strptime(result, "%Y-%m-%d")
assert parsed is not None
# Should be within 2 days of now
now = datetime.now(timezone.utc)
delta = now.replace(tzinfo=None) - parsed
assert 0 <= delta.days <= 2
def test_should_return_date_string_for_30d(self):
result = _budget_window_start_date("30d")
parsed = datetime.strptime(result, "%Y-%m-%d")
now = datetime.now(timezone.utc)
delta = now.replace(tzinfo=None) - parsed
assert 29 <= delta.days <= 31
def test_should_return_at_least_1_day_for_short_durations(self):
"""Even for durations shorter than 1 day (e.g. 1h), we should
look back at least 1 day since daily tables are per-day."""
result = _budget_window_start_date("1h")
parsed = datetime.strptime(result, "%Y-%m-%d")
now = datetime.now(timezone.utc)
delta = now.replace(tzinfo=None) - parsed
assert delta.days >= 1
# ---------------------------------------------------------------------------
# Unit tests: _query_db_and_cache
# ---------------------------------------------------------------------------
class TestQueryDbAndCache:
@pytest.mark.asyncio
async def test_should_query_db(self, budget_limiter):
"""Should call db_lookup and return the result."""
db_lookup = AsyncMock(return_value=25.0)
result = await budget_limiter._query_db_and_cache(
cache_key="miss-key",
db_lookup=db_lookup,
entity_id="ent-1",
model="gpt-4",
budget_duration="1d",
)
assert result == 25.0
db_lookup.assert_awaited_once_with(
entity_id="ent-1", model="gpt-4", budget_duration="1d"
)
@pytest.mark.asyncio
async def test_should_populate_cache_after_db_hit(self, budget_limiter):
"""After a DB query returns a value, it should be cached."""
db_lookup = AsyncMock(return_value=25.0)
await budget_limiter._query_db_and_cache(
cache_key="populate-key",
db_lookup=db_lookup,
entity_id="ent-1",
model="gpt-4",
budget_duration="1d",
)
# Now verify the cache was populated
cached = await budget_limiter.dual_cache.async_get_cache(
key="populate-key"
)
assert cached == 25.0
@pytest.mark.asyncio
async def test_should_not_populate_cache_on_db_none(self, budget_limiter):
"""When DB returns None (e.g. no prisma_client), cache should not be
populated with None."""
db_lookup = AsyncMock(return_value=None)
result = await budget_limiter._query_db_and_cache(
cache_key="none-key",
db_lookup=db_lookup,
entity_id="ent-1",
model="gpt-4",
budget_duration="1d",
)
assert result is None
cached = await budget_limiter.dual_cache.async_get_cache(
key="none-key"
)
assert cached is None
@pytest.mark.asyncio
async def test_should_return_none_on_db_exception(self, budget_limiter):
"""DB errors should fail open — return None, don't block."""
db_lookup = AsyncMock(side_effect=Exception("connection refused"))
result = await budget_limiter._query_db_and_cache(
cache_key="error-key",
db_lookup=db_lookup,
entity_id="ent-1",
model="gpt-4",
budget_duration="1d",
)
assert result is None
@pytest.mark.asyncio
async def test_should_return_zero_from_db(self, budget_limiter):
"""When DB returns 0.0 (no spend yet), cache it and return 0.0."""
db_lookup = AsyncMock(return_value=0.0)
result = await budget_limiter._query_db_and_cache(
cache_key="zero-key",
db_lookup=db_lookup,
entity_id="ent-1",
model="gpt-4",
budget_duration="1d",
)
assert result == 0.0
# 0.0 should be cached
cached = await budget_limiter.dual_cache.async_get_cache(
key="zero-key"
)
assert cached == 0.0
# ---------------------------------------------------------------------------
# Unit tests: Postgres query methods
# ---------------------------------------------------------------------------
class TestQueryEndUserModelSpend:
@pytest.mark.asyncio
async def test_should_return_none_when_no_prisma(self):
"""When prisma_client is None, return None."""
with patch(
"litellm.proxy.proxy_server.prisma_client",
None,
):
result = (
await _PROXY_VirtualKeyModelMaxBudgetLimiter._query_end_user_model_spend(
entity_id="user-1",
model="gpt-4",
budget_duration="1d",
)
)
assert result is None
@pytest.mark.asyncio
async def test_should_sum_spend_from_rows(self):
"""Should sum the spend field from returned rows."""
row1 = SimpleNamespace(spend=10.0)
row2 = SimpleNamespace(spend=5.5)
mock_prisma = MagicMock()
mock_prisma.db.litellm_dailyenduserspend.find_many = AsyncMock(
return_value=[row1, row2]
)
with patch(
"litellm.proxy.proxy_server.prisma_client",
mock_prisma,
):
result = (
await _PROXY_VirtualKeyModelMaxBudgetLimiter._query_end_user_model_spend(
entity_id="user-1",
model="gpt-4",
budget_duration="1d",
)
)
assert result == 15.5
@pytest.mark.asyncio
async def test_should_return_zero_when_no_rows(self):
"""No rows means zero spend."""
mock_prisma = MagicMock()
mock_prisma.db.litellm_dailyenduserspend.find_many = AsyncMock(
return_value=[]
)
with patch(
"litellm.proxy.proxy_server.prisma_client",
mock_prisma,
):
result = (
await _PROXY_VirtualKeyModelMaxBudgetLimiter._query_end_user_model_spend(
entity_id="user-1",
model="gpt-4",
budget_duration="1d",
)
)
assert result == 0.0
@pytest.mark.asyncio
async def test_should_pass_correct_where_clause(self):
"""Verify the Prisma query uses model_group and date filter."""
mock_prisma = MagicMock()
mock_find = AsyncMock(return_value=[])
mock_prisma.db.litellm_dailyenduserspend.find_many = mock_find
with patch(
"litellm.proxy.proxy_server.prisma_client",
mock_prisma,
):
await _PROXY_VirtualKeyModelMaxBudgetLimiter._query_end_user_model_spend(
entity_id="user-1",
model="openai/gpt-4",
budget_duration="7d",
)
mock_find.assert_awaited_once()
where = mock_find.call_args.kwargs["where"]
assert where["end_user_id"] == "user-1"
assert "gte" in where["date"]
# Should include both model variants in OR clause
or_clause = where["OR"]
model_groups = [c["model_group"] for c in or_clause]
assert "openai/gpt-4" in model_groups
assert "gpt-4" in model_groups
class TestQueryVirtualKeyModelSpend:
@pytest.mark.asyncio
async def test_should_return_none_when_no_prisma(self):
with patch(
"litellm.proxy.proxy_server.prisma_client",
None,
):
result = await _PROXY_VirtualKeyModelMaxBudgetLimiter._query_virtual_key_model_spend(
entity_id="hash-1",
model="gpt-4",
budget_duration="1d",
)
assert result is None
@pytest.mark.asyncio
async def test_should_sum_spend_from_rows(self):
row1 = SimpleNamespace(spend=20.0)
row2 = SimpleNamespace(spend=3.0)
mock_prisma = MagicMock()
mock_prisma.db.litellm_dailyuserspend.find_many = AsyncMock(
return_value=[row1, row2]
)
with patch(
"litellm.proxy.proxy_server.prisma_client",
mock_prisma,
):
result = await _PROXY_VirtualKeyModelMaxBudgetLimiter._query_virtual_key_model_spend(
entity_id="hash-1",
model="gpt-4",
budget_duration="1d",
)
assert result == 23.0
@pytest.mark.asyncio
async def test_should_use_api_key_in_where_clause(self):
mock_prisma = MagicMock()
mock_find = AsyncMock(return_value=[])
mock_prisma.db.litellm_dailyuserspend.find_many = mock_find
with patch(
"litellm.proxy.proxy_server.prisma_client",
mock_prisma,
):
await _PROXY_VirtualKeyModelMaxBudgetLimiter._query_virtual_key_model_spend(
entity_id="hash-abc",
model="gpt-4",
budget_duration="30d",
)
where = mock_find.call_args.kwargs["where"]
assert where["api_key"] == "hash-abc"
# ---------------------------------------------------------------------------
# Unit tests: two-key cache lookup (provider prefix stripping)
# ---------------------------------------------------------------------------
class TestTwoKeyCacheLookup:
@pytest.mark.asyncio
async def test_should_hit_cache_with_stripped_prefix(
self, budget_limiter, budget_config_1d
):
"""If cache was written with bare model name but read is requested
with provider prefix, the stripped-key fallback should find it."""
# Write path uses bare model name
bare_key = f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:key-hash:gpt-4:1d"
await budget_limiter.dual_cache.async_set_cache(
key=bare_key, value=30.0
)
with patch.object(
budget_limiter, "_query_virtual_key_model_spend"
) as mock_db:
result = await budget_limiter._get_virtual_key_spend_for_model(
user_api_key_hash="key-hash",
model="openai/gpt-4", # request uses provider prefix
key_budget_config=budget_config_1d,
)
assert result == 30.0
mock_db.assert_not_called()
@pytest.mark.asyncio
async def test_should_hit_cache_with_primary_key_first(
self, budget_limiter, budget_config_1d
):
"""If cache has BOTH keys, the primary (with prefix) should win."""
primary_key = f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:key-hash:openai/gpt-4:1d"
bare_key = f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:key-hash:gpt-4:1d"
await budget_limiter.dual_cache.async_set_cache(
key=primary_key, value=10.0
)
await budget_limiter.dual_cache.async_set_cache(
key=bare_key, value=99.0
)
result = await budget_limiter._get_virtual_key_spend_for_model(
user_api_key_hash="key-hash",
model="openai/gpt-4",
key_budget_config=budget_config_1d,
)
assert result == 10.0 # primary wins
@pytest.mark.asyncio
async def test_should_not_try_stripped_key_when_no_prefix(
self, budget_limiter, budget_config_1d
):
"""When model has no prefix, stripped == model, so only one lookup."""
with patch.object(
budget_limiter.dual_cache,
"async_get_cache",
new_callable=AsyncMock,
return_value=None,
) as mock_cache, patch.object(
budget_limiter,
"_query_virtual_key_model_spend",
new_callable=AsyncMock,
return_value=0.0,
):
await budget_limiter._get_virtual_key_spend_for_model(
user_api_key_hash="key-hash",
model="gpt-4", # no provider prefix
key_budget_config=budget_config_1d,
)
# Should only call cache once (no alt key when stripped == model)
assert mock_cache.call_count == 1
@pytest.mark.asyncio
async def test_end_user_should_hit_stripped_cache(
self, budget_limiter, budget_config_1d
):
"""End-user path should also support the two-key lookup."""
bare_key = f"{END_USER_SPEND_CACHE_KEY_PREFIX}:eu-1:gpt-4:1d"
await budget_limiter.dual_cache.async_set_cache(
key=bare_key, value=7.5
)
with patch.object(
budget_limiter, "_query_end_user_model_spend"
) as mock_db:
result = await budget_limiter._get_end_user_spend_for_model(
end_user_id="eu-1",
model="openai/gpt-4",
key_budget_config=budget_config_1d,
)
assert result == 7.5
mock_db.assert_not_called()
# ---------------------------------------------------------------------------
# Integration tests: read path with DB fallback
# ---------------------------------------------------------------------------
class TestVirtualKeySpendWithDbFallback:
@pytest.mark.asyncio
async def test_should_use_cache_when_available(self, budget_limiter, budget_config_1d):
"""Cache hit should return without querying DB."""
cache_key = f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:key-hash:gpt-4:1d"
await budget_limiter.dual_cache.async_set_cache(
key=cache_key, value=42.0
)
with patch.object(
budget_limiter, "_query_virtual_key_model_spend"
) as mock_db:
result = await budget_limiter._get_virtual_key_spend_for_model(
user_api_key_hash="key-hash",
model="gpt-4",
key_budget_config=budget_config_1d,
)
assert result == 42.0
mock_db.assert_not_called()
@pytest.mark.asyncio
async def test_should_fall_back_to_db_on_cache_miss(
self, budget_limiter, budget_config_1d
):
"""Cache miss should query DB and return the result."""
with patch.object(
budget_limiter,
"_query_virtual_key_model_spend",
new_callable=AsyncMock,
return_value=33.0,
):
result = await budget_limiter._get_virtual_key_spend_for_model(
user_api_key_hash="key-hash-2",
model="gpt-4",
key_budget_config=budget_config_1d,
)
assert result == 33.0
@pytest.mark.asyncio
async def test_should_enforce_budget_with_db_spend(self, budget_limiter):
"""End-to-end: DB reports spend over budget -> BudgetExceededError."""
user_api_key = UserAPIKeyAuth(
token="test-key-hash",
key_alias="test-alias",
model_max_budget={"gpt-4": {"budget_limit": 50.0, "time_period": "1d"}},
)
with patch.object(
budget_limiter,
"_query_virtual_key_model_spend",
new_callable=AsyncMock,
return_value=75.0,
):
with pytest.raises(litellm.BudgetExceededError):
await budget_limiter.is_key_within_model_budget(
user_api_key, "gpt-4"
)
@pytest.mark.asyncio
async def test_should_pass_when_db_spend_within_budget(self, budget_limiter):
"""End-to-end: DB reports spend within budget -> no error."""
user_api_key = UserAPIKeyAuth(
token="test-key-hash",
key_alias="test-alias",
model_max_budget={"gpt-4": {"budget_limit": 50.0, "time_period": "1d"}},
)
with patch.object(
budget_limiter,
"_query_virtual_key_model_spend",
new_callable=AsyncMock,
return_value=25.0,
):
result = await budget_limiter.is_key_within_model_budget(
user_api_key, "gpt-4"
)
assert result is True
class TestEndUserSpendWithDbFallback:
@pytest.mark.asyncio
async def test_should_use_cache_when_available(self, budget_limiter, budget_config_1d):
cache_key = f"{END_USER_SPEND_CACHE_KEY_PREFIX}:eu-1:gpt-4:1d"
await budget_limiter.dual_cache.async_set_cache(
key=cache_key, value=18.0
)
with patch.object(
budget_limiter, "_query_end_user_model_spend"
) as mock_db:
result = await budget_limiter._get_end_user_spend_for_model(
end_user_id="eu-1",
model="gpt-4",
key_budget_config=budget_config_1d,
)
assert result == 18.0
mock_db.assert_not_called()
@pytest.mark.asyncio
async def test_should_fall_back_to_db_on_cache_miss(
self, budget_limiter, budget_config_1d
):
with patch.object(
budget_limiter,
"_query_end_user_model_spend",
new_callable=AsyncMock,
return_value=12.0,
):
result = await budget_limiter._get_end_user_spend_for_model(
end_user_id="eu-2",
model="gpt-4",
key_budget_config=budget_config_1d,
)
assert result == 12.0
@pytest.mark.asyncio
async def test_should_enforce_budget_with_db_spend(self, budget_limiter):
"""End-to-end: DB reports spend over budget -> BudgetExceededError."""
with patch.object(
budget_limiter,
"_query_end_user_model_spend",
new_callable=AsyncMock,
return_value=200.0,
):
with pytest.raises(litellm.BudgetExceededError):
await budget_limiter.is_end_user_within_model_budget(
end_user_id="eu-3",
end_user_model_max_budget={
"gpt-4": {"budget_limit": 100.0, "time_period": "1d"}
},
model="gpt-4",
)
@pytest.mark.asyncio
async def test_should_pass_when_db_spend_within_budget(self, budget_limiter):
with patch.object(
budget_limiter,
"_query_end_user_model_spend",
new_callable=AsyncMock,
return_value=50.0,
):
result = await budget_limiter.is_end_user_within_model_budget(
end_user_id="eu-4",
end_user_model_max_budget={
"gpt-4": {"budget_limit": 100.0, "time_period": "1d"}
},
model="gpt-4",
)
assert result is True
# ---------------------------------------------------------------------------
# Tests: DB failure does not block requests
# ---------------------------------------------------------------------------
class TestFailOpen:
@pytest.mark.asyncio
async def test_should_not_block_when_db_fails_for_virtual_key(self, budget_limiter):
"""If DB query raises, budget check should pass (fail open)."""
user_api_key = UserAPIKeyAuth(
token="test-key-hash",
key_alias="test-alias",
model_max_budget={"gpt-4": {"budget_limit": 50.0, "time_period": "1d"}},
)
with patch.object(
budget_limiter,
"_query_virtual_key_model_spend",
new_callable=AsyncMock,
side_effect=Exception("DB connection refused"),
):
# Should NOT raise — fail open
result = await budget_limiter.is_key_within_model_budget(
user_api_key, "gpt-4"
)
assert result is True
@pytest.mark.asyncio
async def test_should_not_block_when_db_fails_for_end_user(self, budget_limiter):
with patch.object(
budget_limiter,
"_query_end_user_model_spend",
new_callable=AsyncMock,
side_effect=Exception("DB connection refused"),
):
result = await budget_limiter.is_end_user_within_model_budget(
end_user_id="eu-5",
end_user_model_max_budget={
"gpt-4": {"budget_limit": 100.0, "time_period": "1d"}
},
model="gpt-4",
)
assert result is True
@pytest.mark.asyncio
async def test_should_not_block_when_no_prisma_client(self, budget_limiter):
"""No DB connection at all should fail open."""
user_api_key = UserAPIKeyAuth(
token="test-key-hash",
key_alias="test-alias",
model_max_budget={"gpt-4": {"budget_limit": 50.0, "time_period": "1d"}},
)
with patch(
"litellm.proxy.proxy_server.prisma_client",
None,
):
result = await budget_limiter.is_key_within_model_budget(
user_api_key, "gpt-4"
)
assert result is True
# ---------------------------------------------------------------------------
# Tests: write path still uses parent's _increment_spend_for_key
# ---------------------------------------------------------------------------
class TestWritePathPreserved:
@pytest.mark.asyncio
async def test_should_still_call_increment_spend_for_key(self, budget_limiter):
"""async_log_success_event should still call the inherited
_increment_spend_for_key so single-instance/no-DB setups keep
working."""
virtual_key = "test-key-hash"
model = "gpt-4"
budget_duration = "1d"
kwargs = {
"standard_logging_object": {
"response_cost": 0.05,
"model": model,
"metadata": {"user_api_key_hash": virtual_key},
},
"litellm_params": {
"metadata": {
"user_api_key_model_max_budget": {
model: {
"budget_limit": 100.0,
"time_period": budget_duration,
},
}
},
},
}
with patch.object(
budget_limiter,
"_increment_spend_for_key",
new_callable=AsyncMock,
) as mock_increment:
await budget_limiter.async_log_success_event(
kwargs, response_obj=None, start_time=None, end_time=None
)
mock_increment.assert_awaited_once()
call_kwargs = mock_increment.call_args.kwargs
assert call_kwargs["response_cost"] == 0.05
assert f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:{virtual_key}:{model}:{budget_duration}" in call_kwargs["spend_key"]