mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
81717e3582
commit
c5f22f5fdf
4 changed files with 126 additions and 10 deletions
|
|
@ -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
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue