mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
refactor(team): batch bulk member update writes behind PATCH /v2/team/{team_id}/members
Replace the per-member loop (which re-ran budget upserts and rewrote the full members_with_roles JSON once per member) with set-based writes: one litellm_budgettable.update_many for every member that already owns a private budget, cloned budget creates only for members on the shared team default or with no budget row, and a single team row update for role changes, all grouped in one prisma batch transaction. team_id moves to the path and out of the request body. The patch building and budget_reset_at handling are shared with /team/member_update via _build_member_budget_patch and the extracted _budget_patch_to_write_data
This commit is contained in:
parent
8263add5ff
commit
1611b8b1ea
8 changed files with 503 additions and 211 deletions
|
|
@ -600,7 +600,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/team/permissions_list",
|
||||
"/team/permissions_update",
|
||||
"/team/permissions_bulk_update",
|
||||
"/team/member/bulk_update",
|
||||
"/v2/team/{team_id}/members",
|
||||
"/team/daily/activity",
|
||||
# model
|
||||
"/model/new",
|
||||
|
|
@ -742,6 +742,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/team/member_add",
|
||||
"/team/member_delete",
|
||||
"/team/member_update",
|
||||
"/v2/team/{team_id}/members",
|
||||
"/team/permissions_list",
|
||||
"/team/permissions_update",
|
||||
"/team/daily/activity",
|
||||
|
|
@ -3721,7 +3722,6 @@ class TeamMemberBulkUpdateFields(LiteLLMPydanticObjectBase):
|
|||
|
||||
|
||||
class BulkTeamMemberUpdateRequest(LiteLLMPydanticObjectBase):
|
||||
team_id: str
|
||||
user_ids: list[str] | None = None
|
||||
all_members_in_team: bool = False
|
||||
update_fields: TeamMemberBulkUpdateFields
|
||||
|
|
|
|||
|
|
@ -33,7 +33,6 @@ _PROXY_ADMIN_VIEW_ONLY_BLOCKED_ROUTES = frozenset(
|
|||
"/team/unblock",
|
||||
"/team/permissions_update",
|
||||
"/team/permissions_bulk_update",
|
||||
"/team/member/bulk_update",
|
||||
# model
|
||||
"/model/new",
|
||||
"/model/update",
|
||||
|
|
@ -735,7 +734,6 @@ class RouteChecks:
|
|||
"/team/new",
|
||||
"/team/update",
|
||||
"/team/delete",
|
||||
"/team/member/bulk_update",
|
||||
"/model/new",
|
||||
"/model/update",
|
||||
"/model/delete",
|
||||
|
|
|
|||
|
|
@ -423,6 +423,20 @@ def _has_meaningful_budget_limit(budget_values: Dict[str, Any]) -> bool:
|
|||
return any(_is_set_budget_value(budget_values.get(field)) for field in _TEAM_MEMBER_BUDGET_LIMIT_FIELDS)
|
||||
|
||||
|
||||
def _budget_patch_to_write_data(budget_patch: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Turn an RFC 7396-style budget patch into the budget-table write payload:
|
||||
setting budget_duration also recomputes budget_reset_at, clearing the
|
||||
duration clears budget_reset_at, and a patch that never mentions the
|
||||
duration leaves the reset timestamp alone."""
|
||||
if "budget_duration" not in budget_patch:
|
||||
return dict(budget_patch)
|
||||
duration = budget_patch["budget_duration"]
|
||||
return {
|
||||
**budget_patch,
|
||||
"budget_reset_at": get_budget_reset_time(budget_duration=duration) if duration is not None else None,
|
||||
}
|
||||
|
||||
|
||||
async def _upsert_budget_and_membership(
|
||||
tx,
|
||||
*,
|
||||
|
|
@ -450,12 +464,7 @@ async def _upsert_budget_and_membership(
|
|||
if not budget_patch:
|
||||
return
|
||||
|
||||
write_data = dict(budget_patch)
|
||||
if "budget_duration" in write_data:
|
||||
duration = write_data["budget_duration"]
|
||||
write_data["budget_reset_at"] = (
|
||||
get_budget_reset_time(budget_duration=duration) if duration is not None else None
|
||||
)
|
||||
write_data = _budget_patch_to_write_data(budget_patch)
|
||||
|
||||
is_shared_default = (
|
||||
existing_budget_id is not None
|
||||
|
|
|
|||
|
|
@ -34,6 +34,7 @@ from litellm.proxy._types import (
|
|||
DeleteTeamRequest,
|
||||
FailedTeamMemberUpdate,
|
||||
LiteLLM_AuditLogs,
|
||||
LiteLLM_BudgetTable,
|
||||
LiteLLM_DeletedTeamTable,
|
||||
LiteLLM_ManagementEndpoint_MetadataFields,
|
||||
LiteLLM_ManagementEndpoint_MetadataFields_Premium,
|
||||
|
|
@ -60,6 +61,7 @@ from litellm.proxy._types import (
|
|||
TeamInfoResponseObjectTeamTable,
|
||||
TeamListResponseObject,
|
||||
TeamMemberAddRequest,
|
||||
TeamMemberBulkUpdateFields,
|
||||
TeamMemberDeleteRequest,
|
||||
TeamMemberUpdateRequest,
|
||||
TeamMemberUpdateResponse,
|
||||
|
|
@ -81,7 +83,10 @@ 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,
|
||||
|
|
@ -2816,7 +2821,7 @@ _MEMBER_BUDGET_PATCH_FIELDS = {
|
|||
}
|
||||
|
||||
|
||||
def _build_member_budget_patch(data: TeamMemberUpdateRequest) -> Dict[str, Any]:
|
||||
def _build_member_budget_patch(data: TeamMemberUpdateRequest | TeamMemberBulkUpdateFields) -> dict[str, Any]:
|
||||
"""Map the budget fields the request actually set (merge-patch: a sent
|
||||
value updates, an explicit null clears, an absent field is left untouched)
|
||||
to their budget-table columns."""
|
||||
|
|
@ -2828,6 +2833,21 @@ def _build_member_budget_patch(data: TeamMemberUpdateRequest) -> Dict[str, Any]:
|
|||
}
|
||||
|
||||
|
||||
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
|
||||
|
|
@ -3022,14 +3042,15 @@ async def _apply_team_member_update(
|
|||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/team/member/bulk_update",
|
||||
@router.patch(
|
||||
"/v2/team/{team_id}/members",
|
||||
tags=["team management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=BulkTeamMemberUpdateResponse,
|
||||
)
|
||||
@management_endpoint_wrapper
|
||||
async def bulk_update_team_members(
|
||||
team_id: str,
|
||||
data: BulkTeamMemberUpdateRequest,
|
||||
http_request: Request,
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
|
|
@ -3047,11 +3068,11 @@ async def bulk_update_team_members(
|
|||
|
||||
_validate_budget_duration(data.update_fields.budget_duration)
|
||||
|
||||
existing_team_row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": data.team_id})
|
||||
existing_team_row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id})
|
||||
if existing_team_row is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": "Team id={} does not exist in db".format(data.team_id)},
|
||||
detail={"error": "Team id={} does not exist in db".format(team_id)},
|
||||
)
|
||||
existing_team = LiteLLM_TeamTable(**existing_team_row.model_dump())
|
||||
if (
|
||||
|
|
@ -3063,7 +3084,7 @@ async def bulk_update_team_members(
|
|||
status_code=403,
|
||||
detail={
|
||||
"error": "Call not allowed. User not proxy admin OR team admin. route={}, team_id={}".format(
|
||||
"/team/member/bulk_update", data.team_id
|
||||
"/v2/team/{team_id}/members", team_id
|
||||
)
|
||||
},
|
||||
)
|
||||
|
|
@ -3084,35 +3105,112 @@ async def bulk_update_team_members(
|
|||
},
|
||||
)
|
||||
|
||||
returned_team_info: TeamInfoResponseObject = await team_info(
|
||||
http_request=http_request,
|
||||
team_id=data.team_id,
|
||||
key_limit=None,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
member_user_ids = {member.user_id for member in existing_team.members_with_roles if member.user_id is not None}
|
||||
valid_user_ids = [user_id for user_id in user_ids if user_id in member_user_ids]
|
||||
failed_updates = [
|
||||
FailedTeamMemberUpdate(
|
||||
user_id=user_id, failed_reason="User id={} is not a member of team {}".format(user_id, team_id)
|
||||
)
|
||||
for user_id in user_ids
|
||||
if user_id not in member_user_ids
|
||||
]
|
||||
|
||||
budget_patch = _build_member_budget_patch(data.update_fields)
|
||||
updated_by = user_api_key_dict.user_id or ""
|
||||
|
||||
budget_target_user_ids = valid_user_ids if budget_patch else []
|
||||
raw_memberships = (
|
||||
await prisma_client.db.litellm_teammembership.find_many(
|
||||
where={"team_id": team_id, "user_id": {"in": budget_target_user_ids}}
|
||||
)
|
||||
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}),
|
||||
}
|
||||
|
||||
new_role = data.update_fields.role
|
||||
valid_user_id_set = frozenset(valid_user_ids)
|
||||
updated_members_with_roles = (
|
||||
[
|
||||
Member(user_id=member.user_id, role=new_role, user_email=member.user_email)
|
||||
if member.user_id in valid_user_id_set
|
||||
else member
|
||||
for member in existing_team.members_with_roles
|
||||
]
|
||||
if new_role is not None
|
||||
else None
|
||||
)
|
||||
|
||||
update_fields = data.update_fields.model_dump(exclude_unset=True)
|
||||
successful_updates: list[TeamMemberUpdateResponse] = []
|
||||
failed_updates: list[FailedTeamMemberUpdate] = []
|
||||
for user_id in user_ids:
|
||||
try:
|
||||
response = await _apply_team_member_update(
|
||||
data=TeamMemberUpdateRequest(team_id=data.team_id, user_id=user_id, **update_fields),
|
||||
returned_team_info=returned_team_info,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
successful_updates.append(response)
|
||||
except HTTPException as exc:
|
||||
detail = exc.detail
|
||||
failed_reason = detail.get("error", str(detail)) if isinstance(detail, dict) else str(detail)
|
||||
failed_updates.append(FailedTeamMemberUpdate(user_id=user_id, failed_reason=failed_reason))
|
||||
except Exception as exc:
|
||||
verbose_proxy_logger.exception("Failed to bulk update team member %s in team %s", user_id, data.team_id)
|
||||
failed_updates.append(FailedTeamMemberUpdate(user_id=user_id, failed_reason=str(exc)))
|
||||
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])
|
||||
},
|
||||
)
|
||||
|
||||
successful_updates = [
|
||||
TeamMemberUpdateResponse(
|
||||
team_id=team_id,
|
||||
user_id=user_id,
|
||||
max_budget_in_team=data.update_fields.max_budget_in_team,
|
||||
tpm_limit=data.update_fields.tpm_limit,
|
||||
rpm_limit=data.update_fields.rpm_limit,
|
||||
budget_duration=data.update_fields.budget_duration,
|
||||
allowed_models=data.update_fields.allowed_models,
|
||||
)
|
||||
for user_id in valid_user_ids
|
||||
]
|
||||
return BulkTeamMemberUpdateResponse(
|
||||
team_id=data.team_id,
|
||||
team_id=team_id,
|
||||
total_requested=len(user_ids),
|
||||
successful_updates=successful_updates,
|
||||
failed_updates=failed_updates,
|
||||
|
|
|
|||
|
|
@ -131,7 +131,6 @@ def test_proxy_admin_viewer_config_update_route_rejected():
|
|||
"/team/unblock",
|
||||
"/team/permissions_update",
|
||||
"/team/permissions_bulk_update",
|
||||
"/team/member/bulk_update",
|
||||
# JWT key mapping write routes
|
||||
"/jwt/key/mapping/new",
|
||||
"/jwt/key/mapping/update",
|
||||
|
|
@ -2883,6 +2882,46 @@ def test_patch_team_gate_rejects_view_only_admin():
|
|||
)
|
||||
|
||||
|
||||
def test_bulk_member_update_route_has_same_reach_as_member_update():
|
||||
"""PATCH /v2/team/{team_id}/members must be reachable by the same coarse gate
|
||||
as /team/member_update (self_managed_routes; the endpoint enforces proxy /
|
||||
team / org admin itself), without the resolved path colliding with static
|
||||
siblings like /v2/team/list."""
|
||||
from litellm.proxy._types import LiteLLMRoutes
|
||||
|
||||
assert RouteChecks.check_route_access(
|
||||
route="/v2/team/team-1/members", allowed_routes=LiteLLMRoutes.self_managed_routes.value
|
||||
)
|
||||
assert RouteChecks.check_route_access(
|
||||
route="/v2/team/team-1/members", allowed_routes=LiteLLMRoutes.management_routes.value
|
||||
)
|
||||
assert not RouteChecks.check_route_access(
|
||||
route="/v2/team/list", allowed_routes=LiteLLMRoutes.self_managed_routes.value
|
||||
)
|
||||
|
||||
|
||||
def test_bulk_member_update_gate_rejects_view_only_admin():
|
||||
"""A view-only proxy admin cannot PATCH /v2/team/{team_id}/members: the
|
||||
templated path never exact-matches the write blocklists, so the unsafe-method
|
||||
default-deny is what has to catch it."""
|
||||
user_obj = LiteLLM_UserTable(
|
||||
user_id="viewer",
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value,
|
||||
)
|
||||
valid_token = UserAPIKeyAuth(user_id="viewer", user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
RouteChecks.non_proxy_admin_allowed_routes_check(
|
||||
user_obj=user_obj,
|
||||
_user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value,
|
||||
route="/v2/team/team-1/members",
|
||||
request=_patch_team_request(),
|
||||
valid_token=valid_token,
|
||||
request_data={},
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_pass_through_registers_wildcard_for_auth_subpath():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import json
|
||||
import types
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
|
|
@ -9,6 +10,8 @@ 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,
|
||||
Member,
|
||||
|
|
@ -86,7 +89,7 @@ def happy_path_upsert(monkeypatch):
|
|||
AsyncMock(
|
||||
return_value={
|
||||
"team_info": team_row,
|
||||
"team_memberships": [types.SimpleNamespace(user_id="user-1", budget_id="bud-1")],
|
||||
"team_memberships": [LiteLLM_TeamMembership(user_id="user-1", team_id="team-1234", budget_id="bud-1")],
|
||||
}
|
||||
),
|
||||
)
|
||||
|
|
@ -175,90 +178,76 @@ async def test_team_member_update_rejects_invalid_budget_duration(monkeypatch, b
|
|||
upsert_mock.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_team_member_update_applies_patch_and_returns_member_failures(monkeypatch):
|
||||
team_row = LiteLLM_TeamTable(
|
||||
team_id="team-1234",
|
||||
members_with_roles=[Member(user_id="user-1", role="user")],
|
||||
)
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma_client)
|
||||
monkeypatch.setattr(
|
||||
team_endpoints,
|
||||
"team_info",
|
||||
AsyncMock(return_value={"team_info": team_row, "team_memberships": []}),
|
||||
)
|
||||
update_mock = AsyncMock(
|
||||
side_effect=[
|
||||
team_endpoints.TeamMemberUpdateResponse(team_id="team-1234", user_id="user-1", tpm_limit=42),
|
||||
HTTPException(status_code=404, detail={"error": "User is not a team member"}),
|
||||
]
|
||||
)
|
||||
monkeypatch.setattr(team_endpoints, "_apply_team_member_update", update_mock)
|
||||
class _RecordedWrites:
|
||||
def __init__(self):
|
||||
self.budget_update_many: list = []
|
||||
self.budget_creates: list = []
|
||||
self.team_updates: list = []
|
||||
|
||||
response = await bulk_update_team_members(
|
||||
data=BulkTeamMemberUpdateRequest(
|
||||
team_id="team-1234",
|
||||
user_ids=["user-1", "user-2", "user-1"],
|
||||
update_fields=TeamMemberBulkUpdateFields(tpm_limit=42),
|
||||
),
|
||||
http_request=Request({"type": "http", "method": "POST", "path": "/team/member/bulk_update"}),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN.value, user_id="admin"),
|
||||
)
|
||||
|
||||
assert response.total_requested == 2
|
||||
assert [member.user_id for member in response.successful_updates] == ["user-1"]
|
||||
assert response.failed_updates[0].user_id == "user-2"
|
||||
# a dict HTTPException detail must surface the nested error string, not a
|
||||
# python dict repr like "{'error': 'User is not a team member'}"
|
||||
assert response.failed_updates[0].failed_reason == "User is not a team member"
|
||||
assert update_mock.await_args_list[0].kwargs["data"].model_dump(exclude_unset=True) == {
|
||||
"team_id": "team-1234",
|
||||
"user_id": "user-1",
|
||||
"tpm_limit": 42,
|
||||
}
|
||||
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 __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *args):
|
||||
return False
|
||||
|
||||
|
||||
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()
|
||||
self.membership_find_many_wheres: list = []
|
||||
self.budget_find_unique_wheres: list = []
|
||||
|
||||
async def _team_find_unique(where):
|
||||
return team_row
|
||||
|
||||
async def _membership_find_many(where):
|
||||
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 _bulk_setup(monkeypatch, team_row, memberships, default_budget=None):
|
||||
db = _FakeBulkDb(team_row, memberships, default_budget)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", types.SimpleNamespace(db=db))
|
||||
monkeypatch.setattr(proxy_server, "premium_user", False)
|
||||
return db
|
||||
|
||||
|
||||
def _bulk_request():
|
||||
return Request({"type": "http", "method": "PATCH", "path": "/v2/team/team-1234/members"})
|
||||
|
||||
|
||||
_ADMIN = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN.value, user_id="admin")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_team_member_update_returns_unexpected_member_failure(monkeypatch):
|
||||
team_row = LiteLLM_TeamTable(
|
||||
team_id="team-1234",
|
||||
members_with_roles=[Member(user_id="user-1", role="user")],
|
||||
)
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma_client)
|
||||
monkeypatch.setattr(
|
||||
team_endpoints,
|
||||
"team_info",
|
||||
AsyncMock(return_value={"team_info": team_row, "team_memberships": []}),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
team_endpoints, "_apply_team_member_update", AsyncMock(side_effect=RuntimeError("database unavailable"))
|
||||
)
|
||||
|
||||
response = await bulk_update_team_members(
|
||||
data=BulkTeamMemberUpdateRequest(
|
||||
team_id="team-1234",
|
||||
user_ids=["user-1"],
|
||||
update_fields=TeamMemberBulkUpdateFields(tpm_limit=42),
|
||||
),
|
||||
http_request=Request({"type": "http", "method": "POST", "path": "/team/member/bulk_update"}),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN.value, user_id="admin"),
|
||||
)
|
||||
|
||||
assert response.successful_updates == []
|
||||
assert response.failed_updates == [
|
||||
team_endpoints.FailedTeamMemberUpdate(user_id="user-1", failed_reason="database unavailable")
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_team_member_update_resolves_team_info_once(monkeypatch):
|
||||
"""The whole batch must resolve team_info a single time and still upsert every
|
||||
member; resolving it per member re-scans the team, its keys, and all
|
||||
memberships on each iteration, which times out large teams."""
|
||||
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."""
|
||||
team_row = LiteLLM_TeamTable(
|
||||
team_id="team-1234",
|
||||
members_with_roles=[
|
||||
|
|
@ -266,55 +255,215 @@ async def test_bulk_team_member_update_resolves_team_info_once(monkeypatch):
|
|||
Member(user_id="user-2", role="user"),
|
||||
Member(user_id="user-3", role="user"),
|
||||
],
|
||||
metadata={},
|
||||
)
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row)
|
||||
|
||||
class _FakeTx:
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *args):
|
||||
return False
|
||||
|
||||
prisma_client.db.tx = MagicMock(return_value=_FakeTx())
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma_client)
|
||||
monkeypatch.setattr(proxy_server, "premium_user", False)
|
||||
|
||||
team_info_mock = AsyncMock(
|
||||
return_value={
|
||||
"team_info": team_row,
|
||||
"team_memberships": [
|
||||
types.SimpleNamespace(user_id="user-1", budget_id="bud-1"),
|
||||
types.SimpleNamespace(user_id="user-2", budget_id="bud-2"),
|
||||
types.SimpleNamespace(user_id="user-3", budget_id="bud-3"),
|
||||
],
|
||||
}
|
||||
db = _bulk_setup(
|
||||
monkeypatch,
|
||||
team_row,
|
||||
memberships=[
|
||||
LiteLLM_TeamMembership(user_id="user-1", team_id="team-1234", budget_id="bud-1"),
|
||||
LiteLLM_TeamMembership(user_id="user-2", team_id="team-1234", budget_id="bud-2"),
|
||||
],
|
||||
)
|
||||
monkeypatch.setattr(team_endpoints, "team_info", team_info_mock)
|
||||
upsert_mock = AsyncMock()
|
||||
monkeypatch.setattr(team_endpoints, "_upsert_budget_and_membership", upsert_mock)
|
||||
|
||||
response = await bulk_update_team_members(
|
||||
team_id="team-1234",
|
||||
data=BulkTeamMemberUpdateRequest(
|
||||
team_id="team-1234",
|
||||
all_members_in_team=True,
|
||||
user_ids=["user-1", "user-2", "user-1"],
|
||||
update_fields=TeamMemberBulkUpdateFields(tpm_limit=42),
|
||||
),
|
||||
http_request=Request({"type": "http", "method": "POST", "path": "/team/member/bulk_update"}),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN.value, user_id="admin"),
|
||||
http_request=_bulk_request(),
|
||||
user_api_key_dict=_ADMIN,
|
||||
)
|
||||
|
||||
assert team_info_mock.await_count == 1
|
||||
assert upsert_mock.await_count == 3
|
||||
assert [member.user_id for member in response.successful_updates] == ["user-1", "user-2", "user-3"]
|
||||
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"]}}]
|
||||
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."""
|
||||
team_row = LiteLLM_TeamTable(
|
||||
team_id="team-1234",
|
||||
members_with_roles=[
|
||||
Member(user_id="user-1", role="user"),
|
||||
Member(user_id="user-2", role="user"),
|
||||
],
|
||||
metadata={"team_member_budget_id": "default-bud"},
|
||||
)
|
||||
db = _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(
|
||||
team_id="team-1234",
|
||||
data=BulkTeamMemberUpdateRequest(
|
||||
user_ids=["user-1", "user-2"],
|
||||
update_fields=TeamMemberBulkUpdateFields(tpm_limit=42),
|
||||
),
|
||||
http_request=_bulk_request(),
|
||||
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"}]},
|
||||
}
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
@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."""
|
||||
team_row = LiteLLM_TeamTable(
|
||||
team_id="team-1234",
|
||||
members_with_roles=[
|
||||
Member(user_id="user-1", role="admin"),
|
||||
Member(user_id="user-2", role="user", user_email="two@example.com"),
|
||||
Member(user_id="user-3", role="user"),
|
||||
],
|
||||
)
|
||||
db = _bulk_setup(monkeypatch, team_row, memberships=[])
|
||||
|
||||
await bulk_update_team_members(
|
||||
team_id="team-1234",
|
||||
data=BulkTeamMemberUpdateRequest(
|
||||
user_ids=["user-1", "user-2"],
|
||||
update_fields=TeamMemberBulkUpdateFields(role="user"),
|
||||
),
|
||||
http_request=_bulk_request(),
|
||||
user_api_key_dict=_ADMIN,
|
||||
)
|
||||
|
||||
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]
|
||||
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] == [
|
||||
("user-1", "user"),
|
||||
("user-2", "user"),
|
||||
("user-3", "user"),
|
||||
]
|
||||
assert members[1]["user_email"] == "two@example.com"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_reports_non_members_as_failed(monkeypatch):
|
||||
team_row = LiteLLM_TeamTable(
|
||||
team_id="team-1234",
|
||||
members_with_roles=[Member(user_id="user-1", role="user")],
|
||||
)
|
||||
db = _bulk_setup(
|
||||
monkeypatch,
|
||||
team_row,
|
||||
memberships=[LiteLLM_TeamMembership(user_id="user-1", team_id="team-1234", budget_id="bud-1")],
|
||||
)
|
||||
|
||||
response = await bulk_update_team_members(
|
||||
team_id="team-1234",
|
||||
data=BulkTeamMemberUpdateRequest(
|
||||
user_ids=["user-1", "ghost-user"],
|
||||
update_fields=TeamMemberBulkUpdateFields(tpm_limit=42),
|
||||
),
|
||||
http_request=_bulk_request(),
|
||||
user_api_key_dict=_ADMIN,
|
||||
)
|
||||
|
||||
assert response.total_requested == 2
|
||||
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 == []
|
||||
|
||||
|
||||
@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."""
|
||||
team_row = LiteLLM_TeamTable(
|
||||
team_id="team-1234",
|
||||
members_with_roles=[Member(user_id="user-1", role="user")],
|
||||
)
|
||||
db = _bulk_setup(
|
||||
monkeypatch,
|
||||
team_row,
|
||||
memberships=[LiteLLM_TeamMembership(user_id="user-1", team_id="team-1234", budget_id="bud-1")],
|
||||
)
|
||||
|
||||
await bulk_update_team_members(
|
||||
team_id="team-1234",
|
||||
data=BulkTeamMemberUpdateRequest(
|
||||
user_ids=["user-1"],
|
||||
update_fields=TeamMemberBulkUpdateFields(budget_duration=None),
|
||||
),
|
||||
http_request=_bulk_request(),
|
||||
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},
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_update_admin_role_requires_premium(monkeypatch):
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", object())
|
||||
monkeypatch.setattr(proxy_server, "premium_user", False)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await bulk_update_team_members(
|
||||
team_id="team-1234",
|
||||
data=BulkTeamMemberUpdateRequest(
|
||||
user_ids=["user-1"],
|
||||
update_fields=TeamMemberBulkUpdateFields(role="admin"),
|
||||
),
|
||||
http_request=_bulk_request(),
|
||||
user_api_key_dict=_ADMIN,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "premium feature" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
def test_bulk_team_member_update_requires_exactly_one_member_selector():
|
||||
with pytest.raises(ValueError, match="either user_ids or all_members_in_team"):
|
||||
BulkTeamMemberUpdateRequest(
|
||||
team_id="team-1234",
|
||||
user_ids=["user-1"],
|
||||
all_members_in_team=True,
|
||||
update_fields=TeamMemberBulkUpdateFields(tpm_limit=42),
|
||||
|
|
|
|||
|
|
@ -21,10 +21,9 @@ export const teamMemberBulkUpdateCall = async (
|
|||
userIds: string[],
|
||||
updateFields: TeamMemberBulkUpdateFields,
|
||||
) =>
|
||||
apiClient.post(`/team/member/bulk_update`, {
|
||||
apiClient.patch(`/v2/team/${encodeURIComponent(teamId)}/members`, {
|
||||
accessToken,
|
||||
body: {
|
||||
team_id: teamId,
|
||||
user_ids: userIds,
|
||||
update_fields: updateFields,
|
||||
},
|
||||
|
|
|
|||
104
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
104
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -13434,23 +13434,6 @@ export interface paths {
|
|||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/team/member/bulk_update": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
get?: never;
|
||||
put?: never;
|
||||
/** Bulk Update Team Members */
|
||||
post: operations["bulk_update_team_members_team_member_bulk_update_post"];
|
||||
delete?: never;
|
||||
options?: never;
|
||||
head?: never;
|
||||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/team/member_add": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
@ -18959,6 +18942,23 @@ export interface paths {
|
|||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/v2/team/{team_id}/members": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
get?: never;
|
||||
put?: never;
|
||||
post?: never;
|
||||
delete?: never;
|
||||
options?: never;
|
||||
head?: never;
|
||||
/** Bulk Update Team Members */
|
||||
patch: operations["bulk_update_team_members_v2_team__team_id__members_patch"];
|
||||
trace?: never;
|
||||
};
|
||||
"/v2/user/info": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
@ -21500,8 +21500,6 @@ export interface components {
|
|||
* @default false
|
||||
*/
|
||||
all_members_in_team: boolean;
|
||||
/** Team Id */
|
||||
team_id: string;
|
||||
update_fields: components["schemas"]["TeamMemberBulkUpdateFields"];
|
||||
/** User Ids */
|
||||
user_ids?: string[] | null;
|
||||
|
|
@ -50209,39 +50207,6 @@ export interface operations {
|
|||
};
|
||||
};
|
||||
};
|
||||
bulk_update_team_members_team_member_bulk_update_post: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody: {
|
||||
content: {
|
||||
"application/json": components["schemas"]["BulkTeamMemberUpdateRequest"];
|
||||
};
|
||||
};
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["BulkTeamMemberUpdateResponse"];
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
team_member_add_team_member_add_post: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
@ -57622,6 +57587,41 @@ export interface operations {
|
|||
};
|
||||
};
|
||||
};
|
||||
bulk_update_team_members_v2_team__team_id__members_patch: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path: {
|
||||
team_id: string;
|
||||
};
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody: {
|
||||
content: {
|
||||
"application/json": components["schemas"]["BulkTeamMemberUpdateRequest"];
|
||||
};
|
||||
};
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["BulkTeamMemberUpdateResponse"];
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
user_info_v2_v2_user_info_get: {
|
||||
parameters: {
|
||||
query?: {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue