feat(proxy): limit which models an end user can call (#43904)

* feat(proxy): add models column to the end user table

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* feat(proxy): enforce the end user models allowlist in model access checks

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* feat(proxy): accept and return models on the customer endpoints

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): cover the customer models allowlist

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* feat(proxy): resolve team aliases before the end user model check

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: kerry <kerry@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-05 21:58:03 +00:00 • committed by GitHub
parent a3a52466a5
commit a3e15774ad
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
17 changed files with 729 additions and 39 deletions

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_EndUserTable" ADD COLUMN IF NOT EXISTS "models" TEXT[] DEFAULT ARRAY[]::TEXT[];

View file

@ -656,6 +656,7 @@ model LiteLLM_EndUserTable {
spend Float @default(0.0)
allowed_model_region String? // require all user requests to use models in this specific region
default_model String? // use along with 'allowed_model_region'. if no available model in region, default to this model.
models String[] @default([])
budget_id String?
object_permission_id String?
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])

View file

@ -103,6 +103,7 @@ _PROXY_ERROR_TYPE_MAP: Final[Mapping[str, str]] = MappingProxyType(
"org_model_access_denied": MODEL_ACCESS_DENIED,
"project_model_access_denied": MODEL_ACCESS_DENIED,
"agent_model_access_denied": MODEL_ACCESS_DENIED,
"customer_model_access_denied": MODEL_ACCESS_DENIED,
"key_vector_store_access_denied": PERMISSION_DENIED,
"team_vector_store_access_denied": PERMISSION_DENIED,
"org_vector_store_access_denied": PERMISSION_DENIED,

View file

@ -7,7 +7,7 @@ Canonical definition for ``litellm_endusertable``. Re-exported from
from typing import Literal
from pydantic import ConfigDict, model_validator
from pydantic import ConfigDict, Field, model_validator
from litellm.models.budget import LiteLLM_BudgetTable
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
@ -21,6 +21,7 @@ class LiteLLM_EndUserTable(LiteLLMPydanticObjectBase):
spend: float = 0.0
allowed_model_region: Literal["eu", "us"] | None = None
default_model: str | None = None
models: list[str] = Field(default_factory=list)
budget_id: str | None = None
litellm_budget_table: LiteLLM_BudgetTable | None = None
object_permission_id: str | None = None

View file

@ -2130,6 +2130,7 @@ class NewCustomerRequest(BudgetNewRequest):
None # require all user requests to use models in this specific region
)
default_model: str | None = None # if no equivalent model in allowed region - default all requests to this model
models: list[str] | None = None
object_permission: LiteLLM_ObjectPermissionBase | None = None
@model_validator(mode="before")
@ -2156,6 +2157,7 @@ class UpdateCustomerRequest(LiteLLMPydanticObjectBase):
None # require all user requests to use models in this specific region
)
default_model: str | None = None # if no equivalent model in allowed region - default all requests to this model
models: list[str] | None = None
object_permission: LiteLLM_ObjectPermissionBase | None = None
@ -4464,6 +4466,11 @@ class ProxyErrorTypes(str, enum.Enum):
User does not have access to the model
"""
customer_model_access_denied = "customer_model_access_denied"
"""
Customer does not have access to the model
"""
org_model_access_denied = "org_model_access_denied"
"""
Organization does not have access to the model
@ -4553,7 +4560,7 @@ class ProxyErrorTypes(str, enum.Enum):
@classmethod
def get_model_access_error_type_for_object(
cls, object_type: Literal["key", "user", "team", "org", "project", "agent"]
cls, object_type: Literal["key", "user", "customer", "team", "org", "project", "agent"]
) -> "ProxyErrorTypes":
"""
Get the model access error type for object_type
@ -4564,6 +4571,8 @@ class ProxyErrorTypes(str, enum.Enum):
return cls.team_model_access_denied
elif object_type == "user":
return cls.user_model_access_denied
elif object_type == "customer":
return cls.customer_model_access_denied
elif object_type == "org":
return cls.org_model_access_denied
elif object_type == "project":

View file

@ -84,7 +84,10 @@ from litellm.proxy.auth.budget_throttle import (
budget_throttle_percentage,
should_throttle_budget_exceeded,
)
from litellm.proxy.auth.model_access_denied import model_access_denied_client_message
from litellm.proxy.auth.model_access_denied import (
customer_model_access_denied_client_message,
model_access_denied_client_message,
)
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import publish_auth_cache_invalidation
from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec
@ -151,7 +154,7 @@ from .auth_checks_organization import (
add_team_org_context_to_request_body,
organization_role_based_access_check,
)
from .auth_utils import get_model_from_request, get_request_route_template
from .auth_utils import get_model_from_request, get_request_route_template, request_fallback_model_names
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
@ -1068,7 +1071,9 @@ async def common_checks(
if not isinstance(managed_models, (list, tuple)) or not managed_models:
raise HTTPException(403, "This agent has no model grants")
_can_object_call_model(
model=_resolve_team_alias(_model, valid_token.team_model_aliases, valid_token.team_id, llm_router),
model=_resolve_team_alias(
_model, team_model_aliases_for_auth_check(valid_token), valid_token.team_id, llm_router
),
llm_router=llm_router,
models=list(managed_models),
team_id=valid_token.team_id,
@ -1096,6 +1101,23 @@ async def common_checks(
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
)
if end_user_object is not None and end_user_object.models:
with tracer.trace("litellm.proxy.auth.common_checks.can_customer_call_model"):
if _model:
can_customer_access_model(
model=_model,
end_user_object=end_user_object,
llm_router=llm_router,
valid_token=valid_token,
)
for fallback_model in request_fallback_model_names(_typed_request_body(request_body)):
can_customer_access_model(
model=fallback_model,
end_user_object=end_user_object,
llm_router=llm_router,
valid_token=valid_token,
)
# 1.1 - 2.2 - 3.0.2 - 3.0.3: Project checks (blocked, model access, budget)
with tracer.trace("litellm.proxy.auth.common_checks.run_project_checks"):
await _run_project_checks(
@ -1436,6 +1458,7 @@ def get_actual_routes(allowed_routes: list) -> list:
KEY_END_USER_BUDGET_ID_METADATA_FIELD: Final = "end_user_budget_id"
_KEY_METADATA_ADAPTER: Final[TypeAdapter[Mapping[str, object] | None]] = TypeAdapter(Mapping[str, object] | None)
def get_key_end_user_budget_id(key_metadata: Mapping[str, object] | None) -> str | None:
@ -1698,9 +1721,15 @@ def _column_is_set(column: str) -> Mapping[str, object]:
return {column: {"not": None}}
def _array_is_not_empty(column: str) -> Mapping[str, object]:
"""``column`` holds at least one element, as a plain dict for prisma's builder."""
return {column: {"is_empty": False}}
def _restricted_end_user_where() -> Mapping[str, object]:
"""Prisma filter selecting every end-user row that carries a restriction auth enforces."""
return {"OR": [{"blocked": True}, *map(_column_is_set, _RESTRICTED_COLUMNS)]}
restrictions: Final = (*map(_column_is_set, _RESTRICTED_COLUMNS), _array_is_not_empty("models"))
return {"OR": [{"blocked": True}, *restrictions]}
class _RegistryNotCached:
@ -1857,8 +1886,8 @@ async def _end_user_is_known_unrestricted(
True when the cached registry proves the id restricts nothing, so its row need not be read.
Every field ``get_end_user_object`` callers consume (budget, spend under that budget, region,
default model, object permission, blocked) is part of the registry predicate, so an id outside
it is indistinguishable from one with no row at all. The skip is off whenever mere existence of
default model, models, object permission, blocked) is part of the registry predicate, so an id
outside it is indistinguishable from one with no row at all. The skip is off whenever mere existence of
the row is meaningful: ``max_end_user_budget_id`` or the key's ``end_user_budget_id`` grafts a
default budget onto any row that exists, ``validate_end_user_id_in_db`` rejects ids that resolve
to no row, and a token-supplied ``end_user_max_budget`` (a ``user_custom_auth`` callable can set
@ -4463,7 +4492,7 @@ def _can_object_call_model(
team_model_aliases: dict[str, str] | None = None,
team_id: str | None = None,
key_model_aliases: Mapping[str, str] | None = None,
object_type: Literal["user", "team", "key", "org", "project", "agent"] = "user",
object_type: Literal["user", "customer", "team", "key", "org", "project", "agent"] = "user",
fallback_depth: int = 0,
) -> Literal[True]:
"""
@ -4544,7 +4573,11 @@ def _can_object_call_model(
f"Tried to access {model}"
)
raise ModelAccessDeniedProxyException(
message=model_access_denied_client_message(model=model),
message=(
customer_model_access_denied_client_message(model=model)
if object_type == "customer"
else model_access_denied_client_message(model=model)
),
internal_message=internal_message,
type=ProxyErrorTypes.get_model_access_error_type_for_object(object_type=object_type),
param="model",
@ -4554,7 +4587,7 @@ def _can_object_call_model(
def _resolve_team_alias(
model: str | list[str],
team_model_aliases: dict[str, str] | None,
team_model_aliases: Mapping[str, str] | None,
team_id: str | None,
llm_router: Router | None,
) -> str | list[str]:
@ -4566,7 +4599,7 @@ def _resolve_team_alias(
def _live_team_alias_target(
model: str, team_model_aliases: dict[str, str], team_id: str | None, llm_router: Router | None
model: str, team_model_aliases: Mapping[str, str], team_id: str | None, llm_router: Router | None
) -> str:
target: Final = team_model_aliases.get(model)
if target is None:
@ -4600,7 +4633,9 @@ async def _check_agent_access_group_model_access(
if unmanaged is not None
else ()
)
dispatched: Final = _resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router)
dispatched: Final = _resolve_team_alias(
model, team_model_aliases_for_auth_check(valid_token), valid_token.team_id, llm_router
)
for ceiling in ceilings:
if not ceiling.models:
raise ModelAccessDeniedProxyException(
@ -4698,6 +4733,10 @@ def key_model_aliases_for_auth_check(valid_token: UserAPIKeyAuth | None) -> Mapp
return alias_map(valid_token.aliases) if valid_token is not None and valid_token.aliases else None
def team_model_aliases_for_auth_check(valid_token: UserAPIKeyAuth) -> Mapping[str, str] | None:
return alias_map(valid_token.team_model_aliases) if valid_token.team_model_aliases else None
def _resolve_key_models_for_auth_check(valid_token: UserAPIKeyAuth) -> list[str]:
"""
Expand key model sentinels before auth checks.
@ -5135,6 +5174,24 @@ async def can_key_call_resolved_model(
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
)
if valid_token.end_user_id is not None and prisma_client is not None:
key_metadata: Final = _KEY_METADATA_ADAPTER.validate_python(valid_token.metadata)
end_user_object: Final = await get_end_user_object(
end_user_id=valid_token.end_user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
token_end_user_max_budget=valid_token.end_user_max_budget,
key_end_user_budget_id=get_key_end_user_budget_id(key_metadata),
)
if end_user_object is not None and end_user_object.models:
can_customer_access_model(
model=model,
end_user_object=end_user_object,
llm_router=llm_router,
valid_token=valid_token,
)
def can_org_access_model(
model: str,
@ -5302,6 +5359,29 @@ def can_project_access_model(
)
def can_customer_access_model(
model: str | list[str],
end_user_object: LiteLLM_EndUserTable,
llm_router: Router | None,
valid_token: UserAPIKeyAuth | None,
) -> Literal[True]:
team_model_aliases: Final = team_model_aliases_for_auth_check(valid_token) if valid_token is not None else None
listed_names: Final = frozenset(end_user_object.models or ())
unlisted_aliases: Final = (
MappingProxyType({alias: target for alias, target in team_model_aliases.items() if alias not in listed_names})
if team_model_aliases
else None
)
team_id: Final = valid_token.team_id if valid_token is not None else None
return _can_object_call_model(
model=_resolve_team_alias(model, unlisted_aliases, team_id, llm_router),
llm_router=llm_router,
models=end_user_object.models,
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
object_type="customer",
)
async def can_user_call_model(
model: str | list[str],
llm_router: Router | None,

View file

@ -522,6 +522,26 @@ def iter_request_fallback_targets(request_body: Mapping[str, object]) -> Iterato
yield from _iter_fallback_targets(value, 0)
def fallback_target_model_name(target: object) -> str | None:
if isinstance(target, str):
return target
if isinstance(target, Mapping):
model: Final = target.get("model")
if isinstance(model, str):
return model
return None
def request_fallback_model_names(request_body: Mapping[str, object]) -> tuple[str, ...]:
return tuple(
dict.fromkeys(
name
for target in iter_request_fallback_targets(request_body)
if (name := fallback_target_model_name(target)) is not None
)
)
def _reject_url_valued_fallback_target(value: str) -> None:
allowed_hosts: Final = getattr(litellm, "provider_url_destination_allowed_hosts", []) or []
for candidate in provider_url_destination_candidates(value):

View file

@ -7,11 +7,20 @@ MODEL_ACCESS_DENIED_CLIENT_MESSAGE: Final = (
"Check the models available to you and try again."
)
CUSTOMER_MODEL_ACCESS_DENIED_CLIENT_MESSAGE: Final = (
"The requested model '{model}' is not in the allowed models for this customer. "
"Check the models this customer can use and try again."
)
def model_access_denied_client_message(model: str | list[str]) -> str:
return MODEL_ACCESS_DENIED_CLIENT_MESSAGE.format(model=model)
def customer_model_access_denied_client_message(model: str | list[str]) -> str:
return CUSTOMER_MODEL_ACCESS_DENIED_CLIENT_MESSAGE.format(model=model)
class ModelAccessDeniedHTTPException(HTTPException):
def __init__(self, internal_message: str, status_code: int, detail: str | dict[str, str]) -> None:
super().__init__(status_code=status_code, detail=detail)

View file

@ -82,6 +82,7 @@ from litellm.proxy.auth.auth_object_prefetch import (
)
from litellm.proxy.auth.auth_utils import (
abbreviate_api_key,
fallback_target_model_name,
get_end_user_id_from_request_body,
get_model_from_request,
get_request_route,
@ -3796,7 +3797,7 @@ async def _enforce_key_and_fallback_model_access(
fallback_names: Final = tuple(
name
for target in iter_request_fallback_targets(request_data)
if (name := _fallback_target_model_name(target)) is not None
if (name := fallback_target_model_name(target)) is not None
)
for _name in dict.fromkeys(fallback_names): # dedupe, preserve order
@ -3813,16 +3814,6 @@ async def _enforce_key_and_fallback_model_access(
)
def _fallback_target_model_name(target: object) -> str | None:
if isinstance(target, str):
return target
if isinstance(target, dict):
model: Final = target.get("model")
if isinstance(model, str):
return model
return None
async def _run_post_custom_auth_checks(
valid_token: UserAPIKeyAuth,
request: Request,

View file

@ -11,6 +11,7 @@ All /customer management endpoints
#### END-USER/CUSTOMER MANAGEMENT ####
from collections.abc import Mapping, Sequence
from collections.abc import Set as AbstractSet
from datetime import datetime, timedelta
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, NamedTuple, Protocol, TypeVar, overload
@ -56,6 +57,18 @@ from litellm.types.proxy.management_endpoints.customer_endpoints import (
_RowT_co: Final = TypeVar("_RowT_co", covariant=True)
_STR_OBJECT_DICT: Final = TypeAdapter(dict[str, object])
_CLEARABLE_LIST_FIELDS: Final = frozenset({"models"})
def _should_update_field(field: str, value: object, sent_fields: AbstractSet[str]) -> bool:
if value is None:
return False
if field in sent_fields and (isinstance(value, bool) or field in _CLEARABLE_LIST_FIELDS):
return True
if isinstance(value, (list, dict)) and not value:
return False
return value != 0
if TYPE_CHECKING:
@ -333,6 +346,7 @@ async def new_end_user(
- budget_id: Optional[str] - The identifier for an existing budget allocated to the user. Either 'max_budget' or 'budget_id' should be provided, not both.
- allowed_model_region: Optional[Union[Literal["eu"], Literal["us"]]] - Require all user requests to use models in this specific region.
- default_model: Optional[str] - If no equivalent model in the allowed region, default all requests to this model.
- models: Optional[list[str]] - Restrict this customer's access to the listed models.
- metadata: Optional[dict] = Metadata for customer, store information for customer. Example metadata = {"data_training_opt_out": True}
- budget_duration: Optional[str] - Budget is reset at the end of specified duration. If not set, budget is never reset. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d").
- tpm_limit: Optional[int] - [Not Implemented Yet] Specify tpm limit for a given customer (Tokens per minute)
@ -367,6 +381,7 @@ async def new_end_user(
"user_id" : "ishaan-jaff-3",
"allowed_region": "eu",
"budget_id": "free_tier",
"models": ["gpt-4o-mini"],
"default_model": "azure/gpt-3.5-turbo-eu"
}'
@ -608,6 +623,7 @@ async def update_end_user(
- default_model: Optional[str] = (
None # if no equivalent model in allowed region - default all requests to this model
)
- models: Optional[list[str]] = None # omitted or null leaves the allowlist unchanged; an empty list clears it
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - Customer-specific object permissions to control access to resources.
Supported fields:
* mcp_servers: List[str] - List of allowed MCP server IDs
@ -626,7 +642,8 @@ async def update_end_user(
--header 'Content-Type: application/json' \
--data '{
"user_id": "test-litellm-user-4",
"budget_id": "paid_tier"
"budget_id": "paid_tier",
"models": ["gpt-4o-mini"]
}'
# Updating object permissions
@ -653,11 +670,10 @@ async def update_end_user(
if prisma_client is None:
raise Exception("Not connected to DB!")
# get non default values for key
non_default_values: Final = dict[str, object]()
for k, v in data_json.items():
if v is not None and ((isinstance(v, bool) and k in data.fields_set()) or v not in ([], {}, 0)):
non_default_values[k] = v
sent_fields: Final = data.fields_set()
non_default_values: Final[dict[str, object]] = {
k: v for k, v in data_json.items() if _should_update_field(k, v, sent_fields)
}
## Get end user table data ##
end_user_table_data: Final = await _typed_table(EndUserRepository(prisma_client)).find_first(

View file

@ -656,6 +656,7 @@ model LiteLLM_EndUserTable {
spend Float @default(0.0)
allowed_model_region String? // require all user requests to use models in this specific region
default_model String? // use along with 'allowed_model_region'. if no available model in region, default to this model.
models String[] @default([])
budget_id String?
object_permission_id String?
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])

View file

@ -656,6 +656,7 @@ model LiteLLM_EndUserTable {
spend Float @default(0.0)
allowed_model_region String? // require all user requests to use models in this specific region
default_model String? // use along with 'allowed_model_region'. if no available model in region, default to this model.
models String[] @default([])
budget_id String?
object_permission_id String?
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])

View file

@ -0,0 +1,293 @@
import json
import uuid
from collections.abc import Mapping
from pathlib import Path
from typing import Final, Literal
import httpx
import pytest
import yaml
from pydantic import JsonValue, TypeAdapter
from tests.integration._support.client import Gateway, Scenario, object_value, string_value
from tests.integration._support.process import owned_proxy
from tests.integration._support.wire import Reply, Request, Wire, wire_server
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_CHAT_RESPONSE: Final = {
"object": "chat.completion",
"created": 1700000000,
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "scripted response"},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 2, "completion_tokens": 3, "total_tokens": 5},
}
def _customer(scenario: Scenario, models: list[str] | None = None) -> str:
identity: Final = f"integration-customer-{uuid.uuid4().hex}"
body: Final[dict[str, JsonValue]] = {
"user_id": identity,
**({"models": models} if models is not None else {}),
}
scenario.gateway.post("/customer/new", body)
scenario.cleanups.callback(scenario.gateway.post, "/customer/delete", {"user_ids": [identity]})
return identity
def _reply(request: Request) -> Reply:
if request.method != "POST" or not request.body:
return Reply(body=b'{"object":"list","data":[]}')
body: Final = _JSON_OBJECT.validate_json(request.body)
response: Final = {
"id": f"chatcmpl-{uuid.uuid4().hex}",
"model": string_value(body["model"]),
**_CHAT_RESPONSE,
}
return Reply(body=json.dumps(response).encode())
def _fallback_reply(request: Request) -> Reply:
if request.method == "POST" and request.body:
body: Final = _JSON_OBJECT.validate_json(request.body)
if "scripted-primary-" in string_value(body["model"]):
return Reply(status=500, body=b'{"error":{"message":"scripted primary failure","type":"server_error"}}')
return _reply(request)
def _chat(
gateway: Gateway,
key: str,
model: str,
marker: str,
*,
customer: str | None = None,
headers: Mapping[str, str] | None = None,
fallbacks: tuple[str, ...] | None = None,
) -> httpx.Response:
body: Final[dict[str, JsonValue]] = {
"model": model,
"messages": [{"role": "user", "content": marker}],
**({"user": customer} if customer is not None else {}),
**({"fallbacks": list(fallbacks)} if fallbacks is not None else {}),
}
return gateway.request("POST", "/v1/chat/completions", body, key=key, headers=headers)
def _error_type(response: httpx.Response) -> str:
error: Final = object_value(_JSON_OBJECT.validate_json(response.content)["error"])
return string_value(error["type"])
def _assert_upstream(wire: Wire, marker: str, expected: int) -> tuple[Request, ...]:
matching: Final = tuple(request for request in wire.drain() if marker.encode() in request.body)
assert len(matching) == expected, matching
return matching
@pytest.mark.parametrize(
"customer_state",
("without_models", "empty_models", "no_customer_id", "unknown_customer_id"),
)
def test_unrestricted_customers_keep_key_model_access(
gateway: Gateway,
customer_state: Literal["without_models", "empty_models", "no_customer_id", "unknown_customer_id"],
) -> None:
with wire_server(_reply) as wire, gateway.scenario() as scenario:
first: Final = scenario.model(api_base=f"{wire.url}/v1")
second: Final = scenario.model(api_base=f"{wire.url}/v1")
key: Final = scenario.key(models=[first, second])
customer: Final = (
_customer(scenario)
if customer_state == "without_models"
else _customer(scenario, [])
if customer_state == "empty_models"
else f"missing-{uuid.uuid4().hex}"
if customer_state == "unknown_customer_id"
else None
)
for model in (first, second):
marker: Final = uuid.uuid4().hex
response: Final = _chat(gateway, key, model, marker, customer=customer)
assert response.status_code == 200, response.text
assert len(_assert_upstream(wire, marker, 1)) == 1
def test_customer_and_key_model_lists_are_both_enforced(gateway: Gateway) -> None:
with wire_server(_reply) as wire, gateway.scenario() as scenario:
model_a: Final = scenario.model(api_base=f"{wire.url}/v1")
model_b: Final = scenario.model(api_base=f"{wire.url}/v1")
model_c: Final = scenario.model(api_base=f"{wire.url}/v1")
key: Final = scenario.key(models=[model_a, model_b])
headers: Final = {"x-litellm-end-user-id": _customer(scenario, [model_b, model_c])}
marker_a: Final = uuid.uuid4().hex
denied_by_customer: Final = _chat(gateway, key, model_a, marker_a, headers=headers)
assert denied_by_customer.status_code == 403, denied_by_customer.text
assert _error_type(denied_by_customer) == "customer_model_access_denied"
assert _assert_upstream(wire, marker_a, 0) == ()
marker_b: Final = uuid.uuid4().hex
allowed: Final = _chat(gateway, key, model_b, marker_b, headers=headers)
assert allowed.status_code == 200, allowed.text
assert len(_assert_upstream(wire, marker_b, 1)) == 1
marker_c: Final = uuid.uuid4().hex
denied_by_key: Final = _chat(gateway, key, model_c, marker_c, headers=headers)
assert denied_by_key.status_code == 403, denied_by_key.text
assert _error_type(denied_by_key) == "key_model_access_denied"
assert _assert_upstream(wire, marker_c, 0) == ()
def test_request_body_fallback_outside_customer_allowlist_is_denied(gateway: Gateway) -> None:
with wire_server(_fallback_reply) as wire, gateway.scenario() as scenario:
primary: Final = scenario.model(api_base=f"{wire.url}/v1")
fallback: Final = scenario.model(api_base=f"{wire.url}/v1")
key: Final = scenario.key(models=[primary, fallback])
customer: Final = _customer(scenario, [primary])
marker: Final = uuid.uuid4().hex
response: Final = _chat(
gateway,
key,
primary,
marker,
customer=customer,
fallbacks=(fallback,),
)
assert response.status_code == 403, response.text
assert _error_type(response) == "customer_model_access_denied"
assert _assert_upstream(wire, marker, 0) == ()
def _router_config(
directory: Path,
wire: Wire,
primary: str,
fallback: str,
allowed_fallback: str,
primary_provider: str,
fallback_provider: str,
allowed_provider: str,
*,
fallback_target: str,
enforce: bool,
) -> Path:
models: Final = (
(primary, primary_provider),
(fallback, fallback_provider),
(allowed_fallback, allowed_provider),
)
config: Final = {
"model_list": [
{
"model_name": name,
"litellm_params": {
"model": provider_model,
"api_key": "integration-provider-key",
"api_base": f"{wire.url}/v1",
},
}
for name, provider_model in models
],
"general_settings": {
"master_key": "os.environ/LITELLM_MASTER_KEY",
"store_model_in_db": False,
"enforce_fallback_model_access": enforce,
},
"litellm_settings": {"cache": False},
"router_settings": {
"num_retries": 0,
"disable_cooldowns": True,
"fallbacks": [{primary: [fallback_target]}],
},
}
path: Final = directory / f"customer-fallback-{uuid.uuid4().hex}.yaml"
path.write_text(yaml.safe_dump(config))
return path
@pytest.mark.parametrize(
("enforce", "allow_fallback", "expected_status", "expected_upstream_count"),
(
(True, False, 500, 1),
(False, False, 200, 2),
(True, True, 200, 2),
),
)
def test_router_config_fallback_customer_allowlist(
gateway: Gateway,
tmp_path: Path,
enforce: bool,
allow_fallback: bool,
expected_status: int,
expected_upstream_count: int,
) -> None:
with wire_server(_fallback_reply) as wire, gateway.scenario() as model_scenario:
primary_provider: Final = f"openai/scripted-primary-{uuid.uuid4().hex}"
fallback_provider: Final = f"openai/scripted-fallback-{uuid.uuid4().hex}"
allowed_provider: Final = f"openai/scripted-allowed-{uuid.uuid4().hex}"
primary: Final = model_scenario.model(api_base=f"{wire.url}/v1", model=primary_provider)
fallback: Final = model_scenario.model(api_base=f"{wire.url}/v1", model=fallback_provider)
allowed_fallback: Final = model_scenario.model(api_base=f"{wire.url}/v1", model=allowed_provider)
fallback_target: Final = allowed_fallback if allow_fallback else fallback
config: Final = _router_config(
tmp_path,
wire,
primary,
fallback,
allowed_fallback,
primary_provider,
fallback_provider,
allowed_provider,
fallback_target=fallback_target,
enforce=enforce,
)
with owned_proxy(gateway, tmp_path, {}, config=config) as candidate:
with candidate.scenario() as scenario:
key: Final = scenario.key(models=[primary, fallback, allowed_fallback])
customer_models: Final = [primary, fallback_target] if allow_fallback else [primary]
customer: Final = _customer(scenario, customer_models)
marker: Final = uuid.uuid4().hex
response: Final = _chat(candidate, key, primary, marker, customer=customer)
assert response.status_code == expected_status, response.text
requests: Final = _assert_upstream(wire, marker, expected_upstream_count)
models: Final = tuple(
string_value(_JSON_OBJECT.validate_json(request.body)["model"]) for request in requests
)
assert "scripted-primary-" in models[0]
if expected_upstream_count == 2:
assert ("scripted-allowed-" in models[1]) is allow_fallback
assert ("scripted-fallback-" in models[1]) is not allow_fallback
def test_customer_crud_sets_and_clears_models_immediately(gateway: Gateway) -> None:
with wire_server(_reply) as wire, gateway.scenario() as scenario:
model_a: Final = scenario.model(api_base=f"{wire.url}/v1")
model_b: Final = scenario.model(api_base=f"{wire.url}/v1")
key: Final = scenario.key(models=[model_a, model_b])
customer: Final = _customer(scenario)
assert gateway.get("/customer/info", {"end_user_id": customer})["models"] == []
gateway.post("/customer/update", {"user_id": customer, "models": [model_a]})
assert gateway.get("/customer/info", {"end_user_id": customer})["models"] == [model_a]
denied_marker: Final = uuid.uuid4().hex
denied: Final = _chat(gateway, key, model_b, denied_marker, customer=customer)
assert denied.status_code == 403, denied.text
assert _error_type(denied) == "customer_model_access_denied"
assert _assert_upstream(wire, denied_marker, 0) == ()
gateway.post("/customer/update", {"user_id": customer, "models": []})
assert gateway.get("/customer/info", {"end_user_id": customer})["models"] == []
allowed_marker: Final = uuid.uuid4().hex
allowed: Final = _chat(gateway, key, model_b, allowed_marker, customer=customer)
assert allowed.status_code == 200, allowed.text
assert len(_assert_upstream(wire, allowed_marker, 1)) == 1

View file

@ -7157,6 +7157,7 @@ _RESTRICTED_END_USER_WHERE = {
{"allowed_model_region": {"not": None}},
{"default_model": {"not": None}},
{"object_permission_id": {"not": None}},
{"models": {"is_empty": False}},
]
}
@ -8875,6 +8876,156 @@ async def _run_common_checks(
)
async def _common_checks_for_customer_model(
*,
model: str,
customer_models: list[str],
request_overrides: Mapping[str, object] | None = None,
team_model_aliases: dict[str, str] | None = None,
) -> bool:
from litellm.proxy.auth.auth_checks import common_checks
return await common_checks(
request_body={
"model": model,
"messages": [{"role": "user", "content": "hi"}],
**(request_overrides or {}),
},
team_object=None,
user_object=None,
end_user_object=LiteLLM_EndUserTable(user_id="customer-1", blocked=False, models=customer_models),
global_proxy_spend=None,
general_settings={},
route="/chat/completions",
llm_router=None,
proxy_logging_obj=MagicMock(),
valid_token=UserAPIKeyAuth(token="test-token", team_model_aliases=team_model_aliases),
request=MagicMock(spec=Request),
skip_budget_checks=True,
)
@pytest.mark.asyncio
async def test_common_checks_allows_model_in_customer_allowlist() -> None:
assert await _common_checks_for_customer_model(model="A", customer_models=["A"]) is True
@pytest.mark.asyncio
async def test_common_checks_denies_model_outside_customer_allowlist() -> None:
with pytest.raises(ModelAccessDeniedProxyException) as exc_info:
await _common_checks_for_customer_model(model="B", customer_models=["A"])
assert exc_info.value.type == ProxyErrorTypes.customer_model_access_denied
@pytest.mark.asyncio
async def test_common_checks_denies_request_fallback_outside_customer_allowlist() -> None:
with pytest.raises(ModelAccessDeniedProxyException) as exc_info:
await _common_checks_for_customer_model(
model="A",
customer_models=["A"],
request_overrides={"fallbacks": ["B"]},
)
assert exc_info.value.type == ProxyErrorTypes.customer_model_access_denied
@pytest.mark.asyncio
async def test_common_checks_allows_model_with_empty_customer_allowlist() -> None:
assert await _common_checks_for_customer_model(model="B", customer_models=[]) is True
@pytest.mark.asyncio
async def test_common_checks_matches_team_alias_target_against_customer_allowlist() -> None:
team_model_aliases: Final = {"fast": "m1", "slow": "gpt-4o"}
for customer_models in (["m1"], ["fast"]):
assert (
await _common_checks_for_customer_model(
model="fast", customer_models=customer_models, team_model_aliases=team_model_aliases
)
is True
)
with pytest.raises(ModelAccessDeniedProxyException) as exc_info:
await _common_checks_for_customer_model(
model="slow", customer_models=["m1"], team_model_aliases=team_model_aliases
)
assert exc_info.value.type == ProxyErrorTypes.customer_model_access_denied
with pytest.raises(ModelAccessDeniedProxyException):
await _common_checks_for_customer_model(
model="m1", customer_models=["fast"], team_model_aliases=team_model_aliases
)
@pytest.mark.parametrize(
("model", "customer_models", "denied"),
(
("A", ["A"], False),
("B", ["A"], True),
("B", [], False),
),
)
@pytest.mark.asyncio
async def test_can_key_call_resolved_model_checks_customer_allowlist(
monkeypatch: pytest.MonkeyPatch,
model: str,
customer_models: list[str],
denied: bool,
) -> None:
from litellm.proxy.auth import auth_checks
monkeypatch.setattr(proxy_server, "prisma_client", MagicMock())
monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock())
monkeypatch.setattr(proxy_server, "user_api_key_cache", MagicMock())
customer_lookup: Final = AsyncMock(
return_value=LiteLLM_EndUserTable(user_id="customer-1", blocked=False, models=customer_models)
)
monkeypatch.setattr(auth_checks, "get_end_user_object", customer_lookup)
valid_token: Final = UserAPIKeyAuth(end_user_id="customer-1", models=[])
if denied:
with pytest.raises(ModelAccessDeniedProxyException) as exc_info:
await auth_checks.can_key_call_resolved_model(
model=model,
llm_model_list=None,
valid_token=valid_token,
llm_router=None,
)
assert exc_info.value.type == ProxyErrorTypes.customer_model_access_denied
else:
await auth_checks.can_key_call_resolved_model(
model=model,
llm_model_list=None,
valid_token=valid_token,
llm_router=None,
)
customer_lookup.assert_awaited_once()
@pytest.mark.asyncio
async def test_can_key_call_resolved_model_skips_customer_lookup_without_customer_id(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from litellm.proxy.auth import auth_checks
monkeypatch.setattr(proxy_server, "prisma_client", MagicMock())
monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock())
monkeypatch.setattr(proxy_server, "user_api_key_cache", MagicMock())
customer_lookup: Final = AsyncMock()
monkeypatch.setattr(auth_checks, "get_end_user_object", customer_lookup)
await auth_checks.can_key_call_resolved_model(
model="B",
llm_model_list=None,
valid_token=UserAPIKeyAuth(models=[]),
llm_router=None,
)
customer_lookup.assert_not_awaited()
@pytest.mark.asyncio
async def test_common_checks_blocks_unpriced_model_when_enabled(monkeypatch):
monkeypatch.setattr(litellm, "block_requests_for_models_without_pricing", True)

View file

@ -11,11 +11,8 @@ from unittest.mock import AsyncMock, patch
import pytest
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import iter_request_fallback_targets
from litellm.proxy.auth.user_api_key_auth import (
_enforce_key_and_fallback_model_access,
_fallback_target_model_name,
)
from litellm.proxy.auth.auth_utils import fallback_target_model_name, iter_request_fallback_targets
from litellm.proxy.auth.user_api_key_auth import _enforce_key_and_fallback_model_access
def _fallback_model_names(fallbacks):
@ -23,7 +20,7 @@ def _fallback_model_names(fallbacks):
return [
name
for target in iter_request_fallback_targets({"fallbacks": fallbacks})
if (name := _fallback_target_model_name(target)) is not None
if (name := fallback_target_model_name(target)) is not None
]

View file

@ -13,7 +13,7 @@ from litellm.proxy._types import (
ProxyException,
)
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.management_endpoints.customer_endpoints import router
from litellm.proxy.management_endpoints.customer_endpoints import _should_update_field, router
from litellm.types.proxy.management_endpoints.common_daily_activity import (
SpendAnalyticsPaginatedResponse,
)
@ -42,6 +42,48 @@ app.include_router(router)
client = TestClient(app)
@pytest.mark.parametrize(
("field", "value", "sent_fields", "expected"),
(
("models", None, frozenset({"models"}), False),
("blocked", False, frozenset({"blocked"}), True),
("blocked", False, frozenset(), False),
("models", [], frozenset({"models"}), True),
("models", [], frozenset(), False),
("metadata", {}, frozenset({"metadata"}), False),
("object_permission", {"mcp_servers": ["s1"]}, frozenset(), True),
("object_permission", {"mcp_servers": ["s1"]}, frozenset({"object_permission"}), True),
("metadata", ["m1"], frozenset(), True),
("models", ["m1"], frozenset(), True),
("max_budget", 0, frozenset({"max_budget"}), False),
("max_budget", 5.0, frozenset(), True),
("alias", "a", frozenset(), True),
),
ids=(
"null-models",
"explicit-false",
"omitted-false",
"clear-models",
"omitted-empty-models",
"empty-metadata",
"nonempty-object-permission-omitted",
"nonempty-object-permission-sent",
"nonempty-list-other-field",
"nonempty-models",
"zero-budget",
"nonzero-budget",
"alias",
),
)
def test_should_update_field(
field: str,
value: object,
sent_fields: frozenset[str],
expected: bool,
) -> None:
assert _should_update_field(field, value, sent_fields) is expected
@pytest.fixture
def mock_prisma_client():
with patch("litellm.proxy.proxy_server.prisma_client") as mock:
@ -749,6 +791,7 @@ _FULL_DB_ROW = {
"spend": 1.5,
"allowed_model_region": None,
"default_model": None,
"models": ["allowed-model"],
"budget_id": "b1",
"object_permission_id": "p1",
"litellm_budget_table": {
@ -794,6 +837,7 @@ _EXPECTED_CUSTOMER = {
"spend": 1.5,
"allowed_model_region": None,
"default_model": None,
"models": ["allowed-model"],
"budget_id": "b1",
"litellm_budget_table": {
"budget_id": "b1",
@ -857,6 +901,19 @@ def test_char_new_body(mock_prisma_client, mock_user_api_key_auth):
assert response.json() == _EXPECTED_CUSTOMER
def test_customer_new_forwards_models_to_db(mock_prisma_client, mock_user_api_key_auth):
mock_prisma_client.db.litellm_endusertable.create = AsyncMock(return_value=_row(_FULL_DB_ROW))
response = client.post(
"/customer/new",
json={"user_id": "c1", "models": ["allowed-model"]},
headers={"Authorization": "Bearer k"},
)
assert response.status_code == 200, response.text
assert mock_prisma_client.db.litellm_endusertable.create.call_args.kwargs["data"]["models"] == ["allowed-model"]
@pytest.mark.parametrize("bad_duration", ["0s", "-5m"])
def test_customer_new_rejects_a_duration_that_never_advances(
mock_prisma_client, mock_user_api_key_auth, bad_duration
@ -903,6 +960,50 @@ def test_char_update_body(mock_prisma_client, mock_user_api_key_auth):
)
assert response.status_code == 200
assert response.json() == _EXPECTED_CUSTOMER
assert "models" not in mock_prisma_client.db.litellm_endusertable.update.call_args.kwargs["data"]
def test_customer_update_clears_models_allowlist(mock_prisma_client, mock_user_api_key_auth):
mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(
return_value=_row({"user_id": "c1", "blocked": False})
)
mock_prisma_client.db.litellm_endusertable.update = AsyncMock(return_value=_row(_FULL_DB_ROW))
response = client.post(
"/customer/update",
json={"user_id": "c1", "models": []},
headers={"Authorization": "Bearer k"},
)
assert response.status_code == 200, response.text
assert mock_prisma_client.db.litellm_endusertable.update.call_args.kwargs["data"]["models"] == []
def test_customer_update_applies_nonempty_object_permission(mock_prisma_client, mock_user_api_key_auth):
mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(
return_value=_row({"user_id": "c1", "blocked": False, "object_permission_id": None})
)
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=None)
updated_permission = MagicMock()
updated_permission.object_permission_id = "permission-1"
mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock(return_value=updated_permission)
mock_prisma_client.db.litellm_endusertable.update = AsyncMock(
return_value=LiteLLM_EndUserTable(user_id="c1", blocked=False, object_permission_id="permission-1")
)
response = client.post(
"/customer/update",
json={"user_id": "c1", "object_permission": {"mcp_servers": ["s1"]}},
headers={"Authorization": "Bearer k"},
)
assert response.status_code == 200, response.text
mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_awaited_once()
permission_upsert = mock_prisma_client.db.litellm_objectpermissiontable.upsert.call_args.kwargs
assert permission_upsert["data"]["create"]["mcp_servers"] == ["s1"]
assert mock_prisma_client.db.litellm_endusertable.update.call_args.kwargs["data"]["object_permission_id"] == (
"permission-1"
)
def test_char_delete_body(mock_prisma_client, mock_user_api_key_auth):

View file

@ -4301,6 +4301,7 @@ export interface paths {
* - budget_id: Optional[str] - The identifier for an existing budget allocated to the user. Either 'max_budget' or 'budget_id' should be provided, not both.
* - allowed_model_region: Optional[Union[Literal["eu"], Literal["us"]]] - Require all user requests to use models in this specific region.
* - default_model: Optional[str] - If no equivalent model in the allowed region, default all requests to this model.
* - models: Optional[list[str]] - Restrict this customer's access to the listed models.
* - metadata: Optional[dict] = Metadata for customer, store information for customer. Example metadata = {"data_training_opt_out": True}
* - budget_duration: Optional[str] - Budget is reset at the end of specified duration. If not set, budget is never reset. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d").
* - tpm_limit: Optional[int] - [Not Implemented Yet] Specify tpm limit for a given customer (Tokens per minute)
@ -4332,6 +4333,7 @@ export interface paths {
* "user_id" : "ishaan-jaff-3",
* "allowed_region": "eu",
* "budget_id": "free_tier",
* "models": ["gpt-4o-mini"],
* "default_model": "azure/gpt-3.5-turbo-eu"
* }'
*
@ -4411,6 +4413,7 @@ export interface paths {
* - default_model: Optional[str] = (
* None # if no equivalent model in allowed region - default all requests to this model
* )
* - models: Optional[list[str]] = None # omitted or null leaves the allowlist unchanged; an empty list clears it
* - object_permission: Optional[LiteLLM_ObjectPermissionBase] - Customer-specific object permissions to control access to resources.
* Supported fields:
* * mcp_servers: List[str] - List of allowed MCP server IDs
@ -4426,7 +4429,8 @@ export interface paths {
* ```
* curl --location 'http://0.0.0.0:4000/customer/update' --header 'Authorization: Bearer sk-1234' --header 'Content-Type: application/json' --data '{
* "user_id": "test-litellm-user-4",
* "budget_id": "paid_tier"
* "budget_id": "paid_tier",
* "models": ["gpt-4o-mini"]
* }'
*
* # Updating object permissions
@ -4963,6 +4967,7 @@ export interface paths {
* - budget_id: Optional[str] - The identifier for an existing budget allocated to the user. Either 'max_budget' or 'budget_id' should be provided, not both.
* - allowed_model_region: Optional[Union[Literal["eu"], Literal["us"]]] - Require all user requests to use models in this specific region.
* - default_model: Optional[str] - If no equivalent model in the allowed region, default all requests to this model.
* - models: Optional[list[str]] - Restrict this customer's access to the listed models.
* - metadata: Optional[dict] = Metadata for customer, store information for customer. Example metadata = {"data_training_opt_out": True}
* - budget_duration: Optional[str] - Budget is reset at the end of specified duration. If not set, budget is never reset. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d").
* - tpm_limit: Optional[int] - [Not Implemented Yet] Specify tpm limit for a given customer (Tokens per minute)
@ -4994,6 +4999,7 @@ export interface paths {
* "user_id" : "ishaan-jaff-3",
* "allowed_region": "eu",
* "budget_id": "free_tier",
* "models": ["gpt-4o-mini"],
* "default_model": "azure/gpt-3.5-turbo-eu"
* }'
*
@ -5073,6 +5079,7 @@ export interface paths {
* - default_model: Optional[str] = (
* None # if no equivalent model in allowed region - default all requests to this model
* )
* - models: Optional[list[str]] = None # omitted or null leaves the allowlist unchanged; an empty list clears it
* - object_permission: Optional[LiteLLM_ObjectPermissionBase] - Customer-specific object permissions to control access to resources.
* Supported fields:
* * mcp_servers: List[str] - List of allowed MCP server IDs
@ -5088,7 +5095,8 @@ export interface paths {
* ```
* curl --location 'http://0.0.0.0:4000/customer/update' --header 'Authorization: Bearer sk-1234' --header 'Content-Type: application/json' --data '{
* "user_id": "test-litellm-user-4",
* "budget_id": "paid_tier"
* "budget_id": "paid_tier",
* "models": ["gpt-4o-mini"]
* }'
*
* # Updating object permissions
@ -30906,6 +30914,8 @@ export interface components {
/** Default Model */
default_model?: string | null;
litellm_budget_table?: components["schemas"]["LiteLLM_BudgetTableFull"] | null;
/** Models */
models?: string[];
object_permission?: components["schemas"]["LiteLLM_ObjectPermissionTable"] | null;
/** Object Permission Id */
object_permission_id?: string | null;
@ -34364,6 +34374,8 @@ export interface components {
/** Default Model */
default_model?: string | null;
litellm_budget_table?: components["schemas"]["LiteLLM_BudgetTable"] | null;
/** Models */
models?: string[];
object_permission?: components["schemas"]["LiteLLM_ObjectPermissionTable"] | null;
/** Object Permission Id */
object_permission_id?: string | null;
@ -38378,6 +38390,8 @@ export interface components {
model_max_budget?: {
[key: string]: components["schemas"]["BudgetConfig"];
} | null;
/** Models */
models?: string[] | null;
object_permission?: components["schemas"]["LiteLLM_ObjectPermissionBase"] | null;
/**
* Rpm Limit
@ -47367,6 +47381,8 @@ export interface components {
default_model?: string | null;
/** Max Budget */
max_budget?: number | null;
/** Models */
models?: string[] | null;
object_permission?: components["schemas"]["LiteLLM_ObjectPermissionBase"] | null;
/** User Id */
user_id: string;