feat: team override model

This commit is contained in:
Harshit28j 2026-03-22 03:48:55 +05:30
parent 7b31ea40a9
commit 82fecaa9a2
17 changed files with 1597 additions and 400 deletions

View file

@ -113,6 +113,151 @@ curl --location 'http://0.0.0.0:4000/chat/completions' \
### [API Reference](https://litellm-api.up.railway.app/#/team%20management/new_team_team_new_post)
## **Per-Member Model Overrides (Team-Scoped Defaults)**
:::info
Requires `TEAM_MODEL_OVERRIDES=true` environment variable or `litellm.team_model_overrides_enabled = True`.
:::
By default, every team member can access all models in `team.models`. With per-member model overrides, you can:
- Set **`default_models`** on a team — the models every member gets by default
- Set **`models`** on individual team members — additional models only they can access
A member's **effective models** = `default_models` `member.models`. If neither is set, falls back to `team.models` (full backward compatibility).
### Enable the Feature
Add to your `config.yaml`:
```yaml
environment_variables:
TEAM_MODEL_OVERRIDES: "true"
```
### 1. Create a Team with Default Models
```shell
curl -L 'http://localhost:4000/team/new' \
-H 'Authorization: Bearer <your-master-key>' \
-H 'Content-Type: application/json' \
-d '{
"team_alias": "engineering",
"models": ["gpt-4", "gpt-4o-mini", "gpt-4o"],
"default_models": ["gpt-4o-mini"]
}'
```
- `models` — the full pool of models the team is allowed to use
- `default_models` — the subset every member gets by default (must be a subset of `models`)
### 2. Add Members with Per-User Overrides
```shell
# Alice gets the default (gpt-4o-mini only)
curl -L 'http://localhost:4000/team/member_add' \
-H 'Authorization: Bearer <your-master-key>' \
-H 'Content-Type: application/json' \
-d '{
"team_id": "<team-id>",
"member": {"role": "user", "user_id": "alice"}
}'
# Bob gets gpt-4o in addition to the default
curl -L 'http://localhost:4000/team/member_add' \
-H 'Authorization: Bearer <your-master-key>' \
-H 'Content-Type: application/json' \
-d '{
"team_id": "<team-id>",
"member": {"role": "user", "user_id": "bob", "models": ["gpt-4o"]}
}'
```
| Member | Override | Effective Models |
|--------|----------|-----------------|
| Alice | none | `["gpt-4o-mini"]` |
| Bob | `["gpt-4o"]` | `["gpt-4o-mini", "gpt-4o"]` |
### 3. Generate Keys and Test
```shell
# Generate key for Bob
curl -L 'http://localhost:4000/key/generate' \
-H 'Authorization: Bearer <your-master-key>' \
-H 'Content-Type: application/json' \
-d '{"team_id": "<team-id>", "user_id": "bob"}'
```
<Tabs>
<TabItem label="Allowed (Bob → gpt-4o)" value="allowed">
```shell
curl -L 'http://localhost:4000/chat/completions' \
-H 'Authorization: Bearer <bob-key>' \
-H 'Content-Type: application/json' \
-d '{"model": "gpt-4o", "messages": [{"role": "user", "content": "Hello"}]}'
```
Returns `200 OK``gpt-4o` is in Bob's effective set.
</TabItem>
<TabItem label="Denied (Bob → gpt-4)" value="denied">
```shell
curl -L 'http://localhost:4000/chat/completions' \
-H 'Authorization: Bearer <bob-key>' \
-H 'Content-Type: application/json' \
-d '{"model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}]}'
```
Returns `401 Unauthorized``gpt-4` is in the team pool but not in Bob's effective set.
</TabItem>
</Tabs>
### 4. Update Member Overrides
```shell
# Add gpt-4 to Bob's overrides
curl -L 'http://localhost:4000/team/member_update' \
-H 'Authorization: Bearer <your-master-key>' \
-H 'Content-Type: application/json' \
-d '{
"team_id": "<team-id>",
"user_id": "bob",
"models": ["gpt-4o", "gpt-4"]
}'
# Remove all overrides (Bob falls back to default_models only)
curl -L 'http://localhost:4000/team/member_update' \
-H 'Authorization: Bearer <your-master-key>' \
-H 'Content-Type: application/json' \
-d '{
"team_id": "<team-id>",
"user_id": "bob",
"models": []
}'
```
### Validation Rules
| Rule | Error |
|------|-------|
| `default_models` must be a subset of `team.models` | `400` on `/team/new` and `/team/update` |
| Member `models` must be a subset of `team.models` | `400` on `/team/member_add` and `/team/member_update` |
| Key `models` must be a subset of effective models | `403` on `/key/generate` |
| Narrowing `team.models` auto-prunes stale `default_models` | Automatic on `/team/update` |
### Backward Compatibility
When the feature flag is off **or** when neither `default_models` nor member `models` is configured:
- `get_effective_team_models()` returns `team.models` unchanged
- All existing teams and keys work exactly as before
- Zero extra database queries on the auth hot path
## **View Available Fallback Models**

View file

@ -0,0 +1,5 @@
-- AlterTable: Add default_models to LiteLLM_TeamTable
ALTER TABLE "LiteLLM_TeamTable" ADD COLUMN IF NOT EXISTS "default_models" TEXT[] DEFAULT ARRAY[]::TEXT[];
-- AlterTable: Add models to LiteLLM_TeamMembership
ALTER TABLE "LiteLLM_TeamMembership" ADD COLUMN IF NOT EXISTS "models" TEXT[] DEFAULT ARRAY[]::TEXT[];

View file

@ -142,6 +142,7 @@ model LiteLLM_TeamTable {
policies String[] @default([])
model_id Int? @unique // id for LiteLLM_ModelTable -> stores team-level model aliases
allow_team_guardrail_config Boolean @default(false) // if true, team admin can configure guardrails for this team
default_models String[] @default([]) // NEW: team-wide defaults
litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id])
litellm_model_table LiteLLM_ModelTable? @relation(fields: [model_id], references: [id])
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
@ -595,6 +596,7 @@ model LiteLLM_TeamMembership {
team_id String
spend Float @default(0.0)
budget_id String?
models String[] @default([]) // NEW: per-user model overrides
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
@@id([user_id, team_id])
}

View file

@ -218,6 +218,7 @@ use_chat_completions_url_for_anthropic_messages: bool = bool(
os.getenv("LITELLM_USE_CHAT_COMPLETIONS_URL_FOR_ANTHROPIC_MESSAGES", False)
) # When True, routes OpenAI /v1/messages requests to chat/completions instead of the Responses API
retry = True
team_model_overrides_enabled = os.getenv("TEAM_MODEL_OVERRIDES", "").lower() == "true"
### AUTH ###
api_key: Optional[str] = None
openai_key: Optional[str] = None

View file

@ -1611,6 +1611,16 @@ class Member(MemberBase):
] = Field(
description="The role of the user within the team. 'admin' users can manage team settings and members, 'user' is a regular team member"
)
models: Optional[List[str]] = Field(
default=None,
description="Specific models this member can access within the team. If provided, these will be used in addition to the team's default models.",
)
tpm_limit: Optional[int] = Field(
default=None, description="Tokens per minute limit for this team member"
)
rpm_limit: Optional[int] = Field(
default=None, description="Requests per minute limit for this team member"
)
class OrgMember(MemberBase):
@ -1642,6 +1652,7 @@ class TeamBase(LiteLLMPydanticObjectBase):
blocked: bool = False
router_settings: Optional[dict] = None
access_group_ids: Optional[List[str]] = None
default_models: List[str] = []
class NewTeamRequest(TeamBase):
@ -1719,6 +1730,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
model_aliases: Optional[dict] = None
guardrails: Optional[List[str]] = None
policies: Optional[List[str]] = None
default_models: Optional[List[str]] = None
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
team_member_budget: Optional[float] = None
team_member_budget_duration: Optional[str] = None
@ -2387,6 +2399,8 @@ class LiteLLM_VerificationTokenView(LiteLLM_VerificationToken):
team_alias: Optional[str] = None
team_tpm_limit: Optional[int] = None
team_rpm_limit: Optional[int] = None
team_member_models: Optional[List[str]] = None
team_default_models: Optional[List[str]] = None
team_max_budget: Optional[float] = None
team_soft_budget: Optional[float] = None
team_models: List = []
@ -3607,6 +3621,7 @@ class LiteLLM_TeamMembership(LiteLLMPydanticObjectBase):
team_id: str
budget_id: Optional[str] = None
spend: Optional[float] = 0.0
models: List[str] = []
litellm_budget_table: Optional[LiteLLM_BudgetTable]
def safe_get_team_member_rpm_limit(self) -> Optional[int]:
@ -3729,6 +3744,7 @@ class TeamMemberDeleteRequest(MemberDeleteRequest):
class TeamMemberUpdateRequest(TeamMemberDeleteRequest):
max_budget_in_team: Optional[float] = None
role: Optional[Literal["admin", "user"]] = None
models: Optional[List[str]] = None
tpm_limit: Optional[int] = Field(
default=None, description="Tokens per minute limit for this team member"
)
@ -3739,6 +3755,7 @@ class TeamMemberUpdateRequest(TeamMemberDeleteRequest):
class TeamMemberUpdateResponse(MemberUpdateResponse):
team_id: str
models: Optional[List[str]] = None
max_budget_in_team: Optional[float] = None
tpm_limit: Optional[int] = None
rpm_limit: Optional[int] = None

View file

@ -9,6 +9,7 @@ Run checks for:
3. If end_user ('user' passed to /chat/completions, /embeddings endpoint) is in budget
"""
import asyncio
import os
import re
import time
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast
@ -409,20 +410,16 @@ async def common_checks( # noqa: PLR0915
# 2. If team can call model
if _model and team_object:
with tracer.trace("litellm.proxy.auth.common_checks.can_team_access_model"):
if not await can_team_access_model(
# can_team_access_model returns Literal[True] or raises ProxyException
await can_team_access_model(
model=_model,
team_object=team_object,
llm_router=llm_router,
team_model_aliases=valid_token.team_model_aliases
if valid_token
else None,
):
raise ProxyException(
message=f"Team not allowed to access model. Team={team_object.team_id}, Model={_model}. Allowed team models = {team_object.models}",
type=ProxyErrorTypes.team_model_access_denied,
param="model",
code=status.HTTP_401_UNAUTHORIZED,
)
valid_token=valid_token,
)
# Require trace id for agent keys when agent has require_trace_id_on_calls_by_agent
if valid_token is not None and valid_token.agent_id:
@ -2690,11 +2687,76 @@ def can_org_access_model(
)
def compute_effective_models(
team_defaults: List[str],
member_models: List[str],
team_pool: List[str],
) -> List[str]:
"""
Core computation shared by the auth hot-path and key-generation.
effective = union(team_defaults, member_models), capped by team_pool.
- If neither defaults nor overrides are set, falls back to team_pool (backward compat).
- If cap empties the list (all stale), falls back to team_pool (NOT [] which = allow-all).
- team_pool=[] means "allow all" cap is skipped.
"""
effective = list(set(team_defaults + member_models))
if not effective:
return team_pool
if team_pool:
effective = [m for m in effective if m in set(team_pool)]
if not effective:
return team_pool
return effective
def get_effective_team_models(
team_object: Optional[LiteLLM_TeamTable],
valid_token: Optional[UserAPIKeyAuth] = None,
) -> List[str]:
"""
Returns the effective list of models for a team member.
The union of:
- team_object.default_models (OR valid_token.team_default_models if available)
- team_membership.models (OR valid_token.team_member_models if available)
Capped by team_object.models. Falls back to team_object.models when empty.
"""
if not (
litellm.team_model_overrides_enabled
or os.getenv("TEAM_MODEL_OVERRIDES", "").lower() == "true"
):
return team_object.models if team_object else []
# Get from team defaults — prefer team_object (authoritative, fresh from DB/cache)
# over valid_token (snapshot from key creation time, may be stale).
# Use `is not None` instead of truthiness so that an explicit empty list []
# (meaning "no defaults") is not confused with "field missing".
team_defaults: List[str] = []
if team_object and team_object.default_models is not None:
team_defaults = team_object.default_models
elif valid_token and valid_token.team_default_models is not None:
team_defaults = valid_token.team_default_models
# Get from member specific overrides
member_models: List[str] = []
if valid_token and valid_token.team_member_models is not None:
member_models = valid_token.team_member_models
team_pool = team_object.models if team_object else []
return compute_effective_models(team_defaults, member_models, team_pool)
async def can_team_access_model(
model: Union[str, List[str]],
team_object: Optional[LiteLLM_TeamTable],
llm_router: Optional[Router],
team_model_aliases: Optional[Dict[str, str]] = None,
valid_token: Optional[UserAPIKeyAuth] = None,
) -> Literal[True]:
"""
Returns True if the team can access a specific model.
@ -2702,11 +2764,12 @@ async def can_team_access_model(
1. First checks native team-level model permissions (current implementation)
2. If not allowed natively, falls back to access_group_ids on the team
"""
effective_models = get_effective_team_models(team_object, valid_token)
try:
return _can_object_call_model(
model=model,
llm_router=llm_router,
models=team_object.models if team_object else [],
models=effective_models,
team_model_aliases=team_model_aliases,
team_id=team_object.team_id if team_object else None,
object_type="team",

View file

@ -1,4 +1,4 @@
from typing import TYPE_CHECKING, Any, Dict, Optional, Union
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from litellm._logging import verbose_proxy_logger
from litellm.caching import DualCache
@ -354,6 +354,7 @@ async def _upsert_budget_and_membership(
user_api_key_dict: UserAPIKeyAuth,
tpm_limit: Optional[int] = None,
rpm_limit: Optional[int] = None,
models: Optional[List[str]] = None,
):
"""
Helper function to Create/Update or Delete the budget within the team membership
@ -366,35 +367,69 @@ async def _upsert_budget_and_membership(
user_api_key_dict: User API Key dictionary containing user information
tpm_limit: Tokens per minute limit for the team member
rpm_limit: Requests per minute limit for the team member
models: Specific models this member can access within the team.
If max_budget, tpm_limit, and rpm_limit are all None, the user's budget is removed from the team membership.
If any of these values exist, a budget is updated or created and linked to the team membership.
If max_budget, tpm_limit, rpm_limit, and models are all None, the budget is disconnected
(but existing model overrides are preserved models=None means "not specified").
If any of these values exist, a budget is updated or created and linked to the team membership, and models are updated.
To explicitly clear model overrides, pass models=[].
"""
if max_budget is None and tpm_limit is None and rpm_limit is None:
# disconnect the budget since all limits are None
await tx.litellm_teammembership.update(
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}},
data={"litellm_budget_table": {"disconnect": True}},
)
if (
max_budget is None
and tpm_limit is None
and rpm_limit is None
and models is None
):
# Nothing to change — only disconnect budget if one was actually linked.
# Do NOT touch models (models=None means "not specified", not "clear").
# Use upsert (not update) because members added without budget/models
# may not have a LiteLLM_TeamMembership row yet.
if existing_budget_id is not None:
await tx.litellm_teammembership.upsert(
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}},
data={
"create": {"user_id": user_id, "team_id": team_id},
"update": {"litellm_budget_table": {"disconnect": True}},
},
)
return
# create a new budget
create_data: Dict[str, Any] = {
"created_by": user_api_key_dict.user_id or "",
"updated_by": user_api_key_dict.user_id or "",
}
if max_budget is not None:
create_data["max_budget"] = max_budget
if tpm_limit is not None:
create_data["tpm_limit"] = tpm_limit
if rpm_limit is not None:
create_data["rpm_limit"] = rpm_limit
_budget_id = existing_budget_id
if max_budget is not None or tpm_limit is not None or rpm_limit is not None:
# create a new budget
create_data: Dict[str, Any] = {
"created_by": user_api_key_dict.user_id or "",
"updated_by": user_api_key_dict.user_id or "",
}
if max_budget is not None:
create_data["max_budget"] = max_budget
if tpm_limit is not None:
create_data["tpm_limit"] = tpm_limit
if rpm_limit is not None:
create_data["rpm_limit"] = rpm_limit
new_budget = await tx.litellm_budgettable.create(
data=create_data,
include={"team_membership": True},
)
_budget_id = new_budget.budget_id
# upsert the team membership with the new/updated budget and models
membership_create_data: Dict[str, Any] = {
"user_id": user_id,
"team_id": team_id,
}
membership_update_data: Dict[str, Any] = {}
if _budget_id:
budget_connect = {
"litellm_budget_table": {"connect": {"budget_id": _budget_id}}
}
membership_create_data.update(budget_connect)
membership_update_data.update(budget_connect)
if models is not None:
membership_create_data["models"] = models
membership_update_data["models"] = models
new_budget = await tx.litellm_budgettable.create(
data=create_data,
include={"team_membership": True},
)
# upsert the team membership with the new/updated budget
await tx.litellm_teammembership.upsert(
where={
"user_id_team_id": {
@ -403,18 +438,8 @@ async def _upsert_budget_and_membership(
}
},
data={
"create": {
"user_id": user_id,
"team_id": team_id,
"litellm_budget_table": {
"connect": {"budget_id": new_budget.budget_id},
},
},
"update": {
"litellm_budget_table": {
"connect": {"budget_id": new_budget.budget_id},
},
},
"create": membership_create_data,
"update": membership_update_data,
},
)

View file

@ -41,6 +41,7 @@ from litellm.proxy._experimental.mcp_server.db import (
from litellm.proxy._types import *
from litellm.proxy._types import LiteLLM_VerificationToken
from litellm.proxy.auth.auth_checks import (
compute_effective_models,
_delete_cache_key_object,
can_team_access_model,
get_org_object,
@ -636,6 +637,54 @@ async def _common_key_generation_helper( # noqa: PLR0915
data_json = data.model_dump(exclude_unset=True, exclude_none=True) # type: ignore
# [TEAM MODEL OVERRIDES] Handle effective team models for the key
if (
litellm.team_model_overrides_enabled
or os.getenv("TEAM_MODEL_OVERRIDES", "").lower() == "true"
) and team_table is not None:
# Read member models from LiteLLM_TeamMembership table (authoritative source),
# NOT from members_with_roles JSON blob which can be stale after /team/member_update.
# Note: This is a management endpoint (/key/generate), not the hot auth path.
# The hot path (/chat/completions) uses the SQL view join with zero extra queries.
member_models: List[str] = []
if data.user_id and prisma_client is not None:
_membership = await prisma_client.db.litellm_teammembership.find_unique(
where={
"user_id_team_id": {
"user_id": data.user_id,
"team_id": team_table.team_id,
}
}
)
if _membership is not None:
member_models = _membership.models or []
team_default_models = getattr(team_table, "default_models", None) or []
team_pool = team_table.models or []
effective_models = compute_effective_models(
team_defaults=team_default_models,
member_models=member_models,
team_pool=team_pool,
)
if effective_models:
# if 'all-team-models' was requested, restrict it to the effective models
if "all-team-models" in (data.models or []):
data_json["models"] = effective_models
# if explicit models were requested, validate they're a subset of effective set
elif data.models:
disallowed = set(data.models) - set(effective_models)
if disallowed:
raise HTTPException(
status_code=403,
detail={
"error": f"Requested models not in user's effective team models. "
f"Disallowed: {sorted(disallowed)}. "
f"Effective models: {sorted(effective_models)}"
},
)
# if NO models was requested, runtime auth will compute effective models
# from the SQL view join (tm.models + t.default_models), so nothing to store here
data_json = handle_key_type(data, data_json)
# if we get max_budget passed to /key/generate, then use it as key_max_budget. Since generate_key_helper_fn is used to make new users

View file

@ -813,6 +813,22 @@ async def new_team( # noqa: PLR0915
},
)
# Validate default_models is a subset of team models (prevent privilege escalation).
# When data.models is empty/[] (unrestricted team), skip validation — by design,
# team.models=[] means "allow all" so any default_models is a valid subset.
# At runtime, compute_effective_models caps effective set by team.models.
if data.default_models and data.models:
disallowed = set(data.default_models) - set(data.models)
if disallowed:
raise HTTPException(
status_code=400,
detail={
"error": f"default_models must be a subset of team models. "
f"Disallowed: {sorted(disallowed)}. "
f"Team models: {sorted(data.models)}"
},
)
# Check if license is over limit
total_teams = await prisma_client.db.litellm_teamtable.count()
if total_teams and _license_check.is_team_count_over_limit(
@ -1407,6 +1423,25 @@ async def update_team( # noqa: PLR0915
detail={"error": f"Team not found, passed team_id={data.team_id}"},
)
# Validate default_models is a subset of team models (prevent privilege escalation)
if data.default_models:
team_models = (
data.models
if data.models is not None
else (existing_team_row.models or [])
)
if team_models:
disallowed = set(data.default_models) - set(team_models)
if disallowed:
raise HTTPException(
status_code=400,
detail={
"error": f"default_models must be a subset of team models. "
f"Disallowed: {sorted(disallowed)}. "
f"Team models: {sorted(team_models)}"
},
)
if data.soft_budget is not None:
max_budget_to_check = (
data.max_budget
@ -1494,6 +1529,23 @@ async def update_team( # noqa: PLR0915
updated_kv = data.json(exclude_unset=True)
# When team.models is being changed, prune existing default_models
# to prevent stale over-permissive defaults (privilege escalation).
# Must inject into updated_kv directly (not data) because
# data.json(exclude_unset=True) skips fields not set during __init__.
if "models" in updated_kv and "default_models" not in updated_kv:
existing_defaults = existing_team_row.default_models or []
if existing_defaults:
new_models = updated_kv["models"] or []
if not new_models:
# team.models=[] means "allow all" — clear default_models
# so get_effective_team_models falls back to [] (allow all)
updated_kv["default_models"] = []
else:
pruned = [m for m in existing_defaults if m in set(new_models)]
if pruned != existing_defaults:
updated_kv["default_models"] = pruned
# Check budget_duration and budget_reset_at
_set_budget_reset_at(data, updated_kv)
@ -1764,6 +1816,29 @@ async def _process_team_members(
else None
)
# Validate member model overrides are within team.models (prevent privilege escalation)
team_models = (
complete_team_data.models
if hasattr(complete_team_data, "models")
and isinstance(complete_team_data.models, list)
else []
)
members_to_validate = (
[data.member] if isinstance(data.member, Member) else data.member
)
for member in members_to_validate:
if member.models and team_models:
disallowed = set(member.models) - set(team_models)
if disallowed:
raise HTTPException(
status_code=400,
detail={
"error": f"Member model overrides must be a subset of team models. "
f"Disallowed: {sorted(disallowed)}. "
f"Team models: {sorted(team_models)}"
},
)
if isinstance(data.member, Member):
try:
updated_user, updated_tm = await add_new_member(
@ -1774,6 +1849,7 @@ async def _process_team_members(
litellm_proxy_admin_name=litellm_proxy_admin_name,
team_id=data.team_id,
default_team_budget_id=default_team_budget_id,
team_models=team_models,
)
except Exception as e:
raise HTTPException(
@ -1798,6 +1874,7 @@ async def _process_team_members(
litellm_proxy_admin_name=litellm_proxy_admin_name,
team_id=data.team_id,
default_team_budget_id=default_team_budget_id,
team_models=team_models,
)
except Exception as e:
raise HTTPException(
@ -2331,7 +2408,7 @@ async def team_member_delete(
response_model=TeamMemberUpdateResponse,
)
@management_endpoint_wrapper
async def team_member_update(
async def team_member_update( # noqa: PLR0915
data: TeamMemberUpdateRequest,
http_request: Request,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
@ -2424,6 +2501,19 @@ async def team_member_update(
identified_budget_id = tm.budget_id
break
### validate member model overrides are within team.models (when restricted)
if data.models and existing_team_row.models:
disallowed = set(data.models) - set(existing_team_row.models)
if disallowed:
raise HTTPException(
status_code=400,
detail={
"error": f"Member model overrides must be a subset of team models. "
f"Disallowed: {sorted(disallowed)}. "
f"Team models: {sorted(existing_team_row.models)}"
},
)
### upsert new budget
async with prisma_client.db.tx() as tx:
await _upsert_budget_and_membership(
@ -2435,18 +2525,37 @@ async def team_member_update(
user_api_key_dict=user_api_key_dict,
tpm_limit=data.tpm_limit,
rpm_limit=data.rpm_limit,
models=data.models,
)
### update team member role
if data.role is not None:
# Resolve the effective models for this member (from the authoritative
# LiteLLM_TeamMembership table) so we can: (a) keep the members_with_roles
# JSON in sync, and (b) return the actual stored state in the response.
stored_models = data.models
if stored_models is None:
_tm_row = await prisma_client.db.litellm_teammembership.find_unique(
where={
"user_id_team_id": {
"user_id": received_user_id,
"team_id": data.team_id,
}
}
)
stored_models = (_tm_row.models or []) if _tm_row is not None else []
if data.role is not None or data.models is not None:
team_members: List[Member] = []
for member in team_table.members_with_roles:
if member.user_id == received_user_id:
team_members.append(
Member(
user_id=member.user_id,
role=data.role,
role=data.role or member.role,
user_email=data.user_email or member.user_email,
models=stored_models,
tpm_limit=data.tpm_limit if data.tpm_limit is not None else getattr(member, "tpm_limit", None),
rpm_limit=data.rpm_limit if data.rpm_limit is not None else getattr(member, "rpm_limit", None),
)
)
else:
@ -2464,6 +2573,7 @@ async def team_member_update(
team_id=data.team_id,
user_id=received_user_id,
user_email=data.user_email,
models=stored_models,
max_budget_in_team=data.max_budget_in_team,
tpm_limit=data.tpm_limit,
rpm_limit=data.rpm_limit,

View file

@ -2,7 +2,7 @@
## Helper utils for the management endpoints (keys/users/teams)
from datetime import datetime
from functools import wraps
from typing import Optional, Tuple
from typing import List, Optional, Tuple
from fastapi import HTTPException, Request
@ -148,6 +148,7 @@ async def add_new_member(
user_api_key_dict: UserAPIKeyAuth,
litellm_proxy_admin_name: str,
default_team_budget_id: Optional[str] = None,
team_models: Optional[List[str]] = None,
) -> Tuple[LiteLLM_UserTable, Optional[LiteLLM_TeamMembership]]:
"""
Add a new member to a team
@ -208,28 +209,60 @@ async def add_new_member(
# Check if trying to set a budget for team member
if max_budget_in_team is not None:
if (
max_budget_in_team is not None
or new_member.tpm_limit is not None
or new_member.rpm_limit is not None
):
# create a new budget item for this member
_budget_create_data = {
"created_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
"updated_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
}
if max_budget_in_team is not None:
_budget_create_data["max_budget"] = max_budget_in_team
if new_member.tpm_limit is not None:
_budget_create_data["tpm_limit"] = new_member.tpm_limit
if new_member.rpm_limit is not None:
_budget_create_data["rpm_limit"] = new_member.rpm_limit
response = await prisma_client.db.litellm_budgettable.create(
data={
"max_budget": max_budget_in_team,
"created_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
"updated_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
}
data=_budget_create_data # type: ignore
)
_budget_id = response.budget_id
else:
_budget_id = default_team_budget_id
if _budget_id and returned_user is not None and returned_user.user_id is not None:
if (
(_budget_id or new_member.models)
and returned_user is not None
and returned_user.user_id is not None
):
membership_create_data = {
"team_id": team_id,
"user_id": returned_user.user_id,
}
if _budget_id:
membership_create_data["budget_id"] = _budget_id
if new_member.models:
# Defense-in-depth: validate member models are within team models
if team_models:
disallowed = set(new_member.models) - set(team_models)
if disallowed:
raise HTTPException(
status_code=400,
detail={
"error": f"Member model overrides must be a subset of team models. "
f"Disallowed: {sorted(disallowed)}. "
f"Team models: {sorted(team_models)}"
},
)
membership_create_data["models"] = new_member.models
_returned_team_membership = (
await prisma_client.db.litellm_teammembership.create(
data={
"team_id": team_id,
"user_id": returned_user.user_id,
"budget_id": _budget_id,
},
data=membership_create_data, # type: ignore
include={"litellm_budget_table": True},
)
)

View file

@ -142,6 +142,7 @@ model LiteLLM_TeamTable {
policies String[] @default([])
model_id Int? @unique // id for LiteLLM_ModelTable -> stores team-level model aliases
allow_team_guardrail_config Boolean @default(false) // if true, team admin can configure guardrails for this team
default_models String[] @default([]) // NEW: team-wide defaults
litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id])
litellm_model_table LiteLLM_ModelTable? @relation(fields: [model_id], references: [id])
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
@ -599,6 +600,7 @@ model LiteLLM_TeamMembership {
team_id String
spend Float @default(0.0)
budget_id String?
models String[] @default([]) // NEW: per-user model overrides
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
@@id([user_id, team_id])
}

View file

@ -2961,6 +2961,7 @@ class PrismaClient:
t.tpm_limit AS team_tpm_limit,
t.rpm_limit AS team_rpm_limit,
t.models AS team_models,
t.default_models AS team_default_models,
t.metadata AS team_metadata,
t.blocked AS team_blocked,
t.team_alias AS team_alias,
@ -2969,6 +2970,7 @@ class PrismaClient:
t.object_permission_id AS team_object_permission_id,
t.organization_id as org_id,
tm.spend AS team_member_spend,
tm.models AS team_member_models,
m.aliases AS team_model_aliases,
-- Added comma to separate b.* columns
b.max_budget AS litellm_budget_table_max_budget,

View file

@ -142,6 +142,7 @@ model LiteLLM_TeamTable {
policies String[] @default([])
model_id Int? @unique // id for LiteLLM_ModelTable -> stores team-level model aliases
allow_team_guardrail_config Boolean @default(false) // if true, team admin can configure guardrails for this team
default_models String[] @default([]) // NEW: team-wide defaults
litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id])
litellm_model_table LiteLLM_ModelTable? @relation(fields: [model_id], references: [id])
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
@ -590,6 +591,7 @@ model LiteLLM_TeamMembership {
team_id String
spend Float @default(0.0)
budget_id String?
models String[] @default([]) // NEW: per-user model overrides
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
@@id([user_id, team_id])
}

View file

@ -0,0 +1,94 @@
import sys
import os
import pytest
# Add the parent directory to the system path to import litellm
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../..")))
import litellm
from litellm.proxy._types import UserAPIKeyAuth, LiteLLM_TeamTable
from litellm.proxy.auth.auth_checks import (
can_team_access_model,
get_effective_team_models,
)
@pytest.mark.asyncio
async def test_get_effective_team_models():
original_flag = litellm.team_model_overrides_enabled
original_env = os.environ.pop("TEAM_MODEL_OVERRIDES", None)
try:
litellm.team_model_overrides_enabled = True
# Case 1: No overrides, should return team.models
team = LiteLLM_TeamTable(team_id="t1", models=["m1"])
assert get_effective_team_models(team) == ["m1"]
# Case 2: Team defaults exist (d1 must be in team.models pool)
team = LiteLLM_TeamTable(team_id="t1", models=["m1", "d1"], default_models=["d1"])
assert set(get_effective_team_models(team)) == {"d1"}
# Case 3: Team defaults + Member overrides (all in team.models pool)
team = LiteLLM_TeamTable(
team_id="t1", models=["m1", "d1", "mo1"], default_models=["d1"]
)
token = UserAPIKeyAuth(team_member_models=["mo1"])
assert set(get_effective_team_models(team, token)) == {"d1", "mo1"}
# Case 4: No team object (should use token values if available)
token.team_default_models = ["td1"]
assert set(get_effective_team_models(None, token)) == {"td1", "mo1"}
# Case 5: Feature disabled — also ensure env var is cleared
litellm.team_model_overrides_enabled = False
os.environ.pop("TEAM_MODEL_OVERRIDES", None)
assert get_effective_team_models(team, token) == ["m1", "d1", "mo1"]
finally:
litellm.team_model_overrides_enabled = original_flag
if original_env is not None:
os.environ["TEAM_MODEL_OVERRIDES"] = original_env
@pytest.mark.asyncio
async def test_can_team_access_model_with_overrides():
original_flag = litellm.team_model_overrides_enabled
try:
litellm.team_model_overrides_enabled = True
# Team pool includes m1, d1, g1. default_models=["d1"].
team = LiteLLM_TeamTable(
team_id="t1", models=["m1", "d1", "g1"], default_models=["d1"]
)
# With only defaults, should NOT have access to m1
with pytest.raises(Exception):
await can_team_access_model(model="m1", team_object=team, llm_router=None)
# Should have access to d1 (it's a default)
assert (
await can_team_access_model(model="d1", team_object=team, llm_router=None)
is True
)
# Member has extra access to g1
token = UserAPIKeyAuth(team_member_models=["g1"])
assert (
await can_team_access_model(
model="g1", team_object=team, llm_router=None, valid_token=token
)
is True
)
assert (
await can_team_access_model(
model="d1", team_object=team, llm_router=None, valid_token=token
)
is True
)
# Should NOT have access to m1
with pytest.raises(Exception):
await can_team_access_model(
model="m1", team_object=team, llm_router=None, valid_token=token
)
finally:
litellm.team_model_overrides_enabled = original_flag

View file

@ -0,0 +1,367 @@
"""
Tests for team-scoped default + per-user model overrides.
Covers:
1. defaults_only user can access default models, not others
2. defaults + overrides user can access union of both
3. overrides only (no defaults) user can access override models only
4. neither configured falls back to team.models (backward compat, including [] = allow all)
5. key creation rejects models outside effective set 403
6. key creation with no models defaults to effective set
7. remove override next request for that model 403 (revocation)
8. all-team-models key + overrides restricted to effective set
9. cross-user isolation: User A overrides don't affect User B
10. access_group_ids fallback still works when effective models check fails
11. feature flag off all new fields ignored, team.models used
12. empty default_models + empty member models + team.models=[] allow all (backward compat)
"""
import sys
import os
import pytest
sys.path.insert(
0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../.."))
)
from unittest.mock import AsyncMock, patch
import litellm
from litellm.proxy._types import UserAPIKeyAuth, LiteLLM_TeamTable
from litellm.proxy.auth.auth_checks import (
can_team_access_model,
get_effective_team_models,
)
@pytest.fixture(autouse=True)
def enable_feature_flag():
original = litellm.team_model_overrides_enabled
litellm.team_model_overrides_enabled = True
yield
litellm.team_model_overrides_enabled = original
# ── get_effective_team_models unit tests ─────────────────────────────────────
class TestGetEffectiveTeamModels:
def test_defaults_only(self):
"""1. User with defaults only → can access default models."""
team = LiteLLM_TeamTable(
team_id="t1", models=["m1", "m2", "d1", "d2"], default_models=["d1", "d2"]
)
result = get_effective_team_models(team)
assert set(result) == {"d1", "d2"}
def test_defaults_plus_overrides(self):
"""2. User with defaults + overrides → union of both."""
team = LiteLLM_TeamTable(
team_id="t1", models=["m1", "d1", "mo1"], default_models=["d1"]
)
token = UserAPIKeyAuth(team_member_models=["mo1"])
result = get_effective_team_models(team, token)
assert set(result) == {"d1", "mo1"}
def test_overrides_only_no_defaults(self):
"""3. User with overrides only (no defaults) → can access override models."""
team = LiteLLM_TeamTable(team_id="t1", models=["m1", "mo1", "mo2"])
token = UserAPIKeyAuth(team_member_models=["mo1", "mo2"])
result = get_effective_team_models(team, token)
assert set(result) == {"mo1", "mo2"}
def test_neither_configured_fallback(self):
"""4. Neither configured → falls back to team.models."""
team = LiteLLM_TeamTable(team_id="t1", models=["m1", "m2"])
result = get_effective_team_models(team)
assert result == ["m1", "m2"]
def test_unrestricted_team_with_defaults(self):
"""team.models=[] (allow all) + default_models set → members restricted to defaults.
This is by design: admin wants unrestricted team pool but limited member defaults."""
team = LiteLLM_TeamTable(team_id="t1", models=[], default_models=["gpt-4"])
result = get_effective_team_models(team)
# Cap is skipped (team_pool=[]), so defaults pass through
assert result == ["gpt-4"]
def test_neither_configured_empty_team_models_allows_all(self):
"""12. empty default_models + empty member models + team.models=[] → allow all."""
team = LiteLLM_TeamTable(team_id="t1", models=[])
result = get_effective_team_models(team)
assert result == [] # empty = allow all
def test_cross_user_isolation(self):
"""9. User A overrides don't affect User B."""
team = LiteLLM_TeamTable(
team_id="t1", models=["m1", "d1", "mo_a", "mo_b"], default_models=["d1"]
)
token_a = UserAPIKeyAuth(team_member_models=["mo_a"])
token_b = UserAPIKeyAuth(team_member_models=["mo_b"])
result_a = get_effective_team_models(team, token_a)
result_b = get_effective_team_models(team, token_b)
assert set(result_a) == {"d1", "mo_a"}
assert set(result_b) == {"d1", "mo_b"}
assert "mo_a" not in result_b
assert "mo_b" not in result_a
def test_feature_flag_off(self, monkeypatch):
"""11. Feature flag off → all new fields ignored, team.models used."""
litellm.team_model_overrides_enabled = False
monkeypatch.delenv("TEAM_MODEL_OVERRIDES", raising=False)
team = LiteLLM_TeamTable(team_id="t1", models=["m1"], default_models=["d1"])
token = UserAPIKeyAuth(team_member_models=["mo1"])
result = get_effective_team_models(team, token)
assert result == ["m1"]
def test_deduplication(self):
"""Overlapping models are deduplicated."""
team = LiteLLM_TeamTable(
team_id="t1", models=["m1", "shared", "extra"], default_models=["shared"]
)
token = UserAPIKeyAuth(team_member_models=["shared", "extra"])
result = get_effective_team_models(team, token)
assert set(result) == {"shared", "extra"}
assert len(result) == 2 # no duplicates
def test_no_team_object(self):
"""No team object → empty list."""
assert get_effective_team_models(None) == []
def test_no_team_object_with_token(self):
"""No team object but token has defaults → uses token values."""
token = UserAPIKeyAuth(team_default_models=["td1"], team_member_models=["mo1"])
result = get_effective_team_models(None, token)
assert set(result) == {"td1", "mo1"}
# ── can_team_access_model integration tests ──────────────────────────────────
class TestCanTeamAccessModelWithOverrides:
@pytest.mark.asyncio
async def test_defaults_only_allowed(self):
"""1. User with defaults only → can access default models."""
team = LiteLLM_TeamTable(
team_id="t1", models=["m1", "m2", "d1"], default_models=["d1"]
)
assert await can_team_access_model(
model="d1", team_object=team, llm_router=None
)
@pytest.mark.asyncio
async def test_defaults_only_denied(self):
"""1. User with defaults only → cannot access other team models."""
team = LiteLLM_TeamTable(
team_id="t1", models=["m1", "m2", "d1"], default_models=["d1"]
)
with pytest.raises(Exception):
await can_team_access_model(model="m1", team_object=team, llm_router=None)
@pytest.mark.asyncio
async def test_defaults_plus_overrides_allowed(self):
"""2. User with defaults + overrides → can access union."""
team = LiteLLM_TeamTable(
team_id="t1", models=["m1", "d1", "mo1"], default_models=["d1"]
)
token = UserAPIKeyAuth(team_member_models=["mo1"])
assert await can_team_access_model(
model="d1", team_object=team, llm_router=None, valid_token=token
)
assert await can_team_access_model(
model="mo1", team_object=team, llm_router=None, valid_token=token
)
@pytest.mark.asyncio
async def test_defaults_plus_overrides_denied(self):
"""2. User with overrides → cannot access models outside union."""
team = LiteLLM_TeamTable(
team_id="t1", models=["m1", "m2", "d1", "mo1"], default_models=["d1"]
)
token = UserAPIKeyAuth(team_member_models=["mo1"])
with pytest.raises(Exception):
await can_team_access_model(
model="m2", team_object=team, llm_router=None, valid_token=token
)
@pytest.mark.asyncio
async def test_revocation_after_override_removal(self):
"""7. Remove override → model access denied."""
team = LiteLLM_TeamTable(
team_id="t1", models=["m1", "d1", "mo1"], default_models=["d1"]
)
# With override
token_with = UserAPIKeyAuth(team_member_models=["mo1"])
assert await can_team_access_model(
model="mo1", team_object=team, llm_router=None, valid_token=token_with
)
# After override removal (empty member models)
token_without = UserAPIKeyAuth(team_member_models=[])
with pytest.raises(Exception):
await can_team_access_model(
model="mo1",
team_object=team,
llm_router=None,
valid_token=token_without,
)
@pytest.mark.asyncio
async def test_stale_override_capped_by_team_models(self):
"""Stale member override for model removed from team.models → denied."""
team = LiteLLM_TeamTable(
team_id="t1", models=["m1"], default_models=["m1"]
)
# Member has stale override for "m2" which is no longer in team.models
token = UserAPIKeyAuth(team_member_models=["m2"])
# effective = union(["m1"], ["m2"]) capped by team.models=["m1"] → ["m1"]
with pytest.raises(Exception):
await can_team_access_model(
model="m2", team_object=team, llm_router=None, valid_token=token
)
@pytest.mark.asyncio
async def test_all_overrides_stale_does_not_grant_allow_all(self):
"""P0: When ALL overrides are stale (capped out), must NOT return [] (allow all).
Should fall back to team.models to prevent privilege escalation."""
team = LiteLLM_TeamTable(
team_id="t1", models=["m1"] # no default_models
)
# Member has ONLY stale overrides — none are in team.models
token = UserAPIKeyAuth(team_member_models=["stale1", "stale2"])
result = get_effective_team_models(team, token)
# Should fall back to team.models=["m1"], NOT [] (allow all)
assert result == ["m1"]
# And specifically should NOT be empty (which means allow-all)
assert result != []
# Member should be able to access m1 (team pool fallback)
assert await can_team_access_model(
model="m1", team_object=team, llm_router=None, valid_token=token
)
# But not an arbitrary model
with pytest.raises(Exception):
await can_team_access_model(
model="random-model", team_object=team, llm_router=None, valid_token=token
)
@pytest.mark.asyncio
async def test_backward_compat_no_overrides(self):
"""4. Neither configured → uses team.models as before."""
team = LiteLLM_TeamTable(team_id="t1", models=["m1", "m2"])
assert await can_team_access_model(
model="m1", team_object=team, llm_router=None
)
@pytest.mark.asyncio
async def test_backward_compat_empty_team_models_allows_all(self):
"""12. team.models=[] with no overrides → allow all."""
team = LiteLLM_TeamTable(team_id="t1", models=[])
assert await can_team_access_model(
model="any-model", team_object=team, llm_router=None
)
@pytest.mark.asyncio
async def test_feature_flag_off_uses_team_models(self, monkeypatch):
"""11. Feature flag off → ignores overrides, uses team.models."""
litellm.team_model_overrides_enabled = False
monkeypatch.delenv("TEAM_MODEL_OVERRIDES", raising=False)
team = LiteLLM_TeamTable(team_id="t1", models=["m1"], default_models=["d1"])
token = UserAPIKeyAuth(team_member_models=["mo1"])
# Should use team.models=["m1"], not effective models
assert await can_team_access_model(
model="m1", team_object=team, llm_router=None, valid_token=token
)
with pytest.raises(Exception):
await can_team_access_model(
model="d1", team_object=team, llm_router=None, valid_token=token
)
# ── Key-generation enforcement tests ─────────────────────────────────────────
class TestKeyGenerationEnforcement:
"""Tests 5, 6, 8: key-generation model validation against effective set."""
def _get_effective(self, team, token=None):
"""Helper to compute effective models (same logic as key-gen)."""
return get_effective_team_models(team, token)
def test_key_rejects_models_outside_effective_set(self):
"""5. Key creation with models outside effective set → should be rejected."""
team = LiteLLM_TeamTable(
team_id="t1", models=["m1", "m2", "m3"], default_models=["m1"]
)
token = UserAPIKeyAuth(team_member_models=["m2"])
effective = self._get_effective(team, token)
# Simulate key-gen validation: requested models must be subset of effective
requested = ["m3"] # not in effective set {m1, m2}
disallowed = set(requested) - set(effective)
assert disallowed == {"m3"}, "m3 should be disallowed"
def test_key_defaults_to_effective_set(self):
"""6. Key creation with no models → defaults to effective set."""
team = LiteLLM_TeamTable(
team_id="t1", models=["m1", "m2", "m3"], default_models=["m1"]
)
token = UserAPIKeyAuth(team_member_models=["m2"])
effective = self._get_effective(team, token)
# When no models requested, key should get effective set
assert set(effective) == {"m1", "m2"}
def test_all_team_models_restricted_to_effective_set(self):
"""8. all-team-models key + overrides → restricted to effective set, not full team.models."""
team = LiteLLM_TeamTable(
team_id="t1", models=["m1", "m2", "m3"], default_models=["m1"]
)
token = UserAPIKeyAuth(team_member_models=["m2"])
effective = self._get_effective(team, token)
# all-team-models should resolve to effective set, not team.models
assert set(effective) == {"m1", "m2"}
assert "m3" not in effective # m3 is in team.models but not in effective
# ── Access group fallback test ───────────────────────────────────────────────
class TestAccessGroupFallback:
@pytest.mark.asyncio
async def test_access_group_fallback_when_effective_models_deny(self):
"""10. access_group_ids fallback still works when effective models check fails."""
team = LiteLLM_TeamTable(
team_id="t1",
models=["m1", "m2"],
default_models=["m1"],
access_group_ids=["group-1"],
)
# "m2" is NOT in effective set (only "m1" is default, no member overrides)
# But it should be accessible via access_group_ids fallback
with patch(
"litellm.proxy.auth.auth_checks._get_models_from_access_groups",
new_callable=AsyncMock,
return_value=["m2", "m3"],
):
result = await can_team_access_model(
model="m2", team_object=team, llm_router=None
)
assert result is True
@pytest.mark.asyncio
async def test_access_group_fallback_still_denies_unknown_model(self):
"""10b. access_group_ids fallback does not grant access to models outside groups."""
team = LiteLLM_TeamTable(
team_id="t1",
models=["m1", "m2"],
default_models=["m1"],
access_group_ids=["group-1"],
)
with patch(
"litellm.proxy.auth.auth_checks._get_models_from_access_groups",
new_callable=AsyncMock,
return_value=["m2"], # group only has m2
):
with pytest.raises(Exception):
await can_team_access_model(
model="unknown-model", team_object=team, llm_router=None
)

View file

@ -54,13 +54,109 @@ async def test_upsert_disconnect(mock_tx, fake_user):
user_api_key_dict=fake_user,
)
mock_tx.litellm_teammembership.update.assert_awaited_once_with(
where={"user_id_team_id": {"user_id": "user-1", "team_id": "team-1"}},
data={"litellm_budget_table": {"disconnect": True}},
)
# All None + no existing budget → early return, no DB calls at all
mock_tx.litellm_teammembership.upsert.assert_not_called()
mock_tx.litellm_teammembership.update.assert_not_called()
mock_tx.litellm_budgettable.update.assert_not_called()
mock_tx.litellm_budgettable.create.assert_not_called()
mock_tx.litellm_teammembership.upsert.assert_not_called()
@pytest.mark.asyncio
async def test_upsert_disconnect_with_existing_budget(mock_tx, fake_user):
"""When all params are None but a budget was linked, disconnect it."""
await _upsert_budget_and_membership(
mock_tx,
team_id="team-1",
user_id="user-1",
max_budget=None,
existing_budget_id="budget-existing",
user_api_key_dict=fake_user,
)
mock_tx.litellm_teammembership.upsert.assert_awaited_once_with(
where={"user_id_team_id": {"user_id": "user-1", "team_id": "team-1"}},
data={
"create": {"user_id": "user-1", "team_id": "team-1"},
"update": {"litellm_budget_table": {"disconnect": True}},
},
)
@pytest.mark.asyncio
async def test_upsert_models_only_no_budget(mock_tx, fake_user):
"""Setting models with no budget params → upsert with models only."""
await _upsert_budget_and_membership(
mock_tx,
team_id="team-m",
user_id="user-m",
max_budget=None,
existing_budget_id=None,
user_api_key_dict=fake_user,
models=["gpt-4"],
)
mock_tx.litellm_budgettable.create.assert_not_called()
mock_tx.litellm_teammembership.upsert.assert_awaited_once_with(
where={"user_id_team_id": {"user_id": "user-m", "team_id": "team-m"}},
data={
"create": {"user_id": "user-m", "team_id": "team-m", "models": ["gpt-4"]},
"update": {"models": ["gpt-4"]},
},
)
@pytest.mark.asyncio
async def test_upsert_models_empty_list_clears(mock_tx, fake_user):
"""Setting models=[] explicitly clears overrides."""
await _upsert_budget_and_membership(
mock_tx,
team_id="team-c",
user_id="user-c",
max_budget=None,
existing_budget_id=None,
user_api_key_dict=fake_user,
models=[],
)
mock_tx.litellm_budgettable.create.assert_not_called()
mock_tx.litellm_teammembership.upsert.assert_awaited_once_with(
where={"user_id_team_id": {"user_id": "user-c", "team_id": "team-c"}},
data={
"create": {"user_id": "user-c", "team_id": "team-c", "models": []},
"update": {"models": []},
},
)
@pytest.mark.asyncio
async def test_upsert_models_plus_budget(mock_tx, fake_user):
"""Setting models alongside a budget → both written in same upsert."""
await _upsert_budget_and_membership(
mock_tx,
team_id="team-mb",
user_id="user-mb",
max_budget=50.0,
existing_budget_id=None,
user_api_key_dict=fake_user,
models=["gpt-4o"],
)
new_budget_id = mock_tx.litellm_budgettable.create.return_value.budget_id
mock_tx.litellm_teammembership.upsert.assert_awaited_once_with(
where={"user_id_team_id": {"user_id": "user-mb", "team_id": "team-mb"}},
data={
"create": {
"user_id": "user-mb",
"team_id": "team-mb",
"litellm_budget_table": {"connect": {"budget_id": new_budget_id}},
"models": ["gpt-4o"],
},
"update": {
"litellm_budget_table": {"connect": {"budget_id": new_budget_id}},
"models": ["gpt-4o"],
},
},
)
# TEST: existing budget id, creates new budget (current behavior)
@ -316,3 +412,25 @@ async def test_upsert_rpm_only_creates_new_budget(mock_tx, fake_user):
},
},
)
@pytest.mark.asyncio
async def test_upsert_budget_change_preserves_models(mock_tx, fake_user):
"""Updating budget with models=None should NOT touch existing models."""
await _upsert_budget_and_membership(
mock_tx,
team_id="team-bp",
user_id="user-bp",
max_budget=100.0,
existing_budget_id=None,
user_api_key_dict=fake_user,
# models=None (default) — should not appear in upsert data
)
call_args = mock_tx.litellm_teammembership.upsert.call_args
create_data = call_args.kwargs["data"]["create"]
update_data = call_args.kwargs["data"]["update"]
# models should NOT be in create or update data when models=None
assert "models" not in create_data, "models=None should not appear in create"
assert "models" not in update_data, "models=None should not appear in update"