diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 5a258d74181..9341afdf240 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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}" ), ) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index f37d3e64a9b..466c9b30458 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -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", "", "user_id", "", "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() diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index cef45139412..740e09ae050 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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( diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 8612d243c41..f921fbe80fc 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -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 diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index efd1abbb383..9bd39db3f6f 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -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