mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(proxy): align team-member budget keys and validate model budget types
This commit is contained in:
parent
ad6e88547d
commit
4a87c6babe
6 changed files with 112 additions and 12 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue