mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(proxy): simplify team-member model budget keying and alias checks
This commit is contained in:
parent
f38f5d48b0
commit
ad6e88547d
5 changed files with 160 additions and 44 deletions
|
|
@ -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}"
|
||||
),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue