Merge pull request #41632 from BerriAI/litellm_bulk_team_member_budget_update

feat(management_v1): bulk update team member budgets
This commit is contained in:
ryan-crabbe-berri 2026-09-17 15:58:32 -07:00 • committed by GitHub
commit 29a959b3e8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 1309 additions and 46 deletions

View file

@ -850,6 +850,7 @@ class LiteLLMRoutes(enum.Enum):
"/team/member_add",
"/team/member_delete",
"/management/v1/teams/{team_id}/members/bulk_delete",
"/management/v1/teams/{team_id}/members/bulk_update",
"/team/member_update",
"/team/{team_id}/member/{user_id}/reset_spend",
"/team/permissions_list",

View file

@ -31,6 +31,7 @@ _PROXY_ADMIN_VIEW_ONLY_BLOCKED_ROUTES: Final = frozenset(
# team
"/team/new",
"/management/v1/teams/{team_id}/members/bulk_delete",
"/management/v1/teams/{team_id}/members/bulk_update",
"/team/update",
"/team/delete",
"/team/block",
@ -767,6 +768,7 @@ class RouteChecks:
"/user/bulk_update",
"/team/new",
"/management/v1/teams/{team_id}/members/bulk_delete",
"/management/v1/teams/{team_id}/members/bulk_update",
"/team/update",
"/team/delete",
"/model/new",

View file

@ -78,3 +78,27 @@ def get_budget_reset_time(budget_duration: str) -> datetime:
`BudgetResetSettings` by injection (creation/update endpoints, startup backfill).
"""
return compute_budget_reset_at(budget_duration, get_budget_reset_settings())
def _is_persistable_budget_duration(budget_duration: str) -> bool:
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
try:
if duration_in_seconds(budget_duration) <= 0:
return False
get_budget_reset_time(budget_duration=budget_duration)
except (ValueError, OverflowError):
return False
return True
def budget_duration_error(budget_duration: str | None) -> str | None:
"""Why `budget_duration` cannot be persisted, or None when it is usable.
A non-positive duration resolves to a reset time of "now", which leaves the row
permanently due: the reset job re-reads it every tick and, once enough of them
exist, they fill each batch and starve every other tenant's reset.
"""
if budget_duration is None or _is_persistable_budget_duration(budget_duration):
return None
return f"Invalid budget_duration '{budget_duration}'. Use a format like '1h', '24h', '7d', or '30d'."

View file

@ -1,5 +1,6 @@
import math
from collections.abc import Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Optional, Union
from fastapi import HTTPException, status
@ -33,23 +34,11 @@ def validate_budget_duration(budget_duration: str | None, status_code: int = 400
enough of them exist, they fill each batch and starve every other tenant's
reset.
"""
if budget_duration is None:
return
from litellm.proxy.common_utils.timezone_utils import budget_duration_error
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
try:
if duration_in_seconds(budget_duration) <= 0:
raise ValueError("budget_duration must be positive")
get_budget_reset_time(budget_duration=budget_duration)
except (ValueError, OverflowError):
raise HTTPException(
status_code=status_code,
detail={
"error": f"Invalid budget_duration '{budget_duration}'. Use a format like '1h', '24h', '7d', or '30d'."
},
)
error: Final = budget_duration_error(budget_duration)
if error is not None:
raise HTTPException(status_code=status_code, detail={"error": error})
from litellm._logging import verbose_proxy_logger
@ -490,6 +479,33 @@ _TEAM_MEMBER_BUDGET_LIMIT_FIELDS: Final = (
)
MEMBER_BUDGET_PATCH_FIELDS: Final = MappingProxyType(
{
"max_budget_in_team": "max_budget",
"tpm_limit": "tpm_limit",
"rpm_limit": "rpm_limit",
"budget_duration": "budget_duration",
"allowed_models": "allowed_models",
}
)
def _prisma_value(value: object) -> object:
return list(value) if isinstance(value, tuple) else value
def member_budget_patch(source: BaseModel) -> dict[str, Any]:
"""Map the per-member limit fields a request actually set to their budget-table
columns (merge-patch: a sent value updates, an explicit null clears, an absent
field is left untouched)."""
provided: Final = source.model_dump(exclude_unset=True)
return {
column: _prisma_value(provided[request_field])
for request_field, column in MEMBER_BUDGET_PATCH_FIELDS.items()
if request_field in provided
}
def _is_set_budget_value(value: object) -> bool:
if value is None:
return False
@ -513,6 +529,7 @@ async def _upsert_budget_and_membership(
user_api_key_dict: UserAPIKeyAuth,
budget_patch: dict[str, Any],
team_default_budget_id: str | None = None,
shared_budget_ids: frozenset[str] | None = None,
):
"""
Apply a merge-patch of per-member budget fields to a team membership.
@ -527,6 +544,10 @@ async def _upsert_budget_and_membership(
(from team metadata.team_member_budget_id). When the membership still
points at it, we clone-on-write so editing one member's budget does not
mutate the shared default that every other member points at.
``shared_budget_ids`` extends that protection to any other row more than one
membership points at, which a caller patching several members at once has
already counted; a row listed there is cloned rather than written in place.
"""
if not budget_patch:
return
@ -538,10 +559,8 @@ async def _upsert_budget_and_membership(
get_budget_reset_time(budget_duration=duration) if duration is not None else None
)
is_shared_default: Final = (
existing_budget_id is not None
and team_default_budget_id is not None
and existing_budget_id == team_default_budget_id
is_shared_default: Final = existing_budget_id is not None and (
existing_budget_id == team_default_budget_id or existing_budget_id in (shared_budget_ids or frozenset())
)
async def _disconnect():

View file

@ -1,4 +1,4 @@
"""`POST /management/v1/teams/{team_id}/members/bulk_delete`."""
"""`POST /management/v1/teams/{team_id}/members/bulk_delete` and `.../members/bulk_update`."""
from typing import Annotated, Final
@ -9,12 +9,15 @@ from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.list_api.common import PROBLEM_TYPE_BASE, ManagementProblem, reject_unknown_query_params
from litellm.proxy.management_endpoints.management_v1.common import MANAGEMENT_V1_PREFIX
from litellm.proxy.management_helpers.bulk_team_member_budgets import bulk_update_team_member_budgets
from litellm.proxy.management_helpers.bulk_user_deletion import bulk_remove_team_members
from litellm.proxy.management_helpers.utils import (
management_endpoint_wrapper, # pyright: ignore[reportUnknownVariableType] # legacy decorator is untyped
)
from litellm.types.proxy.management_endpoints.management_v1 import ProblemDetail
from litellm.types.proxy.management_endpoints.team_endpoints import (
BulkTeamMemberBudgetUpdateRequest,
BulkTeamMemberBudgetUpdateResponse,
BulkTeamMemberDeleteRequest,
BulkTeamMemberDeleteResponse,
)
@ -92,3 +95,80 @@ async def bulk_delete_team_members_action(
detail="Failed to remove team members.",
)
)
@router.post(
"/teams/{team_id}/members/bulk_update",
tags=["team management"], # mutable-ok: FastAPI types `tags` as list[str], not Sequence
dependencies=(Depends(user_api_key_auth), Depends(reject_unknown_query_params)),
response_model=BulkTeamMemberBudgetUpdateResponse,
)
@management_endpoint_wrapper
async def bulk_update_team_member_budgets_action(
team_id: str,
data: BulkTeamMemberBudgetUpdateRequest,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
) -> BulkTeamMemberBudgetUpdateResponse:
"""
Set per-member limits for up to 500 members of one team in one call. Same
authorization and member addressing as `/team/member_update`: proxy admins, the team's
admins, and admins of the team's organization, with each member named by exactly one of
`user_id` or `user_email`. Unknown body fields are a 422 and an unknown team is a 404.
Each row is a merge patch of that member's limits: a field left out is untouched, a
field sent as null is cleared, and clearing the last limit drops the member back to the
team default. A budget row shared by several memberships, the team default included, is
copied for the member being patched rather than written in place, so one member's new
cap never lands on anybody else.
`data` holds one result per requested member, in request order, carrying the limits in
force after the write. A row is `success: false` with an `error` when it names nobody on
the team or repeats an earlier row. Roles are not part of this route; `/team/member_update`
still owns them.
Example curl:
```
curl --location 'http://0.0.0.0:4000/management/v1/teams/team-1/members/bulk_update' \
--header 'Authorization: Bearer sk-1234' \
--header 'Content-Type: application/json' \
--data '{"members": [{"user_id": "user-1", "max_budget_in_team": 10}, {"user_email": "user-2@example.com", "max_budget_in_team": 10, "budget_duration": "30d"}]}'
```
"""
try:
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
if prisma_client is None:
raise ManagementProblem(
ProblemDetail(
type=f"{PROBLEM_TYPE_BASE}database-not-connected",
title="Database not connected",
status=503,
detail=CommonProxyErrors.db_not_connected_error.value,
)
)
results: Final = await bulk_update_team_member_budgets(
team_id=team_id,
data=data,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
return BulkTeamMemberBudgetUpdateResponse(data=results)
except ManagementProblem:
raise
except Exception as e: # noqa: BLE001 # a driver error answers as a problem document, not the OpenAI error shape
verbose_proxy_logger.exception(
"litellm.proxy.management_endpoints.management_v1.teams.bulk_update_team_member_budgets_action(): "
"Exception occured - %s",
e,
)
raise ManagementProblem(
ProblemDetail(
type=f"{PROBLEM_TYPE_BASE}internal-server-error",
title="Internal server error",
status=500,
detail="Failed to update team member budgets.",
)
)

View file

@ -129,6 +129,7 @@ from litellm.proxy.management_endpoints.common_utils import (
_update_metadata_fields,
_upsert_budget_and_membership,
_user_has_admin_view,
member_budget_patch,
validate_budget_duration,
validate_team_model_max_budget,
)
@ -3686,27 +3687,6 @@ async def team_member_delete(
return existing_team_row
_MEMBER_BUDGET_PATCH_FIELDS: Final = {
"max_budget_in_team": "max_budget",
"tpm_limit": "tpm_limit",
"rpm_limit": "rpm_limit",
"budget_duration": "budget_duration",
"allowed_models": "allowed_models",
}
def _build_member_budget_patch(data: TeamMemberUpdateRequest) -> dict[str, object]:
"""Map the budget fields the request actually set (merge-patch: a sent
value updates, an explicit null clears, an absent field is left untouched)
to their budget-table columns."""
provided: Final = data.model_dump(exclude_unset=True)
return {
column: provided[request_field]
for request_field, column in _MEMBER_BUDGET_PATCH_FIELDS.items()
if request_field in provided
}
@router.post(
"/team/member_update",
tags=["team management"],
@ -3812,7 +3792,7 @@ async def team_member_update(
team_default_budget_id = raw_default_budget_id
### upsert new budget
budget_patch: Final = _build_member_budget_patch(data)
budget_patch: Final = member_budget_patch(data)
async with prisma_client.tx() as tx:
await _upsert_budget_and_membership(
tx=tx,

View file

@ -0,0 +1,192 @@
"""Batched per-member limit writes behind `POST /management/v1/teams/{team_id}/members/bulk_update`.
Every read runs on the writer inside the batch transaction, so the write plan can never be
built from a lagging read replica. Any budget row that more than one membership points at,
the team's shared default included, is cloned before it is written, so raising one member's
cap never moves another member's.
"""
from collections.abc import Sequence
from datetime import timedelta
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from litellm.proxy._types import LiteLLM_TeamTable, LitellmUserRoles, Member, UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import invalidate_team_member_spend_state
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient
from litellm.proxy.management_endpoints.common_utils import (
_is_user_org_admin_for_team, # pyright: ignore[reportPrivateUsage] # same check /team/member_update uses
_is_user_team_admin, # pyright: ignore[reportPrivateUsage] # same check /team/member_update uses
_upsert_budget_and_membership, # pyright: ignore[reportPrivateUsage] # the single-member write, shared so the two surfaces cannot drift
member_budget_patch,
)
from litellm.proxy.management_helpers.bulk_user_deletion import (
_duplicate_member_indexes, # pyright: ignore[reportPrivateUsage] # same duplicate rule as members/bulk_delete
_eq_filter, # pyright: ignore[reportPrivateUsage] # same prisma filter shape as members/bulk_delete
_forbidden, # pyright: ignore[reportPrivateUsage] # same problem shape as members/bulk_delete
_in_filter, # pyright: ignore[reportPrivateUsage] # same prisma filter shape as members/bulk_delete
_team_not_found, # pyright: ignore[reportPrivateUsage] # same problem shape as members/bulk_delete
_team_users_filter, # pyright: ignore[reportPrivateUsage] # same prisma filter shape as members/bulk_delete
)
from litellm.proxy.utils import PrismaClient
from litellm.repositories.team_repository import TeamRepository
from litellm.types.proxy.management_endpoints.team_endpoints import (
BulkTeamMemberBudgetUpdateRequest,
TeamMemberBudgetPatch,
TeamMemberBudgetUpdateResult,
)
if TYPE_CHECKING:
from prisma import Prisma
from prisma import models as prisma_models
from litellm.repositories.prisma_protocols import TableActions
_BATCH_TX_TIMEOUT: Final = timedelta(seconds=60)
_NO_METADATA: Final = MappingProxyType({})
_WITH_BUDGET: Final = MappingProxyType({"litellm_budget_table": True})
def _membership_tx_db(tx: "Prisma") -> "TableActions[prisma_models.LiteLLM_TeamMembership]":
return tx.litellm_teammembership # pyright: ignore[reportReturnType] # TableActions widens the generated inputs to Mapping, as the repositories do
def _budget_tx_db(tx: "Prisma") -> "TableActions[prisma_models.LiteLLM_BudgetTable]":
return tx.litellm_budgettable # pyright: ignore[reportReturnType] # TableActions widens the generated inputs to Mapping, as the repositories do
def _roster_user_id(member: TeamMemberBudgetPatch, roster: Sequence[Member]) -> str | None:
"""The team member this row addresses, or None when it names nobody on the team."""
if member.user_id is not None:
return member.user_id if any(m.user_id == member.user_id for m in roster) else None
return next((m.user_id for m in roster if m.user_email is not None and m.user_email == member.user_email), None)
def _team_default_budget_id(team: LiteLLM_TeamTable) -> str | None:
raw: Final = (team.metadata or _NO_METADATA).get("team_member_budget_id")
return raw if isinstance(raw, str) else None
async def _shared_budget_ids(tx: "Prisma", budget_ids: frozenset[str]) -> frozenset[str]:
"""The rows in ``budget_ids`` more than one membership points at, counted across every
team so a row shared with another team is protected too."""
if not budget_ids:
return frozenset()
rows: Final = await _membership_tx_db(tx).find_many(where=_in_filter("budget_id", budget_ids))
return frozenset(budget_id for budget_id in budget_ids if sum(1 for row in rows if row.budget_id == budget_id) > 1)
def _result(
member: TeamMemberBudgetPatch,
user_id: str | None,
error: str | None,
budget_of: "MappingProxyType[str, prisma_models.LiteLLM_BudgetTable | None]",
team_default_max_budget: float | None,
) -> TeamMemberBudgetUpdateResult:
if error is not None or user_id is None:
return TeamMemberBudgetUpdateResult(
user_id=member.user_id,
user_email=member.user_email,
success=False,
error=error or "User not found in team",
)
budget: Final = budget_of.get(user_id)
own_max_budget: Final = budget.max_budget if budget is not None else None
inherits: Final = own_max_budget is None and team_default_max_budget is not None and team_default_max_budget > 0
return TeamMemberBudgetUpdateResult(
user_id=user_id,
user_email=member.user_email,
success=True,
budget_id=budget.budget_id if budget is not None else None,
max_budget=team_default_max_budget if inherits else own_max_budget,
max_budget_source=("team_default" if inherits else "member" if own_max_budget is not None else None),
tpm_limit=budget.tpm_limit if budget is not None else None,
rpm_limit=budget.rpm_limit if budget is not None else None,
budget_duration=budget.budget_duration if budget is not None else None,
allowed_models=tuple(budget.allowed_models) if budget is not None else None,
)
async def bulk_update_team_member_budgets(
team_id: str,
data: BulkTeamMemberBudgetUpdateRequest,
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
) -> tuple[TeamMemberBudgetUpdateResult, ...]:
"""Apply one merge patch of per-member limits per requested member, in one transaction."""
team: Final = await TeamRepository(WriterPinnedClient(prisma_client.db)).find_by_id(team_id)
if team is None:
raise _team_not_found(team_id)
if (
user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value
and not _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team)
and not await _is_user_org_admin_for_team(user_api_key_dict=user_api_key_dict, team_obj=team)
):
raise _forbidden(
"Call not allowed. User not proxy admin OR team admin OR org admin for this team. "
f"route='/management/v1/teams/{team_id}/members/bulk_update'"
)
roster: Final = team.members_with_roles or ()
named: Final = tuple(_roster_user_id(member, roster) for member in data.members)
duplicates: Final = _duplicate_member_indexes(data.members) | frozenset(
index for index, user_id in enumerate(named) if user_id is not None and user_id in named[:index]
)
applied: Final = tuple(
(index, user_id) for index, user_id in enumerate(named) if user_id is not None and index not in duplicates
)
if not applied:
return tuple(
_result(
member, None, "Duplicate member in request" if index in duplicates else None, MappingProxyType({}), None
)
for index, member in enumerate(data.members)
)
user_ids: Final = sorted(user_id for _, user_id in applied)
default_budget_id: Final = _team_default_budget_id(team)
team_members_filter: Final = _team_users_filter(team_id, user_ids)
async with prisma_client.tx(timeout=_BATCH_TX_TIMEOUT) as tx:
memberships: Final = await _membership_tx_db(tx).find_many(where=team_members_filter)
budget_id_of: Final = MappingProxyType({m.user_id: m.budget_id for m in memberships})
shared: Final = await _shared_budget_ids(
tx, frozenset(budget_id for budget_id in budget_id_of.values() if budget_id is not None)
)
for index, user_id in applied:
await _upsert_budget_and_membership(
tx=tx,
team_id=team_id,
user_id=user_id,
existing_budget_id=budget_id_of.get(user_id),
user_api_key_dict=user_api_key_dict,
budget_patch=member_budget_patch(data.members[index]),
team_default_budget_id=default_budget_id,
shared_budget_ids=shared,
)
written: Final = await _membership_tx_db(tx).find_many(where=team_members_filter, include=_WITH_BUDGET)
team_default: Final = (
await _budget_tx_db(tx).find_unique(where=_eq_filter("budget_id", default_budget_id))
if default_budget_id is not None
else None
)
for user_id in user_ids:
await invalidate_team_member_spend_state(
user_id=user_id, team_id=team_id, user_api_key_cache=user_api_key_cache
)
budget_of: Final = MappingProxyType({m.user_id: m.litellm_budget_table for m in written})
return tuple(
_result(
member,
named[index],
"Duplicate member in request" if index in duplicates else None,
budget_of,
team_default.max_budget if team_default is not None else None,
)
for index, member in enumerate(data.members)
)

View file

@ -1,6 +1,6 @@
from typing import Any, Final, Literal
from pydantic import BaseModel, ConfigDict, Field, model_validator
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from litellm.proxy._types import (
KeyManagementRoutes,
@ -10,12 +10,15 @@ from litellm.proxy._types import (
Member,
MemberDeleteRequest,
)
from litellm.proxy.common_utils.timezone_utils import budget_duration_error
from litellm.types.proxy.management_endpoints.management_v1 import ResourceResponse
TeamIdSearchMatch = Literal["exact", "prefix"]
MAX_BULK_TEAM_MEMBER_DELETES: Final = 500
MAX_BULK_TEAM_MEMBER_BUDGET_UPDATES: Final = 500
class GetTeamMemberPermissionsRequest(BaseModel):
"""Request to get the team member permissions for a team"""
@ -123,7 +126,7 @@ class BulkTeamMemberAddResponse(BaseModel):
class TeamMemberRef(MemberDeleteRequest):
"""One member to remove, named by exactly one of `user_id` or `user_email`."""
"""One member, named by exactly one of `user_id` or `user_email`."""
model_config = ConfigDict(extra="forbid")
@ -155,6 +158,55 @@ class BulkTeamMemberDeleteResponse(ResourceResponse[tuple[TeamMemberDeleteResult
"""`{data: [...]}` with one `TeamMemberDeleteResult` per requested member, in request order."""
class TeamMemberBudgetPatch(TeamMemberRef):
"""One member's per-member limits, merge-patch style: a field left out of the row is
untouched, a field sent as null is cleared, and clearing the last limit drops the
member back to the team default."""
max_budget_in_team: float | None = None
tpm_limit: int | None = None
rpm_limit: int | None = None
budget_duration: str | None = None
allowed_models: tuple[str, ...] | None = None
@field_validator("budget_duration")
@classmethod
def persistable_budget_duration(cls, value: str | None) -> str | None:
error: Final = budget_duration_error(value)
if error is not None:
raise ValueError(error)
return value
class BulkTeamMemberBudgetUpdateRequest(BaseModel):
"""Body of `POST /management/v1/teams/{team_id}/members/bulk_update`."""
model_config = ConfigDict(extra="forbid")
members: tuple[TeamMemberBudgetPatch, ...] = Field(min_length=1, max_length=MAX_BULK_TEAM_MEMBER_BUDGET_UPDATES)
class TeamMemberBudgetUpdateResult(BaseModel):
"""Outcome for one requested member, in request order, carrying the limits in force
after the write rather than the ones that were asked for."""
user_id: str | None = None
user_email: str | None = None
success: bool
error: str | None = None
budget_id: str | None = None
max_budget: float | None = None
max_budget_source: Literal["member", "team_default"] | None = None
tpm_limit: int | None = None
rpm_limit: int | None = None
budget_duration: str | None = None
allowed_models: tuple[str, ...] | None = None
class BulkTeamMemberBudgetUpdateResponse(ResourceResponse[tuple[TeamMemberBudgetUpdateResult, ...]]):
"""`{data: [...]}` with one `TeamMemberBudgetUpdateResult` per requested member, in request order."""
class TeamMemberInfoResponse(LiteLLM_TeamMembership):
"""Response for GET /team/{team_id}/members/me — caller's own membership row."""

View file

@ -0,0 +1,772 @@
"""`POST /management/v1/teams/{team_id}/members/bulk_update`: the per-member limit writes and the
HTTP contract around them.
The in-memory Prisma here follows the one in
`tests/test_litellm/proxy/management_helpers/test_bulk_user_deletion.py`, extended with the budget
table and the membership/budget relation the bulk budget writer needs.
"""
import copy
from collections.abc import Mapping, Sequence
from contextlib import asynccontextmanager
from datetime import datetime, timedelta, timezone
from typing import Final
import pytest
from fastapi import FastAPI, Request
from fastapi.exceptions import RequestValidationError
from fastapi.testclient import TestClient
from pydantic import BaseModel, ConfigDict, Field
from litellm.proxy._types import LiteLLM_TeamTable, LitellmUserRoles, Member, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.user_api_key_cache import (
UserApiKeyCache,
team_membership_auth_cache_key,
team_membership_reservation_cache_key,
)
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
from litellm.proxy.list_api.common import ManagementProblem, problem_response, request_validation_problem
from litellm.proxy.management_endpoints.management_v1 import router
from litellm.proxy.management_endpoints.management_v1.common import MANAGEMENT_V1_PREFIX
from litellm.proxy.management_helpers.bulk_team_member_budgets import bulk_update_team_member_budgets
from litellm.types.proxy.management_endpoints.team_endpoints import (
MAX_BULK_TEAM_MEMBER_BUDGET_UPDATES,
BulkTeamMemberBudgetUpdateRequest,
TeamMemberBudgetUpdateResult,
)
ADMIN: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin")
OUTSIDER: Final = UserAPIKeyAuth(user_id="outsider", user_role=LitellmUserRoles.INTERNAL_USER)
TEAM_ID: Final = "t1"
class _BudgetRow(BaseModel):
"""A `LiteLLM_BudgetTable` row, carrying every column the merge patch reads or writes."""
model_config = ConfigDict(extra="allow")
budget_id: str
max_budget: float | None = None
soft_budget: float | None = None
max_parallel_requests: int | None = None
tpm_limit: int | None = None
rpm_limit: int | None = None
model_max_budget: Mapping[str, object] | None = None
budget_duration: str | None = None
budget_reset_at: datetime | None = None
allowed_models: list[str] = Field(default_factory=list)
created_by: str | None = None
updated_by: str | None = None
class _MembershipRow(BaseModel):
"""A `LiteLLM_TeamMembership` row; `litellm_budget_table` is only filled on an `include` read."""
model_config = ConfigDict(extra="allow")
user_id: str
team_id: str
budget_id: str | None = None
litellm_budget_table: _BudgetRow | None = None
def _wanted(where: Mapping[str, object], field: str) -> set[str] | None:
clause: Final = where.get(field)
if isinstance(clause, dict) and "in" in clause:
return set(clause["in"])
if isinstance(clause, str):
return {clause}
return None
def _matches(row: Mapping[str, object], where: Mapping[str, object]) -> bool:
return all((wanted := _wanted(where, field)) is not None and row.get(field) in wanted for field in where)
class _BudgetTable:
def __init__(self, budgets: Sequence[_BudgetRow]) -> None:
self.rows: dict[str, _BudgetRow] = {b.budget_id: b for b in budgets}
async def find_unique(self, where: Mapping[str, str]) -> _BudgetRow | None:
return self.rows.get(where["budget_id"])
async def update(self, where: Mapping[str, str], data: Mapping[str, object]) -> _BudgetRow:
row: Final = self.rows[where["budget_id"]]
updated: Final = row.model_copy(update=dict(data))
self.rows[row.budget_id] = updated
return updated
async def create(self, data: Mapping[str, object], include: Mapping[str, bool] | None = None) -> _BudgetRow:
budget_id: Final = f"new-budget-{len(self.rows) + 1}"
row: Final = _BudgetRow.model_validate({**data, "budget_id": budget_id})
self.rows[budget_id] = row
return row
class _MembershipTable:
def __init__(self, budgets: _BudgetTable, memberships: Sequence[_MembershipRow]) -> None:
self._budgets = budgets
self.rows: list[_MembershipRow] = list(memberships)
def _index_of(self, user_id: str, team_id: str) -> int | None:
return next(
(i for i, r in enumerate(self.rows) if r.user_id == user_id and r.team_id == team_id),
None,
)
async def find_many(
self, where: Mapping[str, object], include: Mapping[str, bool] | None = None
) -> list[_MembershipRow]:
matched: Final = [r for r in self.rows if _matches(r.model_dump(), where)]
if not include:
return matched
return [
r.model_copy(update={"litellm_budget_table": self._budgets.rows.get(r.budget_id or "")}) for r in matched
]
async def update(self, where: Mapping[str, Mapping[str, str]], data: Mapping[str, object]) -> _MembershipRow:
key: Final = where["user_id_team_id"]
index: Final = self._index_of(key["user_id"], key["team_id"])
assert index is not None, f"no membership row for {key}"
relation: Final = data.get("litellm_budget_table")
if isinstance(relation, dict) and relation.get("disconnect"):
self.rows[index] = self.rows[index].model_copy(update={"budget_id": None})
return self.rows[index]
async def upsert(self, where: Mapping[str, Mapping[str, str]], data: Mapping[str, object]) -> _MembershipRow:
key: Final = where["user_id_team_id"]
budget_id: Final = data["update"]["litellm_budget_table"]["connect"]["budget_id"]
index: Final = self._index_of(key["user_id"], key["team_id"])
if index is None:
self.rows.append(_MembershipRow(user_id=key["user_id"], team_id=key["team_id"], budget_id=budget_id))
return self.rows[-1]
self.rows[index] = self.rows[index].model_copy(update={"budget_id": budget_id})
return self.rows[index]
class _TeamTable:
"""`find_many` and `create` are what `RoutingPrismaWrapper` keys read routing off, so a fake
table without them would silently never route and pass a reader-staleness test on the writer."""
def __init__(self, teams: Sequence[LiteLLM_TeamTable]) -> None:
self.rows: dict[str, LiteLLM_TeamTable] = {t.team_id: t for t in teams}
async def find_unique(self, where: Mapping[str, str]) -> LiteLLM_TeamTable | None:
return self.rows.get(where["team_id"])
async def find_many(self, where: Mapping[str, object] | None = None) -> list[LiteLLM_TeamTable]:
return [t for t in self.rows.values() if where is None or _matches(t.model_dump(), where)]
async def create(self, data: Mapping[str, object]) -> LiteLLM_TeamTable:
row: Final = LiteLLM_TeamTable.model_validate(dict(data))
self.rows[row.team_id] = row
return row
class _Db:
def __init__(
self,
teams: Sequence[LiteLLM_TeamTable],
memberships: Sequence[_MembershipRow],
budgets: Sequence[_BudgetRow],
) -> None:
self.litellm_teamtable = _TeamTable(teams)
self.litellm_budgettable = _BudgetTable(budgets)
self.litellm_teammembership = _MembershipTable(self.litellm_budgettable, memberships)
class _FakePrisma:
def __init__(
self,
teams: Sequence[LiteLLM_TeamTable] = (),
memberships: Sequence[_MembershipRow] = (),
budgets: Sequence[_BudgetRow] = (),
) -> None:
self.db = _Db(teams, memberships, budgets)
@asynccontextmanager
async def tx(self, *, timeout: object = None):
snapshot: Final = copy.deepcopy(self.db)
try:
yield self.db
except BaseException:
self.db = snapshot
raise
class _ReplicatedPrisma:
"""A client whose reads route to a lagging replica, as a proxy with `DATABASE_URL_READ_REPLICA` does."""
def __init__(self, writer: _FakePrisma, reader: _FakePrisma) -> None:
self._writer = writer
self.db = RoutingPrismaWrapper(writer=writer.db, reader=reader.db) # pyright: ignore[reportArgumentType] # fake dbs stand in for PrismaWrapper
def tx(self, *, timeout: object = None):
return self._writer.tx(timeout=timeout)
class _UnreachableDb:
"""A `.db` whose every table access fails, as one behind a dropped connection does."""
def __getattr__(self, name: str) -> object:
raise RuntimeError("connection reset by peer")
class _UnreachablePrisma:
def __init__(self) -> None:
self.db = _UnreachableDb()
def _team(
*members: str,
team_id: str = TEAM_ID,
default_budget_id: str | None = None,
admins: Sequence[str] = (),
) -> LiteLLM_TeamTable:
return LiteLLM_TeamTable(
team_id=team_id,
metadata={"team_member_budget_id": default_budget_id} if default_budget_id else {},
members_with_roles=[
Member(user_id=m, user_email=f"{m}@example.com", role="admin" if m in admins else "user") for m in members
],
)
def _membership(user_id: str, budget_id: str | None = None, team_id: str = TEAM_ID) -> _MembershipRow:
return _MembershipRow(user_id=user_id, team_id=team_id, budget_id=budget_id)
def _budget(
budget_id: str,
*,
max_budget: float | None = None,
tpm_limit: int | None = None,
rpm_limit: int | None = None,
budget_duration: str | None = None,
) -> _BudgetRow:
return _BudgetRow(
budget_id=budget_id,
max_budget=max_budget,
tpm_limit=tpm_limit,
rpm_limit=rpm_limit,
budget_duration=budget_duration,
)
async def _bulk_update(
prisma: _FakePrisma | _ReplicatedPrisma,
members: Sequence[Mapping[str, object]],
team_id: str = TEAM_ID,
caller: UserAPIKeyAuth = ADMIN,
cache: UserApiKeyCache | None = None,
) -> tuple[TeamMemberBudgetUpdateResult, ...]:
return await bulk_update_team_member_budgets(
team_id=team_id,
data=BulkTeamMemberBudgetUpdateRequest.model_validate({"members": list(members)}),
user_api_key_dict=caller,
prisma_client=prisma, # pyright: ignore[reportArgumentType] # fake stands in for PrismaClient
user_api_key_cache=cache or UserApiKeyCache(),
)
def _budget_id_of(prisma: _FakePrisma, user_id: str, team_id: str = TEAM_ID) -> str | None:
row: Final = next(r for r in prisma.db.litellm_teammembership.rows if r.user_id == user_id and r.team_id == team_id)
return row.budget_id
def _budget_of(prisma: _FakePrisma, user_id: str, team_id: str = TEAM_ID) -> _BudgetRow:
budget_id: Final = _budget_id_of(prisma, user_id, team_id)
assert budget_id is not None, f"{user_id} has no budget"
return prisma.db.litellm_budgettable.rows[budget_id]
def _seeded_cache(*user_ids: str, team_id: str = TEAM_ID) -> UserApiKeyCache:
cache: Final = UserApiKeyCache()
for user_id in user_ids:
cache.set_cache(key=team_membership_auth_cache_key(team_id=team_id, user_id=user_id), value={"cap": "old"})
cache.set_cache(
key=team_membership_reservation_cache_key(user_id=user_id, team_id=team_id), value={"cap": "old"}
)
return cache
def _cached_keys(cache: UserApiKeyCache, user_id: str, team_id: str = TEAM_ID) -> tuple[object, object]:
return (
cache.get_cache(key=team_membership_auth_cache_key(team_id=team_id, user_id=user_id)),
cache.get_cache(key=team_membership_reservation_cache_key(user_id=user_id, team_id=team_id)),
)
@pytest.mark.asyncio
async def test_patching_one_member_of_a_shared_budget_row_forks_it_and_leaves_the_other_member_untouched():
prisma = _FakePrisma(
teams=[_team("m1", "m2")],
memberships=[_membership("m1", "shared-b"), _membership("m2", "shared-b")],
budgets=[_budget("shared-b", max_budget=100.0, tpm_limit=900)],
)
results = await _bulk_update(prisma, [{"user_id": "m1", "max_budget_in_team": 50}])
assert [(r.user_id, r.success, r.max_budget) for r in results] == [("m1", True, 50.0)]
assert _budget_id_of(prisma, "m1") not in (None, "shared-b")
assert (_budget_of(prisma, "m1").max_budget, _budget_of(prisma, "m1").tpm_limit) == (50.0, 900)
assert _budget_id_of(prisma, "m2") == "shared-b"
assert prisma.db.litellm_budgettable.rows["shared-b"].max_budget == 100.0
assert results[0].budget_id == _budget_id_of(prisma, "m1")
@pytest.mark.asyncio
async def test_patching_members_of_the_team_default_budget_gives_each_their_own_row_and_leaves_the_default_alone():
prisma = _FakePrisma(
teams=[_team("m1", "m2", "m3", default_budget_id="team-default")],
memberships=[
_membership("m1", "team-default"),
_membership("m2", "team-default"),
_membership("m3", "team-default"),
],
budgets=[_budget("team-default", max_budget=25.0, tpm_limit=1000)],
)
results = await _bulk_update(
prisma,
[{"user_id": "m1", "max_budget_in_team": 5}, {"user_id": "m2", "max_budget_in_team": 7}],
)
assert [r.success for r in results] == [True, True]
default = prisma.db.litellm_budgettable.rows["team-default"]
assert (default.max_budget, default.tpm_limit) == (25.0, 1000)
assert _budget_id_of(prisma, "m3") == "team-default"
patched = (_budget_id_of(prisma, "m1"), _budget_id_of(prisma, "m2"))
assert len(set(patched)) == 2 and "team-default" not in patched
assert (_budget_of(prisma, "m1").max_budget, _budget_of(prisma, "m1").tpm_limit) == (5.0, 1000)
assert (_budget_of(prisma, "m2").max_budget, _budget_of(prisma, "m2").tpm_limit) == (7.0, 1000)
@pytest.mark.asyncio
async def test_the_team_default_row_is_forked_even_when_only_one_membership_points_at_it():
prisma = _FakePrisma(
teams=[_team("m1", "m2", default_budget_id="team-default")],
memberships=[_membership("m1", "team-default")],
budgets=[_budget("team-default", max_budget=25.0, tpm_limit=1000)],
)
results = await _bulk_update(prisma, [{"user_id": "m1", "max_budget_in_team": 5}])
assert [(r.success, r.max_budget, r.tpm_limit) for r in results] == [(True, 5.0, 1000)]
default = prisma.db.litellm_budgettable.rows["team-default"]
assert (default.max_budget, default.tpm_limit) == (25.0, 1000)
assert _budget_id_of(prisma, "m1") not in (None, "team-default")
@pytest.mark.asyncio
async def test_a_budget_row_only_one_member_points_at_is_updated_in_place():
prisma = _FakePrisma(
teams=[_team("m1", "m2", default_budget_id="team-default")],
memberships=[_membership("m1", "priv-m1"), _membership("m2", "team-default")],
budgets=[_budget("team-default", max_budget=25.0), _budget("priv-m1", max_budget=10.0, tpm_limit=5)],
)
results = await _bulk_update(prisma, [{"user_id": "m1", "max_budget_in_team": 20}])
assert [(r.success, r.budget_id, r.max_budget) for r in results] == [(True, "priv-m1", 20.0)]
assert set(prisma.db.litellm_budgettable.rows) == {"team-default", "priv-m1"}
assert _budget_id_of(prisma, "m1") == "priv-m1"
assert (_budget_of(prisma, "m1").max_budget, _budget_of(prisma, "m1").tpm_limit) == (20.0, 5)
@pytest.mark.asyncio
async def test_an_omitted_field_is_kept_an_explicit_null_clears_it_and_clearing_the_last_limit_disconnects():
prisma = _FakePrisma(
teams=[_team("m1")],
memberships=[_membership("m1", "priv-m1")],
budgets=[_budget("priv-m1", max_budget=10.0, tpm_limit=5, rpm_limit=7)],
)
kept = await _bulk_update(prisma, [{"user_id": "m1", "rpm_limit": 9}])
assert (kept[0].max_budget, kept[0].tpm_limit, kept[0].rpm_limit) == (10.0, 5, 9)
cleared = await _bulk_update(prisma, [{"user_id": "m1", "tpm_limit": None}])
assert (cleared[0].max_budget, cleared[0].tpm_limit, cleared[0].rpm_limit) == (10.0, None, 9)
assert _budget_id_of(prisma, "m1") == "priv-m1"
emptied = await _bulk_update(prisma, [{"user_id": "m1", "max_budget_in_team": None, "rpm_limit": None}])
assert (emptied[0].success, emptied[0].budget_id, emptied[0].max_budget) == (True, None, None)
assert _budget_id_of(prisma, "m1") is None
@pytest.mark.asyncio
async def test_budget_duration_seeds_a_reset_time_derived_from_the_duration_and_clearing_it_clears_the_reset():
prisma = _FakePrisma(
teams=[_team("m1", "m2")],
memberships=[_membership("m1", "priv-m1"), _membership("m2", "priv-m2")],
budgets=[_budget("priv-m1", max_budget=10.0), _budget("priv-m2", max_budget=10.0)],
)
before = datetime.now(timezone.utc)
await _bulk_update(
prisma,
[{"user_id": "m1", "budget_duration": "2d"}, {"user_id": "m2", "budget_duration": "5d"}],
)
two_day = _budget_of(prisma, "m1").budget_reset_at
five_day = _budget_of(prisma, "m2").budget_reset_at
assert two_day is not None and five_day is not None
assert before < two_day <= before + timedelta(days=2)
assert before + timedelta(days=4) - timedelta(seconds=1) < five_day <= before + timedelta(days=5)
assert five_day - two_day == timedelta(days=3)
await _bulk_update(prisma, [{"user_id": "m1", "budget_duration": None}])
assert _budget_of(prisma, "m1").budget_reset_at is None
assert _budget_of(prisma, "m1").budget_duration is None
assert _budget_of(prisma, "m1").max_budget == 10.0
@pytest.mark.asyncio
async def test_a_member_named_twice_is_written_once_and_the_later_rows_report_the_duplicate():
prisma = _FakePrisma(
teams=[_team("m1")],
memberships=[_membership("m1", "priv-m1")],
budgets=[_budget("priv-m1", max_budget=1.0)],
)
results = await _bulk_update(
prisma,
[
{"user_id": "m1", "max_budget_in_team": 10},
{"user_id": "m1", "max_budget_in_team": 20},
{"user_email": "m1@example.com", "max_budget_in_team": 30},
],
)
assert [(r.success, r.error) for r in results] == [
(True, None),
(False, "Duplicate member in request"),
(False, "Duplicate member in request"),
]
assert _budget_of(prisma, "m1").max_budget == 10.0
@pytest.mark.asyncio
async def test_a_row_naming_somebody_off_the_team_fails_without_writing_while_the_rest_of_the_batch_lands():
prisma = _FakePrisma(
teams=[_team("m1")],
memberships=[_membership("m1", "priv-m1"), _membership("elsewhere", "priv-other")],
budgets=[_budget("priv-m1", max_budget=1.0), _budget("priv-other", max_budget=2.0)],
)
results = await _bulk_update(
prisma,
[
{"user_id": "elsewhere", "max_budget_in_team": 99},
{"user_email": "nobody@example.com", "max_budget_in_team": 99},
{"user_id": "m1", "max_budget_in_team": 10},
],
)
assert [(r.success, r.error) for r in results] == [
(False, "User not found in team"),
(False, "User not found in team"),
(True, None),
]
assert prisma.db.litellm_budgettable.rows["priv-other"].max_budget == 2.0
assert _budget_of(prisma, "m1").max_budget == 10.0
assert set(prisma.db.litellm_budgettable.rows) == {"priv-m1", "priv-other"}
@pytest.mark.asyncio
async def test_each_result_carries_the_limits_read_back_after_the_write_in_request_order():
prisma = _FakePrisma(
teams=[_team("m1", "m2")],
memberships=[_membership("m1", "priv-m1"), _membership("m2", "priv-m2")],
budgets=[
_budget("priv-m1", tpm_limit=100, budget_duration="7d"),
_budget("priv-m2", rpm_limit=3),
],
)
results = await _bulk_update(
prisma,
[{"user_id": "m2", "rpm_limit": 8}, {"user_id": "m1", "max_budget_in_team": 42}],
)
assert [r.user_id for r in results] == ["m2", "m1"]
assert (results[1].max_budget, results[1].tpm_limit, results[1].budget_duration) == (42.0, 100, "7d")
assert (results[0].rpm_limit, results[0].max_budget) == (8, None)
@pytest.mark.asyncio
async def test_every_written_member_is_evicted_from_both_team_membership_cache_keys():
prisma = _FakePrisma(
teams=[_team("m1", "m2", "m3")],
memberships=[_membership("m1", "priv-m1"), _membership("m2", "priv-m2"), _membership("m3", "priv-m3")],
budgets=[_budget("priv-m1", max_budget=1.0), _budget("priv-m2", max_budget=2.0), _budget("priv-m3")],
)
cache = _seeded_cache("m1", "m2", "m3")
await _bulk_update(
prisma,
[{"user_id": "m1", "max_budget_in_team": 10}, {"user_id": "m2", "max_budget_in_team": 20}],
cache=cache,
)
assert _cached_keys(cache, "m1") == (None, None)
assert _cached_keys(cache, "m2") == (None, None)
assert _cached_keys(cache, "m3") == ({"cap": "old"}, {"cap": "old"})
@pytest.mark.asyncio
async def test_a_member_with_no_cap_of_their_own_reports_the_team_default_cap_but_only_their_own_rate_limits():
prisma = _FakePrisma(
teams=[_team("m1", default_budget_id="team-default")],
memberships=[],
budgets=[_budget("team-default", max_budget=25.0, tpm_limit=1000)],
)
results = await _bulk_update(prisma, [{"user_id": "m1", "tpm_limit": 7}])
assert [(r.success, r.max_budget, r.max_budget_source, r.tpm_limit) for r in results] == [
(True, 25.0, "team_default", 7)
]
assert _budget_of(prisma, "m1").max_budget is None
default = prisma.db.litellm_budgettable.rows["team-default"]
assert (default.max_budget, default.tpm_limit) == (25.0, 1000)
@pytest.mark.asyncio
async def test_an_explicit_cap_reports_as_the_members_own_while_clearing_one_falls_back_to_the_team_default():
prisma = _FakePrisma(
teams=[_team("m1", "m2", default_budget_id="team-default")],
memberships=[_membership("m1", "priv-m1"), _membership("m2", "priv-m2")],
budgets=[
_budget("team-default", max_budget=25.0),
_budget("priv-m1", max_budget=5.0),
_budget("priv-m2", max_budget=9.0),
],
)
results = await _bulk_update(
prisma,
[{"user_id": "m1", "max_budget_in_team": 50}, {"user_id": "m2", "max_budget_in_team": None}],
)
assert [(r.user_id, r.max_budget, r.max_budget_source) for r in results] == [
("m1", 50.0, "member"),
("m2", 25.0, "team_default"),
]
assert results[1].budget_id is None
assert _budget_id_of(prisma, "m2") is None
assert prisma.db.litellm_budgettable.rows["team-default"].max_budget == 25.0
@pytest.mark.asyncio
async def test_a_team_with_no_default_budget_reports_no_effective_cap_for_a_member_without_one():
prisma = _FakePrisma(
teams=[_team("m1")],
memberships=[_membership("m1", "priv-m1")],
budgets=[_budget("priv-m1", tpm_limit=5)],
)
results = await _bulk_update(prisma, [{"user_id": "m1", "rpm_limit": 3}])
assert [(r.success, r.max_budget, r.max_budget_source) for r in results] == [(True, None, None)]
assert (results[0].tpm_limit, results[0].rpm_limit) == (5, 3)
@pytest.mark.asyncio
async def test_a_zero_team_default_reports_no_cap_because_enforcement_reads_zero_there_as_uncapped():
prisma = _FakePrisma(
teams=[_team("m1", default_budget_id="team-default")],
memberships=[_membership("m1", None)],
budgets=[_budget("team-default", max_budget=0.0)],
)
results = await _bulk_update(prisma, [{"user_id": "m1", "tpm_limit": 9}])
assert [(r.success, r.max_budget, r.max_budget_source) for r in results] == [(True, None, None)]
assert results[0].tpm_limit == 9
@pytest.mark.asyncio
async def test_a_row_that_names_nobody_on_the_team_reports_no_cap_and_no_source():
prisma = _FakePrisma(
teams=[_team("m1", default_budget_id="team-default")],
memberships=[_membership("m1", "priv-m1")],
budgets=[_budget("team-default", max_budget=25.0), _budget("priv-m1", max_budget=5.0)],
)
results = await _bulk_update(
prisma,
[{"user_id": "ghost", "max_budget_in_team": 1}, {"user_id": "m1", "max_budget_in_team": 6}],
)
assert [(r.success, r.max_budget, r.max_budget_source) for r in results] == [
(False, None, None),
(True, 6.0, "member"),
]
@pytest.mark.asyncio
async def test_the_roster_authz_read_runs_on_the_writer_so_a_lagging_replica_cannot_let_a_demoted_admin_write():
writer = _FakePrisma(
teams=[_team("lead", "m1")],
memberships=[_membership("m1", "priv-m1")],
budgets=[_budget("priv-m1", max_budget=1.0)],
)
replica = _FakePrisma(teams=[_team("lead", "m1", admins=("lead",))])
demoted = UserAPIKeyAuth(user_id="lead", user_role=LitellmUserRoles.INTERNAL_USER)
with pytest.raises(ManagementProblem) as raised:
await _bulk_update(
_ReplicatedPrisma(writer=writer, reader=replica),
[{"user_id": "m1", "max_budget_in_team": 99}],
caller=demoted,
)
assert raised.value.problem.status == 403
assert writer.db.litellm_budgettable.rows["priv-m1"].max_budget == 1.0
app = FastAPI()
@app.exception_handler(ManagementProblem)
async def management_problem_exception_handler(request: Request, exc: ManagementProblem):
return problem_response(exc.problem)
@app.exception_handler(RequestValidationError)
async def validation_exception_handler(request: Request, exc: RequestValidationError):
return problem_response(request_validation_problem(exc.errors()))
app.include_router(router)
client = TestClient(app)
BULK_UPDATE_PATH: Final = f"{MANAGEMENT_V1_PREFIX}/teams/{TEAM_ID}/members/bulk_update"
@pytest.fixture
def as_proxy_admin():
app.dependency_overrides[user_api_key_auth] = lambda: ADMIN
yield
app.dependency_overrides.clear()
@pytest.fixture
def as_outsider():
app.dependency_overrides[user_api_key_auth] = lambda: OUTSIDER
yield
app.dependency_overrides.clear()
@pytest.fixture
def prisma(monkeypatch):
fake = _FakePrisma(
teams=[_team("m1", "m2")],
memberships=[_membership("m1", "priv-m1")],
budgets=[_budget("priv-m1", max_budget=1.0)],
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", fake)
return fake
def _post(body: object, path: str = BULK_UPDATE_PATH):
return client.post(path, json=body, headers={"Authorization": "Bearer sk-1234"})
def test_unknown_fields_empty_and_oversized_batches_are_422_problem_documents(prisma, as_proxy_admin):
bodies = (
{"members": [{"user_id": "m1", "max_budget": 10}]},
{"members": [{"user_id": "m1"}], "team_id": TEAM_ID},
{"members": []},
{"members": [{"user_id": f"u{i}"} for i in range(MAX_BULK_TEAM_MEMBER_BUDGET_UPDATES + 1)]},
)
for body in bodies:
response = _post(body)
assert response.status_code == 422, body
assert response.headers["content-type"] == "application/problem+json"
assert response.json()["type"] == "urn:litellm:error:invalid-request-body"
assert prisma.db.litellm_budgettable.rows["priv-m1"].max_budget == 1.0
def test_an_unknown_team_is_a_404_problem_document(prisma, as_proxy_admin):
response = _post(
{"members": [{"user_id": "m1", "max_budget_in_team": 10}]},
path=f"{MANAGEMENT_V1_PREFIX}/teams/nope/members/bulk_update",
)
assert response.status_code == 404
assert response.headers["content-type"] == "application/problem+json"
assert response.json()["type"] == "urn:litellm:error:team-not-found"
assert prisma.db.litellm_budgettable.rows["priv-m1"].max_budget == 1.0
def test_a_caller_who_administers_neither_the_team_nor_its_org_is_a_403_problem_document(prisma, as_outsider):
response = _post({"members": [{"user_id": "m1", "max_budget_in_team": 10}]})
assert response.status_code == 403
assert response.headers["content-type"] == "application/problem+json"
assert response.json()["type"] == "urn:litellm:error:forbidden"
assert prisma.db.litellm_budgettable.rows["priv-m1"].max_budget == 1.0
def test_a_team_admin_may_bulk_update_their_own_teams_members(prisma, monkeypatch):
prisma.db.litellm_teamtable.rows[TEAM_ID] = _team("lead", "m1", admins=("lead",))
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="lead", user_role=LitellmUserRoles.INTERNAL_USER
)
try:
response = _post({"members": [{"user_id": "m1", "max_budget_in_team": 10}]})
finally:
app.dependency_overrides.clear()
assert response.status_code == 200
assert [(r["user_id"], r["success"], r["max_budget"]) for r in response.json()["data"]] == [("m1", True, 10.0)]
@pytest.mark.parametrize("duration", ("0d", "nonsense"))
def test_a_budget_duration_no_reset_can_be_scheduled_from_is_a_422_naming_its_row_and_writes_nothing(
prisma, as_proxy_admin, duration
):
response = _post(
{
"members": [
{"user_id": "m1", "max_budget_in_team": 10},
{"user_id": "m2", "budget_duration": duration},
]
}
)
assert response.status_code == 422
assert response.headers["content-type"] == "application/problem+json"
assert response.json()["type"] == "urn:litellm:error:invalid-request-body"
assert "members.1.budget_duration" in response.json()["detail"]
assert prisma.db.litellm_budgettable.rows["priv-m1"].max_budget == 1.0
def test_an_unconnected_database_is_a_503_problem_document(monkeypatch, as_proxy_admin):
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
response = _post({"members": [{"user_id": "m1", "max_budget_in_team": 10}]})
assert response.status_code == 503
assert response.headers["content-type"] == "application/problem+json"
assert response.json()["type"] == "urn:litellm:error:database-not-connected"
def test_a_driver_error_answers_as_a_problem_document_without_leaking_the_exception(monkeypatch, as_proxy_admin):
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", _UnreachablePrisma())
response = _post({"members": [{"user_id": "m1", "max_budget_in_team": 10}]})
assert response.status_code == 500
assert response.headers["content-type"] == "application/problem+json"
assert response.json()["type"] == "urn:litellm:error:internal-server-error"
assert "connection reset by peer" not in response.text

View file

@ -8544,6 +8544,45 @@ export interface paths {
patch?: never;
trace?: never;
};
"/management/v1/teams/{team_id}/members/bulk_update": {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
get?: never;
put?: never;
/**
* Bulk Update Team Member Budgets Action
* @description Set per-member limits for up to 500 members of one team in one call. Same
* authorization and member addressing as `/team/member_update`: proxy admins, the team's
* admins, and admins of the team's organization, with each member named by exactly one of
* `user_id` or `user_email`. Unknown body fields are a 422 and an unknown team is a 404.
*
* Each row is a merge patch of that member's limits: a field left out is untouched, a
* field sent as null is cleared, and clearing the last limit drops the member back to the
* team default. A budget row shared by several memberships, the team default included, is
* copied for the member being patched rather than written in place, so one member's new
* cap never lands on anybody else.
*
* `data` holds one result per requested member, in request order, carrying the limits in
* force after the write. A row is `success: false` with an `error` when it names nobody on
* the team or repeats an earlier row. Roles are not part of this route; `/team/member_update`
* still owns them.
*
* Example curl:
* ```
* curl --location 'http://0.0.0.0:4000/management/v1/teams/team-1/members/bulk_update' --header 'Authorization: Bearer sk-1234' --header 'Content-Type: application/json' --data '{"members": [{"user_id": "user-1", "max_budget_in_team": 10}, {"user_email": "user-2@example.com", "max_budget_in_team": 10, "budget_duration": "30d"}]}'
* ```
*/
post: operations["bulk_update_team_member_budgets_action_management_v1_teams__team_id__members_bulk_update_post"];
delete?: never;
options?: never;
head?: never;
patch?: never;
trace?: never;
};
"/management/v1/users/bulk": {
parameters: {
query?: never;
@ -24928,6 +24967,22 @@ export interface components {
[key: string]: unknown;
} | null;
};
/**
* BulkTeamMemberBudgetUpdateRequest
* @description Body of `POST /management/v1/teams/{team_id}/members/bulk_update`.
*/
BulkTeamMemberBudgetUpdateRequest: {
/** Members */
members: components["schemas"]["TeamMemberBudgetPatch"][];
};
/**
* BulkTeamMemberBudgetUpdateResponse
* @description `{data: [...]}` with one `TeamMemberBudgetUpdateResult` per requested member, in request order.
*/
BulkTeamMemberBudgetUpdateResponse: {
/** Data */
data: components["schemas"]["TeamMemberBudgetUpdateResult"][];
};
/**
* BulkTeamMemberDeleteRequest
* @description Body of `POST /management/v1/teams/{team_id}/members/bulk_delete`.
@ -37927,6 +37982,57 @@ export interface components {
/** User Id */
user_id?: string | null;
};
/**
* TeamMemberBudgetPatch
* @description One member's per-member limits, merge-patch style: a field left out of the row is
* untouched, a field sent as null is cleared, and clearing the last limit drops the
* member back to the team default.
*/
TeamMemberBudgetPatch: {
/** Allowed Models */
allowed_models?: string[] | null;
/** Budget Duration */
budget_duration?: string | null;
/** Max Budget In Team */
max_budget_in_team?: number | null;
/** Rpm Limit */
rpm_limit?: number | null;
/** Tpm Limit */
tpm_limit?: number | null;
/** User Email */
user_email?: string | null;
/** User Id */
user_id?: string | null;
};
/**
* TeamMemberBudgetUpdateResult
* @description Outcome for one requested member, in request order, carrying the limits in force
* after the write rather than the ones that were asked for.
*/
TeamMemberBudgetUpdateResult: {
/** Allowed Models */
allowed_models?: string[] | null;
/** Budget Duration */
budget_duration?: string | null;
/** Budget Id */
budget_id?: string | null;
/** Error */
error?: string | null;
/** Max Budget */
max_budget?: number | null;
/** Max Budget Source */
max_budget_source?: ("member" | "team_default") | null;
/** Rpm Limit */
rpm_limit?: number | null;
/** Success */
success: boolean;
/** Tpm Limit */
tpm_limit?: number | null;
/** User Email */
user_email?: string | null;
/** User Id */
user_id?: string | null;
};
/** TeamMemberDeleteRequest */
TeamMemberDeleteRequest: {
/** Team Id */
@ -37982,7 +38088,7 @@ export interface components {
};
/**
* TeamMemberRef
* @description One member to remove, named by exactly one of `user_id` or `user_email`.
* @description One member, named by exactly one of `user_id` or `user_email`.
*/
TeamMemberRef: {
/** User Email */
@ -52077,6 +52183,41 @@ export interface operations {
};
};
};
bulk_update_team_member_budgets_action_management_v1_teams__team_id__members_bulk_update_post: {
parameters: {
query?: never;
header?: never;
path: {
team_id: string;
};
cookie?: never;
};
requestBody: {
content: {
"application/json": components["schemas"]["BulkTeamMemberBudgetUpdateRequest"];
};
};
responses: {
/** @description Successful Response */
200: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["BulkTeamMemberBudgetUpdateResponse"];
};
};
/** @description Validation Error */
422: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["HTTPValidationError"];
};
};
};
};
bulk_create_users_route_management_v1_users_bulk_post: {
parameters: {
query?: never;