Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/management_helpers/bulk_team_member_budgets.py: 57%
94 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""Batched per-member limit writes behind `POST /management/v1/teams/{team_id}/members/bulk_update`.
3Every read runs on the writer inside the batch transaction, so the write plan can never be
4built from a lagging read replica. Any budget row that more than one membership points at,
5the team's shared default included, is cloned before it is written, so raising one member's
6cap never moves another member's.
7"""
9from collections.abc import Sequence
10from datetime import datetime, timedelta
11from types import MappingProxyType
12from typing import TYPE_CHECKING, Final
14from pydantic import BaseModel, ConfigDict
16from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
17from litellm.proxy._types import (
18 LiteLLM_TeamTable,
19 LitellmTableNames,
20 LitellmUserRoles,
21 Member,
22 UserAPIKeyAuth,
23)
24from litellm.proxy.auth.auth_checks import invalidate_team_member_spend_state
25from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
26from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient
27from litellm.proxy.management_endpoints.common_utils import (
28 _is_user_org_admin_for_team, # pyright: ignore[reportPrivateUsage] # same check /team/member_update uses
29 _is_user_team_admin, # pyright: ignore[reportPrivateUsage] # same check /team/member_update uses
30 _upsert_budget_and_membership, # pyright: ignore[reportPrivateUsage] # the single-member write, shared so the two surfaces cannot drift
31 member_budget_patch,
32)
33from litellm.proxy.management_helpers.audit_logs import create_object_audit_log
34from litellm.proxy.management_helpers.bulk_user_deletion import (
35 _duplicate_member_indexes, # pyright: ignore[reportPrivateUsage] # same duplicate rule as members/bulk_delete
36 _eq_filter, # pyright: ignore[reportPrivateUsage] # same prisma filter shape as members/bulk_delete
37 _forbidden, # pyright: ignore[reportPrivateUsage] # same problem shape as members/bulk_delete
38 _in_filter, # pyright: ignore[reportPrivateUsage] # same prisma filter shape as members/bulk_delete
39 _team_not_found, # pyright: ignore[reportPrivateUsage] # same problem shape as members/bulk_delete
40 _team_users_filter, # pyright: ignore[reportPrivateUsage] # same prisma filter shape as members/bulk_delete
41)
42from litellm.proxy.utils import PrismaClient
43from litellm.repositories.team_repository import TeamRepository
44from litellm.types.proxy.management_endpoints.team_endpoints import (
45 BulkTeamMemberBudgetUpdateRequest,
46 TeamMemberBudgetPatch,
47 TeamMemberBudgetUpdateResult,
48)
50if TYPE_CHECKING: 50 ↛ 51line 50 didn't jump to line 51 because the condition on line 50 was never true
51 from prisma import Prisma
52 from prisma import models as prisma_models
54 from litellm.repositories.prisma_protocols import TableActions
56_BATCH_TX_TIMEOUT: Final = timedelta(seconds=60)
57_NO_METADATA: Final = MappingProxyType({})
58_WITH_BUDGET: Final = MappingProxyType({"litellm_budget_table": True})
61def _membership_tx_db(tx: "Prisma") -> "TableActions[prisma_models.LiteLLM_TeamMembership]":
62 return tx.litellm_teammembership # pyright: ignore[reportReturnType] # TableActions widens the generated inputs to Mapping, as the repositories do
65def _budget_tx_db(tx: "Prisma") -> "TableActions[prisma_models.LiteLLM_BudgetTable]":
66 return tx.litellm_budgettable # pyright: ignore[reportReturnType] # TableActions widens the generated inputs to Mapping, as the repositories do
69def _roster_user_id(member: TeamMemberBudgetPatch, roster: Sequence[Member]) -> str | None:
70 """The team member this row addresses, or None when it names nobody on the team."""
71 if member.user_id is not None:
72 return member.user_id if any(m.user_id == member.user_id for m in roster) else None
73 return next((m.user_id for m in roster if m.user_email is not None and m.user_email == member.user_email), None)
76def _team_default_budget_id(team: LiteLLM_TeamTable) -> str | None:
77 raw: Final = (team.metadata or _NO_METADATA).get("team_member_budget_id")
78 return raw if isinstance(raw, str) else None
81async def _shared_budget_ids(tx: "Prisma", budget_ids: frozenset[str]) -> frozenset[str]:
82 """The rows in ``budget_ids`` more than one membership points at, counted across every
83 team so a row shared with another team is protected too."""
84 if not budget_ids:
85 return frozenset()
86 rows: Final = await _membership_tx_db(tx).find_many(where=_in_filter("budget_id", budget_ids))
87 return frozenset(budget_id for budget_id in budget_ids if sum(1 for row in rows if row.budget_id == budget_id) > 1)
90class _AuditedMemberBudget(BaseModel):
91 """One member's limits as the audit log's before/after values record them."""
93 model_config = ConfigDict(frozen=True)
95 user_id: str
96 budget_id: str | None = None
97 max_budget: float | None = None
98 tpm_limit: int | None = None
99 rpm_limit: int | None = None
100 budget_duration: str | None = None
101 budget_reset_at: datetime | None = None
102 allowed_models: tuple[str, ...] | None = None
105class _AuditedMemberBudgets(BaseModel):
106 """The audit-log columns hold a JSON object, so the per-member list is nested under a key."""
108 model_config = ConfigDict(frozen=True)
110 team_member_budgets: tuple[_AuditedMemberBudget, ...]
113def _audited_member_budget(row: "prisma_models.LiteLLM_TeamMembership") -> _AuditedMemberBudget:
114 budget: Final = row.litellm_budget_table
115 if budget is None:
116 return _AuditedMemberBudget(user_id=row.user_id, budget_id=row.budget_id)
117 return _AuditedMemberBudget(
118 user_id=row.user_id,
119 budget_id=row.budget_id,
120 max_budget=budget.max_budget,
121 tpm_limit=budget.tpm_limit,
122 rpm_limit=budget.rpm_limit,
123 budget_duration=budget.budget_duration,
124 budget_reset_at=budget.budget_reset_at,
125 allowed_models=tuple(budget.allowed_models),
126 )
129def _limits_audit_value(rows: "Sequence[prisma_models.LiteLLM_TeamMembership]") -> str:
130 """Serialize the members' limits for an audit-log value, dropping the limits they do not set."""
131 return safe_dumps(
132 _AuditedMemberBudgets(
133 team_member_budgets=tuple(_audited_member_budget(row) for row in sorted(rows, key=lambda row: row.user_id))
134 ).model_dump(exclude_none=True, mode="json")
135 )
138def _result(
139 member: TeamMemberBudgetPatch,
140 user_id: str | None,
141 error: str | None,
142 budget_of: "MappingProxyType[str, prisma_models.LiteLLM_BudgetTable | None]",
143 team_default_max_budget: float | None,
144) -> TeamMemberBudgetUpdateResult:
145 if error is not None or user_id is None: 145 ↛ 152line 145 didn't jump to line 152 because the condition on line 145 was always true
146 return TeamMemberBudgetUpdateResult(
147 user_id=member.user_id,
148 user_email=member.user_email,
149 success=False,
150 error=error or "User not found in team",
151 )
152 budget: Final = budget_of.get(user_id)
153 own_max_budget: Final = budget.max_budget if budget is not None else None
154 inherits: Final = own_max_budget is None and team_default_max_budget is not None and team_default_max_budget > 0
155 return TeamMemberBudgetUpdateResult(
156 user_id=user_id,
157 user_email=member.user_email,
158 success=True,
159 budget_id=budget.budget_id if budget is not None else None,
160 max_budget=team_default_max_budget if inherits else own_max_budget,
161 max_budget_source=("team_default" if inherits else "member" if own_max_budget is not None else None),
162 tpm_limit=budget.tpm_limit if budget is not None else None,
163 rpm_limit=budget.rpm_limit if budget is not None else None,
164 budget_duration=budget.budget_duration if budget is not None else None,
165 allowed_models=tuple(budget.allowed_models) if budget is not None else None,
166 )
169async def bulk_update_team_member_budgets(
170 team_id: str,
171 data: BulkTeamMemberBudgetUpdateRequest,
172 user_api_key_dict: UserAPIKeyAuth,
173 prisma_client: PrismaClient,
174 user_api_key_cache: UserApiKeyCache,
175 litellm_proxy_admin_name: str,
176 litellm_changed_by: str | None = None,
177) -> tuple[TeamMemberBudgetUpdateResult, ...]:
178 """Apply one merge patch of per-member limits per requested member, in one transaction."""
179 team: Final = await TeamRepository(WriterPinnedClient(prisma_client.db)).find_by_id(team_id)
180 if team is None:
181 raise _team_not_found(team_id)
183 if ( 183 ↛ 188line 183 didn't jump to line 188 because the condition on line 183 was never true
184 user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value
185 and not _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team)
186 and not await _is_user_org_admin_for_team(user_api_key_dict=user_api_key_dict, team_obj=team)
187 ):
188 raise _forbidden(
189 "Call not allowed. User not proxy admin OR team admin OR org admin for this team. "
190 f"route='/management/v1/teams/{team_id}/members/bulk_update'"
191 )
193 roster: Final = team.members_with_roles or ()
194 named: Final = tuple(_roster_user_id(member, roster) for member in data.members)
195 duplicates: Final = _duplicate_member_indexes(data.members) | frozenset(
196 index for index, user_id in enumerate(named) if user_id is not None and user_id in named[:index]
197 )
198 applied: Final = tuple(
199 (index, user_id) for index, user_id in enumerate(named) if user_id is not None and index not in duplicates
200 )
201 if not applied: 201 ↛ 209line 201 didn't jump to line 209 because the condition on line 201 was always true
202 return tuple(
203 _result(
204 member, None, "Duplicate member in request" if index in duplicates else None, MappingProxyType({}), None
205 )
206 for index, member in enumerate(data.members)
207 )
209 user_ids: Final = sorted(user_id for _, user_id in applied)
210 default_budget_id: Final = _team_default_budget_id(team)
211 team_members_filter: Final = _team_users_filter(team_id, user_ids)
213 async with prisma_client.tx(timeout=_BATCH_TX_TIMEOUT) as tx:
214 memberships: Final = await _membership_tx_db(tx).find_many(where=team_members_filter, include=_WITH_BUDGET)
215 budget_id_of: Final = MappingProxyType({m.user_id: m.budget_id for m in memberships})
216 shared: Final = await _shared_budget_ids(
217 tx, frozenset(budget_id for budget_id in budget_id_of.values() if budget_id is not None)
218 )
219 for index, user_id in applied:
220 await _upsert_budget_and_membership(
221 tx=tx,
222 team_id=team_id,
223 user_id=user_id,
224 existing_budget_id=budget_id_of.get(user_id),
225 user_api_key_dict=user_api_key_dict,
226 budget_patch=member_budget_patch(data.members[index]),
227 team_default_budget_id=default_budget_id,
228 shared_budget_ids=shared,
229 )
230 written: Final = await _membership_tx_db(tx).find_many(where=team_members_filter, include=_WITH_BUDGET)
231 team_default: Final = (
232 await _budget_tx_db(tx).find_unique(where=_eq_filter("budget_id", default_budget_id))
233 if default_budget_id is not None
234 else None
235 )
237 for user_id in user_ids:
238 await invalidate_team_member_spend_state(
239 user_id=user_id, team_id=team_id, user_api_key_cache=user_api_key_cache
240 )
242 await create_object_audit_log(
243 object_id=team_id,
244 action="updated",
245 litellm_changed_by=litellm_changed_by,
246 user_api_key_dict=user_api_key_dict,
247 litellm_proxy_admin_name=litellm_proxy_admin_name,
248 table_name=LitellmTableNames.TEAM_TABLE_NAME,
249 before_value=_limits_audit_value(memberships),
250 after_value=_limits_audit_value(written),
251 )
253 budget_of: Final = MappingProxyType({m.user_id: m.litellm_budget_table for m in written})
254 return tuple(
255 _result(
256 member,
257 named[index],
258 "Duplicate member in request" if index in duplicates else None,
259 budget_of,
260 team_default.max_budget if team_default is not None else None,
261 )
262 for index, member in enumerate(data.members)
263 )