From 4a87c6babe54ba60bc64c05eeb0019321925ad67 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Sat, 25 Apr 2026 10:38:38 -0700 Subject: [PATCH] fix(proxy): align team-member budget keys and validate model budget types --- litellm/proxy/auth/auth_checks.py | 21 ++++++++-- litellm/proxy/auth/user_api_key_auth.py | 10 ++++- .../proxy/common_utils/reset_budget_job.py | 10 ++++- litellm/proxy/proxy_server.py | 3 +- .../proxy/auth/test_auth_checks.py | 38 +++++++++++++++-- .../common_utils/test_reset_budget_job.py | 42 +++++++++++++++++++ 6 files changed, 112 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 9341afdf240..8adf4055cf8 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -3325,10 +3325,16 @@ async def _check_team_member_budget( ) or 0.0 # Read from cross-pod counter (Redis-first) if available - from litellm.proxy.proxy_server import get_current_spend + from litellm.proxy.proxy_server import ( + _get_team_member_counter_key, + get_current_spend, + ) team_member_spend = await get_current_spend( - counter_key=f"spend:team_member:{valid_token.user_id}:{team_object.team_id}", + counter_key=_get_team_member_counter_key( + user_id=valid_token.user_id, + team_id=team_object.team_id, + ), fallback_spend=team_member_spend, ) @@ -3400,12 +3406,19 @@ async def _check_team_member_model_budget( if model_budget_config is None: continue - max_budget = ( + raw_max_budget = ( model_budget_config.get("max_budget") if isinstance(model_budget_config, dict) else None ) - if max_budget is None: + if isinstance(raw_max_budget, (int, float)): + max_budget = float(raw_max_budget) + elif isinstance(raw_max_budget, str): + try: + max_budget = float(raw_max_budget) + except ValueError: + continue + else: continue counter_key = _get_team_member_model_counter_key( diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index edb1efc9706..ce78cf62881 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1343,7 +1343,10 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 ) if team_member_budget is not None and team_member_budget > 0: # Read from cross-pod counter (Redis-first) if available - from litellm.proxy.proxy_server import get_current_spend + from litellm.proxy.proxy_server import ( + _get_team_member_counter_key, + get_current_spend, + ) team_member_spend = valid_token.team_member_spend if ( @@ -1351,7 +1354,10 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 and valid_token.team_id is not None ): team_member_spend = await get_current_spend( - counter_key=f"spend:team_member:{valid_token.user_id}:{valid_token.team_id}", + counter_key=_get_team_member_counter_key( + user_id=valid_token.user_id, + team_id=valid_token.team_id, + ), fallback_spend=team_member_spend, ) if team_member_spend > team_member_budget: diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index e486336cec0..d50ff48fee4 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -68,13 +68,19 @@ class ResetBudgetJob: # Reset Redis directly so a transient failure doesn't leave stale # counters that get_current_spend would read as authoritative. try: - from litellm.proxy.proxy_server import spend_counter_cache + from litellm.proxy.proxy_server import ( + _get_team_member_counter_key, + spend_counter_cache, + ) memberships = await self.prisma_client.db.litellm_teammembership.find_many( where={"budget_id": {"in": budget_ids}} ) for m in memberships: - counter_key = f"spend:team_member:{m.user_id}:{m.team_id}" + counter_key = _get_team_member_counter_key( + user_id=m.user_id, + team_id=m.team_id, + ) # Always reset in-memory spend_counter_cache.in_memory_cache.set_cache( key=counter_key, value=0.0 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 740e09ae050..1da6bfe3bec 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1949,7 +1949,8 @@ async def _reseed_spend_from_db(counter_key: str) -> float: spend:key:{token} -> LiteLLM_VerificationToken.spend spend:team:{team_id} -> LiteLLM_TeamTable.spend - spend:team_member:{uid}:{tid} -> LiteLLM_TeamMembership.spend + spend:team_member:team_id::{tid}::user_id::{uid} + -> LiteLLM_TeamMembership.spend spend:user:{user_id} -> LiteLLM_UserTable.spend spend:org:{org_id} -> LiteLLM_OrganizationTable.spend diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index f921fbe80fc..1064f4614bb 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -1968,7 +1968,7 @@ async def test_team_member_budget_check_reads_from_spend_counter(): proxy_logging_obj = ProxyLogging(user_api_key_cache=None) async def mock_get_current_spend(counter_key, fallback_spend): - if counter_key == "spend:team_member:test-user:test-team": + if counter_key == "spend:team_member:team_id::test-team::user_id::test-user": return 1.5 return fallback_spend @@ -2174,7 +2174,7 @@ async def test_team_member_budget_check_falls_back_to_team_default_budget_id(): ) async def mock_get_current_spend(counter_key, fallback_spend): - if counter_key == "spend:team_member:test-user:test-team": + if counter_key == "spend:team_member:team_id::test-team::user_id::test-user": return 70.0 return fallback_spend @@ -2271,7 +2271,7 @@ async def test_team_member_budget_check_per_member_override_wins_over_team_defau mocked_spend = 70.0 async def mock_get_current_spend(counter_key, fallback_spend): - if counter_key == "spend:team_member:test-user:test-team": + if counter_key == "spend:team_member:team_id::test-team::user_id::test-user": return mocked_spend return fallback_spend @@ -2355,3 +2355,35 @@ async def test_team_member_model_budget_uses_model_group_key_for_alias(): ) assert exc_info.value.current_cost == 11.0 assert exc_info.value.max_budget == 10.0 + + +@pytest.mark.asyncio +async def test_team_member_model_budget_accepts_numeric_string_max_budget(): + 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", + ) + + async def mock_get_current_spend(counter_key, fallback_spend): + if ( + counter_key + == "spend:team_member:team_id::test-team::user_id::test-user::model::gpt-4o" + ): + 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="gpt-4o", + llm_router=None, + ) + assert exc_info.value.current_cost == 11.0 + assert exc_info.value.max_budget == 10.0 diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index 379ccf4d9af..60a1ba8b5ff 100644 --- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -823,6 +823,48 @@ def test_reset_budget_for_team_members_preserves_total_spend(): assert "total_spend" not in call_kwargs["data"] +def test_reset_budget_for_team_members_resets_new_counter_key_format(): + from unittest.mock import patch + + expired_budget = type( + "LiteLLM_BudgetTableFull", + (), + {"budget_id": "budget-1"}, + ) + + membership = MagicMock() + membership.user_id = "user-1" + membership.team_id = "team-1" + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teammembership.find_many = AsyncMock( + return_value=[membership] + ) + mock_prisma_client.db.litellm_teammembership.update_many = AsyncMock( + return_value={"count": 1} + ) + + fake_module = types.ModuleType("litellm.proxy.proxy_server") + spend_counter_cache = MagicMock() + spend_counter_cache.in_memory_cache.set_cache = MagicMock() + spend_counter_cache.redis_cache = None + fake_module.spend_counter_cache = spend_counter_cache + fake_module._get_team_member_counter_key = ( + lambda user_id, team_id: f"spend:team_member:team_id::{team_id}::user_id::{user_id}" + ) + + with patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_module}): + job = ResetBudgetJob( + proxy_logging_obj=MagicMock(), prisma_client=mock_prisma_client + ) + asyncio.run(job.reset_budget_for_litellm_team_members([expired_budget])) + + spend_counter_cache.in_memory_cache.set_cache.assert_any_call( + key="spend:team_member:team_id::team-1::user_id::user-1", + value=0.0, + ) + + # --------------------------------------------------------------------------- # reset_budget_windows (per-key / per-team concurrent window resets) # ---------------------------------------------------------------------------