From c5f22f5fdf218b05fc821846233106e17dfdf0aa Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Thu, 16 Jul 2026 14:51:26 -0700 Subject: [PATCH] fix(team): address greptile on bulk member budget writes Echo only fields present in model_fields_set on successful bulk updates, skip redundant budget_reset_at when write_data already set it, and dedupe UpdateBudget ops that share one private budget id --- .../management_endpoints/team_endpoints.py | 34 +++++++++++--- .../team_member_budget_writes.py | 29 ++++++++++-- .../test_team_member_budget_writes.py | 44 +++++++++++++++++++ .../proxy/test_team_member_update.py | 29 ++++++++++++ 4 files changed, 126 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index b4f9f27af65..74109ef8358 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -3064,11 +3064,37 @@ async def _apply_team_member_update( ) +_BULK_MEMBER_RESPONSE_FIELDS = frozenset( + { + "max_budget_in_team", + "tpm_limit", + "rpm_limit", + "budget_duration", + "allowed_models", + } +) + + +def _successful_member_update_response( + *, + team_id: str, + user_id: str, + update_fields: TeamMemberBulkUpdateFields, +) -> TeamMemberUpdateResponse: + fields_set = update_fields.model_fields_set + return TeamMemberUpdateResponse( + team_id=team_id, + user_id=user_id, + **{field: getattr(update_fields, field) for field in _BULK_MEMBER_RESPONSE_FIELDS if field in fields_set}, + ) + + @router.patch( "/v2/team/{team_id}/members", tags=["team management"], dependencies=[Depends(user_api_key_auth)], response_model=BulkTeamMemberUpdateResponse, + response_model_exclude_unset=True, ) @management_endpoint_wrapper async def bulk_update_team_members( @@ -3210,14 +3236,10 @@ async def bulk_update_team_members( ) successful_updates = [ - TeamMemberUpdateResponse( + _successful_member_update_response( 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, + update_fields=data.update_fields, ) for user_id in valid_user_ids ] diff --git a/litellm/proxy/management_endpoints/team_member_budget_writes.py b/litellm/proxy/management_endpoints/team_member_budget_writes.py index 8521f042702..d6342bda3d9 100644 --- a/litellm/proxy/management_endpoints/team_member_budget_writes.py +++ b/litellm/proxy/management_endpoints/team_member_budget_writes.py @@ -1,6 +1,7 @@ from __future__ import annotations from dataclasses import dataclass +from functools import reduce from typing import Any, Callable, Mapping, Protocol, Sequence from uuid import uuid4 @@ -75,13 +76,33 @@ def _build_create_data( 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: + if create_data.get("budget_duration") is None: create_data.pop("budget_reset_at", None) + elif "budget_reset_at" not in create_data: + create_data["budget_reset_at"] = get_budget_reset_time(budget_duration=create_data["budget_duration"]) return create_data +def _dedupe_update_budget_writes(writes: Sequence[MemberBudgetWrite]) -> tuple[MemberBudgetWrite, ...]: + last_update_by_id = {write.budget_id: write for write in writes if isinstance(write, UpdateBudget)} + if len(last_update_by_id) == sum(1 for write in writes if isinstance(write, UpdateBudget)): + return tuple(writes) + + def step( + acc: tuple[tuple[MemberBudgetWrite, ...], frozenset[str]], + write: MemberBudgetWrite, + ) -> tuple[tuple[MemberBudgetWrite, ...], frozenset[str]]: + out, seen = acc + if not isinstance(write, UpdateBudget): + return (*out, write), seen + if write.budget_id in seen: + return acc + return (*out, last_update_by_id[write.budget_id]), seen | {write.budget_id} + + result, _ = reduce(step, writes, ((), frozenset())) + return result + + def plan_member_budget_writes( *, memberships: Sequence[MembershipBudgetSnapshot], @@ -145,7 +166,7 @@ def plan_member_budget_writes( ) ) - return MemberBudgetWritePlan(writes=tuple(writes)) + return MemberBudgetWritePlan(writes=_dedupe_update_budget_writes(writes)) class TeamMemberBudgetDb(Protocol): 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 index 13da65f9c62..6a0c8fe501f 100644 --- 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 @@ -80,6 +80,50 @@ def test_plan_empty_patch_is_noop(): assert plan.writes == () +def test_plan_dedupes_update_budget_writes_for_shared_private_budget_id(): + plan = plan_member_budget_writes( + memberships=( + MembershipBudgetSnapshot(user_id="u1", budget_id="shared-private"), + MembershipBudgetSnapshot(user_id="u2", budget_id="shared-private"), + MembershipBudgetSnapshot(user_id="u3", budget_id="other"), + ), + budgets_by_id={ + "shared-private": BudgetFieldSnapshot(budget_id="shared-private", fields={"tpm_limit": 1}), + "other": BudgetFieldSnapshot(budget_id="other", fields={"tpm_limit": 2}), + }, + budget_patch={"tpm_limit": 9}, + team_default_budget_id=None, + actor_user_id="admin", + ) + + updates = tuple(write for write in plan.writes if isinstance(write, UpdateBudget)) + assert len(updates) == 2 + assert {write.budget_id for write in updates} == {"shared-private", "other"} + + +def test_plan_clone_inherits_shared_duration_and_sets_reset_at(): + plan = plan_member_budget_writes( + memberships=(MembershipBudgetSnapshot(user_id="on-default", budget_id="default"),), + budgets_by_id={ + "default": BudgetFieldSnapshot( + budget_id="default", + fields={"budget_duration": "30d", "max_budget": 50.0}, + ), + }, + budget_patch={"tpm_limit": 3}, + team_default_budget_id="default", + actor_user_id="admin", + new_budget_id_factory=lambda: "nb-1", + ) + + assert len(plan.writes) == 1 + write = plan.writes[0] + assert isinstance(write, CreateAndAttachBudget) + assert write.create_data["budget_duration"] == "30d" + assert write.create_data["budget_reset_at"] is not None + assert write.create_data["tpm_limit"] == 3 + + class _RecordingDb: def __init__(self): self.created = [] diff --git a/tests/test_litellm/proxy/test_team_member_update.py b/tests/test_litellm/proxy/test_team_member_update.py index e05f289f0e1..4097ed8ed76 100644 --- a/tests/test_litellm/proxy/test_team_member_update.py +++ b/tests/test_litellm/proxy/test_team_member_update.py @@ -304,6 +304,35 @@ async def test_bulk_update_plans_budget_writes_without_per_member_inline_prisma( 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"] + assert response.successful_updates[0].model_dump(exclude_unset=True) == { + "team_id": "team-1234", + "user_id": "user-1", + "tpm_limit": 42, + } + + +@pytest.mark.asyncio +async def test_bulk_update_role_only_response_omits_unset_budget_fields(monkeypatch): + team_row = LiteLLM_TeamTable( + team_id="team-1234", + members_with_roles=[Member(user_id="user-1", role="admin")], + ) + _bulk_setup(monkeypatch, team_row) + + response = await bulk_update_team_members( + team_id="team-1234", + data=BulkTeamMemberUpdateRequest( + user_ids=["user-1"], + update_fields=TeamMemberBulkUpdateFields(role="user"), + ), + http_request=_bulk_request(), + user_api_key_dict=_ADMIN, + ) + + assert response.successful_updates[0].model_dump(exclude_unset=True) == { + "team_id": "team-1234", + "user_id": "user-1", + } @pytest.mark.asyncio