fix(proxy): simplify team-member model budget keying and alias checks

This commit is contained in:
Ishaan Jaffer 2026-04-25 09:37:34 -07:00
parent f38f5d48b0
commit ad6e88547d
No known key found for this signature in database
5 changed files with 160 additions and 44 deletions

View file

@ -658,6 +658,7 @@ async def common_checks( # noqa: PLR0915
team_object=team_object,
valid_token=valid_token,
model=_model,
llm_router=llm_router,
)
# 5. If end_user ('user' passed to /chat/completions, /embeddings endpoint) is in budget
@ -3343,6 +3344,7 @@ async def _check_team_member_model_budget(
team_object: Optional[LiteLLM_TeamTable],
valid_token: Optional[UserAPIKeyAuth],
model: Optional[Union[str, List[str]]],
llm_router: Optional[Router],
):
"""
Check if a team member has exceeded their per-model budget for this team.
@ -3350,7 +3352,7 @@ async def _check_team_member_model_budget(
The per-model limits are set at the team level via team_member_model_max_budget
(stored in team metadata) and apply to each member independently.
Spend is tracked via the spend counter cache keyed by
spend:team_member:{user_id}:{team_id}:model:{model}.
spend:team_member:team_id::{team_id}::user_id::{user_id}::model::{model}.
"""
if (
team_object is None
@ -3374,13 +3376,27 @@ async def _check_team_member_model_budget(
return
from litellm.proxy.proxy_server import (
_get_team_member_model_counter_key,
_reseed_spend_from_db,
get_current_spend,
spend_counter_cache,
)
for model_str in models_to_check:
model_budget_config = team_member_model_max_budget.get(model_str)
for requested_model_str in models_to_check:
counter_model = requested_model_str
if requested_model_str in litellm.model_alias_map:
counter_model = litellm.model_alias_map[requested_model_str]
elif (
llm_router is not None
and isinstance(llm_router.model_group_alias, dict)
and requested_model_str in llm_router.model_group_alias
):
resolved_model = llm_router._get_model_from_alias(requested_model_str)
if resolved_model:
counter_model = resolved_model
model_budget_config = team_member_model_max_budget.get(counter_model)
if model_budget_config is None:
model_budget_config = team_member_model_max_budget.get(requested_model_str)
if model_budget_config is None:
continue
@ -3392,7 +3408,11 @@ async def _check_team_member_model_budget(
if max_budget is None:
continue
counter_key = f"spend:team_member:{valid_token.user_id}:{team_object.team_id}:model:{model_str}"
counter_key = _get_team_member_model_counter_key(
user_id=valid_token.user_id,
team_id=team_object.team_id,
model=counter_model,
)
# -1.0 is the sentinel for "cache cold". Negative spend is impossible,
# so this distinguishes "not cached yet" from a real zero balance.
model_spend = await get_current_spend(
@ -3401,15 +3421,8 @@ async def _check_team_member_model_budget(
)
if model_spend < 0:
# Cache cold (pod restart / Redis flush) — fetch authoritative value from DB.
# Use async_set_cache (idempotent overwrite) not async_increment_cache:
# two concurrent requests that both see a miss would each call INCRBYFLOAT,
# doubling the seed value and producing false budget-exceeded 429s.
# Seed unconditionally (including zero spend) so new users don't hit DB
# on every auth request until their first request's cost callback fires.
# Do not write cache from auth path to avoid clobbering in-flight increments.
model_spend = await _reseed_spend_from_db(counter_key)
await spend_counter_cache.async_set_cache(
key=counter_key, value=model_spend
)
if model_spend >= max_budget:
raise litellm.BudgetExceededError(
@ -3418,7 +3431,7 @@ async def _check_team_member_model_budget(
message=(
f"ExceededBudget: Team member model budget exceeded. "
f"User={valid_token.user_id}, Team={team_object.team_id}, "
f"Model={model_str}. Spend=${model_spend:.6f}, Budget=${max_budget:.6f}"
f"Model={counter_model}. Spend=${model_spend:.6f}, Budget=${max_budget:.6f}"
),
)

View file

@ -1824,19 +1824,23 @@ class DBSpendUpdateWriter:
from litellm.proxy.utils import _raise_failed_update_spend_exception
for key, cost in team_member_model_list_transactions.items():
parts = key.split("::")
# Expected: ["team_id", "<tid>", "user_id", "<uid>", "model", "<model>"]
if (
len(parts) != 6
or parts[0] != "team_id"
or parts[2] != "user_id"
or parts[4] != "model"
not key.startswith("team_id::")
or "::user_id::" not in key
or "::model::" not in key
):
verbose_proxy_logger.debug(
"Skipping malformed team_member_model key: %s", key
)
continue
team_id, user_id, model_name = parts[1], parts[3], parts[5]
membership_part, model_name = key.split("::model::", 1)
team_and_user = membership_part[len("team_id::") :]
team_id, user_id = team_and_user.split("::user_id::", 1)
if not team_id or not user_id or not model_name:
verbose_proxy_logger.debug(
"Skipping malformed team_member_model key: %s", key
)
continue
for attempt in range(n_retry_times + 1):
start_time = time.time()

View file

@ -1801,6 +1801,29 @@ async def get_current_spend(counter_key: str, fallback_spend: float) -> float:
return fallback_spend
def _get_team_member_counter_key(user_id: str, team_id: str) -> str:
return f"spend:team_member:team_id::{team_id}::user_id::{user_id}"
def _get_team_member_model_counter_key(user_id: str, team_id: str, model: str) -> str:
return f"spend:team_member:team_id::{team_id}::user_id::{user_id}::model::{model}"
def _parse_team_member_counter_suffix(
suffix: str,
) -> Optional[Tuple[str, str, Optional[str]]]:
if "::model::" in suffix:
membership_part, model_name = suffix.split("::model::", 1)
else:
membership_part, model_name = suffix, None
parts = membership_part.split("::")
if len(parts) != 4 or parts[0] != "team_id" or parts[2] != "user_id":
return None
team_id = parts[1]
user_id = parts[3]
return user_id, team_id, model_name
async def increment_spend_counters(
token: Optional[str],
team_id: Optional[str],
@ -1889,13 +1912,17 @@ async def increment_spend_counters(
if user_id is not None and team_id is not None:
await _init_and_increment_spend_counter(
counter_key=f"spend:team_member:{user_id}:{team_id}",
counter_key=_get_team_member_counter_key(user_id=user_id, team_id=team_id),
source_cache_key=f"team_membership:{user_id}:{team_id}",
increment=response_cost,
)
if model is not None:
await _init_and_increment_spend_counter(
counter_key=f"spend:team_member:{user_id}:{team_id}:model:{model}",
counter_key=_get_team_member_model_counter_key(
user_id=user_id,
team_id=team_id,
model=model,
),
source_cache_key=None,
increment=response_cost,
)
@ -1945,13 +1972,11 @@ async def _reseed_spend_from_db(counter_key: str) -> float:
)
elif counter_key.startswith("spend:team_member:"):
suffix = counter_key[len("spend:team_member:") :]
if ":model:" in suffix:
# format: {user_id}:{team_id}:model:{model_name}
# Read from dedicated atomic table — no read-modify-write race
membership_part, model_name = suffix.rsplit(":model:", 1)
if ":" not in membership_part:
return 0.0
member_user_id, member_team_id = membership_part.rsplit(":", 1)
parsed_team_member_key = _parse_team_member_counter_suffix(suffix=suffix)
if parsed_team_member_key is None:
return 0.0
member_user_id, member_team_id, model_name = parsed_team_member_key
if model_name is not None:
model_row = (
await prisma_client.db.litellm_teammembermodelspend.find_unique(
where={
@ -1964,19 +1989,14 @@ async def _reseed_spend_from_db(counter_key: str) -> float:
)
)
return float(model_row.spend if model_row is not None else 0.0)
else:
membership_part = suffix
if ":" not in membership_part:
return 0.0
member_user_id, member_team_id = membership_part.rsplit(":", 1)
row = await prisma_client.db.litellm_teammembership.find_unique(
where={
"user_id_team_id": {
"user_id": member_user_id,
"team_id": member_team_id,
}
row = await prisma_client.db.litellm_teammembership.find_unique(
where={
"user_id_team_id": {
"user_id": member_user_id,
"team_id": member_team_id,
}
)
}
)
elif counter_key.startswith("spend:team:"):
team_id = counter_key[len("spend:team:") :]
row = await prisma_client.db.litellm_teamtable.find_unique(

View file

@ -30,6 +30,7 @@ from litellm.proxy.auth.auth_checks import (
ExperimentalUIJWTToken,
_can_object_call_vector_stores,
_check_team_member_budget,
_check_team_member_model_budget,
_get_fuzzy_user_object,
_get_team_db_check,
_log_budget_lookup_failure,
@ -1692,6 +1693,7 @@ async def test_virtual_key_max_budget_alert_check_global_fallback():
)
import litellm
original = litellm.default_key_max_budget_alert_emails
try:
litellm.default_key_max_budget_alert_emails = global_config
@ -1731,6 +1733,7 @@ async def test_virtual_key_max_budget_alert_check_per_key_merges_with_global():
)
import litellm
original = litellm.default_key_max_budget_alert_emails
try:
litellm.default_key_max_budget_alert_emails = global_config
@ -2315,3 +2318,40 @@ async def test_team_member_budget_check_per_member_override_wins_over_team_defau
)
assert exc_info.value.current_cost == 250.0
assert exc_info.value.max_budget == 200.0
@pytest.mark.asyncio
async def test_team_member_model_budget_uses_model_group_key_for_alias():
team_object = LiteLLM_TeamTable(
team_id="test-team",
metadata={"team_member_model_max_budget": {"gpt-4o": {"max_budget": 10.0}}},
)
valid_token = UserAPIKeyAuth(
token="test-token",
user_id="test-user",
team_id="test-team",
)
llm_router = MagicMock()
llm_router.model_group_alias = {"alias-gpt4": "gpt-4o"}
llm_router._get_model_from_alias.return_value = "gpt-4o"
expected_counter_key = (
"spend:team_member:team_id::test-team::user_id::test-user::model::gpt-4o"
)
async def mock_get_current_spend(counter_key, fallback_spend):
if counter_key == expected_counter_key:
return 11.0
return fallback_spend
with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend):
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await _check_team_member_model_budget(
team_object=team_object,
valid_token=valid_token,
model="alias-gpt4",
llm_router=llm_router,
)
assert exc_info.value.current_cost == 11.0
assert exc_info.value.max_budget == 10.0

View file

@ -4921,7 +4921,7 @@ async def test_increment_spend_counters_initializes_and_increments():
@pytest.mark.asyncio
async def test_increment_spend_counters_team_and_member():
"""Counter should track team and team member spend separately."""
"""Counter should track team, team member, and team-member-model spend."""
from litellm.caching.dual_cache import DualCache
from litellm.proxy._types import LiteLLM_TeamTable
@ -4953,15 +4953,21 @@ async def test_increment_spend_counters_team_and_member():
team_id="team-1",
user_id="user-1",
response_cost=0.30,
model="gpt-4o",
)
team_counter = counter_cache.in_memory_cache.get_cache(key="spend:team:team-1")
assert team_counter == 2.30
member_counter = counter_cache.in_memory_cache.get_cache(
key="spend:team_member:user-1:team-1"
key="spend:team_member:team_id::team-1::user_id::user-1"
)
assert member_counter == 1.30
member_model_counter = counter_cache.in_memory_cache.get_cache(
key="spend:team_member:team_id::team-1::user_id::user-1::model::gpt-4o"
)
assert member_model_counter == 0.30
finally:
ps.user_api_key_cache = original_key_cache
ps.spend_counter_cache = original_counter_cache
@ -5064,6 +5070,39 @@ async def test_reseed_spend_from_db_user_and_org_prefixes():
ps.prisma_client = orig_prisma
@pytest.mark.asyncio
async def test_reseed_spend_from_db_team_member_model_counter_parses_tagged_key():
import litellm.proxy.proxy_server as ps
from litellm.proxy.proxy_server import _reseed_spend_from_db
model_row = MagicMock()
model_row.spend = 12.34
fake_prisma = MagicMock()
fake_prisma.db.litellm_teammembermodelspend.find_unique = AsyncMock(
return_value=model_row
)
orig_prisma = ps.prisma_client
ps.prisma_client = fake_prisma
try:
spend = await _reseed_spend_from_db(
"spend:team_member:team_id::org:team-1::user_id::google:12345::model::gpt-4o"
)
assert spend == 12.34
fake_prisma.db.litellm_teammembermodelspend.find_unique.assert_awaited_once_with(
where={
"user_id_team_id_model": {
"user_id": "google:12345",
"team_id": "org:team-1",
"model": "gpt-4o",
}
}
)
finally:
ps.prisma_client = orig_prisma
@pytest.mark.asyncio
async def test_reseed_spend_from_db_skips_window_variant_keys():
"""Window counters (spend:*:window:{duration}) share prefixes with