diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260930000000_add_end_user_models/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260930000000_add_end_user_models/migration.sql new file mode 100644 index 00000000000..1e1da38e562 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260930000000_add_end_user_models/migration.sql @@ -0,0 +1,2 @@ +-- AlterTable +ALTER TABLE "LiteLLM_EndUserTable" ADD COLUMN IF NOT EXISTS "models" TEXT[] DEFAULT ARRAY[]::TEXT[]; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index cf76b764350..888b704bc05 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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]) diff --git a/litellm/litellm_core_utils/error_normalization.py b/litellm/litellm_core_utils/error_normalization.py index be0098ec34b..726fdcd1540 100644 --- a/litellm/litellm_core_utils/error_normalization.py +++ b/litellm/litellm_core_utils/error_normalization.py @@ -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, diff --git a/litellm/models/end_user.py b/litellm/models/end_user.py index 8dccf1eb5e7..a5964d0d60a 100644 --- a/litellm/models/end_user.py +++ b/litellm/models/end_user.py @@ -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 diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index b993a5d8d6d..15b3c925a1d 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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": diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index dfdf7dc4e66..6922b5f6cbd 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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, diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index c3c1032a3f0..31e7483b891 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -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): diff --git a/litellm/proxy/auth/model_access_denied.py b/litellm/proxy/auth/model_access_denied.py index ffb73b343cd..b4f0250a42d 100644 --- a/litellm/proxy/auth/model_access_denied.py +++ b/litellm/proxy/auth/model_access_denied.py @@ -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) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index cb39801a8e0..65aa337c057 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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, diff --git a/litellm/proxy/management_endpoints/customer_endpoints.py b/litellm/proxy/management_endpoints/customer_endpoints.py index 7236fd12e9d..2e629185028 100644 --- a/litellm/proxy/management_endpoints/customer_endpoints.py +++ b/litellm/proxy/management_endpoints/customer_endpoints.py @@ -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( diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index cf76b764350..888b704bc05 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -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]) diff --git a/schema.prisma b/schema.prisma index cf76b764350..888b704bc05 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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]) diff --git a/tests/integration/authorization/test_customer_model_allowlist.py b/tests/integration/authorization/test_customer_model_allowlist.py new file mode 100644 index 00000000000..2d88e7b7039 --- /dev/null +++ b/tests/integration/authorization/test_customer_model_allowlist.py @@ -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 diff --git a/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py b/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py index 059cff0c385..6f62e6ed092 100644 --- a/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py +++ b/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py @@ -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) diff --git a/tests/unit/proxy/auth/test_router_override_fallback_auth.py b/tests/unit/proxy/auth/test_router_override_fallback_auth.py index eb1135a240a..c34614a93ed 100644 --- a/tests/unit/proxy/auth/test_router_override_fallback_auth.py +++ b/tests/unit/proxy/auth/test_router_override_fallback_auth.py @@ -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 ] diff --git a/tests/unit/proxy/management_endpoints/test_customer_endpoints.py b/tests/unit/proxy/management_endpoints/test_customer_endpoints.py index 77e52f30bb7..c99db2d8054 100644 --- a/tests/unit/proxy/management_endpoints/test_customer_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_customer_endpoints.py @@ -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): diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 401439f6e63..bb6f75755e2 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -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;