feat(proxy): add POST /user/bulk_new for batched user and team membership creation

Creates up to 500 internal users in one request with set-based validation,
a single create_many for user rows, and one locked write per referenced team.
Rows fail independently, keys are opt-in per row via auto_create_key, and
send_invite_email is rejected for the batch.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
ryan 2026-09-14 01:17:19 +00:00
parent 30f33a949b
commit ff5b59b173
7 changed files with 1321 additions and 26 deletions

View file

@ -656,6 +656,7 @@ class LiteLLMRoutes(enum.Enum):
[
# user
"/user/new",
"/user/bulk_new",
"/user/update",
"/user/bulk_update",
"/user/delete",

View file

@ -24,6 +24,7 @@ _PROXY_ADMIN_VIEW_ONLY_BLOCKED_ROUTES: Final = frozenset(
[
# user
"/user/new",
"/user/bulk_new",
"/user/delete",
"/user/bulk_update",
# team
@ -755,6 +756,7 @@ class RouteChecks:
_ADMIN_VIEWER_BLOCKED_WRITE_ROUTES = frozenset(
[
"/user/new",
"/user/bulk_new",
"/user/delete",
"/user/bulk_update",
"/team/new",

View file

@ -5,6 +5,7 @@ Internal User Management Endpoints
These are members of a Team on LiteLLM
/user/new
/user/bulk_new
/user/update
/user/bulk_update
/user/delete
@ -77,6 +78,8 @@ from litellm.types.proxy.management_endpoints.common_daily_activity import (
SpendAnalyticsPaginatedResponse,
)
from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
BulkNewUserRequest,
BulkNewUserResponse,
BulkUpdateUserRequest,
BulkUpdateUserResponse,
UserListResponse,
@ -637,6 +640,71 @@ async def new_user(
raise handle_exception_on_proxy(e)
@router.post(
"/user/bulk_new",
tags=["Internal User management"],
dependencies=[Depends(user_api_key_auth)],
response_model=BulkNewUserResponse,
)
@management_endpoint_wrapper
async def bulk_new_user(
data: BulkNewUserRequest,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), # noqa: B008 # FastAPI dependency injection
) -> BulkNewUserResponse:
"""
Create up to 500 internal users in one request, optionally adding each one to teams.
Every entry in `users` takes the same fields as `/user/new`, with two differences: `auto_create_key`
defaults to `false` (opt in per user to also get a virtual key back) and `send_invite_email` is not
supported. Rows are validated together (duplicate ids or emails, unknown teams, roles the caller may not
grant), inserted in one statement, and each referenced team is written once for all of its new members.
Rows fail independently: a bad row is reported in `results` with `success: false` and an `error`, and the
other rows still get created. A user that was created but could not be added to one of its teams is
reported with `success: true`, `teams` listing where they did land, and `error` naming the failed team.
The whole request is rejected with 403 only if creating the valid rows would exceed the license seat limit.
Usage Example
```shell
curl -X POST "http://localhost:4000/user/bulk_new" \\
-H "Content-Type: application/json" \\
-H "Authorization: Bearer sk-1234" \\
-d '{
"users": [
{"user_email": "a@example.com", "user_role": "internal_user", "teams": ["team-1"]},
{"user_email": "b@example.com", "user_role": "internal_user", "auto_create_key": true}
]
}'
```
Returns `results` (one entry per input row, in order, with `user_id`, `user_email`, `success`, `teams`,
`key`, `error`), `total_requested`, `successful_creations` and `failed_creations`.
"""
from litellm.proxy.management_helpers.bulk_user_creation import bulk_create_users
from litellm.proxy.proxy_server import (
_license_check, # pyright: ignore[reportPrivateUsage] # same proxy license singleton /user/new reads
litellm_proxy_admin_name,
prisma_client,
user_api_key_cache,
)
if prisma_client is None:
raise HTTPException(status_code=400, detail=CommonProxyErrors.db_not_connected_error.value)
try:
return await bulk_create_users(
users=data.users,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
license_check=_license_check,
litellm_proxy_admin_name=litellm_proxy_admin_name,
user_api_key_cache=user_api_key_cache,
)
except Exception as e: # noqa: BLE001 # normalize every failure to the proxy exception contract
verbose_proxy_logger.exception("/user/bulk_new: Exception occured - %s", e)
raise handle_exception_on_proxy(e)
@router.get(
"/user/available_roles",
tags=["Internal User management"],

View file

@ -4082,6 +4082,40 @@ def _check_model_access_group(models: list[str] | None, llm_router: Router | Non
return True
_NO_METADATA: Final[Mapping[str, object]] = MappingProxyType({})
def metadata_json_with_limits(
metadata: Mapping[str, object] | None,
*,
model_rpm_limit: Mapping[str, object] | None,
model_tpm_limit: Mapping[str, object] | None,
mcp_rpm_limit: Mapping[str, int] | None,
tag_rpm_limit: Mapping[str, int] | None,
guardrails: Sequence[str] | None,
policies: Sequence[str] | None,
prompts: Sequence[str] | None,
) -> str:
"""Serialize the stored metadata blob with the per-model, MCP, tag, guardrail, policy and prompt settings folded in."""
limits: Final = tuple(
(name, value)
for name, value in (
("model_rpm_limit", model_rpm_limit),
("model_tpm_limit", model_tpm_limit),
("mcp_rpm_limit", mcp_rpm_limit),
("tag_rpm_limit", tag_rpm_limit),
("guardrails", guardrails),
("policies", policies),
("prompts", prompts),
)
if value is not None
)
if metadata is None and not limits:
return json.dumps(None)
merged: Final = {**(metadata or _NO_METADATA), **dict(limits)} # mutable-ok: encrypt_callback_vars takes a dict
return json.dumps(encrypt_callback_vars(merged))
async def generate_key_helper_fn(
request_type: Literal["user", "key"], # identifies if this request is from /user/new or /key/generate
duration: str | None = None,
@ -4184,31 +4218,16 @@ async def generate_key_helper_fn(
permissions_json: Final = json.dumps(permissions)
router_settings_json: Final = safe_dumps(router_settings) if router_settings is not None else safe_dumps({})
# Add model_rpm_limit and model_tpm_limit to metadata
if model_rpm_limit is not None:
metadata = metadata or {}
metadata["model_rpm_limit"] = model_rpm_limit
if model_tpm_limit is not None:
metadata = metadata or {}
metadata["model_tpm_limit"] = model_tpm_limit
if mcp_rpm_limit is not None:
metadata = metadata or {}
metadata["mcp_rpm_limit"] = mcp_rpm_limit
if tag_rpm_limit is not None:
metadata = metadata or {}
metadata["tag_rpm_limit"] = tag_rpm_limit
if guardrails is not None:
metadata = metadata or {}
metadata["guardrails"] = guardrails
if policies is not None:
metadata = metadata or {}
metadata["policies"] = policies
if prompts is not None:
metadata = metadata or {}
metadata["prompts"] = prompts
metadata = encrypt_callback_vars(metadata)
metadata_json: Final = json.dumps(metadata)
metadata_json: Final = metadata_json_with_limits(
metadata,
model_rpm_limit=model_rpm_limit,
model_tpm_limit=model_tpm_limit,
mcp_rpm_limit=mcp_rpm_limit,
tag_rpm_limit=tag_rpm_limit,
guardrails=guardrails,
policies=policies,
prompts=prompts,
)
validate_model_max_budget(model_max_budget)
model_max_budget_json: Final = json.dumps(model_max_budget)
budget_fallbacks_json: Final = json.dumps(budget_fallbacks or {})

View file

@ -0,0 +1,820 @@
"""Batched internal user creation behind `/user/bulk_new`.
The batch is validated with set queries, user rows land in one `create_many`, and every
referenced team is written once under its advisory lock instead of once per user.
"""
import asyncio
import json
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Literal, TypeAlias, TypeVar
from fastapi import HTTPException, Request
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.integrations.prometheus import PrometheusLogger
from litellm.proxy._types import (
LiteLLM_TeamTable,
LitellmUserRoles,
Member,
NewUserRequestTeam,
OrganizationMemberAddRequest,
OrgMember,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import invalidate_team_member_spend_state
from litellm.proxy.auth.litellm_license import LicenseCheck
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
from litellm.proxy.hooks.user_management_event_hooks import UserManagementEventHooks
from litellm.proxy.management_endpoints.common_utils import (
_is_user_org_admin_for_team, # pyright: ignore[reportPrivateUsage] # same team-admin check /user/new uses
_is_user_team_admin, # pyright: ignore[reportPrivateUsage] # same team-admin check /user/new uses
validate_budget_duration,
)
from litellm.proxy.management_endpoints.internal_user_endpoints import (
_update_internal_new_user_params, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # /user/new defaults; result validated below
check_if_default_team_set,
)
from litellm.proxy.management_endpoints.key_management_endpoints import (
_check_permissions_caller_permission, # pyright: ignore[reportPrivateUsage] # same permission check /user/new uses
generate_key_helper_fn, # pyright: ignore[reportUnknownVariableType] # legacy untyped helper; result validated by _KEY_RESPONSE
metadata_json_with_limits,
)
from litellm.proxy.management_endpoints.organization_endpoints import organization_member_add
from litellm.proxy.management_helpers.access_group_team_sync import TEAM_ADVISORY_LOCK_SQL
from litellm.proxy.management_helpers.object_permission_utils import (
_set_object_permission, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # shared with /user/new; result validated below
)
from litellm.proxy.management_helpers.utils import (
_resolve_member_budget_id, # pyright: ignore[reportPrivateUsage] # shared with /team/member_add
)
from litellm.proxy.utils import PrismaClient
from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.user_repository import UserRepository
from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
BulkNewUserItem,
BulkNewUserResponse,
UserCreateResult,
)
if TYPE_CHECKING:
from prisma import Prisma
from prisma import models as prisma_models
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
BULK_NEW_USER_CONCURRENCY: Final = 10
TeamRole: TypeAlias = Literal["user", "admin"]
KeyGenerator: TypeAlias = Callable[..., Awaitable[object]]
_T: Final = TypeVar("_T")
@dataclass(frozen=True, slots=True)
class _RowFailure:
index: int
user_id: str | None
user_email: str | None
error: str
@dataclass(frozen=True, slots=True)
class _PendingUser:
index: int
request: BulkNewUserItem
user_id: str
teams: tuple[NewUserRequestTeam, ...]
class _UserRow(BaseModel):
"""The `/user/new` body after defaults and object permission were applied."""
model_config = ConfigDict(extra="ignore")
user_id: str
user_email: str | None = None
user_alias: str | None = None
user_role: str | None = None
team_id: str | None = None
max_budget: float | None = None
spend: float | None = 0.0
models: tuple[str, ...] | None = None
metadata: Mapping[str, object] | None = None
max_parallel_requests: int | None = None
tpm_limit: int | None = None
rpm_limit: int | None = None
budget_duration: str | None = None
allowed_cache_controls: tuple[str, ...] | None = None
sso_user_id: str | None = None
object_permission_id: str | None = None
model_max_budget: Mapping[str, object] | None = None
model_rpm_limit: Mapping[str, object] | None = None
model_tpm_limit: Mapping[str, object] | None = None
mcp_rpm_limit: Mapping[str, int] | None = None
tag_rpm_limit: Mapping[str, int] | None = None
guardrails: tuple[str, ...] | None = None
policies: tuple[str, ...] | None = None
prompts: tuple[str, ...] | None = None
duration: str | None = None
key_alias: str | None = None
organizations: tuple[str, ...] | None = None
_USER_ROW: Final = TypeAdapter(_UserRow)
@dataclass(frozen=True, slots=True)
class _PreparedUser:
pending: _PendingUser
row: _UserRow
@dataclass(frozen=True, slots=True)
class _TeamAssignment:
user_id: str
user_email: str | None
role: TeamRole
max_budget_in_team: float | None
@dataclass(frozen=True, slots=True)
class _TeamWrite:
"""Outcome of one locked roster write. `failed` maps user ids to the reason they were not added."""
team_id: str
after: tuple[Member, ...]
added: frozenset[str]
failed: Mapping[str, str]
@dataclass(frozen=True, slots=True)
class _CreatedUser:
prepared: _PreparedUser
teams: tuple[str, ...]
key: str | None
errors: tuple[str, ...]
_ERROR_DETAIL: Final = TypeAdapter(Mapping[str, object])
_JSON_OBJECT: Final = TypeAdapter(dict[str, object])
class _KeyResponse(BaseModel):
token: str
_KEY_RESPONSE: Final = TypeAdapter(_KeyResponse)
def _error_message(exc: BaseException) -> str:
if not isinstance(exc, HTTPException):
return str(exc)
try:
detail: Final = _ERROR_DETAIL.validate_python(exc.detail)
except ValidationError:
return str(exc.detail)
return str(detail.get("error", detail))
def _requested_teams(item: BulkNewUserItem) -> tuple[NewUserRequestTeam, ...]:
if item.team_id is not None:
return (NewUserRequestTeam(team_id=item.team_id),)
teams: Final = item.teams if item.teams is not None else check_if_default_team_set()
if teams is None:
return ()
return tuple(team if isinstance(team, NewUserRequestTeam) else NewUserRequestTeam(team_id=team) for team in teams)
def _row_error(item: BulkNewUserItem, user_api_key_dict: UserAPIKeyAuth) -> str | None:
if (
item.user_role in (LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY)
and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN
):
return (
"Only proxy admins can create administrative users (proxy_admin, proxy_admin_viewer). "
f"Attempted to create user with role: {item.user_role}. Your role: {user_api_key_dict.user_role}"
)
try:
validate_budget_duration(item.budget_duration)
_check_permissions_caller_permission(data=item, user_api_key_dict=user_api_key_dict)
except Exception as exc: # noqa: BLE001 # any validation failure is reported on this row only
return _error_message(exc)
return None
def _normalized_email(email: str | None) -> str | None:
return email.strip().lower() if email else None
def _partition_rows(
users: Sequence[BulkNewUserItem], user_api_key_dict: UserAPIKeyAuth
) -> tuple[tuple[_PendingUser, ...], tuple[_RowFailure, ...]]:
"""Assign ids, run the per-row checks and fail later rows that repeat an earlier row's id or email."""
user_ids: Final = tuple(item.user_id or str(uuid.uuid4()) for item in users)
first_index_by_id: Final = MappingProxyType(
{user_id: index for index, user_id in reversed(tuple(enumerate(user_ids)))}
)
first_index_by_email: Final = MappingProxyType(
{
email: index
for index, email in reversed(tuple(enumerate(_normalized_email(item.user_email) for item in users)))
if email is not None
}
)
def classify(index: int, item: BulkNewUserItem) -> _PendingUser | _RowFailure:
user_id: Final = user_ids[index]
email: Final = _normalized_email(item.user_email)
if first_index_by_id[user_id] != index:
return _RowFailure(index, user_id, item.user_email, f"Duplicate user_id in request: {user_id}")
if email is not None and first_index_by_email[email] != index:
return _RowFailure(index, user_id, item.user_email, f"Duplicate user_email in request: {item.user_email}")
error: Final = _row_error(item, user_api_key_dict)
if error is not None:
return _RowFailure(index, user_id, item.user_email, error)
return _PendingUser(index, item, user_id, _requested_teams(item))
outcomes: Final = tuple(classify(index, item) for index, item in enumerate(users))
return (
tuple(outcome for outcome in outcomes if isinstance(outcome, _PendingUser)),
tuple(outcome for outcome in outcomes if isinstance(outcome, _RowFailure)),
)
def _user_table(prisma_client: PrismaClient) -> "TableActions[prisma_models.LiteLLM_UserTable]":
return UserRepository(prisma_client).table
async def _existing_user_conflicts(
prisma_client: PrismaClient, pending: Sequence[_PendingUser]
) -> tuple[frozenset[str], frozenset[str]]:
"""Return the requested user ids and (lowercased) emails that already exist, using one query each."""
user_ids: Final = sorted(user.user_id for user in pending)
emails: Final = sorted(frozenset(user.request.user_email for user in pending if user.request.user_email))
if not user_ids:
return frozenset(), frozenset()
table: Final = _user_table(prisma_client)
id_filter: Final = {"user_id": {"in": user_ids}} # mutable-ok: Prisma query filters are dict-shaped
email_filter: Final = {"user_email": {"in": emails, "mode": "insensitive"}} # mutable-ok: Prisma filter
id_rows: Final = await table.find_many(where=id_filter)
email_rows: Final = await table.find_many(where=email_filter) if emails else ()
return (
frozenset(row.user_id for row in id_rows),
frozenset(lowered for row in email_rows if (lowered := _normalized_email(row.user_email)) is not None),
)
async def _load_teams(prisma_client: PrismaClient, team_ids: frozenset[str]) -> Mapping[str, LiteLLM_TeamTable]:
if not team_ids:
return MappingProxyType({})
rows: Final = await TeamRepository(prisma_client).table.find_many(
where={"team_id": {"in": sorted(team_ids)}} # mutable-ok: Prisma query filters are dict-shaped
)
return MappingProxyType({row.team_id: LiteLLM_TeamTable.model_validate(row.model_dump()) for row in rows})
async def _team_permission_error(team: LiteLLM_TeamTable, user_api_key_dict: UserAPIKeyAuth) -> str | None:
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value:
return None
if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team):
return None
if await _is_user_org_admin_for_team(user_api_key_dict=user_api_key_dict, team_obj=team):
return None
return f"Call not allowed. User not proxy admin OR team admin. team_id={team.team_id}"
async def _unusable_teams(
prisma_client: PrismaClient,
pending: Sequence[_PendingUser],
user_api_key_dict: UserAPIKeyAuth,
) -> tuple[Mapping[str, LiteLLM_TeamTable], Mapping[str, str]]:
"""Load every referenced team once and explain, per team id, why rows naming it cannot proceed."""
team_ids: Final = frozenset(team.team_id for user in pending for team in user.teams)
teams: Final = await _load_teams(prisma_client, team_ids)
permission_errors: Final = await asyncio.gather(
*(_team_permission_error(team, user_api_key_dict) for team in teams.values())
)
missing: Final = tuple(
(team_id, f"Team id={team_id} does not exist") for team_id in team_ids if team_id not in teams
)
denied: Final = tuple(
(team.team_id, error)
for team, error in zip(teams.values(), permission_errors, strict=True)
if error is not None
)
return teams, MappingProxyType({team_id: error for team_id, error in (*missing, *denied)})
def _db_failure(
user: _PendingUser,
existing_ids: frozenset[str],
existing_emails: frozenset[str],
team_errors: Mapping[str, str],
) -> _RowFailure | None:
email: Final = _normalized_email(user.request.user_email)
if user.user_id in existing_ids:
return _RowFailure(user.index, user.user_id, user.request.user_email, f"User id={user.user_id} already exists")
if email is not None and email in existing_emails:
return _RowFailure(
user.index, user.user_id, user.request.user_email, f"User email={user.request.user_email} already exists"
)
errors: Final = tuple(team_errors[team.team_id] for team in user.teams if team.team_id in team_errors)
if errors:
return _RowFailure(user.index, user.user_id, user.request.user_email, "; ".join(errors))
return None
async def _prepare_user(user: _PendingUser, prisma_client: PrismaClient) -> _PreparedUser | _RowFailure:
try:
dumped: Final = user.request.model_dump(exclude={"user_id"}) # mutable-ok: pydantic IncEx takes a set
data: Final = {**dumped, "user_id": user.user_id} # mutable-ok: /user/new defaults helper mutates in place
data_json: Final = _JSON_OBJECT.validate_python(_update_internal_new_user_params(data, user.request))
with_permission: Final = _JSON_OBJECT.validate_python(
await _set_object_permission(data_json=data_json, prisma_client=prisma_client) # pyright: ignore[reportUnknownArgumentType] # validated by the adapter
)
return _PreparedUser(user, _USER_ROW.validate_python(with_permission))
except Exception as exc: # noqa: BLE001 # any preparation failure is reported on this row only
verbose_proxy_logger.warning("/user/bulk_new: could not prepare user %s - %s", user.user_id, exc)
return _RowFailure(user.index, user.user_id, user.request.user_email, _error_message(exc))
class _UserCreateData(TypedDict):
"""One `LiteLLM_UserTable` row as `create_many` takes it; JSON columns are pre-serialized."""
user_id: ReadOnly[str]
user_email: ReadOnly[str | None]
user_alias: ReadOnly[str | None]
user_role: ReadOnly[str | None]
team_id: ReadOnly[str | None]
max_budget: ReadOnly[float | None]
spend: ReadOnly[float]
models: ReadOnly[tuple[str, ...]]
metadata: ReadOnly[str]
max_parallel_requests: ReadOnly[int | None]
tpm_limit: ReadOnly[int | None]
rpm_limit: ReadOnly[int | None]
budget_duration: ReadOnly[str | None]
budget_reset_at: ReadOnly[datetime | None]
allowed_cache_controls: ReadOnly[tuple[str, ...]]
sso_user_id: ReadOnly[str | None]
object_permission_id: ReadOnly[str | None]
teams: ReadOnly[tuple[str, ...]]
model_max_budget: ReadOnly[str]
def _user_create_payload(prepared: _PreparedUser) -> _UserCreateData:
row: Final = prepared.row
metadata_json: Final = metadata_json_with_limits(
row.metadata,
model_rpm_limit=row.model_rpm_limit,
model_tpm_limit=row.model_tpm_limit,
mcp_rpm_limit=row.mcp_rpm_limit,
tag_rpm_limit=row.tag_rpm_limit,
guardrails=row.guardrails,
policies=row.policies,
prompts=row.prompts,
)
payload: Final[_UserCreateData] = {
"user_id": row.user_id,
"user_email": row.user_email,
"user_alias": row.user_alias,
"user_role": row.user_role,
"team_id": row.team_id,
"max_budget": row.max_budget,
"spend": row.spend or 0.0,
"models": row.models or (),
"metadata": metadata_json,
"max_parallel_requests": row.max_parallel_requests,
"tpm_limit": row.tpm_limit,
"rpm_limit": row.rpm_limit,
"budget_duration": row.budget_duration,
"budget_reset_at": get_budget_reset_time(row.budget_duration) if row.budget_duration else None,
"allowed_cache_controls": row.allowed_cache_controls or (),
"sso_user_id": row.sso_user_id,
"object_permission_id": row.object_permission_id,
"teams": tuple(team.team_id for team in prepared.pending.teams),
"model_max_budget": json.dumps(row.model_max_budget) if row.model_max_budget else "{}",
}
return payload
async def _bounded(limit: int, awaitables: Sequence[Awaitable[_T]]) -> tuple[_T | BaseException, ...]:
semaphore: Final = asyncio.Semaphore(limit)
async def run(awaitable: Awaitable[_T]) -> _T:
async with semaphore:
return await awaitable
return tuple(await asyncio.gather(*(run(awaitable) for awaitable in awaitables), return_exceptions=True))
async def _insert_users(
prisma_client: PrismaClient, prepared: Sequence[_PreparedUser]
) -> tuple[tuple[_PreparedUser, ...], tuple[_RowFailure, ...]]:
"""Insert every row in one statement. If that fails, retry rows one at a time so the error lands on its row."""
if not prepared:
return (), ()
table: Final = _user_table(prisma_client)
payloads: Final = tuple(_user_create_payload(user) for user in prepared)
try:
await table.create_many(data=payloads)
return tuple(prepared), ()
except Exception as exc: # noqa: BLE001 # fall back to per-row inserts so the failing row can be identified
verbose_proxy_logger.warning("/user/bulk_new: create_many failed, retrying rows individually - %s", exc)
outcomes: Final = await _bounded(
BULK_NEW_USER_CONCURRENCY, tuple(table.create(data=payload) for payload in payloads)
)
return (
tuple(user for user, outcome in zip(prepared, outcomes, strict=True) if not isinstance(outcome, BaseException)),
tuple(
_RowFailure(user.pending.index, user.pending.user_id, user.row.user_email, _error_message(outcome))
for user, outcome in zip(prepared, outcomes, strict=True)
if isinstance(outcome, BaseException)
),
)
def _assignments_by_team(created: Sequence[_PreparedUser]) -> Mapping[str, tuple[_TeamAssignment, ...]]:
team_ids: Final = tuple(dict.fromkeys(team.team_id for user in created for team in user.pending.teams))
return MappingProxyType(
{
team_id: tuple(
_TeamAssignment(user.pending.user_id, user.row.user_email, team.user_role, team.max_budget_in_team)
for user in created
for team in user.pending.teams
if team.team_id == team_id
)
for team_id in team_ids
}
)
class _MembershipData(TypedDict):
team_id: ReadOnly[str]
user_id: ReadOnly[str]
budget_id: ReadOnly[str | None]
class _RosterData(TypedDict):
members_with_roles: ReadOnly[str]
class _TeamsData(TypedDict):
teams: ReadOnly[tuple[str, ...]]
def _default_member_budget_id(team: LiteLLM_TeamTable) -> str | None:
metadata: Final = (
_JSON_OBJECT.validate_python(
team.metadata # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # LiteLLM_TeamTable.metadata is a bare dict; validated by the adapter
)
if team.metadata # pyright: ignore[reportUnknownMemberType] # same bare dict
else None
)
budget_id: Final = metadata.get("team_member_budget_id") if metadata is not None else None
return budget_id if isinstance(budget_id, str) else None
def _team_tx_db(tx: "Prisma") -> "TableActions[prisma_models.LiteLLM_TeamTable]":
return tx.litellm_teamtable # pyright: ignore[reportReturnType] # TableActions widens the generated inputs to Mapping, as the repositories do
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
async def _write_team_roster(
prisma_client: PrismaClient,
team: LiteLLM_TeamTable,
members: Sequence[_TeamAssignment],
user_api_key_dict: UserAPIKeyAuth,
litellm_proxy_admin_name: str,
) -> _TeamWrite:
"""Add every new member to one team under its advisory lock: one roster rewrite and one membership insert."""
try:
async with prisma_client.tx() as tx:
await tx.query_raw(TEAM_ADVISORY_LOCK_SQL, team.team_id)
roster: Final = await TeamRepository(prisma_client).get_members_with_roles_locked(tx, team.team_id)
if roster is None:
raise ValueError(f"Team id={team.team_id} does not exist")
already_present: Final = frozenset(member.user_id for member in roster if member.user_id)
new_members: Final = tuple(member for member in members if member.user_id not in already_present)
budget_ids: Final = tuple(
[ # mutable-ok: budgets are created one at a time on the transaction's single connection
await _resolve_member_budget_id(
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=litellm_proxy_admin_name,
max_budget_in_team=member.max_budget_in_team,
allowed_models=team.default_team_member_models or None,
budget_duration=None,
default_team_budget_id=_default_member_budget_id(team),
tx=tx, # pyright: ignore[reportArgumentType] # MemberWriteTx lags the generated Prisma signatures, same as /team/member_add
)
for member in new_members
]
)
await _membership_tx_db(tx).create_many(
data=tuple(
_MembershipData(team_id=team.team_id, user_id=member.user_id, budget_id=budget_id)
for member, budget_id in zip(new_members, budget_ids, strict=True)
),
skip_duplicates=True,
)
after: Final = (
*roster,
*(Member(user_id=m.user_id, user_email=m.user_email, role=m.role) for m in new_members),
)
await _team_tx_db(tx).update(
where={"team_id": team.team_id}, # mutable-ok: Prisma query filters are dict-shaped
data=_RosterData(members_with_roles=json.dumps(tuple(member.model_dump() for member in after))),
)
return _TeamWrite(
team_id=team.team_id,
after=after,
added=frozenset(member.user_id for member in new_members),
failed=MappingProxyType({}),
)
except Exception as exc: # noqa: BLE001 # the team write failure is reported on each affected row
verbose_proxy_logger.exception("/user/bulk_new: failed to add members to team %s - %s", team.team_id, exc)
message: Final = f"Failed to add user to team {team.team_id}: {_error_message(exc)}"
return _TeamWrite(
team_id=team.team_id,
after=(),
added=frozenset(),
failed=MappingProxyType({member.user_id: message for member in members}),
)
async def _detach_failed_teams(
prisma_client: PrismaClient, created: Sequence[_PreparedUser], writes: Mapping[str, _TeamWrite]
) -> None:
"""Users are inserted with `teams` already set; drop the teams whose roster write did not take them."""
table: Final = _user_table(prisma_client)
updates: Final = tuple(
table.update(
where={"user_id": user.row.user_id}, # mutable-ok: Prisma query filters are dict-shaped
data=_TeamsData(teams=landed),
)
for user in created
if (landed := _row_teams(user, writes)[0]) != tuple(team.team_id for team in user.pending.teams)
)
for outcome in await _bounded(BULK_NEW_USER_CONCURRENCY, updates):
if isinstance(outcome, BaseException):
verbose_proxy_logger.warning("/user/bulk_new: could not detach failed teams from user - %s", outcome)
async def _publish_team_writes(writes: Sequence[_TeamWrite], user_api_key_cache: "UserApiKeyCache") -> None:
prometheus_logger: Final = PrometheusLogger.get_instance()
for write in writes:
if prometheus_logger is None or not write.added:
continue
try:
prometheus_logger.set_team_members_metric(
LiteLLM_TeamTable(
team_id=write.team_id,
members_with_roles=write.after, # pyright: ignore[reportArgumentType] # pydantic coerces the tuple into the declared list
)
)
except Exception as exc: # noqa: BLE001 # metrics are best-effort and must not fail the request
verbose_proxy_logger.debug("Prometheus: failed to emit team members metric: %s", exc)
evictions: Final = await _bounded(
BULK_NEW_USER_CONCURRENCY,
tuple(
invalidate_team_member_spend_state(
user_id=user_id, team_id=write.team_id, user_api_key_cache=user_api_key_cache
)
for write in writes
for user_id in write.added
),
)
for eviction in evictions:
if isinstance(eviction, BaseException):
verbose_proxy_logger.warning("/user/bulk_new: cache eviction failed - %s", eviction)
_KEY_FIELDS: Final = MappingProxyType(
{
name: True
for name in (
"user_id",
"team_id",
"duration",
"key_alias",
"models",
"metadata",
"max_parallel_requests",
"tpm_limit",
"rpm_limit",
"allowed_cache_controls",
"model_max_budget",
"model_rpm_limit",
"model_tpm_limit",
"mcp_rpm_limit",
"tag_rpm_limit",
"guardrails",
"policies",
"prompts",
"object_permission_id",
)
}
)
async def _generate_key(prepared: _PreparedUser, generate_key: KeyGenerator) -> str:
response: Final = _KEY_RESPONSE.validate_python(
await generate_key(
request_type="key", table_name="key", **prepared.row.model_dump(include=_KEY_FIELDS, exclude_none=True)
)
)
return response.token
async def _add_to_organizations(
prepared: _PreparedUser, organizations: Sequence[str], user_api_key_dict: UserAPIKeyAuth
) -> None:
for organization_id in organizations:
await organization_member_add(
data=OrganizationMemberAddRequest(
organization_id=organization_id,
member=OrgMember(user_id=prepared.row.user_id, role=LitellmUserRoles.INTERNAL_USER),
),
http_request=Request(scope={"type": "http", "path": "/user/bulk_new"}), # mutable-ok: ASGI scopes are dicts
user_api_key_dict=user_api_key_dict,
)
async def _run_per_user(
created: Sequence[_PreparedUser],
select: Callable[[_PreparedUser], bool],
action: Callable[[_PreparedUser], Awaitable[_T]],
) -> Mapping[str, _T | BaseException]:
chosen: Final = tuple(user for user in created if select(user))
outcomes: Final = await _bounded(BULK_NEW_USER_CONCURRENCY, tuple(action(user) for user in chosen))
return MappingProxyType({user.row.user_id: outcome for user, outcome in zip(chosen, outcomes, strict=True)})
async def _write_audit_logs(
prisma_client: PrismaClient,
created: Sequence[_PreparedUser],
user_api_key_dict: UserAPIKeyAuth,
litellm_proxy_admin_name: str,
) -> None:
if not created:
return
created_ids: Final = sorted(user.row.user_id for user in created)
created_filter: Final = {"user_id": {"in": created_ids}} # mutable-ok: Prisma query filters are dict-shaped
rows: Final = await _user_table(prisma_client).find_many(where=created_filter)
outcomes: Final = await _bounded(
BULK_NEW_USER_CONCURRENCY,
tuple(
UserManagementEventHooks.create_internal_user_audit_log(
user_id=row.user_id,
action="created",
litellm_changed_by=user_api_key_dict.user_id,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=litellm_proxy_admin_name,
before_value=None,
after_value=row.model_dump_json(exclude_none=True),
)
for row in rows
),
)
for outcome in outcomes:
if isinstance(outcome, BaseException):
verbose_proxy_logger.warning("Unable to create audit log for user on `/user/bulk_new` - %s", outcome)
def _row_teams(prepared: _PreparedUser, writes: Mapping[str, _TeamWrite]) -> tuple[tuple[str, ...], tuple[str, ...]]:
"""Split a user's requested teams into the ones they landed in and the errors for the ones they did not."""
requested: Final = tuple(team.team_id for team in prepared.pending.teams)
return (
tuple(team_id for team_id in requested if prepared.row.user_id in writes[team_id].added),
tuple(
writes[team_id].failed[prepared.row.user_id]
for team_id in requested
if prepared.row.user_id in writes[team_id].failed
),
)
def _to_result(created: _CreatedUser) -> UserCreateResult:
return UserCreateResult(
user_id=created.prepared.row.user_id,
user_email=created.prepared.row.user_email,
success=True,
teams=created.teams,
key=created.key,
error="; ".join(created.errors) if created.errors else None,
)
def _failure_result(failure: _RowFailure) -> UserCreateResult:
return UserCreateResult(user_id=failure.user_id, user_email=failure.user_email, success=False, error=failure.error)
async def bulk_create_users(
users: Sequence[BulkNewUserItem],
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient,
license_check: LicenseCheck,
litellm_proxy_admin_name: str,
user_api_key_cache: "UserApiKeyCache",
generate_key: KeyGenerator = generate_key_helper_fn,
) -> BulkNewUserResponse:
"""Create every valid row in `users`; rows that fail validation or a write are reported, not raised.
Raises `HTTPException(403)` only when the whole batch would push the deployment over its license seat limit.
"""
pending, request_failures = _partition_rows(users, user_api_key_dict)
existing_ids, existing_emails = await _existing_user_conflicts(prisma_client, pending)
teams, team_errors = await _unusable_teams(prisma_client, pending, user_api_key_dict)
db_failures: Final = tuple(
failure
for user in pending
if (failure := _db_failure(user, existing_ids, existing_emails, team_errors)) is not None
)
failed_indexes: Final = frozenset(failure.index for failure in db_failures)
creatable: Final = tuple(user for user in pending if user.index not in failed_indexes)
billable_users: Final = await UserRepository(prisma_client).count_billable_users()
if creatable and license_check.is_over_limit(total_users=billable_users + len(creatable)):
raise HTTPException(
status_code=403,
detail="License is over limit. Please contact support@berri.ai to upgrade your license.",
)
prepared_outcomes: Final = tuple([await _prepare_user(user, prisma_client) for user in creatable])
prepare_failures: Final = tuple(o for o in prepared_outcomes if isinstance(o, _RowFailure))
created, insert_failures = await _insert_users(
prisma_client, tuple(o for o in prepared_outcomes if isinstance(o, _PreparedUser))
)
team_writes: Final = MappingProxyType(
{
team_id: await _write_team_roster(
prisma_client, teams[team_id], members, user_api_key_dict, litellm_proxy_admin_name
)
for team_id, members in _assignments_by_team(created).items()
}
)
await _detach_failed_teams(prisma_client, created, team_writes)
await _publish_team_writes(tuple(team_writes.values()), user_api_key_cache)
keys: Final = await _run_per_user(
created, lambda user: user.pending.request.auto_create_key, lambda user: _generate_key(user, generate_key)
)
org_outcomes: Final = await _run_per_user(
created,
lambda user: bool(user.row.organizations),
lambda user: _add_to_organizations(user, user.row.organizations or (), user_api_key_dict),
)
await _write_audit_logs(prisma_client, created, user_api_key_dict, litellm_proxy_admin_name)
def finish(prepared: _PreparedUser) -> _CreatedUser:
landed, team_failures = _row_teams(prepared, team_writes)
key_outcome: Final = keys.get(prepared.row.user_id)
org_outcome: Final = org_outcomes.get(prepared.row.user_id)
return _CreatedUser(
prepared=prepared,
teams=landed,
key=key_outcome if isinstance(key_outcome, str) else None,
errors=(
*team_failures,
*(
(f"Failed to create key: {_error_message(key_outcome)}",)
if isinstance(key_outcome, BaseException)
else ()
),
*(
(f"Failed to add user to organizations: {_error_message(org_outcome)}",)
if isinstance(org_outcome, BaseException)
else ()
),
),
)
failures: Final = MappingProxyType(
{
failure.index: _failure_result(failure)
for failure in (*request_failures, *db_failures, *prepare_failures, *insert_failures)
}
)
successes_by_index: Final = MappingProxyType({user.pending.index: _to_result(finish(user)) for user in created})
results: Final = tuple(
failures[index] if index in failures else successes_by_index[index] for index in range(len(users))
)
successes: Final = sum(1 for result in results if result.success)
return BulkNewUserResponse(
results=results,
total_requested=len(users),
successful_creations=successes,
failed_creations=len(users) - successes,
)

View file

@ -1,15 +1,18 @@
from collections.abc import Mapping
from typing import Any, Final, Literal
from pydantic import BaseModel, field_validator
from pydantic import BaseModel, Field, field_validator
from typing_extensions import ReadOnly, TypedDict
from litellm.proxy._types import (
LiteLLM_UserTableWithKeyCount,
NewUserRequest,
UpdateUserRequest,
UpdateUserRequestNoUserIDorEmail,
)
MAX_BULK_NEW_USERS: Final = 500
class InsensitiveContains(TypedDict):
contains: ReadOnly[str]
@ -83,3 +86,38 @@ class BulkUpdateUserResponse(BaseModel):
total_requested: int
successful_updates: int
failed_updates: int
class BulkNewUserItem(NewUserRequest):
"""One row of `/user/bulk_new`: the `/user/new` body, with keys opt-in and invite emails unsupported."""
auto_create_key: bool = False
@field_validator("send_invite_email")
@classmethod
def reject_invite_email(cls, value: bool | None) -> bool | None:
if value:
raise ValueError("send_invite_email is not supported on /user/bulk_new; invite users separately")
return value
class BulkNewUserRequest(BaseModel):
users: tuple[BulkNewUserItem, ...] = Field(min_length=1, max_length=MAX_BULK_NEW_USERS)
class UserCreateResult(BaseModel):
"""Outcome for one row of `/user/bulk_new`. `teams` lists the teams the user was actually added to."""
user_id: str | None = None
user_email: str | None = None
success: bool
teams: tuple[str, ...] | None = None
key: str | None = None
error: str | None = None
class BulkNewUserResponse(BaseModel):
results: tuple[UserCreateResult, ...]
total_requested: int
successful_creations: int
failed_creations: int

View file

@ -0,0 +1,347 @@
import json
from contextlib import asynccontextmanager
from typing import Final
import pytest
from fastapi import HTTPException
from pydantic import BaseModel, ConfigDict, ValidationError
from litellm.caching.caching import DualCache
from litellm.proxy._types import LiteLLM_TeamTable, LitellmUserRoles, Member, UserAPIKeyAuth
from litellm.proxy.management_helpers.bulk_user_creation import bulk_create_users
from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
BulkNewUserItem,
BulkNewUserRequest,
)
ADMIN: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
INTERNAL: Final = UserAPIKeyAuth(user_id="someone", user_role=LitellmUserRoles.INTERNAL_USER)
class _UserRow(BaseModel):
model_config = ConfigDict(extra="allow")
user_id: str
user_email: str | None = None
user_role: str | None = None
teams: list[str] = []
max_budget: float | None = None
class _UserTable:
"""Enough of the Prisma user table for the bulk path: set lookups, one create_many and per-row fallbacks."""
def __init__(self, fail_ids: frozenset[str] = frozenset()) -> None:
self.rows: dict[str, _UserRow] = {}
self.fail_ids = fail_ids
self.create_many_calls = 0
async def count(self, where: object = None) -> int:
return 0 if where is not None else len(self.rows)
async def find_many(self, where: dict[str, dict[str, object]]) -> list[_UserRow]:
if "user_id" in where:
wanted = where["user_id"]["in"]
return [row for row in self.rows.values() if row.user_id in wanted]
wanted_emails = {str(e).lower() for e in where["user_email"]["in"]}
return [row for row in self.rows.values() if (row.user_email or "").lower() in wanted_emails]
async def create(self, data: dict[str, object]) -> _UserRow:
row = _UserRow.model_validate(data)
if row.user_id in self.fail_ids or row.user_id in self.rows:
raise RuntimeError(f"insert failed for {row.user_id}")
self.rows[row.user_id] = row
return row
async def create_many(self, data: list[dict[str, object]]) -> int:
self.create_many_calls += 1
rows = [_UserRow.model_validate(d) for d in data]
if any(row.user_id in self.fail_ids for row in rows):
raise RuntimeError("batch insert failed")
for row in rows:
self.rows[row.user_id] = row
return len(rows)
async def update(self, where: dict[str, str], data: dict[str, object]) -> _UserRow:
row = self.rows[where["user_id"]]
updated = _UserRow.model_validate({**row.model_dump(), **data})
self.rows[row.user_id] = updated
return updated
class _TeamTable:
def __init__(self, teams: list[LiteLLM_TeamTable]) -> None:
self.rows = {team.team_id: team for team in teams}
self.update_calls = 0
async def find_many(self, where: dict[str, dict[str, list[str]]]) -> list[LiteLLM_TeamTable]:
return [self.rows[team_id] for team_id in where["team_id"]["in"] if team_id in self.rows]
async def update(self, where: dict[str, str], data: dict[str, str]) -> LiteLLM_TeamTable:
self.update_calls += 1
team = self.rows[where["team_id"]]
team.members_with_roles = [Member(**m) for m in json.loads(data["members_with_roles"])]
return team
class _MembershipTable:
def __init__(self) -> None:
self.rows: list[dict[str, object]] = []
async def create_many(self, data: list[dict[str, object]], skip_duplicates: bool = False) -> int:
self.rows.extend(data)
return len(data)
class _Tx:
def __init__(self, db: "_Db") -> None:
self.litellm_teamtable = db.litellm_teamtable
self.litellm_teammembership = db.litellm_teammembership
self.locks: list[str] = []
async def query_raw(self, sql: str, *args: object) -> list[dict[str, object]]:
if "pg_advisory_xact_lock" in sql:
self.locks.append(str(args[0]))
return []
team = self.litellm_teamtable.rows.get(str(args[0]))
if team is None:
return []
return [{"members_with_roles": [m.model_dump() for m in team.members_with_roles]}]
class _Db:
def __init__(self, teams: list[LiteLLM_TeamTable], fail_ids: frozenset[str] = frozenset()) -> None:
self.litellm_usertable = _UserTable(fail_ids)
self.litellm_teamtable = _TeamTable(teams)
self.litellm_teammembership = _MembershipTable()
class _FakePrisma:
def __init__(self, teams: list[LiteLLM_TeamTable] | None = None, fail_ids: frozenset[str] = frozenset()) -> None:
self.db = _Db(teams or [], fail_ids)
self.tx_count = 0
self.locks: list[str] = []
def jsonify_object(self, data: dict[str, object]) -> dict[str, object]:
return data
@asynccontextmanager
async def tx(self):
self.tx_count += 1
tx = _Tx(self.db)
yield tx
self.locks.extend(tx.locks)
class _License:
def __init__(self, max_users: int | None = None) -> None:
self.max_users = max_users
self.seen: list[int] = []
def is_over_limit(self, total_users: int) -> bool:
self.seen.append(total_users)
return self.max_users is not None and total_users > self.max_users
def _team(team_id: str, members: list[Member] | None = None) -> LiteLLM_TeamTable:
return LiteLLM_TeamTable(team_id=team_id, members_with_roles=members or [])
async def _no_keys(**kwargs: object) -> dict[str, object]:
raise AssertionError(f"key generation was not requested: {kwargs}")
async def _run(prisma, users, caller=ADMIN, license=None, generate_key=_no_keys):
return await bulk_create_users(
users=[BulkNewUserItem(**u) for u in users],
user_api_key_dict=caller,
prisma_client=prisma,
license_check=license or _License(),
litellm_proxy_admin_name="default_user_id",
user_api_key_cache=DualCache(),
generate_key=generate_key,
)
@pytest.mark.asyncio
async def test_creates_users_and_team_membership_in_every_store():
prisma = _FakePrisma(teams=[_team("t1", [Member(user_id="existing", role="admin")]), _team("t2")])
response = await _run(
prisma,
[
{"user_id": "u1", "user_email": "a@example.com", "teams": ["t1", "t2"], "max_budget": 50},
{"user_id": "u2", "user_email": "b@example.com", "teams": ["t1"]},
{"user_id": "u3", "user_email": "c@example.com"},
],
)
assert (response.total_requested, response.successful_creations, response.failed_creations) == (3, 3, 0)
assert [r.user_id for r in response.results] == ["u1", "u2", "u3"]
assert all(r.success and r.key is None and r.error is None for r in response.results)
assert [r.teams for r in response.results] == [("t1", "t2"), ("t1",), ()]
users = prisma.db.litellm_usertable.rows
assert users["u1"].teams == ["t1", "t2"] and users["u1"].max_budget == 50
assert users["u2"].teams == ["t1"] and users["u3"].teams == []
assert [m.user_id for m in prisma.db.litellm_teamtable.rows["t1"].members_with_roles] == ["existing", "u1", "u2"]
assert [m.user_id for m in prisma.db.litellm_teamtable.rows["t2"].members_with_roles] == ["u1"]
assert sorted((m["team_id"], m["user_id"]) for m in prisma.db.litellm_teammembership.rows) == [
("t1", "u1"),
("t1", "u2"),
("t2", "u1"),
]
@pytest.mark.asyncio
async def test_one_insert_and_one_locked_write_per_team():
prisma = _FakePrisma(teams=[_team("t1"), _team("t2")])
await _run(
prisma,
[{"user_id": f"u{i}", "teams": ["t1"] if i % 2 else ["t1", "t2"]} for i in range(20)],
)
assert prisma.db.litellm_usertable.create_many_calls == 1
assert prisma.tx_count == 2
assert sorted(prisma.locks) == ["t1", "t2"]
assert prisma.db.litellm_teamtable.update_calls == 2
assert len(prisma.db.litellm_teamtable.rows["t1"].members_with_roles) == 20
assert len(prisma.db.litellm_teamtable.rows["t2"].members_with_roles) == 10
@pytest.mark.asyncio
async def test_bad_rows_fail_alone_and_good_rows_still_land():
prisma = _FakePrisma(teams=[_team("t1")])
prisma.db.litellm_usertable.rows["taken"] = _UserRow(user_id="taken", user_email="Taken@Example.com")
response = await _run(
prisma,
[
{"user_id": "u1", "user_email": "a@example.com", "teams": ["t1"]},
{"user_id": "u2", "user_email": "A@EXAMPLE.COM"},
{"user_id": "u1", "user_email": "z@example.com"},
{"user_id": "u3", "user_email": "taken@example.com"},
{"user_id": "taken"},
{"user_id": "u4", "teams": ["missing"]},
{"user_id": "u5", "teams": ["t1", "missing"]},
{"user_id": "u6", "budget_duration": "not-a-duration"},
{"user_id": "u7", "user_email": "ok@example.com", "teams": ["t1"]},
],
)
assert [r.success for r in response.results] == [True, False, False, False, False, False, False, False, True]
assert (response.successful_creations, response.failed_creations) == (2, 7)
errors = [r.error for r in response.results]
assert "Duplicate user_email" in errors[1]
assert "Duplicate user_id" in errors[2]
assert "already exists" in errors[3] and "already exists" in errors[4]
assert "missing" in errors[5] and "does not exist" in errors[5]
assert "missing" in errors[6]
assert errors[7] is not None
assert set(prisma.db.litellm_usertable.rows) == {"taken", "u1", "u7"}
assert [m.user_id for m in prisma.db.litellm_teamtable.rows["t1"].members_with_roles] == ["u1", "u7"]
@pytest.mark.asyncio
async def test_insert_failure_falls_back_to_per_row_and_reports_only_that_row():
prisma = _FakePrisma(teams=[_team("t1")], fail_ids=frozenset({"u2"}))
response = await _run(
prisma,
[{"user_id": "u1", "teams": ["t1"]}, {"user_id": "u2", "teams": ["t1"]}, {"user_id": "u3"}],
)
assert [r.success for r in response.results] == [True, False, True]
assert "insert failed for u2" in (response.results[1].error or "")
assert set(prisma.db.litellm_usertable.rows) == {"u1", "u3"}
assert [m.user_id for m in prisma.db.litellm_teamtable.rows["t1"].members_with_roles] == ["u1"]
@pytest.mark.asyncio
async def test_team_write_failure_keeps_user_and_reports_it_on_the_row():
prisma = _FakePrisma(teams=[_team("t1"), _team("t2")])
async def explode(where, data):
raise RuntimeError("roster write failed")
prisma.db.litellm_teamtable.update = explode
response = await _run(prisma, [{"user_id": "u1", "teams": ["t1", "t2"]}])
result = response.results[0]
assert result.success is True
assert result.teams == ()
assert "t1" in (result.error or "") and "roster write failed" in (result.error or "")
assert prisma.db.litellm_usertable.rows["u1"].teams == []
assert (response.successful_creations, response.failed_creations) == (1, 0)
@pytest.mark.asyncio
async def test_keys_are_opt_in_per_row():
prisma = _FakePrisma()
calls: list[dict[str, object]] = []
async def generate_key(**kwargs: object) -> dict[str, object]:
calls.append(kwargs)
return {"token": f"sk-{kwargs['user_id']}"}
response = await _run(
prisma,
[
{"user_id": "u1"},
{"user_id": "u2", "auto_create_key": True, "models": ["gpt-4o"], "key_alias": "u2-key"},
{"user_id": "u3", "auto_create_key": False},
],
generate_key=generate_key,
)
assert [r.key for r in response.results] == [None, "sk-u2", None]
assert len(calls) == 1
assert calls[0]["user_id"] == "u2" and calls[0]["table_name"] == "key"
assert calls[0]["models"] == ("gpt-4o",) and calls[0]["key_alias"] == "u2-key"
assert set(prisma.db.litellm_usertable.rows) == {"u1", "u2", "u3"}
@pytest.mark.asyncio
async def test_non_admin_cannot_create_admin_users_but_other_rows_proceed():
prisma = _FakePrisma()
response = await _run(
prisma,
[{"user_id": "u1", "user_role": "proxy_admin"}, {"user_id": "u2", "user_role": "internal_user"}],
caller=INTERNAL,
)
assert [r.success for r in response.results] == [False, True]
assert "Only proxy admins" in (response.results[0].error or "")
assert set(prisma.db.litellm_usertable.rows) == {"u2"}
@pytest.mark.asyncio
async def test_license_is_checked_once_against_the_whole_batch():
prisma = _FakePrisma()
prisma.db.litellm_usertable.rows["existing"] = _UserRow(user_id="existing")
license = _License(max_users=3)
with pytest.raises(HTTPException) as exc:
await _run(prisma, [{"user_id": f"u{i}"} for i in range(3)], license=license)
assert exc.value.status_code == 403
assert license.seen == [4]
assert set(prisma.db.litellm_usertable.rows) == {"existing"}
ok = await _run(prisma, [{"user_id": f"u{i}"} for i in range(2)], license=license)
assert ok.successful_creations == 2
resend = await _run(prisma, [{"user_id": f"u{i}"} for i in range(2)], license=license)
assert [r.success for r in resend.results] == [False, False]
assert all("already exists" in (r.error or "") for r in resend.results)
assert license.seen == [4, 3]
assert set(prisma.db.litellm_usertable.rows) == {"existing", "u0", "u1"}
def test_request_rejects_empty_oversized_and_invite_rows():
with pytest.raises(ValidationError):
BulkNewUserRequest(users=[])
with pytest.raises(ValidationError):
BulkNewUserRequest(users=[{"user_email": f"{i}@example.com"} for i in range(501)])
with pytest.raises(ValidationError, match="send_invite_email"):
BulkNewUserItem(user_email="a@example.com", send_invite_email=True)
assert len(BulkNewUserRequest(users=[{"user_email": f"{i}@example.com"} for i in range(500)]).users) == 500
assert BulkNewUserItem(user_email="a@example.com").auto_create_key is False