fix(auth): use cached model-spend fallback without hot-path DB reads
Some checks failed
Unit Tests: Security / security (push) Has been cancelled
Unit Tests: Caching (Redis) / caching-redis (push) Has been cancelled
Unit Tests: Proxy DB Operations / assert-shard-coverage (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-utils (push) Has been cancelled
Unit Tests: Proxy DB Operations / auth-checks (push) Has been cancelled
Unit Tests: Proxy DB Operations / budgets (push) Has been cancelled
Unit Tests: Proxy DB Operations / custom-logging (push) Has been cancelled
Unit Tests: Proxy DB Operations / db-and-spend (push) Has been cancelled
Unit Tests: Proxy DB Operations / endpoints-and-responses (push) Has been cancelled
Unit Tests: Proxy DB Operations / guardrails-hooks (push) Has been cancelled
Unit Tests: Proxy DB Operations / jwt-and-keys (push) Has been cancelled
Unit Tests: Proxy DB Operations / key-generation (push) Has been cancelled
Unit Tests: Proxy DB Operations / logging-misc (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-runtime (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-server-core (push) Has been cancelled
Unit Tests: Proxy DB Operations / schema-migration (push) Has been cancelled

This commit is contained in:
Ishaan Jaffer 2026-04-25 14:07:16 -07:00
parent 41d2cc1012
commit 0610262ca8
No known key found for this signature in database
4 changed files with 36 additions and 27 deletions

View file

@ -1296,6 +1296,19 @@ async def get_team_membership(
return None
async def get_team_member_model_spend(
user_id: str,
team_id: str,
model: str,
user_api_key_cache: DualCache,
) -> float:
_key = f"team_member_model_spend:{user_id}:{team_id}:{model}"
cached_spend = await user_api_key_cache.async_get_cache(key=_key)
if cached_spend is not None:
return float(cached_spend)
return 0.0
def model_in_access_group(
model: str, team_models: Optional[List[str]], llm_router: Optional[Router]
) -> bool:
@ -3385,9 +3398,7 @@ async def _check_team_member_model_budget(
from litellm.proxy.proxy_server import (
_get_team_member_model_counter_key,
_reseed_spend_from_db,
get_current_spend,
spend_counter_cache,
)
for requested_model_str in models_to_check:
@ -3429,15 +3440,18 @@ async def _check_team_member_model_budget(
team_id=team_object.team_id,
model=counter_model,
)
fallback_spend = 0.0
if user_api_key_cache is not None:
fallback_spend = await get_team_member_model_spend(
user_id=valid_token.user_id,
team_id=team_object.team_id,
model=counter_model,
user_api_key_cache=user_api_key_cache,
)
model_spend = await get_current_spend(
counter_key=counter_key,
fallback_spend=-1.0,
fallback_spend=fallback_spend,
)
if model_spend < 0:
model_spend = await _reseed_spend_from_db(counter_key)
await spend_counter_cache.async_set_cache(
key=counter_key, value=model_spend, nx=True
)
if model_spend >= max_budget:
raise litellm.BudgetExceededError(

View file

@ -1926,6 +1926,10 @@ async def increment_spend_counters(
source_cache_key=None,
increment=response_cost,
)
await user_api_key_cache.async_increment_cache(
key=f"team_member_model_spend:{user_id}:{team_id}:{model}",
value=response_cost,
)
if user_id is not None:
await _init_and_increment_spend_counter(

View file

@ -2390,7 +2390,7 @@ async def test_team_member_model_budget_accepts_numeric_string_max_budget():
@pytest.mark.asyncio
async def test_team_member_model_budget_reseeds_on_counter_miss():
async def test_team_member_model_budget_uses_cached_fallback_spend():
team_object = LiteLLM_TeamTable(
team_id="test-team",
metadata={"team_member_model_max_budget": {"gpt-4o": {"max_budget": 10.0}}},
@ -2401,20 +2401,15 @@ async def test_team_member_model_budget_reseeds_on_counter_miss():
team_id="test-team",
)
user_api_key_cache = MagicMock()
user_api_key_cache.async_get_cache = AsyncMock(return_value=11.0)
async def mock_get_current_spend(counter_key, fallback_spend):
assert fallback_spend == 11.0
return fallback_spend
with (
patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend),
patch(
"litellm.proxy.proxy_server._reseed_spend_from_db",
new_callable=AsyncMock,
return_value=11.0,
) as mock_reseed,
patch(
"litellm.proxy.proxy_server.spend_counter_cache.async_set_cache",
new_callable=AsyncMock,
) as mock_set_cache,
):
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await _check_team_member_model_budget(
@ -2422,15 +2417,7 @@ async def test_team_member_model_budget_reseeds_on_counter_miss():
valid_token=valid_token,
model="gpt-4o",
llm_router=None,
user_api_key_cache=None,
user_api_key_cache=user_api_key_cache,
)
assert exc_info.value.current_cost == 11.0
assert exc_info.value.max_budget == 10.0
mock_reseed.assert_awaited_once_with(
"spend:team_member:team_id::test-team::user_id::test-user::model::gpt-4o"
)
mock_set_cache.assert_awaited_once_with(
key="spend:team_member:team_id::test-team::user_id::test-user::model::gpt-4o",
value=11.0,
nx=True,
)

View file

@ -4968,6 +4968,10 @@ async def test_increment_spend_counters_team_and_member():
key="spend:team_member:team_id::team-1::user_id::user-1::model::gpt-4o"
)
assert member_model_counter == 0.30
model_fallback_spend = key_cache.in_memory_cache.get_cache(
key="team_member_model_spend:user-1:team-1:gpt-4o"
)
assert model_fallback_spend == 0.30
finally:
ps.user_api_key_cache = original_key_cache
ps.spend_counter_cache = original_counter_cache