mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
fix(team): correct per-member budget semantics in bulk member update
Reuse the audited single-member _upsert_budget_and_membership per selected member so bulk updates match POST /team/member_update: private budgets are disconnected when a patch clears their last limit, shared-default limits are cloned only for members still on the default (never leaked to null-budget members), all_members_in_team ids are deduplicated, and the cached team is refreshed after a role change so auth sees new roles immediately Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
42851f8c95
commit
937a3cf168
3 changed files with 212 additions and 167 deletions
|
|
@ -34,7 +34,6 @@ from litellm.proxy._types import (
|
|||
DeleteTeamRequest,
|
||||
FailedTeamMemberUpdate,
|
||||
LiteLLM_AuditLogs,
|
||||
LiteLLM_BudgetTable,
|
||||
LiteLLM_DeletedTeamTable,
|
||||
LiteLLM_ManagementEndpoint_MetadataFields,
|
||||
LiteLLM_ManagementEndpoint_MetadataFields_Premium,
|
||||
|
|
@ -83,10 +82,7 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
|||
from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars
|
||||
from litellm.proxy.common_utils.json_merge_patch import apply_json_merge_patch
|
||||
from litellm.proxy.management_endpoints.common_utils import (
|
||||
_TEAM_MEMBER_BUDGET_LIMIT_FIELDS,
|
||||
_budget_patch_to_write_data,
|
||||
_check_passthrough_routes_caller_permission,
|
||||
_is_set_budget_value,
|
||||
_is_user_org_admin_for_team,
|
||||
_is_user_team_admin,
|
||||
_set_object_metadata_field,
|
||||
|
|
@ -2833,21 +2829,6 @@ def _build_member_budget_patch(data: TeamMemberUpdateRequest | TeamMemberBulkUpd
|
|||
}
|
||||
|
||||
|
||||
async def _default_member_budget_fields(prisma_client: PrismaClient, default_budget_id: str) -> dict[str, Any]:
|
||||
"""Fetch the team's shared default member budget and return the limit
|
||||
fields a clone-on-write copy must inherit, so patching a member off the
|
||||
shared default keeps the limits the default was giving them."""
|
||||
default_budget_row = await prisma_client.db.litellm_budgettable.find_unique(where={"budget_id": default_budget_id})
|
||||
if default_budget_row is None:
|
||||
return {}
|
||||
default_budget = LiteLLM_BudgetTable(**default_budget_row.model_dump())
|
||||
return {
|
||||
field: value
|
||||
for field, value in default_budget.model_dump().items()
|
||||
if field in _TEAM_MEMBER_BUDGET_LIMIT_FIELDS and _is_set_budget_value(value)
|
||||
}
|
||||
|
||||
|
||||
def _validate_budget_duration(budget_duration: Optional[str]) -> None:
|
||||
"""Reject budget durations that can't be parsed, are non-positive, or
|
||||
overflow date math, so a bad value can't be persisted and later crash the
|
||||
|
|
@ -3055,7 +3036,12 @@ async def bulk_update_team_members(
|
|||
http_request: Request,
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
):
|
||||
from litellm.proxy.proxy_server import premium_user, prisma_client
|
||||
from litellm.proxy.proxy_server import (
|
||||
premium_user,
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail={"error": "No db connected"})
|
||||
|
|
@ -3090,7 +3076,9 @@ async def bulk_update_team_members(
|
|||
)
|
||||
|
||||
if data.all_members_in_team:
|
||||
user_ids = [member.user_id for member in existing_team.members_with_roles if member.user_id is not None]
|
||||
user_ids = list(
|
||||
dict.fromkeys(member.user_id for member in existing_team.members_with_roles if member.user_id is not None)
|
||||
)
|
||||
else:
|
||||
user_ids = list(dict.fromkeys(data.user_ids or []))
|
||||
|
||||
|
|
@ -3116,7 +3104,9 @@ async def bulk_update_team_members(
|
|||
]
|
||||
|
||||
budget_patch = _build_member_budget_patch(data.update_fields)
|
||||
updated_by = user_api_key_dict.user_id or ""
|
||||
|
||||
raw_default_budget_id = (existing_team.metadata or {}).get("team_member_budget_id")
|
||||
default_budget_id = raw_default_budget_id if isinstance(raw_default_budget_id, str) else None
|
||||
|
||||
budget_target_user_ids = valid_user_ids if budget_patch else []
|
||||
raw_memberships = (
|
||||
|
|
@ -3126,36 +3116,8 @@ async def bulk_update_team_members(
|
|||
if budget_target_user_ids
|
||||
else []
|
||||
)
|
||||
memberships = [LiteLLM_TeamMembership(**membership.model_dump()) for membership in raw_memberships]
|
||||
budget_id_by_user: dict[str, str | None] = {membership.user_id: membership.budget_id for membership in memberships}
|
||||
|
||||
raw_default_budget_id = (existing_team.metadata or {}).get("team_member_budget_id")
|
||||
default_budget_id = raw_default_budget_id if isinstance(raw_default_budget_id, str) else None
|
||||
|
||||
budget_ids_to_update = sorted(
|
||||
{
|
||||
budget_id
|
||||
for budget_id in budget_id_by_user.values()
|
||||
if budget_id is not None and budget_id != default_budget_id
|
||||
}
|
||||
)
|
||||
create_user_ids = [
|
||||
user_id for user_id in budget_target_user_ids if budget_id_by_user.get(user_id) in (None, default_budget_id)
|
||||
]
|
||||
|
||||
needs_default_clone = default_budget_id is not None and any(
|
||||
budget_id_by_user.get(user_id) == default_budget_id for user_id in create_user_ids
|
||||
)
|
||||
inherited_default_fields = (
|
||||
await _default_member_budget_fields(prisma_client, default_budget_id)
|
||||
if needs_default_clone and default_budget_id is not None
|
||||
else {}
|
||||
)
|
||||
write_data = _budget_patch_to_write_data(budget_patch)
|
||||
create_data = {
|
||||
"created_by": updated_by,
|
||||
"updated_by": updated_by,
|
||||
**_budget_patch_to_write_data({**inherited_default_fields, **budget_patch}),
|
||||
budget_id_by_user: dict[str, str | None] = {
|
||||
membership.user_id: membership.budget_id for membership in raw_memberships
|
||||
}
|
||||
|
||||
new_role = data.update_fields.role
|
||||
|
|
@ -3171,31 +3133,36 @@ async def bulk_update_team_members(
|
|||
else None
|
||||
)
|
||||
|
||||
if budget_ids_to_update or create_user_ids or updated_members_with_roles is not None:
|
||||
async with prisma_client.db.batch_() as batcher:
|
||||
if budget_ids_to_update:
|
||||
batcher.litellm_budgettable.update_many(
|
||||
where={"budget_id": {"in": budget_ids_to_update}},
|
||||
data={"updated_by": updated_by, **write_data},
|
||||
)
|
||||
for user_id in create_user_ids:
|
||||
batcher.litellm_budgettable.create(
|
||||
data={
|
||||
**create_data,
|
||||
"team_membership": (
|
||||
{"connect": [{"user_id_team_id": {"user_id": user_id, "team_id": team_id}}]}
|
||||
if user_id in budget_id_by_user
|
||||
else {"create": [{"user_id": user_id, "team_id": team_id}]}
|
||||
),
|
||||
}
|
||||
)
|
||||
if updated_members_with_roles is not None:
|
||||
batcher.litellm_teamtable.update(
|
||||
where={"team_id": team_id},
|
||||
data={
|
||||
"members_with_roles": json.dumps([member.model_dump() for member in updated_members_with_roles])
|
||||
},
|
||||
async def _apply_writes():
|
||||
async with prisma_client.db.tx() as tx:
|
||||
for user_id in budget_target_user_ids:
|
||||
await _upsert_budget_and_membership(
|
||||
tx=tx,
|
||||
team_id=team_id,
|
||||
user_id=user_id,
|
||||
existing_budget_id=budget_id_by_user.get(user_id),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
budget_patch=budget_patch,
|
||||
team_default_budget_id=default_budget_id,
|
||||
)
|
||||
if updated_members_with_roles is None:
|
||||
return None
|
||||
return await tx.litellm_teamtable.update(
|
||||
where={"team_id": team_id},
|
||||
data={"members_with_roles": json.dumps([member.model_dump() for member in updated_members_with_roles])},
|
||||
include={"object_permission": True},
|
||||
)
|
||||
|
||||
refreshed_team_row = (
|
||||
await _apply_writes() if budget_target_user_ids or updated_members_with_roles is not None else None
|
||||
)
|
||||
|
||||
if refreshed_team_row is not None:
|
||||
await _refresh_cached_team(
|
||||
team_row=refreshed_team_row,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
successful_updates = [
|
||||
TeamMemberUpdateResponse(
|
||||
|
|
|
|||
|
|
@ -194,6 +194,91 @@ async def test_bulk_member_update_writes_member_budget_row(
|
|||
assert membership.litellm_budget_table.tpm_limit == 4242
|
||||
|
||||
|
||||
async def test_bulk_member_update_clearing_last_limit_disconnects_private_budget(
|
||||
proxy_client, prisma, scratch, world
|
||||
):
|
||||
member_id = scratch.tag("member")
|
||||
budget_id = scratch.tag("bud")
|
||||
await create_scratch_team(
|
||||
prisma,
|
||||
scratch.prefix,
|
||||
organization_id=world.org_a_id,
|
||||
member_user_ids=[member_id],
|
||||
)
|
||||
await prisma.db.litellm_budgettable.create(
|
||||
data={"budget_id": budget_id, "tpm_limit": 999, "created_by": "t", "updated_by": "t"}
|
||||
)
|
||||
await prisma.db.litellm_teammembership.create(
|
||||
data={"user_id": member_id, "team_id": scratch.prefix, "budget_id": budget_id}
|
||||
)
|
||||
|
||||
resp = await proxy_client.patch(
|
||||
ROUTE.format(team_id=scratch.prefix),
|
||||
headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"},
|
||||
json={"user_ids": [member_id], "update_fields": {"tpm_limit": None}},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
|
||||
membership = await prisma.db.litellm_teammembership.find_unique(
|
||||
where={"user_id_team_id": {"user_id": member_id, "team_id": scratch.prefix}},
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
assert membership is not None
|
||||
assert membership.budget_id is None, "cleared budget must be disconnected"
|
||||
assert membership.litellm_budget_table is None
|
||||
|
||||
|
||||
async def test_bulk_member_update_does_not_leak_default_limits_to_null_budget_members(
|
||||
proxy_client, prisma, scratch, world
|
||||
):
|
||||
on_default = scratch.tag("ondefault")
|
||||
no_budget = scratch.tag("nobudget")
|
||||
default_budget_id = scratch.tag("defbud")
|
||||
await create_scratch_team(
|
||||
prisma,
|
||||
scratch.prefix,
|
||||
organization_id=world.org_a_id,
|
||||
member_user_ids=[on_default, no_budget],
|
||||
metadata={"team_member_budget_id": default_budget_id},
|
||||
)
|
||||
await prisma.db.litellm_budgettable.create(
|
||||
data={"budget_id": default_budget_id, "max_budget": 100.0, "created_by": "t", "updated_by": "t"}
|
||||
)
|
||||
await prisma.db.litellm_teammembership.create(
|
||||
data={"user_id": on_default, "team_id": scratch.prefix, "budget_id": default_budget_id}
|
||||
)
|
||||
await prisma.db.litellm_teammembership.create(
|
||||
data={"user_id": no_budget, "team_id": scratch.prefix}
|
||||
)
|
||||
|
||||
resp = await proxy_client.patch(
|
||||
ROUTE.format(team_id=scratch.prefix),
|
||||
headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"},
|
||||
json={"user_ids": [on_default, no_budget], "update_fields": {"tpm_limit": 42}},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
|
||||
on_default_m = await prisma.db.litellm_teammembership.find_unique(
|
||||
where={"user_id_team_id": {"user_id": on_default, "team_id": scratch.prefix}},
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
assert on_default_m is not None and on_default_m.litellm_budget_table is not None
|
||||
assert on_default_m.budget_id != default_budget_id, "must clone, not patch the shared row"
|
||||
assert on_default_m.litellm_budget_table.tpm_limit == 42
|
||||
assert on_default_m.litellm_budget_table.max_budget == 100.0, "clone inherits the default's limits"
|
||||
|
||||
no_budget_m = await prisma.db.litellm_teammembership.find_unique(
|
||||
where={"user_id_team_id": {"user_id": no_budget, "team_id": scratch.prefix}},
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
assert no_budget_m is not None and no_budget_m.litellm_budget_table is not None
|
||||
assert no_budget_m.litellm_budget_table.tpm_limit == 42
|
||||
assert no_budget_m.litellm_budget_table.max_budget is None, "null-budget member must not inherit default limits"
|
||||
|
||||
default_row = await prisma.db.litellm_budgettable.find_unique(where={"budget_id": default_budget_id})
|
||||
assert default_row is not None and default_row.max_budget == 100.0, "shared default must be untouched"
|
||||
|
||||
|
||||
async def test_bulk_member_update_over_max_batch_is_400(
|
||||
proxy_client, prisma, scratch, world
|
||||
):
|
||||
|
|
|
|||
|
|
@ -10,7 +10,6 @@ import litellm.proxy.proxy_server as proxy_server
|
|||
import litellm.proxy.management_endpoints.team_endpoints as team_endpoints
|
||||
from litellm.proxy._types import (
|
||||
BulkTeamMemberUpdateRequest,
|
||||
LiteLLM_BudgetTable,
|
||||
LiteLLM_TeamMembership,
|
||||
LiteLLM_TeamTable,
|
||||
LitellmUserRoles,
|
||||
|
|
@ -178,20 +177,15 @@ async def test_team_member_update_rejects_invalid_budget_duration(monkeypatch, b
|
|||
upsert_mock.assert_not_called()
|
||||
|
||||
|
||||
class _RecordedWrites:
|
||||
def __init__(self):
|
||||
self.budget_update_many: list = []
|
||||
self.budget_creates: list = []
|
||||
self.team_updates: list = []
|
||||
class _FakeTx:
|
||||
def __init__(self, team_row, recorder):
|
||||
self._team_row = team_row
|
||||
self._recorder = recorder
|
||||
self.litellm_teamtable = types.SimpleNamespace(update=self._team_update)
|
||||
|
||||
|
||||
class _FakeBatcher:
|
||||
def __init__(self, writes: _RecordedWrites):
|
||||
self.litellm_budgettable = types.SimpleNamespace(
|
||||
update_many=lambda **kwargs: writes.budget_update_many.append(kwargs),
|
||||
create=lambda **kwargs: writes.budget_creates.append(kwargs),
|
||||
)
|
||||
self.litellm_teamtable = types.SimpleNamespace(update=lambda **kwargs: writes.team_updates.append(kwargs))
|
||||
async def _team_update(self, **kwargs):
|
||||
self._recorder.team_updates.append(kwargs)
|
||||
return self._team_row
|
||||
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
|
@ -201,14 +195,11 @@ class _FakeBatcher:
|
|||
|
||||
|
||||
class _FakeBulkDb:
|
||||
"""Typed fake for the exact prisma surface the bulk endpoint touches, so the
|
||||
tests assert the real queries issued (one update_many, batched creates, one
|
||||
team update) instead of monkeypatching endpoint internals."""
|
||||
|
||||
def __init__(self, team_row, memberships, default_budget=None):
|
||||
self.writes = _RecordedWrites()
|
||||
def __init__(self, team_row, memberships):
|
||||
self.team_row = team_row
|
||||
self.memberships = memberships
|
||||
self.membership_find_many_wheres: list = []
|
||||
self.budget_find_unique_wheres: list = []
|
||||
self.team_updates: list = []
|
||||
|
||||
async def _team_find_unique(where):
|
||||
return team_row
|
||||
|
|
@ -217,23 +208,24 @@ class _FakeBulkDb:
|
|||
self.membership_find_many_wheres.append(where)
|
||||
return memberships
|
||||
|
||||
async def _budget_find_unique(where):
|
||||
self.budget_find_unique_wheres.append(where)
|
||||
return default_budget
|
||||
|
||||
self.litellm_teamtable = types.SimpleNamespace(find_unique=_team_find_unique)
|
||||
self.litellm_teammembership = types.SimpleNamespace(find_many=_membership_find_many)
|
||||
self.litellm_budgettable = types.SimpleNamespace(find_unique=_budget_find_unique)
|
||||
|
||||
def batch_(self):
|
||||
return _FakeBatcher(self.writes)
|
||||
def tx(self):
|
||||
return _FakeTx(self.team_row, self)
|
||||
|
||||
|
||||
def _bulk_setup(monkeypatch, team_row, memberships, default_budget=None):
|
||||
db = _FakeBulkDb(team_row, memberships, default_budget)
|
||||
def _bulk_setup(monkeypatch, team_row, memberships):
|
||||
db = _FakeBulkDb(team_row, memberships)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", types.SimpleNamespace(db=db))
|
||||
monkeypatch.setattr(proxy_server, "premium_user", False)
|
||||
return db
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", object())
|
||||
monkeypatch.setattr(proxy_server, "proxy_logging_obj", object())
|
||||
upsert_mock = AsyncMock()
|
||||
monkeypatch.setattr(team_endpoints, "_upsert_budget_and_membership", upsert_mock)
|
||||
refresh_mock = AsyncMock()
|
||||
monkeypatch.setattr(team_endpoints, "_refresh_cached_team", refresh_mock)
|
||||
return db, upsert_mock, refresh_mock
|
||||
|
||||
|
||||
def _bulk_request():
|
||||
|
|
@ -243,11 +235,12 @@ def _bulk_request():
|
|||
_ADMIN = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN.value, user_id="admin")
|
||||
|
||||
|
||||
def _upsert_call_by_user(upsert_mock):
|
||||
return {call.kwargs["user_id"]: call.kwargs for call in upsert_mock.await_args_list}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_patches_private_budgets_with_one_update_many(monkeypatch):
|
||||
"""Members that already own a private budget must be covered by a single
|
||||
update_many over their budget ids; a query per member re-introduces the
|
||||
n round trips this endpoint exists to avoid."""
|
||||
async def test_bulk_update_delegates_budget_upsert_per_valid_member(monkeypatch):
|
||||
team_row = LiteLLM_TeamTable(
|
||||
team_id="team-1234",
|
||||
members_with_roles=[
|
||||
|
|
@ -256,7 +249,7 @@ async def test_bulk_update_patches_private_budgets_with_one_update_many(monkeypa
|
|||
Member(user_id="user-3", role="user"),
|
||||
],
|
||||
)
|
||||
db = _bulk_setup(
|
||||
db, upsert_mock, refresh_mock = _bulk_setup(
|
||||
monkeypatch,
|
||||
team_row,
|
||||
memberships=[
|
||||
|
|
@ -275,22 +268,22 @@ async def test_bulk_update_patches_private_budgets_with_one_update_many(monkeypa
|
|||
user_api_key_dict=_ADMIN,
|
||||
)
|
||||
|
||||
assert db.writes.budget_update_many == [
|
||||
{"where": {"budget_id": {"in": ["bud-1", "bud-2"]}}, "data": {"updated_by": "admin", "tpm_limit": 42}}
|
||||
]
|
||||
assert db.writes.budget_creates == []
|
||||
assert db.writes.team_updates == []
|
||||
assert db.membership_find_many_wheres == [{"team_id": "team-1234", "user_id": {"in": ["user-1", "user-2"]}}]
|
||||
calls = _upsert_call_by_user(upsert_mock)
|
||||
assert set(calls) == {"user-1", "user-2"}
|
||||
assert calls["user-1"]["existing_budget_id"] == "bud-1"
|
||||
assert calls["user-2"]["existing_budget_id"] == "bud-2"
|
||||
for kwargs in calls.values():
|
||||
assert kwargs["budget_patch"] == {"tpm_limit": 42}
|
||||
assert kwargs["team_default_budget_id"] is None
|
||||
assert db.team_updates == []
|
||||
refresh_mock.assert_not_awaited()
|
||||
assert response.total_requested == 2
|
||||
assert [member.user_id for member in response.successful_updates] == ["user-1", "user-2"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_clones_default_budget_instead_of_patching_it(monkeypatch):
|
||||
"""A member on the team's shared default budget must get their own cloned
|
||||
budget (default limits + patch); patching the shared row in place would
|
||||
change limits for every member outside the request. A member with no
|
||||
membership row gets a new budget wired to a created membership."""
|
||||
async def test_bulk_update_forwards_shared_default_only_for_members_still_on_it(monkeypatch):
|
||||
team_row = LiteLLM_TeamTable(
|
||||
team_id="team-1234",
|
||||
members_with_roles=[
|
||||
|
|
@ -299,11 +292,10 @@ async def test_bulk_update_clones_default_budget_instead_of_patching_it(monkeypa
|
|||
],
|
||||
metadata={"team_member_budget_id": "default-bud"},
|
||||
)
|
||||
db = _bulk_setup(
|
||||
_db, upsert_mock, _refresh = _bulk_setup(
|
||||
monkeypatch,
|
||||
team_row,
|
||||
memberships=[LiteLLM_TeamMembership(user_id="user-1", team_id="team-1234", budget_id="default-bud")],
|
||||
default_budget=LiteLLM_BudgetTable(budget_id="default-bud", max_budget=100.0),
|
||||
)
|
||||
|
||||
await bulk_update_team_members(
|
||||
|
|
@ -316,34 +308,15 @@ async def test_bulk_update_clones_default_budget_instead_of_patching_it(monkeypa
|
|||
user_api_key_dict=_ADMIN,
|
||||
)
|
||||
|
||||
assert db.writes.budget_update_many == []
|
||||
assert db.budget_find_unique_wheres == [{"budget_id": "default-bud"}]
|
||||
assert db.writes.budget_creates == [
|
||||
{
|
||||
"data": {
|
||||
"created_by": "admin",
|
||||
"updated_by": "admin",
|
||||
"max_budget": 100.0,
|
||||
"tpm_limit": 42,
|
||||
"team_membership": {"connect": [{"user_id_team_id": {"user_id": "user-1", "team_id": "team-1234"}}]},
|
||||
}
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"created_by": "admin",
|
||||
"updated_by": "admin",
|
||||
"max_budget": 100.0,
|
||||
"tpm_limit": 42,
|
||||
"team_membership": {"create": [{"user_id": "user-2", "team_id": "team-1234"}]},
|
||||
}
|
||||
},
|
||||
]
|
||||
calls = _upsert_call_by_user(upsert_mock)
|
||||
assert calls["user-1"]["existing_budget_id"] == "default-bud"
|
||||
assert calls["user-2"]["existing_budget_id"] is None
|
||||
for kwargs in calls.values():
|
||||
assert kwargs["team_default_budget_id"] == "default-bud"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_role_writes_team_row_once(monkeypatch):
|
||||
"""A role-only bulk update must rewrite members_with_roles in a single team
|
||||
update covering every targeted member, and must not touch budgets at all."""
|
||||
async def test_bulk_update_role_writes_team_row_once_and_refreshes_cache(monkeypatch):
|
||||
team_row = LiteLLM_TeamTable(
|
||||
team_id="team-1234",
|
||||
members_with_roles=[
|
||||
|
|
@ -352,7 +325,7 @@ async def test_bulk_update_role_writes_team_row_once(monkeypatch):
|
|||
Member(user_id="user-3", role="user"),
|
||||
],
|
||||
)
|
||||
db = _bulk_setup(monkeypatch, team_row, memberships=[])
|
||||
db, upsert_mock, refresh_mock = _bulk_setup(monkeypatch, team_row, memberships=[])
|
||||
|
||||
await bulk_update_team_members(
|
||||
team_id="team-1234",
|
||||
|
|
@ -365,10 +338,9 @@ async def test_bulk_update_role_writes_team_row_once(monkeypatch):
|
|||
)
|
||||
|
||||
assert db.membership_find_many_wheres == []
|
||||
assert db.writes.budget_update_many == []
|
||||
assert db.writes.budget_creates == []
|
||||
assert len(db.writes.team_updates) == 1
|
||||
update = db.writes.team_updates[0]
|
||||
upsert_mock.assert_not_awaited()
|
||||
assert len(db.team_updates) == 1
|
||||
update = db.team_updates[0]
|
||||
assert update["where"] == {"team_id": "team-1234"}
|
||||
members = json.loads(update["data"]["members_with_roles"])
|
||||
assert [(member["user_id"], member["role"]) for member in members] == [
|
||||
|
|
@ -377,6 +349,35 @@ async def test_bulk_update_role_writes_team_row_once(monkeypatch):
|
|||
("user-3", "user"),
|
||||
]
|
||||
assert members[1]["user_email"] == "two@example.com"
|
||||
refresh_mock.assert_awaited_once()
|
||||
assert refresh_mock.await_args.kwargs["team_row"] is team_row
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_all_members_in_team_dedups_members(monkeypatch):
|
||||
team_row = LiteLLM_TeamTable(
|
||||
team_id="team-1234",
|
||||
members_with_roles=[
|
||||
Member(user_id="user-1", role="user"),
|
||||
Member(user_id="user-1", role="user"),
|
||||
Member(user_id="user-2", role="user"),
|
||||
],
|
||||
)
|
||||
db, upsert_mock, _refresh = _bulk_setup(monkeypatch, team_row, memberships=[])
|
||||
|
||||
response = await bulk_update_team_members(
|
||||
team_id="team-1234",
|
||||
data=BulkTeamMemberUpdateRequest(
|
||||
all_members_in_team=True,
|
||||
update_fields=TeamMemberBulkUpdateFields(tpm_limit=42),
|
||||
),
|
||||
http_request=_bulk_request(),
|
||||
user_api_key_dict=_ADMIN,
|
||||
)
|
||||
|
||||
assert db.membership_find_many_wheres == [{"team_id": "team-1234", "user_id": {"in": ["user-1", "user-2"]}}]
|
||||
assert list(_upsert_call_by_user(upsert_mock)) == ["user-1", "user-2"]
|
||||
assert response.total_requested == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -385,7 +386,7 @@ async def test_bulk_update_reports_non_members_as_failed(monkeypatch):
|
|||
team_id="team-1234",
|
||||
members_with_roles=[Member(user_id="user-1", role="user")],
|
||||
)
|
||||
db = _bulk_setup(
|
||||
_db, upsert_mock, _refresh = _bulk_setup(
|
||||
monkeypatch,
|
||||
team_row,
|
||||
memberships=[LiteLLM_TeamMembership(user_id="user-1", team_id="team-1234", budget_id="bud-1")],
|
||||
|
|
@ -405,19 +406,16 @@ async def test_bulk_update_reports_non_members_as_failed(monkeypatch):
|
|||
assert [member.user_id for member in response.successful_updates] == ["user-1"]
|
||||
assert response.failed_updates[0].user_id == "ghost-user"
|
||||
assert "not a member" in response.failed_updates[0].failed_reason
|
||||
assert db.writes.budget_update_many[0]["where"] == {"budget_id": {"in": ["bud-1"]}}
|
||||
assert db.writes.budget_creates == []
|
||||
assert list(_upsert_call_by_user(upsert_mock)) == ["user-1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_explicit_null_duration_clears_reset_at(monkeypatch):
|
||||
"""budget_duration: null must clear both the duration and budget_reset_at in
|
||||
the same update_many, otherwise stale reset timestamps keep firing."""
|
||||
async def test_bulk_update_explicit_null_duration_forwarded_as_patch(monkeypatch):
|
||||
team_row = LiteLLM_TeamTable(
|
||||
team_id="team-1234",
|
||||
members_with_roles=[Member(user_id="user-1", role="user")],
|
||||
)
|
||||
db = _bulk_setup(
|
||||
_db, upsert_mock, _refresh = _bulk_setup(
|
||||
monkeypatch,
|
||||
team_row,
|
||||
memberships=[LiteLLM_TeamMembership(user_id="user-1", team_id="team-1234", budget_id="bud-1")],
|
||||
|
|
@ -433,12 +431,7 @@ async def test_bulk_update_explicit_null_duration_clears_reset_at(monkeypatch):
|
|||
user_api_key_dict=_ADMIN,
|
||||
)
|
||||
|
||||
assert db.writes.budget_update_many == [
|
||||
{
|
||||
"where": {"budget_id": {"in": ["bud-1"]}},
|
||||
"data": {"updated_by": "admin", "budget_duration": None, "budget_reset_at": None},
|
||||
}
|
||||
]
|
||||
assert _upsert_call_by_user(upsert_mock)["user-1"]["budget_patch"] == {"budget_duration": None}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue