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
This commit is contained in:
mubashir1osmani 2026-07-16 14:51:26 -07:00
parent 81717e3582
commit c5f22f5fdf
4 changed files with 126 additions and 10 deletions

View file

@ -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
]

View file

@ -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):

View file

@ -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 = []

View file

@ -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