fix(proxy): align team-member budget keys and validate model budget types

This commit is contained in:
Ishaan Jaffer 2026-04-25 10:38:38 -07:00
parent ad6e88547d
commit 4a87c6babe
No known key found for this signature in database
6 changed files with 112 additions and 12 deletions

View file

@ -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(

View file

@ -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:

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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)
# ---------------------------------------------------------------------------