Merge pull request #41347 from BerriAI/litellm_team_member_budget_link_default

fix(team): apply team_member_budget updates to members still on the team default
This commit is contained in:
ryan-crabbe-berri 2026-09-18 17:06:25 -07:00 committed by GitHub
commit 4bb1ae115b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 240 additions and 86 deletions

View file

@ -293,9 +293,9 @@ async def _clone_team_default_budget_for_member(
member budget. Returns the new budget_id, or None if the default budget
no longer exists in the DB.
Used when adding a new team member without an explicit per-member budget,
so the member starts with the team default's values but gets their own
private budget row (which can be edited independently).
Used when adding a new team member with a per-member ``budget_duration``
but no other per-member limit, so the member keeps the team default's
values in their own private budget row while the reset window differs.
``budget_duration_override`` replaces the default's reset window for this
member while keeping the default's other limits, so an admin can set a
@ -346,14 +346,21 @@ async def _resolve_member_budget_id(
"""
Resolve the budget a new team member should be linked to.
Explicit per-member limits create a fresh budget. Otherwise the team's
default member budget is cloned (with ``budget_duration`` overriding its
reset window while keeping its other limits). A lone ``budget_duration``
with no team default creates a window-only budget. With nothing set the
member gets no budget, though ``add_new_member`` still writes its membership row.
Explicit per-member limits create a fresh budget. Otherwise the member is
linked to the team's shared default member budget, so later ``/team/update``
changes reach them; ``/team/member_update`` clones that row on first write.
A lone ``budget_duration`` clones the default with the reset window
overridden, or creates a window-only budget when there is no team default.
With nothing set the member gets no budget, though ``add_new_member`` still writes its membership row.
"""
has_explicit_limit: Final = max_budget_in_team is not None or allowed_models is not None
if not has_explicit_limit and default_team_budget_id is not None and budget_duration is None:
default_budget: Final = await _budget_table(prisma_client, tx).find_unique(
where={"budget_id": default_team_budget_id}
)
return default_team_budget_id if default_budget is not None else None
if not has_explicit_limit and default_team_budget_id is not None:
return await _clone_team_default_budget_for_member(
prisma_client=prisma_client,

View file

@ -1,13 +1,17 @@
import json
from collections.abc import Mapping
from datetime import datetime, timezone
from litellm._uuid import uuid
from unittest.mock import AsyncMock, MagicMock
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import litellm
from litellm._uuid import uuid
from litellm.proxy._types import (
LiteLLM_BudgetTable,
LiteLLM_TeamMembership,
LiteLLM_TeamTable,
LiteLLM_UserTable,
Member,
UserAPIKeyAuth,
@ -164,21 +168,12 @@ async def test_management_otel_span_redacts_nested_submission_env_var_secrets(
@pytest.mark.asyncio
async def test_add_new_member_clones_default_team_budget_id():
"""
Test that add_new_member CLONES the team's default member budget when
max_budget_in_team is None and a default_team_budget_id is provided.
Cloning (rather than sharing the same budget row) is what lets admins later
edit one member's budget without mutating every other member's budget.
"""
async def test_add_new_member_links_default_team_budget_id():
from litellm.proxy._types import LitellmUserRoles
# Setup test data
test_user_id = "test_user_123"
test_team_id = "test_team_456"
test_default_budget_id = "default_budget_789"
test_cloned_budget_id = "cloned_budget_xyz"
test_admin_name = "test_admin"
new_member = Member(user_id=test_user_id, role="user")
@ -202,36 +197,19 @@ async def test_add_new_member_clones_default_team_budget_id():
return_value=mock_user_response
)
# Mock the default budget row fetched for cloning.
mock_default_budget_row = MagicMock()
mock_default_budget_row.model_dump.return_value = {
"budget_id": test_default_budget_id,
"max_budget": 100.0,
"soft_budget": None,
"max_parallel_requests": None,
"tpm_limit": 1000,
"rpm_limit": None,
"model_max_budget": None,
"budget_duration": "1d",
"allowed_models": [],
}
mock_default_budget_row.budget_id = test_default_budget_id
mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(
return_value=mock_default_budget_row
)
# Mock the cloned budget row that .create() returns.
mock_cloned_budget_row = MagicMock()
mock_cloned_budget_row.budget_id = test_cloned_budget_id
mock_prisma_client.db.litellm_budgettable.create = AsyncMock(
return_value=mock_cloned_budget_row
)
mock_prisma_client.db.litellm_budgettable.create = AsyncMock()
# Mock the team membership creation
mock_team_membership_response = MagicMock()
mock_team_membership_response.model_dump.return_value = {
"team_id": test_team_id,
"user_id": test_user_id,
"budget_id": test_cloned_budget_id,
"budget_id": test_default_budget_id,
"litellm_budget_table": None,
}
mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock(
@ -251,33 +229,71 @@ async def test_add_new_member_clones_default_team_budget_id():
assert result_user is not None
assert result_user.user_id == test_user_id
# Membership should be linked to the new cloned budget, not the shared default.
assert result_team_membership is not None
assert result_team_membership.budget_id == test_cloned_budget_id
assert result_team_membership.budget_id != test_default_budget_id
assert result_team_membership.budget_id == test_default_budget_id
mock_prisma_client.db.litellm_usertable.upsert.assert_called_once()
mock_prisma_client.db.litellm_teammembership.upsert.assert_called_once()
# The clone must have happened: find_unique on the default, create for the clone.
mock_prisma_client.db.litellm_budgettable.find_unique.assert_called_once_with(
where={"budget_id": test_default_budget_id}
)
mock_prisma_client.db.litellm_budgettable.create.assert_called_once()
cloned_create_data = (
mock_prisma_client.db.litellm_budgettable.create.call_args.kwargs["data"]
)
# Cloned values from the default budget row
assert cloned_create_data["max_budget"] == 100.0
assert cloned_create_data["tpm_limit"] == 1000
assert cloned_create_data["budget_duration"] == "1d"
assert cloned_create_data["created_by"] == user_api_key_dict.user_id
mock_prisma_client.db.litellm_budgettable.create.assert_not_called()
team_membership_call_args = (
mock_prisma_client.db.litellm_teammembership.upsert.call_args
)
create_data = team_membership_call_args.kwargs["data"]["create"]
assert create_data["budget_id"] == test_cloned_budget_id
assert create_data["budget_id"] == test_default_budget_id
@pytest.mark.asyncio
async def test_add_new_member_no_budget_when_default_budget_row_is_missing():
from litellm.proxy._types import LitellmUserRoles
new_member = Member(user_id="missing-default-user", role="user")
user_api_key_dict = UserAPIKeyAuth(
user_id="admin_user", user_role=LitellmUserRoles.PROXY_ADMIN
)
mock_prisma_client = AsyncMock()
mock_user_response = MagicMock()
mock_user_response.model_dump.return_value = {
"user_id": "missing-default-user",
"user_email": None,
"teams": ["team-md"],
"user_role": "internal_user",
}
mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_response)
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
return_value=mock_user_response
)
mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_budgettable.create = AsyncMock()
mock_membership = MagicMock()
mock_membership.model_dump.return_value = {
"team_id": "team-md",
"user_id": "missing-default-user",
"budget_id": None,
"litellm_budget_table": None,
}
mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock(return_value=mock_membership)
_, result_team_membership = await add_new_member(
new_member=new_member,
max_budget_in_team=None,
prisma_client=mock_prisma_client,
team_id="team-md",
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name="test_admin",
default_team_budget_id="deleted-default",
)
assert result_team_membership is not None
assert result_team_membership.budget_id is None
mock_prisma_client.db.litellm_budgettable.create.assert_not_called()
upsert_kwargs = mock_prisma_client.db.litellm_teammembership.upsert.call_args.kwargs
assert upsert_kwargs["data"]["create"] == {"user_id": "missing-default-user", "team_id": "team-md"}
@pytest.mark.asyncio
@ -661,18 +677,12 @@ async def test_add_new_member_persists_budget_duration_without_max_budget():
@pytest.mark.asyncio
async def test_add_new_member_with_user_email_clones_default_budget():
"""
Test add_new_member with user_email instead of user_id and a team default
budget. The default budget should be CLONED into a new private row for
this user, not shared with other members of the team.
"""
async def test_add_new_member_with_user_email_links_default_budget():
from litellm.proxy._types import LitellmUserRoles
test_user_email = "test@example.com"
test_team_id = "test_team_456"
test_default_budget_id = "default_budget_789"
test_cloned_budget_id = "cloned_budget_for_email_user"
test_admin_name = "test_admin"
new_member = Member(user_email=test_user_email, role="user")
@ -694,35 +704,18 @@ async def test_add_new_member_with_user_email_clones_default_budget():
}
mock_prisma_client.insert_data = AsyncMock(return_value=mock_user_response)
# Default budget that will be cloned
mock_default_budget_row = MagicMock()
mock_default_budget_row.model_dump.return_value = {
"budget_id": test_default_budget_id,
"max_budget": 25.0,
"soft_budget": None,
"max_parallel_requests": None,
"tpm_limit": None,
"rpm_limit": None,
"model_max_budget": None,
"budget_duration": None,
"allowed_models": [],
}
mock_default_budget_row.budget_id = test_default_budget_id
mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(
return_value=mock_default_budget_row
)
# Cloned budget result
mock_cloned_budget_row = MagicMock()
mock_cloned_budget_row.budget_id = test_cloned_budget_id
mock_prisma_client.db.litellm_budgettable.create = AsyncMock(
return_value=mock_cloned_budget_row
)
mock_prisma_client.db.litellm_budgettable.create = AsyncMock()
mock_team_membership_response = MagicMock()
mock_team_membership_response.model_dump.return_value = {
"team_id": test_team_id,
"user_id": "generated_user_id",
"budget_id": test_cloned_budget_id,
"budget_id": test_default_budget_id,
"litellm_budget_table": None,
}
mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock(
@ -742,9 +735,8 @@ async def test_add_new_member_with_user_email_clones_default_budget():
assert result_user is not None
assert result_user.user_email == test_user_email
# Membership should point at the cloned (private) budget, not the shared default.
assert result_team_membership is not None
assert result_team_membership.budget_id == test_cloned_budget_id
assert result_team_membership.budget_id == test_default_budget_id
mock_prisma_client.get_data.assert_called_once_with(
key_val={"user_email": test_user_email},
@ -758,11 +750,166 @@ async def test_add_new_member_with_user_email_clones_default_budget():
assert insert_data["user_email"] == test_user_email
assert insert_data["teams"] == [test_team_id]
# Confirm the clone path ran
mock_prisma_client.db.litellm_budgettable.find_unique.assert_called_once_with(
where={"budget_id": test_default_budget_id}
)
mock_prisma_client.db.litellm_budgettable.create.assert_called_once()
mock_prisma_client.db.litellm_budgettable.create.assert_not_called()
class _FakeBudgetTable:
def __init__(self) -> None:
self.rows: dict[str, dict[str, object]] = {}
def _record(self, budget_id: str) -> LiteLLM_BudgetTable:
row: Final = self.rows[budget_id]
return LiteLLM_BudgetTable(**{k: v for k, v in row.items() if k in LiteLLM_BudgetTable.model_fields})
async def create(
self, *, data: Mapping[str, object], include: Mapping[str, bool] | None = None
) -> LiteLLM_BudgetTable:
budget_id: Final = str(data.get("budget_id") or uuid.uuid4())
self.rows[budget_id] = {**data, "budget_id": budget_id}
return self._record(budget_id)
async def find_unique(self, *, where: Mapping[str, str]) -> LiteLLM_BudgetTable | None:
return self._record(where["budget_id"]) if where["budget_id"] in self.rows else None
async def update(self, *, where: Mapping[str, str], data: Mapping[str, object]) -> LiteLLM_BudgetTable:
self.rows[where["budget_id"]] = {**self.rows[where["budget_id"]], **data}
return self._record(where["budget_id"])
class _FakeMembershipTable:
def __init__(self, budgets: _FakeBudgetTable) -> None:
self.budgets: Final = budgets
self.budget_ids: dict[tuple[str, str], str | None] = {}
def membership(self, team_id: str, user_id: str) -> LiteLLM_TeamMembership:
budget_id: Final = self.budget_ids[(team_id, user_id)]
return LiteLLM_TeamMembership(
user_id=user_id,
team_id=team_id,
budget_id=budget_id,
litellm_budget_table=self.budgets._record(budget_id) if budget_id is not None else None,
)
@staticmethod
def _linked_budget_id(row: Mapping[str, object]) -> str | None:
budget_id: Final = row.get("budget_id")
if isinstance(budget_id, str):
return budget_id
connect: Final = row.get("litellm_budget_table")
if isinstance(connect, dict):
return connect["connect"]["budget_id"]
return None
async def upsert(
self,
*,
where: Mapping[str, Mapping[str, str]],
data: Mapping[str, Mapping[str, object]],
include: Mapping[str, bool] | None = None,
) -> LiteLLM_TeamMembership:
key: Final = where["user_id_team_id"]
membership_key: Final = (key["team_id"], key["user_id"])
if membership_key not in self.budget_ids:
self.budget_ids[membership_key] = self._linked_budget_id(data["create"])
elif "litellm_budget_table" in data["update"]:
self.budget_ids[membership_key] = self._linked_budget_id(data["update"])
return self.membership(*membership_key)
class _FakeUserTable:
async def upsert(self, *, where: Mapping[str, str], data: Mapping[str, Mapping[str, object]]) -> LiteLLM_UserTable:
return LiteLLM_UserTable(user_id=where["user_id"], teams=list(data["create"].get("teams", [])))
async def update_many(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> int:
return 1
class _FakeDb:
def __init__(self) -> None:
self.litellm_budgettable: Final = _FakeBudgetTable()
self.litellm_teammembership: Final = _FakeMembershipTable(self.litellm_budgettable)
self.litellm_usertable: Final = _FakeUserTable()
@pytest.mark.asyncio
async def test_team_update_reaches_inherited_members_but_not_overridden_ones():
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.auth.auth_checks import _check_team_member_budget
from litellm.proxy.management_endpoints.common_utils import _upsert_budget_and_membership
from litellm.proxy.management_endpoints.team_endpoints import TeamMemberBudgetHandler
from litellm.proxy.utils import ProxyLogging
db: Final = _FakeDb()
prisma_client: Final = MagicMock()
prisma_client.db = db
admin: Final = UserAPIKeyAuth(user_id="admin_user", user_role=LitellmUserRoles.PROXY_ADMIN)
team_id: Final = "team-shared-default"
default_budget: Final = await db.litellm_budgettable.create(data={"budget_id": "team-default", "max_budget": 100.0})
team: Final = LiteLLM_TeamTable(team_id=team_id, metadata={"team_member_budget_id": default_budget.budget_id})
for user_id in ("inherits", "overridden"):
await add_new_member(
new_member=Member(user_id=user_id, role="user"),
max_budget_in_team=None,
prisma_client=prisma_client,
team_id=team_id,
user_api_key_dict=admin,
litellm_proxy_admin_name="admin",
default_team_budget_id=default_budget.budget_id,
)
await _upsert_budget_and_membership(
db,
team_id=team_id,
user_id="overridden",
existing_budget_id=default_budget.budget_id,
user_api_key_dict=admin,
budget_patch={"max_budget": 50.0},
team_default_budget_id=default_budget.budget_id,
)
assert db.litellm_teammembership.membership(team_id, "inherits").budget_id == default_budget.budget_id
assert db.litellm_teammembership.membership(team_id, "overridden").budget_id != default_budget.budget_id
assert db.litellm_budgettable.rows[default_budget.budget_id]["max_budget"] == 100.0
with patch( # test-quality-ok: update_budget reads this module global; no dependency injection seam exists
"litellm.proxy.proxy_server.prisma_client", prisma_client
):
await TeamMemberBudgetHandler.upsert_team_member_budget_table(
team_table=team,
user_api_key_dict=admin,
updated_kv={},
team_member_budget=1.0,
)
async def spend_from_membership(counter_key: str, fallback_spend: float, max_budget: float | None = None) -> float:
return fallback_spend
async def check(user_id: str, spend: float) -> None:
membership: Final = db.litellm_teammembership.membership(team_id, user_id).model_copy(update={"spend": spend})
with patch( # test-quality-ok: production auth reads this module global; no dependency injection seam exists
"litellm.proxy.proxy_server.get_current_spend", spend_from_membership
):
await _check_team_member_budget(
team_object=team,
user_object=LiteLLM_UserTable(user_id=user_id),
valid_token=UserAPIKeyAuth(token="tok", user_id=user_id, team_id=team_id),
prisma_client=prisma_client,
user_api_key_cache=MagicMock(),
proxy_logging_obj=ProxyLogging(user_api_key_cache=None),
team_membership=membership,
team_membership_loaded=True,
)
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await check("inherits", spend=2.0)
assert exc_info.value.max_budget == 1.0
await check("overridden", spend=2.0)
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await check("overridden", spend=60.0)
assert exc_info.value.max_budget == 50.0
@pytest.mark.asyncio