diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 5e3ea4b7dcb..44550c80aa9 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -600,6 +600,7 @@ class LiteLLMRoutes(enum.Enum): "/team/permissions_list", "/team/permissions_update", "/team/permissions_bulk_update", + "/v2/team/{team_id}/members", "/team/daily/activity", # model "/model/new", @@ -741,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", @@ -1011,10 +1013,10 @@ class LiteLLM_ObjectPermissionBase(LiteLLMPydanticObjectBase): mcp_tool_search_enabled: Optional[bool] = None +from litellm.models.team import BudgetLimitEntry as BudgetLimitEntry # noqa: E402 from litellm.types.object_permission import ( # noqa: E402 ObjectPermissionDict as ObjectPermissionDict, ) -from litellm.models.team import BudgetLimitEntry as BudgetLimitEntry # noqa: E402 class GenerateRequestBase(LiteLLMPydanticObjectBase): @@ -3715,6 +3717,47 @@ class TeamMemberUpdateResponse(MemberUpdateResponse): allowed_models: Optional[List[str]] = None +class TeamMemberBulkUpdateFields(LiteLLMPydanticObjectBase): + max_budget_in_team: float | None = None + role: Literal["admin", "user"] | None = None + tpm_limit: int | None = None + rpm_limit: int | None = None + budget_duration: str | None = None + allowed_models: list[str] | None = None + + @model_validator(mode="after") + def require_at_least_one_field(self): + if not self.model_fields_set: + raise ValueError("update_fields must specify at least one field to update") + return self + + +class BulkTeamMemberUpdateRequest(LiteLLMPydanticObjectBase): + user_ids: list[str] | None = None + all_members_in_team: bool = False + update_fields: TeamMemberBulkUpdateFields + + @model_validator(mode="after") + def validate_selection(self): + has_user_ids = self.user_ids is not None and len(self.user_ids) > 0 + if has_user_ids and self.all_members_in_team: + raise ValueError("Provide either user_ids or all_members_in_team=True, not both") + if not has_user_ids and not self.all_members_in_team: + raise ValueError("Must provide either user_ids (non-empty) or all_members_in_team=True") + return self + + +class FailedTeamMemberUpdate(MemberUpdateResponse): + failed_reason: str + + +class BulkTeamMemberUpdateResponse(LiteLLMPydanticObjectBase): + team_id: str + total_requested: int + successful_updates: list[TeamMemberUpdateResponse] + failed_updates: list[FailedTeamMemberUpdate] + + class TeamModelAddRequest(BaseModel): """Request to add models to a team""" diff --git a/litellm/proxy/management_endpoints/common_utils.py b/litellm/proxy/management_endpoints/common_utils.py index 8162babef40..b20f494d8a1 100644 --- a/litellm/proxy/management_endpoints/common_utils.py +++ b/litellm/proxy/management_endpoints/common_utils.py @@ -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 diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index d267a3cac69..b4f9f27af65 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -14,7 +14,7 @@ import json import math import traceback from datetime import datetime, timezone -from typing import Annotated, Any, Dict, List, Optional, Tuple, Union, cast +from typing import Annotated, Any, Dict, List, Optional, Sequence, Tuple, Union, cast import fastapi from fastapi import APIRouter, Depends, Header, HTTPException, Request, status @@ -28,8 +28,11 @@ from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.proxy._types import ( UI_TEAM_ID, BlockTeamRequest, + BulkTeamMemberUpdateRequest, + BulkTeamMemberUpdateResponse, CommonProxyErrors, DeleteTeamRequest, + FailedTeamMemberUpdate, LiteLLM_AuditLogs, LiteLLM_DeletedTeamTable, LiteLLM_ManagementEndpoint_MetadataFields, @@ -57,6 +60,7 @@ from litellm.proxy._types import ( TeamInfoResponseObjectTeamTable, TeamListResponseObject, TeamMemberAddRequest, + TeamMemberBulkUpdateFields, TeamMemberDeleteRequest, TeamMemberUpdateRequest, TeamMemberUpdateResponse, @@ -93,6 +97,15 @@ from litellm.proxy.management_endpoints.organization_endpoints import ( from litellm.proxy.management_endpoints.tag_management_endpoints import ( get_daily_activity, ) +from litellm.proxy.management_endpoints.team_member_budget_writes import ( + BudgetFieldSnapshot, + MembershipBudgetSnapshot, + PrismaTeamMemberBudgetDb, + apply_member_budget_write_plan, + budget_snapshot_from_row, + invalidate_team_membership_caches, + plan_member_budget_writes, +) from litellm.proxy.management_helpers.object_permission_utils import ( _set_object_permission, enforce_all_proxy_mcp_servers_grant_is_admin_only, @@ -2814,7 +2827,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.""" @@ -2826,6 +2839,37 @@ def _build_member_budget_patch(data: TeamMemberUpdateRequest) -> Dict[str, Any]: } +async def _load_member_budget_snapshots( + *, + prisma_client: PrismaClient, + team_id: str, + user_ids: Sequence[str], + team_default_budget_id: str | None, +) -> tuple[tuple[MembershipBudgetSnapshot, ...], dict[str, BudgetFieldSnapshot]]: + if not user_ids: + return (), {} + + raw_memberships = await TeamMembershipRepository(prisma_client).table.find_many( + where={"team_id": team_id, "user_id": {"in": list(user_ids)}} + ) + budget_id_by_user = {membership.user_id: membership.budget_id for membership in raw_memberships} + membership_snapshots = tuple( + MembershipBudgetSnapshot(user_id=user_id, budget_id=budget_id_by_user.get(user_id)) for user_id in user_ids + ) + + budget_ids = {budget_id for budget_id in budget_id_by_user.values() if budget_id is not None} + if team_default_budget_id is not None: + budget_ids.add(team_default_budget_id) + if not budget_ids: + return membership_snapshots, {} + + raw_budgets = await BudgetRepository(prisma_client).table.find_many(where={"budget_id": {"in": list(budget_ids)}}) + budgets_by_id = { + budget.budget_id: budget_snapshot_from_row(budget.budget_id, budget.model_dump()) for budget in raw_budgets + } + return membership_snapshots, budgets_by_id + + 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 @@ -2922,6 +2966,23 @@ async def team_member_update( user_api_key_dict=user_api_key_dict, ) + return await _apply_team_member_update( + data=data, + returned_team_info=returned_team_info, + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + ) + + +async def _apply_team_member_update( + data: TeamMemberUpdateRequest, + returned_team_info: TeamInfoResponseObject, + prisma_client: PrismaClient, + user_api_key_dict: UserAPIKeyAuth, +) -> TeamMemberUpdateResponse: + """Apply a single member update against already-fetched team info so a bulk + caller can resolve team_info once and reuse it for every member, instead of + re-scanning the team, its keys, and all memberships per member.""" team_table = returned_team_info["team_info"] ## get user id @@ -2929,7 +2990,7 @@ async def team_member_update( if data.user_id is not None: received_user_id = data.user_id elif data.user_email is not None: - for member in returned_team_info["team_info"].members_with_roles: + for member in team_table.members_with_roles: if member.user_email is not None and member.user_email == data.user_email: received_user_id = member.user_id break @@ -3003,6 +3064,171 @@ async def team_member_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)], +): + 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"}) + + if data.update_fields.role == "admin" and not premium_user: + raise HTTPException( + status_code=400, + detail="Assigning team admins is a premium feature. You must be a LiteLLM Enterprise user to use this feature. If you have a license please set `LITELLM_LICENSE` in your env. Get a 7 day trial key here: https://www.litellm.ai/#trial. Pricing: https://www.litellm.ai/#pricing", + ) + + _validate_budget_duration(data.update_fields.budget_duration) + + 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(team_id)}, + ) + existing_team = LiteLLM_TeamTable(**existing_team_row.model_dump()) + if ( + user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value + and not _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=existing_team) + and not await _is_user_org_admin_for_team(user_api_key_dict=user_api_key_dict, team_obj=existing_team) + ): + raise HTTPException( + status_code=403, + detail={ + "error": "Call not allowed. User not proxy admin OR team admin. route={}, team_id={}".format( + "/v2/team/{team_id}/members", team_id + ) + }, + ) + + if data.all_members_in_team: + 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 [])) + + max_batch_size = 500 + if len(user_ids) > max_batch_size: + raise HTTPException( + status_code=400, + detail={ + "error": "Maximum {} team members can be updated at once. Found {} user_ids.".format( + max_batch_size, len(user_ids) + ) + }, + ) + + 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) + + 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 [] + membership_snapshots, budgets_by_id = await _load_member_budget_snapshots( + prisma_client=prisma_client, + team_id=team_id, + user_ids=budget_target_user_ids, + team_default_budget_id=default_budget_id, + ) + budget_plan = plan_member_budget_writes( + memberships=membership_snapshots, + budgets_by_id=budgets_by_id, + budget_patch=budget_patch, + team_default_budget_id=default_budget_id, + actor_user_id=user_api_key_dict.user_id or "", + ) + + new_role = data.update_fields.role + valid_user_id_set = frozenset(valid_user_ids) + members_with_roles_json = ( + json.dumps( + [ + Member( + user_id=member.user_id, + role=new_role, + user_email=member.user_email, + ).model_dump() + if member.user_id in valid_user_id_set + else member.model_dump() + for member in existing_team.members_with_roles + ] + ) + if new_role is not None + else None + ) + + has_budget_writes = len(budget_plan.writes) > 0 + has_role_writes = members_with_roles_json is not None + refreshed_team_row = None + if has_budget_writes or has_role_writes: + async with prisma_client.db.tx() as tx: + refreshed_team_row = await apply_member_budget_write_plan( + db=PrismaTeamMemberBudgetDb(tx), + team_id=team_id, + plan=budget_plan, + members_with_roles_json=members_with_roles_json, + ) + + 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, + ) + + if has_budget_writes or has_role_writes: + await invalidate_team_membership_caches( + user_ids=valid_user_ids, + team_id=team_id, + user_api_key_cache=user_api_key_cache, + ) + + 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=team_id, + total_requested=len(user_ids), + successful_updates=successful_updates, + failed_updates=failed_updates, + ) + + def _create_results_from_response( members: List[Member], response: TeamAddMemberResponse, diff --git a/litellm/proxy/management_endpoints/team_member_budget_writes.py b/litellm/proxy/management_endpoints/team_member_budget_writes.py new file mode 100644 index 00000000000..8521f042702 --- /dev/null +++ b/litellm/proxy/management_endpoints/team_member_budget_writes.py @@ -0,0 +1,246 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Callable, Mapping, Protocol, Sequence +from uuid import uuid4 + +from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time +from litellm.proxy.management_endpoints.common_utils import ( + _TEAM_MEMBER_BUDGET_LIMIT_FIELDS, + _budget_patch_to_write_data, + _has_meaningful_budget_limit, + _is_set_budget_value, +) + + +@dataclass(frozen=True, slots=True) +class MembershipBudgetSnapshot: + user_id: str + budget_id: str | None + + +@dataclass(frozen=True, slots=True) +class BudgetFieldSnapshot: + budget_id: str + fields: Mapping[str, Any] + + +@dataclass(frozen=True, slots=True) +class DisconnectBudget: + user_id: str + + +@dataclass(frozen=True, slots=True) +class UpdateBudget: + budget_id: str + write_data: Mapping[str, Any] + + +@dataclass(frozen=True, slots=True) +class CreateAndAttachBudget: + user_id: str + budget_id: str + create_data: Mapping[str, Any] + + +MemberBudgetWrite = DisconnectBudget | UpdateBudget | CreateAndAttachBudget + + +@dataclass(frozen=True, slots=True) +class MemberBudgetWritePlan: + writes: tuple[MemberBudgetWrite, ...] + + +def _limit_fields_from_row(row: Mapping[str, Any]) -> dict[str, Any]: + return {field: row.get(field) for field in _TEAM_MEMBER_BUDGET_LIMIT_FIELDS if field in row} + + +def budget_snapshot_from_row(budget_id: str, row: Mapping[str, Any]) -> BudgetFieldSnapshot: + return BudgetFieldSnapshot(budget_id=budget_id, fields=_limit_fields_from_row(row)) + + +def _build_create_data( + *, + actor_user_id: str, + write_data: Mapping[str, Any], + shared_default_fields: Mapping[str, Any] | None, +) -> dict[str, Any]: + create_data: dict[str, Any] = { + "created_by": actor_user_id, + "updated_by": actor_user_id, + } + if shared_default_fields is not None: + for field in _TEAM_MEMBER_BUDGET_LIMIT_FIELDS: + value = shared_default_fields.get(field) + if _is_set_budget_value(value): + create_data[field] = value + create_data.update(write_data) + if create_data.get("budget_duration") is not None: + create_data["budget_reset_at"] = get_budget_reset_time(budget_duration=create_data["budget_duration"]) + else: + create_data.pop("budget_reset_at", None) + return create_data + + +def plan_member_budget_writes( + *, + memberships: Sequence[MembershipBudgetSnapshot], + budgets_by_id: Mapping[str, BudgetFieldSnapshot], + budget_patch: Mapping[str, Any], + team_default_budget_id: str | None, + actor_user_id: str, + new_budget_id_factory: Callable[[], str] = lambda: str(uuid4()), +) -> MemberBudgetWritePlan: + if not budget_patch: + return MemberBudgetWritePlan(writes=()) + + write_data = _budget_patch_to_write_data(dict(budget_patch)) + writes: list[MemberBudgetWrite] = [] + + for membership in memberships: + existing_budget_id = membership.budget_id + is_shared_default = ( + existing_budget_id is not None + and team_default_budget_id is not None + and existing_budget_id == team_default_budget_id + ) + + if existing_budget_id is not None and not is_shared_default: + existing = budgets_by_id.get(existing_budget_id) + merged = dict(existing.fields) if existing is not None else {} + merged.update(write_data) + if not _has_meaningful_budget_limit(merged): + writes.append(DisconnectBudget(user_id=membership.user_id)) + continue + writes.append( + UpdateBudget( + budget_id=existing_budget_id, + write_data={"updated_by": actor_user_id, **write_data}, + ) + ) + continue + + shared_fields = None + if is_shared_default and existing_budget_id is not None: + default_snap = budgets_by_id.get(existing_budget_id) + if default_snap is not None: + shared_fields = default_snap.fields + + create_data = _build_create_data( + actor_user_id=actor_user_id, + write_data=write_data, + shared_default_fields=shared_fields, + ) + if not _has_meaningful_budget_limit(create_data): + if existing_budget_id is not None: + writes.append(DisconnectBudget(user_id=membership.user_id)) + continue + + new_budget_id = new_budget_id_factory() + writes.append( + CreateAndAttachBudget( + user_id=membership.user_id, + budget_id=new_budget_id, + create_data={**create_data, "budget_id": new_budget_id}, + ) + ) + + return MemberBudgetWritePlan(writes=tuple(writes)) + + +class TeamMemberBudgetDb(Protocol): + async def create_budgets(self, rows: Sequence[Mapping[str, Any]]) -> None: ... + + async def update_budget(self, budget_id: str, data: Mapping[str, Any]) -> None: ... + + async def disconnect_membership_budget(self, *, team_id: str, user_id: str) -> None: ... + + async def attach_membership_budget(self, *, team_id: str, user_id: str, budget_id: str) -> None: ... + + async def update_team_members_with_roles(self, *, team_id: str, members_with_roles_json: str) -> Any: ... + + +async def apply_member_budget_write_plan( + *, + db: TeamMemberBudgetDb, + team_id: str, + plan: MemberBudgetWritePlan, + members_with_roles_json: str | None, +) -> Any | None: + creates = tuple(write for write in plan.writes if isinstance(write, CreateAndAttachBudget)) + updates = tuple(write for write in plan.writes if isinstance(write, UpdateBudget)) + disconnects = tuple(write for write in plan.writes if isinstance(write, DisconnectBudget)) + + if creates: + await db.create_budgets(tuple(write.create_data for write in creates)) + for write in creates: + await db.attach_membership_budget( + team_id=team_id, + user_id=write.user_id, + budget_id=write.budget_id, + ) + + for write in updates: + await db.update_budget(write.budget_id, write.write_data) + + for write in disconnects: + await db.disconnect_membership_budget(team_id=team_id, user_id=write.user_id) + + if members_with_roles_json is None: + return None + return await db.update_team_members_with_roles( + team_id=team_id, + members_with_roles_json=members_with_roles_json, + ) + + +class PrismaTeamMemberBudgetDb: + def __init__(self, tx: Any): + self._tx = tx + + async def create_budgets(self, rows: Sequence[Mapping[str, Any]]) -> None: + if not rows: + return + await self._tx.litellm_budgettable.create_many(data=list(rows)) + + async def update_budget(self, budget_id: str, data: Mapping[str, Any]) -> None: + await self._tx.litellm_budgettable.update(where={"budget_id": budget_id}, data=dict(data)) + + async def disconnect_membership_budget(self, *, team_id: str, user_id: str) -> None: + await self._tx.litellm_teammembership.update( + where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, + data={"litellm_budget_table": {"disconnect": True}}, + ) + + async def attach_membership_budget(self, *, team_id: str, user_id: str, budget_id: str) -> None: + await self._tx.litellm_teammembership.upsert( + where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, + data={ + "create": { + "user_id": user_id, + "team_id": team_id, + "litellm_budget_table": {"connect": {"budget_id": budget_id}}, + }, + "update": { + "litellm_budget_table": {"connect": {"budget_id": budget_id}}, + }, + }, + ) + + async def update_team_members_with_roles(self, *, team_id: str, members_with_roles_json: str) -> Any: + return await self._tx.litellm_teamtable.update( + where={"team_id": team_id}, + data={"members_with_roles": members_with_roles_json}, + include={"object_permission": True}, + ) + + +async def invalidate_team_membership_caches( + *, + user_ids: Sequence[str], + team_id: str, + user_api_key_cache: Any, +) -> None: + for user_id in user_ids: + await user_api_key_cache.async_delete_cache(f"team_membership:{user_id}:{team_id}") + await user_api_key_cache.async_delete_cache(f"{team_id}_{user_id}") diff --git a/tests/proxy_behavior/management/test_team_bulk_member_update.py b/tests/proxy_behavior/management/test_team_bulk_member_update.py new file mode 100644 index 00000000000..a14faf15026 --- /dev/null +++ b/tests/proxy_behavior/management/test_team_bulk_member_update.py @@ -0,0 +1,293 @@ +import pytest + +from .actors import Actor +from .conftest import create_scratch_team + +pytestmark = pytest.mark.asyncio(loop_scope="session") + +ROUTE = "/v2/team/{team_id}/members" + + +def _role_of(row, user_id: str): + for m in row.members_with_roles or []: + if m["user_id"] == user_id: + return m["role"] + return None + + +# PATCH /v2/team/{team_id}/members — same coarse route gate and admin checks as +# POST /team/member_update, so the actor x team-shape matrix mirrors it: the +# scratch team is raw-seeded with a "user"-role member and each scenario bulk +# promotes it to "admin". PROXY_ADMIN, the team's team admin, or an org admin of +# the team's org may update; everyone else is 403. (The harness forces +# premium_user, so the admin-role premium gate never decides the outcome.) +_MATRIX = [ + ("alpha/proxy_admin", Actor.PROXY_ADMIN, "alpha", 200), + ("alpha/org_admin", Actor.ORG_ADMIN, "alpha", 200), + ("alpha/team_admin", Actor.TEAM_ADMIN, "alpha", 200), + ("alpha/internal_user", Actor.INTERNAL_USER, "alpha", 403), + ("alpha/owner", Actor.OWNER, "alpha", 403), + ("alpha/unrelated_same_org", Actor.UNRELATED_SAME_ORG, "alpha", 403), + ("alpha/cross_org_user", Actor.CROSS_ORG_USER, "alpha", 403), + ("alpha/service_account", Actor.SERVICE_ACCOUNT, "alpha", 403), + ("alpha/org_b_admin", Actor.ORG_B_ADMIN, "alpha", 403), + ("beta/proxy_admin", Actor.PROXY_ADMIN, "beta", 200), + ("beta/org_admin", Actor.ORG_ADMIN, "beta", 403), + ("beta/team_admin", Actor.TEAM_ADMIN, "beta", 403), + ("beta/internal_user", Actor.INTERNAL_USER, "beta", 403), + ("beta/owner", Actor.OWNER, "beta", 403), + ("beta/unrelated_same_org", Actor.UNRELATED_SAME_ORG, "beta", 403), + ("beta/cross_org_user", Actor.CROSS_ORG_USER, "beta", 403), + ("beta/service_account", Actor.SERVICE_ACCOUNT, "beta", 403), + ("beta/org_b_admin", Actor.ORG_B_ADMIN, "beta", 200), +] + + +async def _seed_target(prisma, world, shape: str, team_id: str, member_id: str) -> None: + if shape == "alpha": + await create_scratch_team( + prisma, + team_id, + organization_id=world.org_a_id, + admin_user_ids=[world.keys[Actor.TEAM_ADMIN].user_id], + member_user_ids=[member_id], + ) + elif shape == "beta": + await create_scratch_team( + prisma, + team_id, + organization_id=world.org_b_id, + member_user_ids=[member_id], + ) + else: # pragma: no cover - guard + pytest.fail(f"unknown shape={shape}") + + +@pytest.mark.parametrize( + "actor,shape,expected_status", + [(a, sh, s) for (_id, a, sh, s) in _MATRIX], + ids=[s[0] for s in _MATRIX], +) +async def test_bulk_member_update_authz_matrix( + actor: Actor, + shape: str, + expected_status: int, + proxy_client, + prisma, + scratch, + world, +): + member_id = scratch.tag("member") + await _seed_target(prisma, world, shape, scratch.prefix, member_id) + caller = world.keys[actor] + + resp = await proxy_client.patch( + ROUTE.format(team_id=scratch.prefix), + headers={"Authorization": f"Bearer {caller.cleartext}"}, + json={"user_ids": [member_id], "update_fields": {"role": "admin"}}, + ) + assert ( + resp.status_code == expected_status + ), f"{actor.value} {shape}: {resp.status_code} {resp.text}" + + row = await prisma.db.litellm_teamtable.find_unique( + where={"team_id": scratch.prefix} + ) + assert row is not None + if expected_status == 200: + assert _role_of(row, member_id) == "admin" + else: + assert _role_of(row, member_id) == "user", "denied but role changed" + + +async def test_bulk_member_update_reports_non_members_without_failing_batch( + proxy_client, prisma, scratch, world +): + """Valid members are all promoted in one write; a user_id that is not a + member of the team lands in failed_updates while the rest still succeed.""" + m1 = scratch.tag("m1") + m2 = scratch.tag("m2") + stranger = scratch.tag("stranger") + await create_scratch_team( + prisma, + scratch.prefix, + organization_id=world.org_a_id, + member_user_ids=[m1, m2], + ) + + resp = await proxy_client.patch( + ROUTE.format(team_id=scratch.prefix), + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={ + "user_ids": [m1, m2, stranger], + "update_fields": {"role": "admin"}, + }, + ) + assert resp.status_code == 200, resp.text + body = resp.json() + assert body["total_requested"] == 3 + assert {u["user_id"] for u in body["successful_updates"]} == {m1, m2} + assert [f["user_id"] for f in body["failed_updates"]] == [stranger] + + row = await prisma.db.litellm_teamtable.find_unique( + where={"team_id": scratch.prefix} + ) + assert row is not None + assert _role_of(row, m1) == "admin" + assert _role_of(row, m2) == "admin" + + +async def test_bulk_member_update_all_members_in_team_promotes_everyone( + proxy_client, prisma, scratch, world +): + """all_members_in_team=True applies the patch to every current member.""" + m1 = scratch.tag("m1") + m2 = scratch.tag("m2") + await create_scratch_team( + prisma, + scratch.prefix, + organization_id=world.org_a_id, + member_user_ids=[m1, m2], + ) + + resp = await proxy_client.patch( + ROUTE.format(team_id=scratch.prefix), + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={"all_members_in_team": True, "update_fields": {"role": "admin"}}, + ) + assert resp.status_code == 200, resp.text + assert resp.json()["total_requested"] == 2 + + row = await prisma.db.litellm_teamtable.find_unique( + where={"team_id": scratch.prefix} + ) + assert row is not None + assert _role_of(row, m1) == "admin" + assert _role_of(row, m2) == "admin" + + +async def test_bulk_member_update_writes_member_budget_row( + proxy_client, prisma, scratch, world +): + """A limit patch on a member with no budget row creates one membership + + budget carrying the patched value; this is the set-based write path.""" + member_id = scratch.tag("member") + await create_scratch_team( + prisma, + scratch.prefix, + organization_id=world.org_a_id, + member_user_ids=[member_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": 4242}}, + ) + 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 and membership.litellm_budget_table is not None + 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 +): + """More than the 500-member cap of user_ids is rejected 400.""" + await create_scratch_team(prisma, scratch.prefix, organization_id=world.org_a_id) + user_ids = [f"{scratch.prefix}-u{i}" for i in range(501)] + resp = await proxy_client.patch( + ROUTE.format(team_id=scratch.prefix), + headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, + json={"user_ids": user_ids, "update_fields": {"role": "user"}}, + ) + assert resp.status_code == 400, resp.text diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index a6d4dc63697..43bd5b7bed3 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -2882,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(): """ diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_member_budget_writes.py b/tests/test_litellm/proxy/management_endpoints/test_team_member_budget_writes.py new file mode 100644 index 00000000000..13da65f9c62 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/test_team_member_budget_writes.py @@ -0,0 +1,134 @@ +import pytest + +from litellm.proxy.management_endpoints.team_member_budget_writes import ( + BudgetFieldSnapshot, + CreateAndAttachBudget, + DisconnectBudget, + MembershipBudgetSnapshot, + UpdateBudget, + apply_member_budget_write_plan, + plan_member_budget_writes, +) + + +def test_plan_updates_private_budget_in_place(): + plan = plan_member_budget_writes( + memberships=(MembershipBudgetSnapshot(user_id="u1", budget_id="b1"),), + budgets_by_id={"b1": BudgetFieldSnapshot(budget_id="b1", fields={"tpm_limit": 1})}, + budget_patch={"tpm_limit": 9}, + team_default_budget_id=None, + actor_user_id="admin", + ) + + assert len(plan.writes) == 1 + write = plan.writes[0] + assert isinstance(write, UpdateBudget) + assert write.budget_id == "b1" + assert write.write_data["tpm_limit"] == 9 + assert write.write_data["updated_by"] == "admin" + + +def test_plan_disconnects_when_last_private_limit_cleared(): + plan = plan_member_budget_writes( + memberships=(MembershipBudgetSnapshot(user_id="u1", budget_id="b1"),), + budgets_by_id={"b1": BudgetFieldSnapshot(budget_id="b1", fields={"tpm_limit": 1})}, + budget_patch={"tpm_limit": None}, + team_default_budget_id=None, + actor_user_id="admin", + ) + + assert plan.writes == (DisconnectBudget(user_id="u1"),) + + +def test_plan_clone_on_write_for_shared_default_only(): + plan = plan_member_budget_writes( + memberships=( + MembershipBudgetSnapshot(user_id="on-default", budget_id="default"), + MembershipBudgetSnapshot(user_id="no-budget", budget_id=None), + ), + budgets_by_id={ + "default": BudgetFieldSnapshot(budget_id="default", fields={"max_budget": 50.0}), + }, + budget_patch={"tpm_limit": 3}, + team_default_budget_id="default", + actor_user_id="admin", + new_budget_id_factory=iter(["nb-1", "nb-2"]).__next__, + ) + + assert len(plan.writes) == 2 + first, second = plan.writes + assert isinstance(first, CreateAndAttachBudget) + assert first.user_id == "on-default" + assert first.budget_id == "nb-1" + assert first.create_data["max_budget"] == 50.0 + assert first.create_data["tpm_limit"] == 3 + assert isinstance(second, CreateAndAttachBudget) + assert second.user_id == "no-budget" + assert second.budget_id == "nb-2" + assert "max_budget" not in second.create_data + assert second.create_data["tpm_limit"] == 3 + + +def test_plan_empty_patch_is_noop(): + plan = plan_member_budget_writes( + memberships=(MembershipBudgetSnapshot(user_id="u1", budget_id="b1"),), + budgets_by_id={"b1": BudgetFieldSnapshot(budget_id="b1", fields={"tpm_limit": 1})}, + budget_patch={}, + team_default_budget_id=None, + actor_user_id="admin", + ) + assert plan.writes == () + + +class _RecordingDb: + def __init__(self): + self.created = [] + self.updated = [] + self.disconnected = [] + self.attached = [] + self.team_role_json = None + + async def create_budgets(self, rows): + self.created.extend(rows) + + async def update_budget(self, budget_id, data): + self.updated.append((budget_id, dict(data))) + + async def disconnect_membership_budget(self, *, team_id, user_id): + self.disconnected.append((team_id, user_id)) + + async def attach_membership_budget(self, *, team_id, user_id, budget_id): + self.attached.append((team_id, user_id, budget_id)) + + async def update_team_members_with_roles(self, *, team_id, members_with_roles_json): + self.team_role_json = members_with_roles_json + return {"team_id": team_id} + + +@pytest.mark.asyncio +async def test_apply_runs_set_oriented_write_groups(): + plan = plan_member_budget_writes( + memberships=( + MembershipBudgetSnapshot(user_id="u1", budget_id="b1"), + MembershipBudgetSnapshot(user_id="u2", budget_id=None), + ), + budgets_by_id={"b1": BudgetFieldSnapshot(budget_id="b1", fields={"tpm_limit": 1})}, + budget_patch={"tpm_limit": 9}, + team_default_budget_id=None, + actor_user_id="admin", + new_budget_id_factory=lambda: "created-1", + ) + db = _RecordingDb() + team_row = await apply_member_budget_write_plan( + db=db, + team_id="team-1", + plan=plan, + members_with_roles_json='[{"user_id":"u1","role":"user"}]', + ) + + assert team_row == {"team_id": "team-1"} + assert db.created[0]["budget_id"] == "created-1" + assert db.attached == [("team-1", "u2", "created-1")] + assert db.updated[0][0] == "b1" + assert db.updated[0][1]["tpm_limit"] == 9 + assert db.team_role_json == '[{"user_id":"u1","role":"user"}]' diff --git a/tests/test_litellm/proxy/test_team_member_update.py b/tests/test_litellm/proxy/test_team_member_update.py index 352c68d491c..e05f289f0e1 100644 --- a/tests/test_litellm/proxy/test_team_member_update.py +++ b/tests/test_litellm/proxy/test_team_member_update.py @@ -1,3 +1,4 @@ +import json import types from unittest.mock import AsyncMock, MagicMock @@ -8,13 +9,23 @@ from starlette.requests import Request 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_TeamMembership, LiteLLM_TeamTable, LitellmUserRoles, Member, + TeamMemberBulkUpdateFields, TeamMemberUpdateRequest, UserAPIKeyAuth, ) -from litellm.proxy.management_endpoints.team_endpoints import team_member_update +from litellm.proxy.management_endpoints.team_endpoints import ( + bulk_update_team_members, + team_member_update, +) +from litellm.proxy.management_endpoints.team_member_budget_writes import ( + BudgetFieldSnapshot, + MembershipBudgetSnapshot, +) @pytest.mark.asyncio @@ -81,9 +92,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")], } ), ) @@ -93,9 +102,7 @@ def happy_path_upsert(monkeypatch): def _member_update_request(**overrides): - data = TeamMemberUpdateRequest( - team_id="team-1234", user_id="user-1", role="user", **overrides - ) + data = TeamMemberUpdateRequest(team_id="team-1234", user_id="user-1", role="user", **overrides) request = Request({"type": "http", "method": "POST", "path": "/team/member_update"}) auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN.value, user_id="admin") return data, request, auth @@ -105,9 +112,7 @@ def _member_update_request(**overrides): async def test_team_member_update_sends_provided_fields_as_patch(happy_path_upsert): """Fields the request sets must reach _upsert_budget_and_membership as a budget patch, otherwise the member budget is never written/reset.""" - data, request, auth = _member_update_request( - max_budget_in_team=10.0, budget_duration="30d" - ) + data, request, auth = _member_update_request(max_budget_in_team=10.0, budget_duration="30d") response = await team_member_update(data, request, auth) @@ -127,9 +132,7 @@ async def test_team_member_update_explicit_null_clears_field(happy_path_upsert): await team_member_update(data, request, auth) - assert happy_path_upsert.await_args.kwargs["budget_patch"] == { - "budget_duration": None - } + assert happy_path_upsert.await_args.kwargs["budget_patch"] == {"budget_duration": None} @pytest.mark.asyncio @@ -153,9 +156,7 @@ async def test_team_member_update_omits_unset_fields_from_patch(happy_path_upser ], ) @pytest.mark.asyncio -async def test_team_member_update_rejects_invalid_budget_duration( - monkeypatch, bad_duration -): +async def test_team_member_update_rejects_invalid_budget_duration(monkeypatch, bad_duration): """An invalid budget_duration must be rejected with a 400 before any DB write, so it can never be persisted and later break the budget reset job.""" monkeypatch.setattr(proxy_server, "prisma_client", object()) @@ -178,3 +179,344 @@ async def test_team_member_update_rejects_invalid_budget_duration( assert exc_info.value.status_code == 400 assert "budget_duration" in str(exc_info.value.detail) upsert_mock.assert_not_called() + + +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) + self.litellm_budgettable = types.SimpleNamespace( + create_many=self._create_many, + update=self._budget_update, + ) + self.litellm_teammembership = types.SimpleNamespace( + update=self._membership_update, + upsert=self._membership_upsert, + ) + + async def _team_update(self, **kwargs): + self._recorder.team_updates.append(kwargs) + return self._team_row + + async def _create_many(self, data): + self._recorder.budget_creates.extend(data) + return {"count": len(data)} + + async def _budget_update(self, **kwargs): + self._recorder.budget_updates.append(kwargs) + return types.SimpleNamespace(budget_id=kwargs["where"]["budget_id"]) + + async def _membership_update(self, **kwargs): + self._recorder.membership_updates.append(kwargs) + return types.SimpleNamespace() + + async def _membership_upsert(self, **kwargs): + self._recorder.membership_upserts.append(kwargs) + return types.SimpleNamespace() + + async def __aenter__(self): + return self + + async def __aexit__(self, *args): + return False + + +class _FakeBulkDb: + def __init__(self, team_row): + self.team_row = team_row + self.team_updates: list = [] + self.budget_creates: list = [] + self.budget_updates: list = [] + self.membership_updates: list = [] + self.membership_upserts: list = [] + self.litellm_teamtable = types.SimpleNamespace(find_unique=AsyncMock(return_value=team_row)) + + def tx(self): + return _FakeTx(self.team_row, self) + + +def _bulk_setup( + monkeypatch, + team_row, + *, + membership_snapshots=(), + budgets_by_id=None, +): + db = _FakeBulkDb(team_row) + cache = types.SimpleNamespace(async_delete_cache=AsyncMock()) + monkeypatch.setattr(proxy_server, "prisma_client", types.SimpleNamespace(db=db)) + monkeypatch.setattr(proxy_server, "premium_user", False) + monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", object()) + load_mock = AsyncMock(return_value=(tuple(membership_snapshots), budgets_by_id or {})) + monkeypatch.setattr(team_endpoints, "_load_member_budget_snapshots", load_mock) + refresh_mock = AsyncMock() + monkeypatch.setattr(team_endpoints, "_refresh_cached_team", refresh_mock) + return db, load_mock, refresh_mock, cache + + +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_update_plans_budget_writes_without_per_member_inline_prisma(monkeypatch): + 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"), + Member(user_id="user-3", role="user"), + ], + ) + db, load_mock, refresh_mock, cache = _bulk_setup( + monkeypatch, + team_row, + membership_snapshots=( + MembershipBudgetSnapshot(user_id="user-1", budget_id="bud-1"), + MembershipBudgetSnapshot(user_id="user-2", budget_id="bud-2"), + ), + budgets_by_id={ + "bud-1": BudgetFieldSnapshot(budget_id="bud-1", fields={"tpm_limit": 1}), + "bud-2": BudgetFieldSnapshot(budget_id="bud-2", fields={"tpm_limit": 2}), + }, + ) + + response = await bulk_update_team_members( + team_id="team-1234", + data=BulkTeamMemberUpdateRequest( + user_ids=["user-1", "user-2", "user-1"], + update_fields=TeamMemberBulkUpdateFields(tpm_limit=42), + ), + http_request=_bulk_request(), + user_api_key_dict=_ADMIN, + ) + + load_mock.assert_awaited_once() + assert load_mock.await_args.kwargs["user_ids"] == ["user-1", "user-2"] + assert {update["where"]["budget_id"] for update in db.budget_updates} == {"bud-1", "bud-2"} + assert db.team_updates == [] + refresh_mock.assert_not_awaited() + assert cache.async_delete_cache.await_count == 4 + 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_shared_default_instead_of_mutating_it(monkeypatch): + 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"}, + ) + ids = iter(["new-bud-1", "new-bud-2"]) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.team_member_budget_writes.uuid4", + lambda: next(ids), + ) + db, _load, _refresh, _cache = _bulk_setup( + monkeypatch, + team_row, + membership_snapshots=( + MembershipBudgetSnapshot(user_id="user-1", budget_id="default-bud"), + MembershipBudgetSnapshot(user_id="user-2", budget_id=None), + ), + budgets_by_id={ + "default-bud": BudgetFieldSnapshot(budget_id="default-bud", fields={"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.budget_updates == [] + created_ids = {row["budget_id"] for row in db.budget_creates} + assert created_ids == {"new-bud-1", "new-bud-2"} + user1_create = next(row for row in db.budget_creates if row["budget_id"] == "new-bud-1") + assert user1_create["max_budget"] == 100.0 + assert user1_create["tpm_limit"] == 42 + user2_create = next(row for row in db.budget_creates if row["budget_id"] == "new-bud-2") + assert "max_budget" not in user2_create + assert user2_create["tpm_limit"] == 42 + assert len(db.membership_upserts) == 2 + + +@pytest.mark.asyncio +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=[ + 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, load_mock, refresh_mock, cache = _bulk_setup(monkeypatch, team_row) + + 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, + ) + + load_mock.assert_awaited_once() + assert load_mock.await_args.kwargs["user_ids"] == [] + 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] == [ + ("user-1", "user"), + ("user-2", "user"), + ("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 + assert cache.async_delete_cache.await_count == 4 + + +@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, load_mock, _refresh, _cache = _bulk_setup( + monkeypatch, + team_row, + membership_snapshots=( + MembershipBudgetSnapshot(user_id="user-1", budget_id="bud-1"), + MembershipBudgetSnapshot(user_id="user-2", budget_id="bud-2"), + ), + budgets_by_id={ + "bud-1": BudgetFieldSnapshot(budget_id="bud-1", fields={"tpm_limit": 1}), + "bud-2": BudgetFieldSnapshot(budget_id="bud-2", fields={"tpm_limit": 2}), + }, + ) + + 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 load_mock.await_args.kwargs["user_ids"] == ["user-1", "user-2"] + assert {update["where"]["budget_id"] for update in db.budget_updates} == {"bud-1", "bud-2"} + assert response.total_requested == 2 + + +@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, load_mock, _refresh, _cache = _bulk_setup( + monkeypatch, + team_row, + membership_snapshots=(MembershipBudgetSnapshot(user_id="user-1", budget_id="bud-1"),), + budgets_by_id={"bud-1": BudgetFieldSnapshot(budget_id="bud-1", fields={"tpm_limit": 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 load_mock.await_args.kwargs["user_ids"] == ["user-1"] + assert {update["where"]["budget_id"] for update in db.budget_updates} == {"bud-1"} + + +@pytest.mark.asyncio +async def test_bulk_update_explicit_null_duration_disconnects_private_budget(monkeypatch): + team_row = LiteLLM_TeamTable( + team_id="team-1234", + members_with_roles=[Member(user_id="user-1", role="user")], + ) + db, _load, _refresh, _cache = _bulk_setup( + monkeypatch, + team_row, + membership_snapshots=(MembershipBudgetSnapshot(user_id="user-1", budget_id="bud-1"),), + budgets_by_id={ + "bud-1": BudgetFieldSnapshot(budget_id="bud-1", fields={"budget_duration": "1d"}), + }, + ) + + 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.budget_updates == [] + assert len(db.membership_updates) == 1 + assert db.membership_updates[0]["data"] == {"litellm_budget_table": {"disconnect": True}} + + +@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( + user_ids=["user-1"], + all_members_in_team=True, + update_fields=TeamMemberBulkUpdateFields(tpm_limit=42), + ) diff --git a/ui/litellm-dashboard/src/components/common_components/MemberTable.tsx b/ui/litellm-dashboard/src/components/common_components/MemberTable.tsx index d58094642d2..9c343ea6d02 100644 --- a/ui/litellm-dashboard/src/components/common_components/MemberTable.tsx +++ b/ui/litellm-dashboard/src/components/common_components/MemberTable.tsx @@ -18,6 +18,8 @@ export interface MemberTableProps { extraColumns?: ColumnsType; showDeleteForMember?: (member: Member) => boolean; emptyText?: string; + rowSelection?: React.ComponentProps>["rowSelection"]; + extraActions?: React.ReactNode; } export default function MemberTable({ @@ -31,6 +33,8 @@ export default function MemberTable({ extraColumns = [], showDeleteForMember, emptyText, + rowSelection, + extraActions, }: MemberTableProps) { const baseColumns: ColumnsType = [ { @@ -106,16 +110,22 @@ export default function MemberTable({ record.user_id ?? record.user_email ?? JSON.stringify(record)} pagination={false} size="small" scroll={{ x: "max-content" }} locale={emptyText ? { emptyText } : undefined} /> - {onAddMember && canEdit && ( - + {(onAddMember || extraActions) && canEdit && ( + + {onAddMember && ( + + )} + {extraActions} + )} ); diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx index 74f201b40c6..faa5b5761d0 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx @@ -896,6 +896,12 @@ const TeamInfoView: React.FC = ({ setSelectedEditMember={setSelectedEditMember} setIsEditMemberModalVisible={setIsEditMemberModalVisible} setIsAddMemberModalVisible={setIsAddMemberModalVisible} + onMembersUpdated={async () => { + if (!accessToken) return; + const updatedTeamData = await teamInfoCall(accessToken, teamId); + setTeamData(updatedTeamData); + onUpdate(updatedTeamData); + }} /> ), }, diff --git a/ui/litellm-dashboard/src/components/team/TeamMemberTab.test.tsx b/ui/litellm-dashboard/src/components/team/TeamMemberTab.test.tsx index a07c57eaa30..016720cbe9e 100644 --- a/ui/litellm-dashboard/src/components/team/TeamMemberTab.test.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamMemberTab.test.tsx @@ -182,6 +182,25 @@ describe("TeamMembersComponent", () => { expect(screen.getByText("Add Member")).toBeInTheDocument(); }); + it("should show checkboxes after Select Members is clicked", async () => { + const user = userEvent.setup(); + renderWithProviders( + , + ); + + await user.click(screen.getByRole("button", { name: "Select Members" })); + + expect(screen.getByRole("button", { name: "Bulk Edit (0 selected)" })).toBeDisabled(); + expect(screen.getAllByRole("checkbox")).toHaveLength(3); + }); + it("should display dash when user email is null", () => { renderWithProviders( void; setIsEditMemberModalVisible: (visible: boolean) => void; setIsAddMemberModalVisible: (visible: boolean) => void; + onMembersUpdated?: () => Promise; } export default function TeamMemberTab({ @@ -26,7 +32,13 @@ export default function TeamMemberTab({ setSelectedEditMember, setIsEditMemberModalVisible, setIsAddMemberModalVisible, + onMembersUpdated, }: TeamMemberTabProps) { + const [isBulkUpdateVisible, setIsBulkUpdateVisible] = useState(false); + const [isBulkUpdating, setIsBulkUpdating] = useState(false); + const [selectionMode, setSelectionMode] = useState(false); + const [selectedMembers, setSelectedMembers] = useState([]); + const [bulkUpdateForm] = Form.useForm(); const formatNumber = (value: number | null): string => { if (value === null || value === undefined) return "0"; @@ -79,7 +91,7 @@ export default function TeamMemberTab({ }; const { data: uiSettingsData } = useUISettings(); - const { userId, userRole } = useAuthorized(); + const { accessToken, userId, userRole } = useAuthorized(); const disableTeamAdminDeleteTeamUser = Boolean(uiSettingsData?.values?.disable_team_admin_delete_team_user); const isUserTeamAdmin = isUserTeamAdminForSingleTeam(teamData.team_info.members_with_roles, userId || ""); const isProxyAdmin = isProxyAdminRole(userRole || ""); @@ -183,31 +195,217 @@ export default function TeamMemberTab({ }, ]; - return ( - { - const membership = teamData.team_memberships.find((tm) => tm.user_id === record.user_id); - const enhancedMember = { - ...record, - max_budget_in_team: membership?.litellm_budget_table?.max_budget || null, - tpm_limit: membership?.litellm_budget_table?.tpm_limit || null, - rpm_limit: membership?.litellm_budget_table?.rpm_limit || null, - budget_duration: membership?.litellm_budget_table?.budget_duration || null, - allowed_models: membership?.litellm_budget_table?.allowed_models || [], - }; - setSelectedEditMember(enhancedMember); - setIsEditMemberModalVisible(true); - }} - onDelete={handleMemberDelete} - onAddMember={() => setIsAddMemberModalVisible(true)} - roleColumnTitle="Team Role" - roleTooltip="This role applies only to this team and is independent from the user's proxy-level role." - extraColumns={extraColumns} - showDeleteForMember={() => - isProxyAdmin || (canEditTeam && !isUserTeamAdmin) || (isUserTeamAdmin && !disableTeamAdminDeleteTeamUser) + const handleBulkUpdate = async (values: { + apply_role?: boolean; + role?: "admin" | "user"; + apply_max_budget?: boolean; + max_budget_in_team?: number | null; + apply_budget_duration?: boolean; + budget_duration?: string | null; + apply_tpm_limit?: boolean; + tpm_limit?: number | null; + apply_rpm_limit?: boolean; + rpm_limit?: number | null; + apply_allowed_models?: boolean; + allowed_models?: string[]; + }) => { + if (!accessToken) return; + const userIds = selectedMembers.flatMap((member) => (member.user_id ? [member.user_id] : [])); + if (userIds.length === 0) { + NotificationsManager.fromBackend("Select at least one team member"); + return; + } + const updateFields: TeamMemberBulkUpdateFields = { + ...(values.apply_role ? { role: values.role } : {}), + ...(values.apply_max_budget ? { max_budget_in_team: values.max_budget_in_team ?? null } : {}), + ...(values.apply_budget_duration ? { budget_duration: values.budget_duration ?? null } : {}), + ...(values.apply_tpm_limit ? { tpm_limit: values.tpm_limit ?? null } : {}), + ...(values.apply_rpm_limit ? { rpm_limit: values.rpm_limit ?? null } : {}), + ...(values.apply_allowed_models ? { allowed_models: values.allowed_models ?? [] } : {}), + }; + if (Object.keys(updateFields).length === 0) { + NotificationsManager.fromBackend("Select at least one field to update"); + return; + } + + setIsBulkUpdating(true); + try { + const response = await teamMemberBulkUpdateCall(accessToken, teamData.team_id, userIds, updateFields); + await onMembersUpdated?.(); + setIsBulkUpdateVisible(false); + setSelectedMembers([]); + setSelectionMode(false); + bulkUpdateForm.resetFields(); + NotificationsManager.success( + `${response.successful_updates.length} team member${response.successful_updates.length === 1 ? "" : "s"} updated`, + ); + if (response.failed_updates.length > 0) { + NotificationsManager.fromBackend(`${response.failed_updates.length} team member updates failed`); } - /> + } catch (error) { + NotificationsManager.fromBackend(error instanceof Error ? error.message : "Failed to bulk update team members"); + } finally { + setIsBulkUpdating(false); + } + }; + + return ( + <> + { + const membership = teamData.team_memberships.find((tm) => tm.user_id === record.user_id); + const enhancedMember = { + ...record, + max_budget_in_team: membership?.litellm_budget_table?.max_budget || null, + tpm_limit: membership?.litellm_budget_table?.tpm_limit || null, + rpm_limit: membership?.litellm_budget_table?.rpm_limit || null, + budget_duration: membership?.litellm_budget_table?.budget_duration || null, + allowed_models: membership?.litellm_budget_table?.allowed_models || [], + }; + setSelectedEditMember(enhancedMember); + setIsEditMemberModalVisible(true); + }} + onDelete={handleMemberDelete} + onAddMember={() => setIsAddMemberModalVisible(true)} + extraActions={ + <> + + {selectionMode && ( + + )} + + } + roleColumnTitle="Team Role" + roleTooltip="This role applies only to this team and is independent from the user's proxy-level role." + extraColumns={extraColumns} + rowSelection={ + selectionMode + ? { + selectedRowKeys: selectedMembers.flatMap((member) => (member.user_id ? [member.user_id] : [])), + onChange: (_selectedRowKeys, selectedRows) => setSelectedMembers(selectedRows), + getCheckboxProps: (member) => ({ disabled: member.user_id === null }), + } + : undefined + } + showDeleteForMember={() => + isProxyAdmin || (canEditTeam && !isUserTeamAdmin) || (isUserTeamAdmin && !disableTeamAdminDeleteTeamUser) + } + /> + setIsBulkUpdateVisible(false)} + onOk={() => bulkUpdateForm.submit()} + okText="Update Members" + confirmLoading={isBulkUpdating} + > +
+ Choose the fields to apply to every selected member. + + Team role + + + {({ getFieldValue }) => + getFieldValue("apply_role") && ( + + ({ label: model, value: model }))} + /> + + ) + } + + +
+ ); } diff --git a/ui/litellm-dashboard/src/components/team/teamMemberBulkUpdate.ts b/ui/litellm-dashboard/src/components/team/teamMemberBulkUpdate.ts new file mode 100644 index 00000000000..1eba49f950b --- /dev/null +++ b/ui/litellm-dashboard/src/components/team/teamMemberBulkUpdate.ts @@ -0,0 +1,30 @@ +import { getGlobalLitellmHeaderName, getProxyBaseUrl } from "@/components/networking"; +import { createApiClient } from "@/lib/http/client"; + +export interface TeamMemberBulkUpdateFields { + role?: "admin" | "user" | null; + max_budget_in_team?: number | null; + tpm_limit?: number | null; + rpm_limit?: number | null; + budget_duration?: string | null; + allowed_models?: string[] | null; +} + +const apiClient = createApiClient({ + getBaseUrl: getProxyBaseUrl, + getAuthHeaderName: getGlobalLitellmHeaderName, +}); + +export const teamMemberBulkUpdateCall = async ( + accessToken: string, + teamId: string, + userIds: string[], + updateFields: TeamMemberBulkUpdateFields, +) => + apiClient.patch(`/v2/team/${encodeURIComponent(teamId)}/members`, { + accessToken, + body: { + user_ids: userIds, + update_fields: updateFields, + }, + }); diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 87c760257d1..7a6088cbc59 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -18943,6 +18943,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; @@ -21481,6 +21498,28 @@ export interface components { [key: string]: unknown; } | null; }; + /** BulkTeamMemberUpdateRequest */ + BulkTeamMemberUpdateRequest: { + /** + * All Members In Team + * @default false + */ + all_members_in_team: boolean; + update_fields: components["schemas"]["TeamMemberBulkUpdateFields"]; + /** User Ids */ + user_ids?: string[] | null; + }; + /** BulkTeamMemberUpdateResponse */ + BulkTeamMemberUpdateResponse: { + /** Failed Updates */ + failed_updates: components["schemas"]["FailedTeamMemberUpdate"][]; + /** Successful Updates */ + successful_updates: components["schemas"]["TeamMemberUpdateResponse"][]; + /** Team Id */ + team_id: string; + /** Total Requested */ + total_requested: number; + }; /** * BulkUpdateKeyRequest * @description Request for bulk key updates @@ -23696,6 +23735,15 @@ export interface components { [key: string]: unknown; } | null; }; + /** FailedTeamMemberUpdate */ + FailedTeamMemberUpdate: { + /** Failed Reason */ + failed_reason: string; + /** User Email */ + user_email?: string | null; + /** User Id */ + user_id: string; + }; /** * FallbackCreateRequest * @description Request model for creating/updating fallbacks @@ -31355,6 +31403,21 @@ export interface components { /** User Id */ user_id?: string | null; }; + /** TeamMemberBulkUpdateFields */ + TeamMemberBulkUpdateFields: { + /** Allowed Models */ + allowed_models?: string[] | null; + /** Budget Duration */ + budget_duration?: string | null; + /** Max Budget In Team */ + max_budget_in_team?: number | null; + /** Role */ + role?: ("admin" | "user") | null; + /** Rpm Limit */ + rpm_limit?: number | null; + /** Tpm Limit */ + tpm_limit?: number | null; + }; /** TeamMemberDeleteRequest */ TeamMemberDeleteRequest: { /** Team Id */ @@ -57552,6 +57615,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?: {