mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
feat: team override model
This commit is contained in:
parent
7b31ea40a9
commit
82fecaa9a2
17 changed files with 1597 additions and 400 deletions
|
|
@ -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**
|
||||
|
||||
|
|
|
|||
|
|
@ -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[];
|
||||
|
|
@ -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])
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
},
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
}
|
||||
|
|
|
|||
94
tests/proxy_unit_tests/test_team_model_overrides.py
Normal file
94
tests/proxy_unit_tests/test_team_model_overrides.py
Normal 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
|
||||
367
tests/test_litellm/proxy/auth/test_team_model_overrides.py
Normal file
367
tests/test_litellm/proxy/auth/test_team_model_overrides.py
Normal 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
|
||||
)
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
Loading…
Add table
Reference in a new issue