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

1"""Batched per-member limit writes behind `POST /management/v1/teams/{team_id}/members/bulk_update`. 

2 

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

8 

9from collections.abc import Sequence 

10from datetime import datetime, timedelta 

11from types import MappingProxyType 

12from typing import TYPE_CHECKING, Final 

13 

14from pydantic import BaseModel, ConfigDict 

15 

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) 

49 

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 

53 

54 from litellm.repositories.prisma_protocols import TableActions 

55 

56_BATCH_TX_TIMEOUT: Final = timedelta(seconds=60) 

57_NO_METADATA: Final = MappingProxyType({}) 

58_WITH_BUDGET: Final = MappingProxyType({"litellm_budget_table": True}) 

59 

60 

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 

63 

64 

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 

67 

68 

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) 

74 

75 

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 

79 

80 

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) 

88 

89 

90class _AuditedMemberBudget(BaseModel): 

91 """One member's limits as the audit log's before/after values record them.""" 

92 

93 model_config = ConfigDict(frozen=True) 

94 

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 

103 

104 

105class _AuditedMemberBudgets(BaseModel): 

106 """The audit-log columns hold a JSON object, so the per-member list is nested under a key.""" 

107 

108 model_config = ConfigDict(frozen=True) 

109 

110 team_member_budgets: tuple[_AuditedMemberBudget, ...] 

111 

112 

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 ) 

127 

128 

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 ) 

136 

137 

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 ) 

167 

168 

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) 

182 

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 ) 

192 

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 ) 

208 

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) 

212 

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 ) 

236 

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 ) 

241 

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 ) 

252 

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 )