mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
a3a52466a5
commit
a3e15774ad
17 changed files with 729 additions and 39 deletions
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_EndUserTable" ADD COLUMN IF NOT EXISTS "models" TEXT[] DEFAULT ARRAY[]::TEXT[];
|
||||
|
|
@ -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])
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
|
|
|
|||
293
tests/integration/authorization/test_customer_model_allowlist.py
Normal file
293
tests/integration/authorization/test_customer_model_allowlist.py
Normal 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
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
20
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
20
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue