From e726e910dfd31df08970279a08eb076229afd921 Mon Sep 17 00:00:00 2001 From: Milan Date: Fri, 1 May 2026 01:44:56 +0300 Subject: [PATCH] refactor(proxy): split team member budget upsert for PLR0915 Extract helpers from _upsert_budget_and_membership and team_member_update so each stays under Ruff max-statements (no noqa). Made-with: Cursor --- .../management_endpoints/common_utils.py | 342 ++++++++++++------ .../management_endpoints/team_endpoints.py | 277 ++++++++------ 2 files changed, 389 insertions(+), 230 deletions(-) diff --git a/litellm/proxy/management_endpoints/common_utils.py b/litellm/proxy/management_endpoints/common_utils.py index 05de65f6601..2fc1e7cf84d 100644 --- a/litellm/proxy/management_endpoints/common_utils.py +++ b/litellm/proxy/management_endpoints/common_utils.py @@ -345,6 +345,198 @@ def _set_object_metadata_field( object_data.metadata[field_name] = value +def _budget_disconnect_all_unset( + max_budget: Optional[float], + tpm_limit: Optional[int], + rpm_limit: Optional[int], + allowed_models: Optional[List[str]], + budget_duration: Optional[str], + budget_duration_explicit: bool, +) -> bool: + return ( + max_budget is None + and tpm_limit is None + and rpm_limit is None + and allowed_models is None + and budget_duration is None + and not budget_duration_explicit + ) + + +async def _disconnect_team_member_budget( + tx, *, team_id: str, user_id: str +) -> None: + await tx.litellm_teammembership.update( + where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, + data={"litellm_budget_table": {"disconnect": True}}, + ) + + +def _non_none_budget_limit_fields( + max_budget: Optional[float], + tpm_limit: Optional[int], + rpm_limit: Optional[int], + allowed_models: Optional[List[str]], +) -> Dict[str, Any]: + out: Dict[str, Any] = {} + if max_budget is not None: + out["max_budget"] = max_budget + if tpm_limit is not None: + out["tpm_limit"] = tpm_limit + if rpm_limit is not None: + out["rpm_limit"] = rpm_limit + if allowed_models is not None: + out["allowed_models"] = allowed_models + return out + + +def _apply_budget_duration_to_update_dict( + update_data: Dict[str, Any], + budget_duration: Optional[str], + budget_duration_explicit: bool, +) -> None: + if budget_duration_explicit: + if budget_duration is not None: + update_data["budget_duration"] = budget_duration + update_data["budget_reset_at"] = get_budget_reset_time( + budget_duration=budget_duration + ) + else: + update_data["budget_duration"] = None + update_data["budget_reset_at"] = None + elif budget_duration is not None: + update_data["budget_duration"] = budget_duration + update_data["budget_reset_at"] = get_budget_reset_time( + budget_duration=budget_duration + ) + + +async def _update_existing_member_budget_in_place( + tx, + *, + existing_budget_id: str, + user_api_key_dict: UserAPIKeyAuth, + max_budget: Optional[float], + tpm_limit: Optional[int], + rpm_limit: Optional[int], + allowed_models: Optional[List[str]], + budget_duration: Optional[str], + budget_duration_explicit: bool, +) -> None: + update_data: Dict[str, Any] = { + "updated_by": user_api_key_dict.user_id or "", + **_non_none_budget_limit_fields( + max_budget, tpm_limit, rpm_limit, allowed_models + ), + } + _apply_budget_duration_to_update_dict( + update_data, budget_duration, budget_duration_explicit + ) + await tx.litellm_budgettable.update( + where={"budget_id": existing_budget_id}, + data=update_data, + ) + + +async def _clone_shared_default_into_create_data( + tx, create_data: Dict[str, Any], *, existing_budget_id: str +) -> None: + default_budget_row = await tx.litellm_budgettable.find_unique( + where={"budget_id": existing_budget_id} + ) + if default_budget_row is None: + return + default_budget_dict = default_budget_row.model_dump() + for field in ( + "max_budget", + "soft_budget", + "max_parallel_requests", + "tpm_limit", + "rpm_limit", + "model_max_budget", + "budget_duration", + "allowed_models", + ): + value = default_budget_dict.get(field) + if value is None: + continue + if isinstance(value, list) and len(value) == 0: + continue + create_data[field] = value + + +def _merge_limit_overrides_into_create_data( + create_data: Dict[str, Any], + *, + max_budget: Optional[float], + tpm_limit: Optional[int], + rpm_limit: Optional[int], + allowed_models: Optional[List[str]], +) -> None: + if max_budget is not None: + create_data["max_budget"] = max_budget + if tpm_limit is not None: + create_data["tpm_limit"] = tpm_limit + if rpm_limit is not None: + create_data["rpm_limit"] = rpm_limit + if allowed_models is not None: + create_data["allowed_models"] = allowed_models + + +def _merge_duration_into_create_data( + create_data: Dict[str, Any], + budget_duration: Optional[str], + budget_duration_explicit: bool, +) -> None: + if budget_duration_explicit: + if budget_duration is not None: + create_data["budget_duration"] = budget_duration + else: + create_data.pop("budget_duration", None) + create_data.pop("budget_reset_at", None) + elif budget_duration is not None: + create_data["budget_duration"] = budget_duration + + +def _set_create_data_budget_reset_from_duration(create_data: Dict[str, Any]) -> None: + bd = create_data.get("budget_duration") + if bd is not None: + create_data["budget_reset_at"] = get_budget_reset_time(budget_duration=bd) + else: + create_data.pop("budget_reset_at", None) + + +async def _create_private_budget_and_link_membership( + tx, *, team_id: str, user_id: str, create_data: Dict[str, Any] +) -> None: + new_budget = await tx.litellm_budgettable.create( + data=create_data, + include={"team_membership": True}, + ) + await 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": new_budget.budget_id}, + }, + }, + "update": { + "litellm_budget_table": { + "connect": {"budget_id": new_budget.budget_id}, + }, + }, + }, + ) + + async def _upsert_budget_and_membership( tx, *, @@ -389,19 +581,15 @@ async def _upsert_budget_and_membership( If any of these values exist, a budget is updated or created and linked to the team membership. """ - if ( - max_budget is None - and tpm_limit is None - and rpm_limit is None - and allowed_models is None - and budget_duration is None - and not budget_duration_explicit + if _budget_disconnect_all_unset( + max_budget, + tpm_limit, + rpm_limit, + allowed_models, + budget_duration, + budget_duration_explicit, ): - # disconnect the budget since all limits are None - await tx.litellm_teammembership.update( - where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, - data={"litellm_budget_table": {"disconnect": True}}, - ) + await _disconnect_team_member_budget(tx, team_id=team_id, user_id=user_id) return is_shared_default = ( @@ -411,121 +599,41 @@ async def _upsert_budget_and_membership( ) if existing_budget_id is not None and not is_shared_default: - # Update the existing budget in-place to preserve fields not being changed. - # Only write fields that the caller explicitly provided (non-None). - update_data: Dict[str, Any] = { - "updated_by": user_api_key_dict.user_id or "", - } - if max_budget is not None: - update_data["max_budget"] = max_budget - if tpm_limit is not None: - update_data["tpm_limit"] = tpm_limit - if rpm_limit is not None: - update_data["rpm_limit"] = rpm_limit - if allowed_models is not None: - update_data["allowed_models"] = allowed_models - if budget_duration_explicit: - if budget_duration is not None: - update_data["budget_duration"] = budget_duration - update_data["budget_reset_at"] = get_budget_reset_time( - budget_duration=budget_duration - ) - else: - update_data["budget_duration"] = None - update_data["budget_reset_at"] = None - elif budget_duration is not None: - update_data["budget_duration"] = budget_duration - update_data["budget_reset_at"] = get_budget_reset_time( - budget_duration=budget_duration - ) - await tx.litellm_budgettable.update( - where={"budget_id": existing_budget_id}, - data=update_data, + await _update_existing_member_budget_in_place( + tx, + existing_budget_id=existing_budget_id, + user_api_key_dict=user_api_key_dict, + max_budget=max_budget, + tpm_limit=tpm_limit, + rpm_limit=rpm_limit, + allowed_models=allowed_models, + budget_duration=budget_duration, + budget_duration_explicit=budget_duration_explicit, ) return - # Either there is no existing budget, OR the membership is still pointing - # at the team's shared default member budget. In both cases we create a - # NEW private budget for this user and (re)link the membership to it. create_data: Dict[str, Any] = { "created_by": user_api_key_dict.user_id or "", "updated_by": user_api_key_dict.user_id or "", } - - # If we're forking off the shared default, seed the new row with the - # default's values so fields the caller did not change carry over. if is_shared_default: - default_budget_row = await tx.litellm_budgettable.find_unique( - where={"budget_id": existing_budget_id} + assert existing_budget_id is not None + await _clone_shared_default_into_create_data( + tx, create_data, existing_budget_id=existing_budget_id ) - if default_budget_row is not None: - default_budget_dict = default_budget_row.model_dump() - for field in ( - "max_budget", - "soft_budget", - "max_parallel_requests", - "tpm_limit", - "rpm_limit", - "model_max_budget", - "budget_duration", - "allowed_models", - ): - value = default_budget_dict.get(field) - if value is None: - continue - if isinstance(value, list) and len(value) == 0: - continue - create_data[field] = value - - # Caller-provided values take precedence over the cloned defaults. - if max_budget is not None: - create_data["max_budget"] = max_budget - if tpm_limit is not None: - create_data["tpm_limit"] = tpm_limit - if rpm_limit is not None: - create_data["rpm_limit"] = rpm_limit - if allowed_models is not None: - create_data["allowed_models"] = allowed_models - if budget_duration_explicit: - if budget_duration is not None: - create_data["budget_duration"] = budget_duration - else: - create_data.pop("budget_duration", None) - create_data.pop("budget_reset_at", None) - elif budget_duration is not None: - create_data["budget_duration"] = budget_duration - - bd = create_data.get("budget_duration") - if bd is not None: - create_data["budget_reset_at"] = get_budget_reset_time(budget_duration=bd) - else: - create_data.pop("budget_reset_at", None) - - new_budget = await tx.litellm_budgettable.create( - data=create_data, - include={"team_membership": True}, + _merge_limit_overrides_into_create_data( + create_data, + max_budget=max_budget, + tpm_limit=tpm_limit, + rpm_limit=rpm_limit, + allowed_models=allowed_models, ) - await 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": new_budget.budget_id}, - }, - }, - "update": { - "litellm_budget_table": { - "connect": {"budget_id": new_budget.budget_id}, - }, - }, - }, + _merge_duration_into_create_data( + create_data, budget_duration, budget_duration_explicit + ) + _set_create_data_budget_reset_from_duration(create_data) + await _create_private_budget_and_link_membership( + tx, team_id=team_id, user_id=user_id, create_data=create_data ) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 9f7c50dfdec..a53019dddf9 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -2614,6 +2614,147 @@ async def team_member_delete( return existing_team_row +def _validate_team_member_update_request( + data: TeamMemberUpdateRequest, premium_user: bool +) -> None: + if data.team_id is None: + raise HTTPException(status_code=400, detail={"error": "No team id passed in"}) + + if data.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", + ) + if data.user_id is None and data.user_email is None: + raise HTTPException( + status_code=400, + detail={"error": "Either user_id or user_email needs to be passed in"}, + ) + + +async def _team_member_update_fetch_team_or_raise( + prisma_client: PrismaClient, team_id: str +) -> LiteLLM_TeamTable: + _existing_team_row = await prisma_client.db.litellm_teamtable.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)}, + ) + return LiteLLM_TeamTable(**_existing_team_row.model_dump()) + + +async def _team_member_update_require_authorized( + user_api_key_dict: UserAPIKeyAuth, existing_team_row: LiteLLM_TeamTable +) -> None: + 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_row + ) + and not await _is_user_org_admin_for_team( + user_api_key_dict=user_api_key_dict, team_obj=existing_team_row + ) + ): + raise HTTPException( + status_code=403, + detail={ + "error": "Call not allowed. User not proxy admin OR team admin. route={}, team_id={}".format( + "/team/member_delete", existing_team_row.team_id + ) + }, + ) + + +def _resolve_team_member_user_id_from_request( + data: TeamMemberUpdateRequest, returned_team_info: TeamInfoResponseObject +) -> str: + if data.user_id is not None: + return data.user_id + assert data.user_email is not None + for member in returned_team_info["team_info"].members_with_roles: + if member.user_email is not None and member.user_email == data.user_email: + if member.user_id is None: + break + return member.user_id + raise HTTPException( + status_code=400, + detail={"error": "User id doesn't exist in team table. Data={}".format(data)}, + ) + + +def _identified_member_budget_id( + team_memberships: List[LiteLLM_TeamMembership], received_user_id: str +) -> Optional[str]: + for tm in team_memberships: + if tm.user_id == received_user_id: + return tm.budget_id + return None + + +def _team_default_member_budget_id( + team_table: TeamInfoResponseObjectTeamTable, +) -> Optional[str]: + if team_table.metadata is None: + return None + raw_default_budget_id = team_table.metadata.get("team_member_budget_id") + if isinstance(raw_default_budget_id, str): + return raw_default_budget_id + return None + + +async def _effective_member_budget_duration( + data: TeamMemberUpdateRequest, + prisma_client: PrismaClient, + team_default_budget_id: Optional[str], +) -> Tuple[bool, Optional[str]]: + budget_duration_explicit = "budget_duration" in data.model_fields_set + if budget_duration_explicit: + return budget_duration_explicit, data.budget_duration + if team_default_budget_id: + _team_budget_row = await prisma_client.db.litellm_budgettable.find_unique( + where={"budget_id": team_default_budget_id} + ) + return ( + budget_duration_explicit, + _team_budget_row.budget_duration if _team_budget_row else None, + ) + return budget_duration_explicit, None + + +async def _team_member_update_apply_role_if_requested( + prisma_client: PrismaClient, + *, + data: TeamMemberUpdateRequest, + team_table: TeamInfoResponseObjectTeamTable, + received_user_id: str, +) -> None: + if data.role is None: + return + team_members: List[Member] = [] + for member in team_table.members_with_roles: + if member.user_id == received_user_id: + team_members.append( + Member( + user_id=member.user_id, + role=data.role, + user_email=data.user_email or member.user_email, + ) + ) + else: + team_members.append(member) + + team_table.members_with_roles = team_members + + _db_team_members: List[dict] = [m.model_dump() for m in team_members] + await prisma_client.db.litellm_teamtable.update( + where={"team_id": data.team_id}, + data={"members_with_roles": json.dumps(_db_team_members)}, # type: ignore + ) + + @router.post( "/team/member_update", tags=["team management"], @@ -2636,51 +2777,13 @@ async def team_member_update( if prisma_client is None: raise HTTPException(status_code=500, detail={"error": "No db connected"}) - if data.team_id is None: - raise HTTPException(status_code=400, detail={"error": "No team id passed in"}) - - if data.role == "admin" and not premium_user: - # exactly the same text your proxy throws for add: - 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", - ) - if data.user_id is None and data.user_email is None: - raise HTTPException( - status_code=400, - detail={"error": "Either user_id or user_email needs to be passed in"}, - ) - - _existing_team_row = await prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": data.team_id} + _validate_team_member_update_request(data, premium_user) + existing_team_row = await _team_member_update_fetch_team_or_raise( + prisma_client, data.team_id + ) + await _team_member_update_require_authorized( + user_api_key_dict, existing_team_row ) - - if _existing_team_row is None: - raise HTTPException( - status_code=400, - detail={"error": "Team id={} does not exist in db".format(data.team_id)}, - ) - existing_team_row = LiteLLM_TeamTable(**_existing_team_row.model_dump()) - - ## CHECK IF USER IS PROXY ADMIN OR TEAM ADMIN OR ORG ADMIN - - 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_row - ) - and not await _is_user_org_admin_for_team( - user_api_key_dict=user_api_key_dict, team_obj=existing_team_row - ) - ): - raise HTTPException( - status_code=403, - detail={ - "error": "Call not allowed. User not proxy admin OR team admin. route={}, team_id={}".format( - "/team/member_delete", existing_team_row.team_id - ) - }, - ) returned_team_info: TeamInfoResponseObject = await team_info( http_request=http_request, @@ -2689,55 +2792,19 @@ async def team_member_update( ) team_table = returned_team_info["team_info"] - - ## get user id - received_user_id: Optional[str] = None - 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: - if member.user_email is not None and member.user_email == data.user_email: - received_user_id = member.user_id - break - - if received_user_id is None: - raise HTTPException( - status_code=400, - detail={ - "error": "User id doesn't exist in team table. Data={}".format(data) - }, + received_user_id = _resolve_team_member_user_id_from_request( + data, returned_team_info + ) + identified_budget_id = _identified_member_budget_id( + returned_team_info["team_memberships"], received_user_id + ) + team_default_budget_id = _team_default_member_budget_id(team_table) + budget_duration_explicit, effective_budget_duration = ( + await _effective_member_budget_duration( + data, prisma_client, team_default_budget_id ) - ## find the relevant team membership - identified_budget_id: Optional[str] = None - for tm in returned_team_info["team_memberships"]: - if tm.user_id == received_user_id: - identified_budget_id = tm.budget_id - break + ) - team_default_budget_id: Optional[str] = None - if team_table.metadata is not None: - raw_default_budget_id = team_table.metadata.get("team_member_budget_id") - if isinstance(raw_default_budget_id, str): - team_default_budget_id = raw_default_budget_id - - ### resolve effective budget_duration - # - Explicit value (including null) takes precedence - # - If omitted, inherit the team's configured team_member_budget_duration - budget_duration_explicit = "budget_duration" in data.model_fields_set - if budget_duration_explicit: - effective_budget_duration = data.budget_duration - else: - if team_default_budget_id: - _team_budget_row = await prisma_client.db.litellm_budgettable.find_unique( - where={"budget_id": team_default_budget_id} - ) - effective_budget_duration = ( - _team_budget_row.budget_duration if _team_budget_row else None - ) - else: - effective_budget_duration = None - - ### upsert new budget async with prisma_client.db.tx() as tx: await _upsert_budget_and_membership( tx=tx, @@ -2754,28 +2821,12 @@ async def team_member_update( budget_duration_explicit=budget_duration_explicit, ) - ### update team member role - if data.role is not None: - team_members: List[Member] = [] - for member in team_table.members_with_roles: - if member.user_id == received_user_id: - team_members.append( - Member( - user_id=member.user_id, - role=data.role, - user_email=data.user_email or member.user_email, - ) - ) - else: - team_members.append(member) - - team_table.members_with_roles = team_members - - _db_team_members: List[dict] = [m.model_dump() for m in team_members] - await prisma_client.db.litellm_teamtable.update( - where={"team_id": data.team_id}, - data={"members_with_roles": json.dumps(_db_team_members)}, # type: ignore - ) + await _team_member_update_apply_role_if_requested( + prisma_client, + data=data, + team_table=team_table, + received_user_id=received_user_id, + ) return TeamMemberUpdateResponse( team_id=data.team_id,