mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(proxy): port the member auto-router write path and grant plumbing to stable/1.100.x
Bugbot on #42668 flagged that the backport's picks left the member auto-router management path unwired on this line. This ports the pieces that make it work, mirroring main: the member write slot in model_management_endpoints (FOR UPDATE lock, team reload with the model table include, identity and name-collision checks, post-commit config publish), StoredAutoRouterIdentity wiring, the license feature helpers, team tpd_limit, Router.config_deployments, the member_auto_router ModelInfo flag, the _TEAM_GRANT_RELATIONS include on team lookups, and the UserAPIKeyAuth fields the team_grants unpack needs. Test files were rebuilt as line content plus the picks' own additions, and ui_sso/test_team_grants carry the pick's grant assertions. Gate-clearing edits stay local to what the picks added: prisma TypedDict arguments replace mutable dict literals, remaining dict/mapping arguments carry reasoned mutable-ok comments, test-quality-ok comments mark the picks' internal-seam patches, and the regenerated dashboard api types are staged. The only remaining make check failure is pre-existing staging drift in untouched tests/test_litellm/test_router.py:2969.
This commit is contained in:
parent
a392f7bc66
commit
54b42e3f05
20 changed files with 1701 additions and 2056 deletions
|
|
@ -71,6 +71,7 @@ class TeamBase(LiteLLMPydanticObjectBase):
|
||||||
metadata: dict | None = None
|
metadata: dict | None = None
|
||||||
tpm_limit: int | None = None
|
tpm_limit: int | None = None
|
||||||
rpm_limit: int | None = None
|
rpm_limit: int | None = None
|
||||||
|
tpd_limit: int | None = None
|
||||||
max_budget: float | None = None
|
max_budget: float | None = None
|
||||||
soft_budget: float | None = None
|
soft_budget: float | None = None
|
||||||
budget_duration: str | None = None
|
budget_duration: str | None = None
|
||||||
|
|
|
||||||
|
|
@ -2779,8 +2779,10 @@ class LiteLLM_VerificationTokenView(LiteLLM_VerificationToken):
|
||||||
team_alias: str | None = None
|
team_alias: str | None = None
|
||||||
team_tpm_limit: int | None = None
|
team_tpm_limit: int | None = None
|
||||||
team_rpm_limit: int | None = None
|
team_rpm_limit: int | None = None
|
||||||
|
team_tpd_limit: int | None = None
|
||||||
team_max_budget: float | None = None
|
team_max_budget: float | None = None
|
||||||
team_soft_budget: float | None = None
|
team_soft_budget: float | None = None
|
||||||
|
team_model_max_budget: dict[str, object] | None = None
|
||||||
team_models: list = []
|
team_models: list = []
|
||||||
team_blocked: bool = False
|
team_blocked: bool = False
|
||||||
soft_budget: float | None = None
|
soft_budget: float | None = None
|
||||||
|
|
|
||||||
|
|
@ -473,6 +473,7 @@ def _is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None
|
||||||
|
|
||||||
|
|
||||||
_NO_MODEL_INFO: Final[Mapping[str, object]] = MappingProxyType({})
|
_NO_MODEL_INFO: Final[Mapping[str, object]] = MappingProxyType({})
|
||||||
|
_TEAM_GRANT_RELATIONS: Final[Mapping[str, object]] = MappingProxyType({"litellm_model_table": True})
|
||||||
|
|
||||||
|
|
||||||
def _has_ptu_flat_cost(model: str, llm_router: "Router") -> bool:
|
def _has_ptu_flat_cost(model: str, llm_router: "Router") -> bool:
|
||||||
|
|
@ -549,11 +550,12 @@ def _model_group_has_pricing(model: str, llm_router: "Router") -> bool:
|
||||||
return True
|
return True
|
||||||
|
|
||||||
model_id = (deployment.get("model_info") or _EMPTY_COST_ENTRY).get("id")
|
model_id = (deployment.get("model_info") or _EMPTY_COST_ENTRY).get("id")
|
||||||
if model_id is None:
|
if not isinstance(model_id, str):
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
model_name = litellm_params.get("model")
|
||||||
model_info = llm_router.get_deployment_model_info(
|
model_info = llm_router.get_deployment_model_info(
|
||||||
model_id=model_id, model_name=litellm_params.get("model") or ""
|
model_id=model_id, model_name=model_name if isinstance(model_name, str) else ""
|
||||||
)
|
)
|
||||||
if model_info is not None and _entry_has_priced_metric(model_info):
|
if model_info is not None and _entry_has_priced_metric(model_info):
|
||||||
return True
|
return True
|
||||||
|
|
@ -2847,7 +2849,10 @@ class TeamNotFoundError(HTTPException):
|
||||||
async def _get_team_db_check(
|
async def _get_team_db_check(
|
||||||
team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None
|
team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None
|
||||||
) -> "_PrismaTeamRow | None":
|
) -> "_PrismaTeamRow | None":
|
||||||
response = await _team_table(TeamRepository(prisma_client)).find_unique(where={"team_id": team_id})
|
response = await _team_table(TeamRepository(prisma_client)).find_unique(
|
||||||
|
where={"team_id": team_id}, # mutable-ok: prisma where clause
|
||||||
|
include=_TEAM_GRANT_RELATIONS,
|
||||||
|
)
|
||||||
|
|
||||||
if response is None and team_id_upsert:
|
if response is None and team_id_upsert:
|
||||||
from litellm.proxy.management_endpoints.team_endpoints import new_team
|
from litellm.proxy.management_endpoints.team_endpoints import new_team
|
||||||
|
|
@ -3147,7 +3152,10 @@ async def get_team_object_by_alias(
|
||||||
|
|
||||||
# Query database by team_alias
|
# Query database by team_alias
|
||||||
try:
|
try:
|
||||||
teams: Final = await _team_table(TeamRepository(prisma_client)).find_many(where={"team_alias": team_alias})
|
teams: Final = await _team_table(TeamRepository(prisma_client)).find_many(
|
||||||
|
where={"team_alias": team_alias}, # mutable-ok: prisma where clause
|
||||||
|
include=_TEAM_GRANT_RELATIONS,
|
||||||
|
)
|
||||||
|
|
||||||
if not teams:
|
if not teams:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
|
|
|
||||||
|
|
@ -52,6 +52,7 @@ from litellm.proxy._types import (
|
||||||
)
|
)
|
||||||
from litellm.proxy.auth.auth_checks import can_team_access_model
|
from litellm.proxy.auth.auth_checks import can_team_access_model
|
||||||
from litellm.proxy.auth.route_checks import RouteChecks
|
from litellm.proxy.auth.route_checks import RouteChecks
|
||||||
|
from litellm.proxy.auth.team_grants import team_model_aliases
|
||||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||||
UserApiKeyCache,
|
UserApiKeyCache,
|
||||||
get_management_object_ttl,
|
get_management_object_ttl,
|
||||||
|
|
@ -1553,7 +1554,9 @@ class JWTAuthManager:
|
||||||
model=requested_model,
|
model=requested_model,
|
||||||
team_object=team_object,
|
team_object=team_object,
|
||||||
llm_router=llm_router,
|
llm_router=llm_router,
|
||||||
team_model_aliases=None,
|
team_model_aliases=dict(aliases)
|
||||||
|
if (aliases := team_model_aliases(team_object)) is not None
|
||||||
|
else None,
|
||||||
)
|
)
|
||||||
):
|
):
|
||||||
is_allowed = allowed_routes_check(
|
is_allowed = allowed_routes_check(
|
||||||
|
|
@ -2090,7 +2093,9 @@ class JWTAuthManager:
|
||||||
model=requested_model,
|
model=requested_model,
|
||||||
team_object=team_object,
|
team_object=team_object,
|
||||||
llm_router=llm_router,
|
llm_router=llm_router,
|
||||||
team_model_aliases=None,
|
team_model_aliases=dict(aliases)
|
||||||
|
if (aliases := team_model_aliases(team_object)) is not None
|
||||||
|
else None,
|
||||||
)
|
)
|
||||||
except ProxyException:
|
except ProxyException:
|
||||||
continue
|
continue
|
||||||
|
|
|
||||||
|
|
@ -15,6 +15,10 @@ from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from litellm.proxy._types import EnterpriseLicenseData
|
from litellm.proxy._types import EnterpriseLicenseData
|
||||||
|
|
||||||
|
AUTO_ROUTER_LICENSE_FEATURE: Final = "auto_router"
|
||||||
|
LICENSE_ALL_FEATURES: Final = "*"
|
||||||
|
AUTO_ROUTER_LICENSE_REMEDY: Final = "A LiteLLM license with the 'auto_router' feature lifts the limit."
|
||||||
|
|
||||||
|
|
||||||
class LicenseCheck:
|
class LicenseCheck:
|
||||||
"""
|
"""
|
||||||
|
|
@ -149,6 +153,24 @@ class LicenseCheck:
|
||||||
return False
|
return False
|
||||||
return team_count > _max_teams_in_license
|
return team_count > _max_teams_in_license
|
||||||
|
|
||||||
|
def grants_feature(self, feature: str) -> bool:
|
||||||
|
if self.airgapped_license_data is None:
|
||||||
|
return False
|
||||||
|
allowed_features: Final = self.airgapped_license_data.get("allowed_features")
|
||||||
|
granted: Final = allowed_features if isinstance(allowed_features, list) else (allowed_features,)
|
||||||
|
return feature in granted or LICENSE_ALL_FEATURES in granted
|
||||||
|
|
||||||
|
def auto_router_capability_limit(self) -> int | None:
|
||||||
|
"""
|
||||||
|
How many auto-routers may claim each gated classifier or customization capability:
|
||||||
|
unlimited (None) only when the signed license lists the auto_router feature or the
|
||||||
|
"*" wildcard that grants every feature, otherwise one per capability. A license verified
|
||||||
|
through the API carries no feature list, so it does not lift the limit either.
|
||||||
|
"""
|
||||||
|
if self.grants_feature(AUTO_ROUTER_LICENSE_FEATURE):
|
||||||
|
return None
|
||||||
|
return 1
|
||||||
|
|
||||||
def verify_license_without_api_request(self, public_key, license_key):
|
def verify_license_without_api_request(self, public_key, license_key):
|
||||||
try:
|
try:
|
||||||
from cryptography.hazmat.primitives import hashes
|
from cryptography.hazmat.primitives import hashes
|
||||||
|
|
|
||||||
|
|
@ -6,7 +6,7 @@ callers kept losing grants (aliases, permissions, limits) one field at a time. B
|
||||||
``team_grants`` and the two paths cannot drift.
|
``team_grants`` and the two paths cannot drift.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from collections.abc import Mapping, Sequence
|
from collections.abc import Mapping
|
||||||
from types import MappingProxyType
|
from types import MappingProxyType
|
||||||
from typing import Annotated, Final
|
from typing import Annotated, Final
|
||||||
|
|
||||||
|
|
@ -61,10 +61,10 @@ class TeamGrants(TypedDict, total=False):
|
||||||
team_soft_budget: ReadOnly[float | None]
|
team_soft_budget: ReadOnly[float | None]
|
||||||
team_model_max_budget: ReadOnly[dict[str, object] | None] # mutable-ok: prisma table field typed loosely
|
team_model_max_budget: ReadOnly[dict[str, object] | None] # mutable-ok: prisma table field typed loosely
|
||||||
team_spend: ReadOnly[float | None]
|
team_spend: ReadOnly[float | None]
|
||||||
team_models: ReadOnly[Sequence[str]]
|
team_models: ReadOnly[list[str]] # mutable-ok: UserAPIKeyAuth declares a list field
|
||||||
team_blocked: ReadOnly[bool]
|
team_blocked: ReadOnly[bool]
|
||||||
team_metadata: ReadOnly[Mapping[str, object] | None]
|
team_metadata: ReadOnly[dict[str, object] | None] # mutable-ok: UserAPIKeyAuth declares a dict field
|
||||||
team_model_aliases: ReadOnly[Mapping[str, str] | None]
|
team_model_aliases: ReadOnly[dict[str, str] | None] # mutable-ok: UserAPIKeyAuth declares a dict field
|
||||||
team_object_permission_id: ReadOnly[str | None]
|
team_object_permission_id: ReadOnly[str | None]
|
||||||
team_object_permission: ReadOnly[LiteLLM_ObjectPermissionTable | None]
|
team_object_permission: ReadOnly[LiteLLM_ObjectPermissionTable | None]
|
||||||
team_member: ReadOnly[Member | None]
|
team_member: ReadOnly[Member | None]
|
||||||
|
|
@ -104,11 +104,14 @@ def team_grants(
|
||||||
team_soft_budget=team_object.soft_budget,
|
team_soft_budget=team_object.soft_budget,
|
||||||
team_model_max_budget=team_object.model_max_budget,
|
team_model_max_budget=team_object.model_max_budget,
|
||||||
team_spend=team_object.spend,
|
team_spend=team_object.spend,
|
||||||
team_models=tuple(team_object.models),
|
team_models=list(team_object.models),
|
||||||
team_blocked=team_object.blocked,
|
team_blocked=team_object.blocked,
|
||||||
team_metadata=json_columns.metadata,
|
team_metadata=dict(json_columns.metadata) if json_columns.metadata is not None else None,
|
||||||
team_model_aliases=(
|
team_model_aliases=(
|
||||||
json_columns.litellm_model_table.model_aliases if json_columns.litellm_model_table is not None else None
|
dict(json_columns.litellm_model_table.model_aliases)
|
||||||
|
if json_columns.litellm_model_table is not None
|
||||||
|
and json_columns.litellm_model_table.model_aliases is not None
|
||||||
|
else None
|
||||||
),
|
),
|
||||||
team_object_permission_id=team_object.object_permission_id,
|
team_object_permission_id=team_object.object_permission_id,
|
||||||
team_object_permission=team_object.object_permission,
|
team_object_permission=team_object.object_permission,
|
||||||
|
|
|
||||||
|
|
@ -77,6 +77,7 @@ from litellm.proxy.auth.oauth2_proxy_hook import handle_oauth2_proxy_request
|
||||||
from litellm.proxy.auth.resolvers import CredentialRef, Principal
|
from litellm.proxy.auth.resolvers import CredentialRef, Principal
|
||||||
from litellm.proxy.auth.resolvers.store import IdentityStore
|
from litellm.proxy.auth.resolvers.store import IdentityStore
|
||||||
from litellm.proxy.auth.route_checks import RouteChecks
|
from litellm.proxy.auth.route_checks import RouteChecks
|
||||||
|
from litellm.proxy.auth.team_grants import team_grants
|
||||||
from litellm.proxy.auth.trusted_proxy_utils import get_trusted_proxy_cidrs
|
from litellm.proxy.auth.trusted_proxy_utils import get_trusted_proxy_cidrs
|
||||||
from litellm.proxy.common_utils.cache_coordinator import EventDrivenCacheCoordinator
|
from litellm.proxy.common_utils.cache_coordinator import EventDrivenCacheCoordinator
|
||||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||||
|
|
@ -1469,24 +1470,16 @@ async def _user_api_key_auth_builder(
|
||||||
user_id=user_id,
|
user_id=user_id,
|
||||||
user_email=user_email,
|
user_email=user_email,
|
||||||
team_id=team_id,
|
team_id=team_id,
|
||||||
team_alias=(team_object.team_alias if team_object is not None else None),
|
|
||||||
team_tpm_limit=(team_object.tpm_limit if team_object is not None else None),
|
|
||||||
team_rpm_limit=(team_object.rpm_limit if team_object is not None else None),
|
|
||||||
team_models=(team_object.models if team_object is not None else []),
|
|
||||||
team_metadata=(team_object.metadata if team_object is not None else None),
|
|
||||||
org_id=org_id,
|
org_id=org_id,
|
||||||
end_user_id=end_user_id,
|
end_user_id=end_user_id,
|
||||||
parent_otel_span=parent_otel_span,
|
parent_otel_span=parent_otel_span,
|
||||||
jwt_claims=jwt_claims,
|
jwt_claims=jwt_claims,
|
||||||
|
**team_grants(team_object=team_object, team_membership=team_membership, user_id=user_id),
|
||||||
)
|
)
|
||||||
|
|
||||||
valid_token = UserAPIKeyAuth(
|
valid_token = UserAPIKeyAuth(
|
||||||
api_key=None,
|
api_key=None,
|
||||||
team_id=team_id,
|
team_id=team_id,
|
||||||
team_alias=(team_object.team_alias if team_object is not None else None),
|
|
||||||
team_tpm_limit=(team_object.tpm_limit if team_object is not None else None),
|
|
||||||
team_rpm_limit=(team_object.rpm_limit if team_object is not None else None),
|
|
||||||
team_models=(team_object.models if team_object is not None else []),
|
|
||||||
user_role=(
|
user_role=(
|
||||||
LitellmUserRoles(user_object.user_role)
|
LitellmUserRoles(user_object.user_role)
|
||||||
if user_object is not None and user_object.user_role is not None
|
if user_object is not None and user_object.user_role is not None
|
||||||
|
|
@ -1500,17 +1493,8 @@ async def _user_api_key_auth_builder(
|
||||||
user_tpm_limit=(user_object.tpm_limit if user_object is not None else None),
|
user_tpm_limit=(user_object.tpm_limit if user_object is not None else None),
|
||||||
user_rpm_limit=(user_object.rpm_limit if user_object is not None else None),
|
user_rpm_limit=(user_object.rpm_limit if user_object is not None else None),
|
||||||
user_model_max_budget=(user_object.model_max_budget if user_object is not None else None),
|
user_model_max_budget=(user_object.model_max_budget if user_object is not None else None),
|
||||||
team_member_rpm_limit=(
|
|
||||||
team_membership.safe_get_team_member_rpm_limit() if team_membership is not None else None
|
|
||||||
),
|
|
||||||
team_member_tpm_limit=(
|
|
||||||
team_membership.safe_get_team_member_tpm_limit() if team_membership is not None else None
|
|
||||||
),
|
|
||||||
team_metadata=(team_object.metadata if team_object is not None else None),
|
|
||||||
jwt_claims=jwt_claims,
|
jwt_claims=jwt_claims,
|
||||||
)
|
**team_grants(team_object=team_object, team_membership=team_membership, user_id=user_id),
|
||||||
valid_token.team_object_permission = (
|
|
||||||
team_object.object_permission if team_object is not None else None
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# AUTO_REGISTER deferred from _resolve_jwt_to_virtual_key.
|
# AUTO_REGISTER deferred from _resolve_jwt_to_virtual_key.
|
||||||
|
|
|
||||||
|
|
@ -13,10 +13,13 @@ model/{model_id}/update - PATCH endpoint for model update.
|
||||||
import asyncio
|
import asyncio
|
||||||
import datetime
|
import datetime
|
||||||
import json
|
import json
|
||||||
from collections.abc import Awaitable, Mapping, Sequence
|
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence
|
||||||
|
from contextlib import AbstractAsyncContextManager, asynccontextmanager
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from fnmatch import fnmatchcase
|
||||||
from json import JSONDecodeError
|
from json import JSONDecodeError
|
||||||
from types import MappingProxyType
|
from types import MappingProxyType
|
||||||
from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, cast
|
from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, TypeVar, cast, runtime_checkable
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
|
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
|
||||||
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError, field_validator
|
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError, field_validator
|
||||||
|
|
@ -54,7 +57,7 @@ from litellm.proxy.common_utils.config_sync_pubsub import (
|
||||||
coordination_redis_cache,
|
coordination_redis_cache,
|
||||||
publish_config_change,
|
publish_config_change,
|
||||||
)
|
)
|
||||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper
|
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper
|
||||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||||
from litellm.proxy.management_endpoints.common_utils import _is_user_team_admin
|
from litellm.proxy.management_endpoints.common_utils import _is_user_team_admin
|
||||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||||
|
|
@ -66,6 +69,13 @@ from litellm.proxy.management_endpoints.team_endpoints import (
|
||||||
update_team as _legacy_update_team,
|
update_team as _legacy_update_team,
|
||||||
)
|
)
|
||||||
from litellm.proxy.management_helpers.audit_logs import create_object_audit_log
|
from litellm.proxy.management_helpers.audit_logs import create_object_audit_log
|
||||||
|
from litellm.proxy.management_helpers.auto_router_permissions import (
|
||||||
|
MemberAutoRouterWrite,
|
||||||
|
StoredAutoRouterIdentity,
|
||||||
|
authorize_member_auto_router_dependencies,
|
||||||
|
authorize_member_auto_router_team,
|
||||||
|
authorize_member_auto_router_write,
|
||||||
|
)
|
||||||
from litellm.proxy.spend_tracking.ptu_feature_flag import (
|
from litellm.proxy.spend_tracking.ptu_feature_flag import (
|
||||||
PTU_COST_ATTRIBUTION_ENV_VAR,
|
PTU_COST_ATTRIBUTION_ENV_VAR,
|
||||||
is_ptu_cost_attribution_enabled,
|
is_ptu_cost_attribution_enabled,
|
||||||
|
|
@ -103,11 +113,13 @@ from litellm.types.router import (
|
||||||
GenericLiteLLMParams,
|
GenericLiteLLMParams,
|
||||||
ModelInfo,
|
ModelInfo,
|
||||||
updateDeployment,
|
updateDeployment,
|
||||||
|
updateLiteLLMParams,
|
||||||
)
|
)
|
||||||
from litellm.utils import get_utc_datetime
|
from litellm.utils import get_utc_datetime
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from prisma import models as prisma_models
|
from prisma import models as prisma_models
|
||||||
|
from prisma import types as prisma_types
|
||||||
|
|
||||||
router: Final = APIRouter()
|
router: Final = APIRouter()
|
||||||
|
|
||||||
|
|
@ -155,8 +167,39 @@ class _ProxyModelTable(Protocol):
|
||||||
def delete_many(self, *, where: Mapping[str, object]) -> Awaitable[int]: ...
|
def delete_many(self, *, where: Mapping[str, object]) -> Awaitable[int]: ...
|
||||||
|
|
||||||
|
|
||||||
|
class _TxTable(Protocol):
|
||||||
|
def find_unique(
|
||||||
|
self, *, where: Mapping[str, object], include: Mapping[str, bool] | None = None
|
||||||
|
) -> Awaitable[BaseModel | None]: ...
|
||||||
|
|
||||||
|
|
||||||
class _TxModelTables(Protocol):
|
class _TxModelTables(Protocol):
|
||||||
litellm_proxymodeltable: _ProxyModelTable
|
litellm_proxymodeltable: _ProxyModelTable
|
||||||
|
litellm_teamtable: _TxTable
|
||||||
|
litellm_teammembership: _TxTable
|
||||||
|
litellm_organizationtable: _TxTable
|
||||||
|
litellm_projecttable: _TxTable
|
||||||
|
|
||||||
|
async def query_raw(self, query: str, *args: object) -> Sequence[Mapping[str, object]]: ...
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
|
class _TransactionFactory(Protocol):
|
||||||
|
def __call__(self, *, timeout: datetime.timedelta = ...) -> AbstractAsyncContextManager[_TxModelTables]: ...
|
||||||
|
|
||||||
|
|
||||||
|
class _ModelTransactionClient(BaseModel):
|
||||||
|
model_config = ConfigDict(arbitrary_types_allowed=True, from_attributes=True)
|
||||||
|
|
||||||
|
tx: _TransactionFactory
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class _TransactionClient:
|
||||||
|
db: _TxModelTables
|
||||||
|
|
||||||
|
|
||||||
|
_RowT = TypeVar("_RowT")
|
||||||
|
|
||||||
|
|
||||||
class _ExistingModelRow(Protocol):
|
class _ExistingModelRow(Protocol):
|
||||||
|
|
@ -241,6 +284,142 @@ def _effective_complexity_router_config(
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _effective_model(
|
||||||
|
incoming_params: GenericLiteLLMParams | None, existing_params: GenericLiteLLMParams | None
|
||||||
|
) -> str | None:
|
||||||
|
"""The model a write leaves on the row, decrypting an existing value only when the patch omits it."""
|
||||||
|
incoming: Final = None if incoming_params is None else incoming_params.model
|
||||||
|
if incoming is not None:
|
||||||
|
return incoming
|
||||||
|
existing: Final = None if existing_params is None else existing_params.model
|
||||||
|
if existing is None:
|
||||||
|
return None
|
||||||
|
decrypted: Final = decrypt_value_helper(
|
||||||
|
value=existing,
|
||||||
|
key="model",
|
||||||
|
exception_type="debug",
|
||||||
|
return_original_value=True,
|
||||||
|
)
|
||||||
|
return decrypted if isinstance(decrypted, str) else None
|
||||||
|
|
||||||
|
|
||||||
|
def _member_auto_router_marker_for_update(
|
||||||
|
*,
|
||||||
|
incoming_params: updateLiteLLMParams | None,
|
||||||
|
existing: Deployment,
|
||||||
|
member_write: MemberAutoRouterWrite | None,
|
||||||
|
) -> bool | None:
|
||||||
|
if member_write is not None:
|
||||||
|
return True
|
||||||
|
if not existing.model_info.member_auto_router:
|
||||||
|
return None
|
||||||
|
if incoming_params is None:
|
||||||
|
return True
|
||||||
|
if any(getattr(incoming_params, field, None) is not None for field in STRATEGY_ROUTER_PARAM_FIELDS):
|
||||||
|
return False
|
||||||
|
return incoming_params.model is None or incoming_params.model == _effective_model(None, existing.litellm_params)
|
||||||
|
|
||||||
|
|
||||||
|
@asynccontextmanager
|
||||||
|
async def _member_auto_router_write_slot(
|
||||||
|
prisma_client: PrismaClient,
|
||||||
|
*,
|
||||||
|
member_write: MemberAutoRouterWrite | None,
|
||||||
|
) -> AsyncGenerator[_ProxyModelTable, None]:
|
||||||
|
"""Hand out the model table a member write goes through.
|
||||||
|
|
||||||
|
Member writes to a team auto router recheck their authorization inside one
|
||||||
|
transaction that locks the row first, so two concurrent member writes
|
||||||
|
cannot both pass the ownership and name checks. Non-member writes keep the
|
||||||
|
direct table. The transaction write bypasses the repository's
|
||||||
|
publish-on-write, so the config change is published once after commit.
|
||||||
|
"""
|
||||||
|
if member_write is None:
|
||||||
|
yield _proxy_model_table(prisma_client)
|
||||||
|
return
|
||||||
|
import litellm
|
||||||
|
from litellm.proxy.auth.team_grants import team_model_aliases
|
||||||
|
from litellm.proxy.proxy_server import llm_router, premium_user
|
||||||
|
|
||||||
|
transaction_client: Final = _ModelTransactionClient.model_validate(prisma_client.db)
|
||||||
|
async with transaction_client.tx(timeout=datetime.timedelta(seconds=30)) as tx_ctx:
|
||||||
|
tables: Final[_TxModelTables] = tx_ctx
|
||||||
|
config_rows: Final = () if llm_router is None else tuple(llm_router.config_deployments())
|
||||||
|
if member_write.model_id is not None:
|
||||||
|
await tx_ctx.query_raw(
|
||||||
|
'SELECT model_id FROM "LiteLLM_ProxyModelTable" WHERE model_id = $1 FOR UPDATE',
|
||||||
|
member_write.model_id,
|
||||||
|
)
|
||||||
|
pinned_client: Final = _TransactionClient(tx_ctx)
|
||||||
|
team_where: Final[prisma_types.LiteLLM_TeamTableWhereUniqueInput] = {"team_id": member_write.team_id}
|
||||||
|
team_row: Final = await tx_ctx.litellm_teamtable.find_unique(
|
||||||
|
where=team_where,
|
||||||
|
include={"litellm_model_table": True}, # mutable-ok: prisma include clause
|
||||||
|
)
|
||||||
|
if team_row is None or llm_router is None:
|
||||||
|
raise HTTPException(status_code=403, detail="The auto router's team or model catalog is unavailable.")
|
||||||
|
team: Final = LiteLLM_TeamTable.model_validate(team_row.model_dump())
|
||||||
|
authorize_member_auto_router_team(user_api_key_dict=member_write.actor, team=team, premium_user=premium_user)
|
||||||
|
if member_write.model_id is not None:
|
||||||
|
model_where: Final[prisma_types.LiteLLM_ProxyModelTableWhereInput] = {"model_id": member_write.model_id}
|
||||||
|
current_row: Final = await tables.litellm_proxymodeltable.find_unique(where=model_where)
|
||||||
|
current_identity: Final = (
|
||||||
|
StoredAutoRouterIdentity.model_validate(current_row.model_dump()) if current_row is not None else None
|
||||||
|
)
|
||||||
|
current_model: Final = (
|
||||||
|
Deployment.model_validate(current_row.model_dump()) if current_row is not None else None
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
current_identity is None
|
||||||
|
or current_identity.created_by != member_write.actor.user_id
|
||||||
|
or current_model is None
|
||||||
|
or current_model.model_info.team_id != member_write.team_id
|
||||||
|
):
|
||||||
|
raise HTTPException(status_code=403, detail="Team members can update only their own auto routers.")
|
||||||
|
if current_identity.updated_at != member_write.updated_at:
|
||||||
|
raise HTTPException(status_code=409, detail="This auto router changed. Reload it before updating.")
|
||||||
|
else:
|
||||||
|
all_models: Final[prisma_types.LiteLLM_ProxyModelTableWhereInput] = {}
|
||||||
|
rows_for_names: Final = await tables.litellm_proxymodeltable.find_many(where=all_models)
|
||||||
|
stored_names: Final = tuple(
|
||||||
|
(
|
||||||
|
row.model_name,
|
||||||
|
model_info_as_mapping(row.model_info),
|
||||||
|
)
|
||||||
|
for row in rows_for_names
|
||||||
|
)
|
||||||
|
config_names: Final = tuple(
|
||||||
|
(str(row.get("model_name", "")), model_info_as_mapping(row.get("model_info"))) for row in config_rows
|
||||||
|
)
|
||||||
|
team_aliases: Final = team_model_aliases(team)
|
||||||
|
aliases: Final = (
|
||||||
|
*(llm_router.model_group_alias or ()),
|
||||||
|
*(litellm.model_alias_map or ()),
|
||||||
|
*(team_aliases or ()),
|
||||||
|
)
|
||||||
|
if member_write.public_name in aliases or any(
|
||||||
|
fnmatchcase(
|
||||||
|
member_write.public_name,
|
||||||
|
str(info.get("team_public_model_name") or name)
|
||||||
|
if info is not None and info.get("team_id") == member_write.team_id
|
||||||
|
else name,
|
||||||
|
)
|
||||||
|
for name, info in (*stored_names, *config_names)
|
||||||
|
if info is None or info.get("team_id") in (None, member_write.team_id)
|
||||||
|
):
|
||||||
|
raise HTTPException(status_code=409, detail="This auto-router name is already used by a model.")
|
||||||
|
await authorize_member_auto_router_dependencies(
|
||||||
|
config=member_write.config,
|
||||||
|
default_model=member_write.default_model,
|
||||||
|
user_api_key_dict=member_write.actor,
|
||||||
|
team=team,
|
||||||
|
prisma_client=pinned_client,
|
||||||
|
llm_router=llm_router,
|
||||||
|
)
|
||||||
|
yield tables.litellm_proxymodeltable
|
||||||
|
await publish_config_change(redis_cache=coordination_redis_cache(), object_type="litellm_proxymodeltable")
|
||||||
|
|
||||||
|
|
||||||
def _strategy_router_write_violation(
|
def _strategy_router_write_violation(
|
||||||
incoming_params: GenericLiteLLMParams | None,
|
incoming_params: GenericLiteLLMParams | None,
|
||||||
existing_params: GenericLiteLLMParams | None,
|
existing_params: GenericLiteLLMParams | None,
|
||||||
|
|
@ -724,11 +903,39 @@ async def patch_model(
|
||||||
param=None,
|
param=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
await ModelManagementAuthChecks.can_user_make_model_call(
|
write_authorization: Final = await ModelManagementAuthChecks.can_user_make_model_call(
|
||||||
model_params=db_model,
|
model_params=db_model,
|
||||||
user_api_key_dict=user_api_key_dict,
|
user_api_key_dict=user_api_key_dict,
|
||||||
prisma_client=prisma_client,
|
prisma_client=prisma_client,
|
||||||
premium_user=premium_user,
|
premium_user=premium_user,
|
||||||
|
member_operation="update",
|
||||||
|
incoming_model_params=patch_data,
|
||||||
|
)
|
||||||
|
member_write: Final = write_authorization if isinstance(write_authorization, MemberAutoRouterWrite) else None
|
||||||
|
member_marker: Final = _member_auto_router_marker_for_update(
|
||||||
|
incoming_params=patch_data.litellm_params, existing=db_model, member_write=member_write
|
||||||
|
)
|
||||||
|
marker_info: Final = (
|
||||||
|
ModelInfo(id=db_model.model_info.id)
|
||||||
|
if member_write is not None
|
||||||
|
else patch_data.model_info or ModelInfo(id=db_model.model_info.id)
|
||||||
|
)
|
||||||
|
effective_info: Final = (
|
||||||
|
marker_info.model_copy(update=MappingProxyType({"member_auto_router": member_marker}))
|
||||||
|
if member_marker is not None
|
||||||
|
else patch_data.model_info
|
||||||
|
)
|
||||||
|
effective_patch: Final = (
|
||||||
|
patch_data.model_copy(
|
||||||
|
update=MappingProxyType(
|
||||||
|
{
|
||||||
|
"model_name": None if member_write is not None else patch_data.model_name,
|
||||||
|
"model_info": effective_info,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if member_marker is not None
|
||||||
|
else patch_data
|
||||||
)
|
)
|
||||||
|
|
||||||
# Pause/resume (`blocked`) is a proxy-admin-only privilege. Team admins
|
# Pause/resume (`blocked`) is a proxy-admin-only privilege. Team admins
|
||||||
|
|
@ -747,22 +954,23 @@ async def patch_model(
|
||||||
existing_params=db_model.litellm_params,
|
existing_params=db_model.litellm_params,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
async def write_row(update_data: PrismaCompatibleUpdateDBModel) -> _ProxyModelRow | None:
|
||||||
|
update_data["updated_by"] = (
|
||||||
|
user_api_key_dict.user_id or litellm_proxy_admin_name
|
||||||
|
) # mutable-ok: prisma update payload is dict-shaped
|
||||||
|
update_data["updated_at"] = cast(
|
||||||
|
str, get_utc_datetime()
|
||||||
|
) # mutable-ok: prisma update payload is dict-shaped
|
||||||
|
async with _member_auto_router_write_slot(prisma_client, member_write=member_write) as table:
|
||||||
|
return await table.update(where={"model_id": model_id}, data=update_data)
|
||||||
|
|
||||||
# Handle team model updates with proper alias management
|
# Handle team model updates with proper alias management
|
||||||
update_data: Final = await _update_team_model_in_db(
|
updated_model: Final = await _update_team_model_in_db(
|
||||||
db_model=db_model,
|
db_model=db_model,
|
||||||
patch_data=patch_data,
|
patch_data=effective_patch,
|
||||||
user_api_key_dict=user_api_key_dict,
|
user_api_key_dict=user_api_key_dict,
|
||||||
prisma_client=prisma_client,
|
prisma_client=prisma_client,
|
||||||
)
|
write_row=write_row,
|
||||||
|
|
||||||
# Add metadata about update
|
|
||||||
update_data["updated_by"] = user_api_key_dict.user_id or litellm_proxy_admin_name
|
|
||||||
update_data["updated_at"] = cast(str, get_utc_datetime())
|
|
||||||
|
|
||||||
# Perform partial update
|
|
||||||
updated_model: Final = await _proxy_model_table(prisma_client).update(
|
|
||||||
where={"model_id": model_id},
|
|
||||||
data=update_data,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if updated_model is None:
|
if updated_model is None:
|
||||||
|
|
@ -993,6 +1201,7 @@ async def _add_model_to_db(
|
||||||
prisma_client: PrismaClient,
|
prisma_client: PrismaClient,
|
||||||
new_encryption_key: str | None = None,
|
new_encryption_key: str | None = None,
|
||||||
should_create_model_in_db: bool = True,
|
should_create_model_in_db: bool = True,
|
||||||
|
slot: AbstractAsyncContextManager[_ProxyModelTable] | None = None,
|
||||||
) -> "prisma_models.LiteLLM_ProxyModelTable | LiteLLM_ProxyModelTable | None":
|
) -> "prisma_models.LiteLLM_ProxyModelTable | LiteLLM_ProxyModelTable | None":
|
||||||
# encrypt litellm params #
|
# encrypt litellm params #
|
||||||
_litellm_params_dict: Final = model_params.litellm_params.dict(exclude_none=True)
|
_litellm_params_dict: Final = model_params.litellm_params.dict(exclude_none=True)
|
||||||
|
|
@ -1011,17 +1220,19 @@ async def _add_model_to_db(
|
||||||
if model_params.model_info.id is not None:
|
if model_params.model_info.id is not None:
|
||||||
_data["model_id"] = model_params.model_info.id
|
_data["model_id"] = model_params.model_info.id
|
||||||
_create_data: Final = cast("Mapping[str, object]", _data) # cast-ok: str-keyed json payload built just above
|
_create_data: Final = cast("Mapping[str, object]", _data) # cast-ok: str-keyed json payload built just above
|
||||||
if should_create_model_in_db:
|
if not should_create_model_in_db:
|
||||||
model_response = await ModelRepository(prisma_client).table.create(data=_create_data)
|
return LiteLLM_ProxyModelTable(**_data)
|
||||||
else:
|
if slot is None:
|
||||||
model_response = LiteLLM_ProxyModelTable(**_data)
|
return await ModelRepository(prisma_client).table.create(data=_create_data)
|
||||||
return model_response
|
async with slot as table:
|
||||||
|
return await table.create(data=_create_data)
|
||||||
|
|
||||||
|
|
||||||
async def _add_team_model_to_db(
|
async def _add_team_model_to_db(
|
||||||
model_params: Deployment,
|
model_params: Deployment,
|
||||||
user_api_key_dict: UserAPIKeyAuth,
|
user_api_key_dict: UserAPIKeyAuth,
|
||||||
prisma_client: PrismaClient,
|
prisma_client: PrismaClient,
|
||||||
|
slot: AbstractAsyncContextManager[_ProxyModelTable] | None = None,
|
||||||
) -> "prisma_models.LiteLLM_ProxyModelTable | LiteLLM_ProxyModelTable | None":
|
) -> "prisma_models.LiteLLM_ProxyModelTable | LiteLLM_ProxyModelTable | None":
|
||||||
"""
|
"""
|
||||||
If 'team_id' is provided,
|
If 'team_id' is provided,
|
||||||
|
|
@ -1053,6 +1264,7 @@ async def _add_team_model_to_db(
|
||||||
model_params=model_params,
|
model_params=model_params,
|
||||||
user_api_key_dict=user_api_key_dict,
|
user_api_key_dict=user_api_key_dict,
|
||||||
prisma_client=prisma_client,
|
prisma_client=prisma_client,
|
||||||
|
slot=slot,
|
||||||
)
|
)
|
||||||
|
|
||||||
if original_model_name:
|
if original_model_name:
|
||||||
|
|
@ -1073,7 +1285,8 @@ async def _update_team_model_in_db(
|
||||||
patch_data: updateDeployment,
|
patch_data: updateDeployment,
|
||||||
user_api_key_dict: UserAPIKeyAuth,
|
user_api_key_dict: UserAPIKeyAuth,
|
||||||
prisma_client: PrismaClient,
|
prisma_client: PrismaClient,
|
||||||
) -> PrismaCompatibleUpdateDBModel:
|
write_row: Callable[[PrismaCompatibleUpdateDBModel], Awaitable[_RowT]],
|
||||||
|
) -> _RowT:
|
||||||
"""
|
"""
|
||||||
Handle team model updates with proper alias management.
|
Handle team model updates with proper alias management.
|
||||||
|
|
||||||
|
|
@ -1081,6 +1294,9 @@ async def _update_team_model_in_db(
|
||||||
- Creates unique internal model_name and team alias
|
- Creates unique internal model_name and team alias
|
||||||
- Adds model to team object
|
- Adds model to team object
|
||||||
- Preserves team_public_model_name for external reference
|
- Preserves team_public_model_name for external reference
|
||||||
|
|
||||||
|
The row is written through ``write_row`` before the team's model list is touched, so a
|
||||||
|
refused or failed write leaves the team as it was (the create path orders itself the same way).
|
||||||
"""
|
"""
|
||||||
# Validate team_id if present in patch_data
|
# Validate team_id if present in patch_data
|
||||||
from litellm.proxy.proxy_server import premium_user
|
from litellm.proxy.proxy_server import premium_user
|
||||||
|
|
@ -1114,7 +1330,7 @@ async def _update_team_model_in_db(
|
||||||
|
|
||||||
# No team_id in patch, proceed with standard update
|
# No team_id in patch, proceed with standard update
|
||||||
if patch_team_id is None:
|
if patch_team_id is None:
|
||||||
return update_db_model(db_model=db_model, updated_patch=patch_data)
|
return await write_row(update_db_model(db_model=db_model, updated_patch=patch_data))
|
||||||
|
|
||||||
# Determine public model name
|
# Determine public model name
|
||||||
public_model_name: Final = _get_public_model_name(
|
public_model_name: Final = _get_public_model_name(
|
||||||
|
|
@ -1133,6 +1349,10 @@ async def _update_team_model_in_db(
|
||||||
db_team_id: Final = db_model.model_info.team_id if db_model.model_info else None
|
db_team_id: Final = db_model.model_info.team_id if db_model.model_info else None
|
||||||
is_new_team_assignment: Final = db_team_id != patch_team_id
|
is_new_team_assignment: Final = db_team_id != patch_team_id
|
||||||
|
|
||||||
|
# Team rows keep their internal UUID-based model_name; the public name lives in model_info
|
||||||
|
patch_data.model_name = f"model_name_{patch_team_id}_{uuid.uuid4()}" if is_new_team_assignment else None
|
||||||
|
row: Final = await write_row(update_db_model(db_model=db_model, updated_patch=patch_data))
|
||||||
|
|
||||||
if is_new_team_assignment:
|
if is_new_team_assignment:
|
||||||
await _setup_new_team_model_assignment(
|
await _setup_new_team_model_assignment(
|
||||||
team_id=patch_team_id,
|
team_id=patch_team_id,
|
||||||
|
|
@ -1150,7 +1370,7 @@ async def _update_team_model_in_db(
|
||||||
prisma_client=prisma_client,
|
prisma_client=prisma_client,
|
||||||
)
|
)
|
||||||
|
|
||||||
return update_db_model(db_model=db_model, updated_patch=patch_data)
|
return row
|
||||||
|
|
||||||
|
|
||||||
def _get_public_model_name(
|
def _get_public_model_name(
|
||||||
|
|
@ -1313,7 +1533,7 @@ async def _get_team_public_model_names(
|
||||||
model_info = model_info_as_mapping(row.model_info)
|
model_info = model_info_as_mapping(row.model_info)
|
||||||
if model_info is not None:
|
if model_info is not None:
|
||||||
public_name = model_info.get("team_public_model_name")
|
public_name = model_info.get("team_public_model_name")
|
||||||
if public_name:
|
if isinstance(public_name, str) and public_name:
|
||||||
public_names.add(public_name)
|
public_names.add(public_name)
|
||||||
return public_names
|
return public_names
|
||||||
|
|
||||||
|
|
@ -1546,7 +1766,14 @@ class ModelManagementAuthChecks:
|
||||||
prisma_client: PrismaClient,
|
prisma_client: PrismaClient,
|
||||||
premium_user: bool,
|
premium_user: bool,
|
||||||
allow_missing_team: bool = False,
|
allow_missing_team: bool = False,
|
||||||
) -> Literal[True]:
|
member_operation: Literal["create", "update"] | None = None,
|
||||||
|
incoming_model_params: updateDeployment | None = None,
|
||||||
|
) -> Literal[True] | MemberAutoRouterWrite:
|
||||||
|
if user_api_key_dict.user_role in (
|
||||||
|
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
|
||||||
|
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY,
|
||||||
|
):
|
||||||
|
raise HTTPException(status_code=403, detail="View-only users cannot manage models.")
|
||||||
## Check team model auth
|
## Check team model auth
|
||||||
if model_params.model_info is not None and model_params.model_info.team_id is not None:
|
if model_params.model_info is not None and model_params.model_info.team_id is not None:
|
||||||
team_obj_row: Final = await _repo_team_table(prisma_client).find_unique(
|
team_obj_row: Final = await _repo_team_table(prisma_client).find_unique(
|
||||||
|
|
@ -1569,6 +1796,27 @@ class ModelManagementAuthChecks:
|
||||||
)
|
)
|
||||||
team_obj: Final = LiteLLM_TeamTable.model_validate(team_obj_row.model_dump())
|
team_obj: Final = LiteLLM_TeamTable.model_validate(team_obj_row.model_dump())
|
||||||
|
|
||||||
|
if (
|
||||||
|
member_operation is not None
|
||||||
|
and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN
|
||||||
|
and not _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj)
|
||||||
|
):
|
||||||
|
from litellm.proxy.proxy_server import llm_router
|
||||||
|
|
||||||
|
if llm_router is None or (member_operation == "update" and incoming_model_params is None):
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400, detail="An auto-router configuration and model catalog are required."
|
||||||
|
)
|
||||||
|
return await authorize_member_auto_router_write(
|
||||||
|
incoming=incoming_model_params if incoming_model_params is not None else model_params,
|
||||||
|
existing=model_params if member_operation == "update" else None,
|
||||||
|
user_api_key_dict=user_api_key_dict,
|
||||||
|
team=team_obj,
|
||||||
|
premium_user=premium_user,
|
||||||
|
prisma_client=prisma_client,
|
||||||
|
llm_router=llm_router,
|
||||||
|
)
|
||||||
|
|
||||||
return ModelManagementAuthChecks.can_user_make_team_model_call(
|
return ModelManagementAuthChecks.can_user_make_team_model_call(
|
||||||
team_id=model_params.model_info.team_id,
|
team_id=model_params.model_info.team_id,
|
||||||
user_api_key_dict=user_api_key_dict,
|
user_api_key_dict=user_api_key_dict,
|
||||||
|
|
@ -1820,11 +2068,17 @@ async def add_new_model(
|
||||||
)
|
)
|
||||||
|
|
||||||
## Auth check
|
## Auth check
|
||||||
await ModelManagementAuthChecks.can_user_make_model_call(
|
write_authorization: Final = await ModelManagementAuthChecks.can_user_make_model_call(
|
||||||
model_params=model_params,
|
model_params=model_params,
|
||||||
user_api_key_dict=user_api_key_dict,
|
user_api_key_dict=user_api_key_dict,
|
||||||
prisma_client=prisma_client,
|
prisma_client=prisma_client,
|
||||||
premium_user=premium_user,
|
premium_user=premium_user,
|
||||||
|
member_operation="create",
|
||||||
|
)
|
||||||
|
member_write: Final = write_authorization if isinstance(write_authorization, MemberAutoRouterWrite) else None
|
||||||
|
if member_write is not None and model_params.model_info is not None:
|
||||||
|
model_params.model_info = model_params.model_info.model_copy( # rebind-ok: downstream team-model handling mutates this same object
|
||||||
|
update=MappingProxyType({"member_auto_router": True})
|
||||||
)
|
)
|
||||||
|
|
||||||
_raise_on_strategy_router_write_violation(
|
_raise_on_strategy_router_write_violation(
|
||||||
|
|
@ -1854,17 +2108,20 @@ async def add_new_model(
|
||||||
reload_outcome: ReconcileOutcome = ReconcileOutcome(still_desired=None, live_after=None)
|
reload_outcome: ReconcileOutcome = ReconcileOutcome(still_desired=None, live_after=None)
|
||||||
try:
|
try:
|
||||||
_original_litellm_model_name: Final = model_params.model_name
|
_original_litellm_model_name: Final = model_params.model_name
|
||||||
|
add_slot: Final = _member_auto_router_write_slot(prisma_client, member_write=member_write)
|
||||||
if model_params.model_info.team_id is None:
|
if model_params.model_info.team_id is None:
|
||||||
model_response = await _add_model_to_db(
|
model_response = await _add_model_to_db(
|
||||||
model_params=priced_model_params,
|
model_params=priced_model_params,
|
||||||
user_api_key_dict=user_api_key_dict,
|
user_api_key_dict=user_api_key_dict,
|
||||||
prisma_client=prisma_client,
|
prisma_client=prisma_client,
|
||||||
|
slot=add_slot,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
model_response = await _add_team_model_to_db(
|
model_response = await _add_team_model_to_db(
|
||||||
model_params=priced_model_params,
|
model_params=priced_model_params,
|
||||||
user_api_key_dict=user_api_key_dict,
|
user_api_key_dict=user_api_key_dict,
|
||||||
prisma_client=prisma_client,
|
prisma_client=prisma_client,
|
||||||
|
slot=add_slot,
|
||||||
)
|
)
|
||||||
reload_outcome = await proxy_config.add_deployment(
|
reload_outcome = await proxy_config.add_deployment(
|
||||||
prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj
|
prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj
|
||||||
|
|
@ -1992,11 +2249,17 @@ async def update_model(
|
||||||
raise Exception("model not found")
|
raise Exception("model not found")
|
||||||
deployment: Final = Deployment(**_existing_litellm_params.model_dump())
|
deployment: Final = Deployment(**_existing_litellm_params.model_dump())
|
||||||
|
|
||||||
await ModelManagementAuthChecks.can_user_make_model_call(
|
write_authorization: Final = await ModelManagementAuthChecks.can_user_make_model_call(
|
||||||
model_params=deployment,
|
model_params=deployment,
|
||||||
user_api_key_dict=user_api_key_dict,
|
user_api_key_dict=user_api_key_dict,
|
||||||
prisma_client=prisma_client,
|
prisma_client=prisma_client,
|
||||||
premium_user=premium_user,
|
premium_user=premium_user,
|
||||||
|
member_operation="update",
|
||||||
|
incoming_model_params=model_params,
|
||||||
|
)
|
||||||
|
member_write: Final = write_authorization if isinstance(write_authorization, MemberAutoRouterWrite) else None
|
||||||
|
member_marker: Final = _member_auto_router_marker_for_update(
|
||||||
|
incoming_params=model_params.litellm_params, existing=deployment, member_write=member_write
|
||||||
)
|
)
|
||||||
|
|
||||||
_raise_on_strategy_router_write_violation(
|
_raise_on_strategy_router_write_violation(
|
||||||
|
|
@ -2038,11 +2301,21 @@ async def update_model(
|
||||||
if value is not None or _existing_litellm_params_dict.get(key) is not None
|
if value is not None or _existing_litellm_params_dict.get(key) is not None
|
||||||
}
|
}
|
||||||
|
|
||||||
_data: Final[dict[str, str]] = {
|
_data: Final[dict[str, str]] = { # mutable-ok: prisma update payload is dict-shaped
|
||||||
"litellm_params": json.dumps(merged_dictionary),
|
"litellm_params": json.dumps(merged_dictionary),
|
||||||
"updated_by": user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME,
|
"updated_by": user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME,
|
||||||
|
**(
|
||||||
|
{
|
||||||
|
"model_info": deployment.model_info.model_copy(
|
||||||
|
update=MappingProxyType({"member_auto_router": member_marker})
|
||||||
|
).model_dump_json(exclude_none=True)
|
||||||
}
|
}
|
||||||
model_response: Final = await _proxy_model_table(prisma_client).update(
|
if member_marker is not None
|
||||||
|
else {}
|
||||||
|
),
|
||||||
|
}
|
||||||
|
async with _member_auto_router_write_slot(prisma_client, member_write=member_write) as update_table:
|
||||||
|
model_response: Final = await update_table.update(
|
||||||
where={"model_id": _model_id},
|
where={"model_id": _model_id},
|
||||||
data=_data,
|
data=_data,
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -4364,7 +4364,7 @@ async def team_info(
|
||||||
try:
|
try:
|
||||||
team_info: BaseModel | None = await _team_db(prisma_client).find_unique(
|
team_info: BaseModel | None = await _team_db(prisma_client).find_unique(
|
||||||
where={"team_id": team_id},
|
where={"team_id": team_id},
|
||||||
include={"litellm_model_table": True, "object_permission": True},
|
include={"litellm_model_table": True, "object_permission": True}, # mutable-ok: prisma include clause
|
||||||
)
|
)
|
||||||
if team_info is None:
|
if team_info is None:
|
||||||
raise Exception
|
raise Exception
|
||||||
|
|
@ -5567,7 +5567,7 @@ async def team_model_add(
|
||||||
updated_team: Final = await _team_db(prisma_client).update(
|
updated_team: Final = await _team_db(prisma_client).update(
|
||||||
where={"team_id": data.team_id},
|
where={"team_id": data.team_id},
|
||||||
data={"updated_at": datetime.now(timezone.utc)},
|
data={"updated_at": datetime.now(timezone.utc)},
|
||||||
include={"object_permission": True},
|
include={"litellm_model_table": True, "object_permission": True}, # mutable-ok: prisma include clause
|
||||||
)
|
)
|
||||||
if updated_team is None:
|
if updated_team is None:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
|
|
@ -5654,7 +5654,7 @@ async def team_model_delete(
|
||||||
updated_team: Final = await _team_db(prisma_client).update(
|
updated_team: Final = await _team_db(prisma_client).update(
|
||||||
where={"team_id": data.team_id},
|
where={"team_id": data.team_id},
|
||||||
data={"models": updated_models},
|
data={"models": updated_models},
|
||||||
include={"object_permission": True},
|
include={"litellm_model_table": True, "object_permission": True}, # mutable-ok: prisma include clause
|
||||||
)
|
)
|
||||||
if updated_team is None:
|
if updated_team is None:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
|
|
|
||||||
|
|
@ -22,7 +22,6 @@ from html import escape
|
||||||
from types import MappingProxyType
|
from types import MappingProxyType
|
||||||
from typing import (
|
from typing import (
|
||||||
TYPE_CHECKING,
|
TYPE_CHECKING,
|
||||||
Annotated,
|
|
||||||
Any,
|
Any,
|
||||||
Final,
|
Final,
|
||||||
Literal,
|
Literal,
|
||||||
|
|
@ -41,7 +40,7 @@ if TYPE_CHECKING:
|
||||||
import jwt
|
import jwt
|
||||||
from fastapi import APIRouter, Depends, Header, HTTPException, Request, Response, status
|
from fastapi import APIRouter, Depends, Header, HTTPException, Request, Response, status
|
||||||
from fastapi.responses import RedirectResponse
|
from fastapi.responses import RedirectResponse
|
||||||
from pydantic import BaseModel, BeforeValidator, ConfigDict, TypeAdapter, ValidationError
|
from pydantic import BaseModel, TypeAdapter, ValidationError
|
||||||
|
|
||||||
import litellm
|
import litellm
|
||||||
from litellm._logging import verbose_proxy_logger
|
from litellm._logging import verbose_proxy_logger
|
||||||
|
|
@ -92,6 +91,7 @@ from litellm.proxy.auth.auth_utils import (
|
||||||
_has_user_setup_sso,
|
_has_user_setup_sso,
|
||||||
)
|
)
|
||||||
from litellm.proxy.auth.handle_jwt import JWTHandler
|
from litellm.proxy.auth.handle_jwt import JWTHandler
|
||||||
|
from litellm.proxy.auth.team_grants import TeamModelAliasTable
|
||||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||||
from litellm.proxy.common_utils.admin_ui_utils import (
|
from litellm.proxy.common_utils.admin_ui_utils import (
|
||||||
admin_ui_disabled,
|
admin_ui_disabled,
|
||||||
|
|
@ -202,31 +202,14 @@ def _team_detail_db(repo: TeamRepository) -> "TableActions[_TeamDetailRow]":
|
||||||
return repo.table
|
return repo.table
|
||||||
|
|
||||||
|
|
||||||
_MODEL_ALIASES_ADAPTER: Final = TypeAdapter(dict[str, str])
|
|
||||||
_SSO_TOKEN_CLAIMS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
_SSO_TOKEN_CLAIMS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
||||||
|
|
||||||
|
|
||||||
def _decode_model_aliases(value: object) -> object:
|
|
||||||
"""``/team/new`` stores team model aliases as a JSON-encoded string in the Json column."""
|
|
||||||
if not isinstance(value, str):
|
|
||||||
return value
|
|
||||||
try:
|
|
||||||
return _MODEL_ALIASES_ADAPTER.validate_json(value)
|
|
||||||
except ValidationError:
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
class _TeamModelAliasTable(BaseModel):
|
|
||||||
model_config = ConfigDict(protected_namespaces=())
|
|
||||||
|
|
||||||
model_aliases: Annotated[Mapping[str, str] | None, BeforeValidator(_decode_model_aliases)] = None
|
|
||||||
|
|
||||||
|
|
||||||
class _TeamRowGrants(BaseModel):
|
class _TeamRowGrants(BaseModel):
|
||||||
team_id: str
|
team_id: str
|
||||||
team_alias: str | None = None
|
team_alias: str | None = None
|
||||||
models: tuple[str, ...] = ()
|
models: tuple[str, ...] = ()
|
||||||
litellm_model_table: _TeamModelAliasTable | None = None
|
litellm_model_table: TeamModelAliasTable | None = None
|
||||||
|
|
||||||
|
|
||||||
class CliSsoTeamDetail(BaseModel):
|
class CliSsoTeamDetail(BaseModel):
|
||||||
|
|
|
||||||
|
|
@ -21,7 +21,7 @@ import time
|
||||||
import traceback
|
import traceback
|
||||||
import weakref
|
import weakref
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Generator, Mapping, Sequence
|
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Generator, Iterator, Mapping, Sequence
|
||||||
from functools import lru_cache
|
from functools import lru_cache
|
||||||
from types import MappingProxyType
|
from types import MappingProxyType
|
||||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias, TypeVar, Union, cast
|
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias, TypeVar, Union, cast
|
||||||
|
|
@ -9446,6 +9446,15 @@ class Router:
|
||||||
if _budget_limiter is not None:
|
if _budget_limiter is not None:
|
||||||
_budget_limiter.register_deployment_budget(deployment=deployment.to_json(exclude_none=True))
|
_budget_limiter.register_deployment_budget(deployment=deployment.to_json(exclude_none=True))
|
||||||
|
|
||||||
|
def config_deployments(self) -> Iterator[Mapping[str, object]]:
|
||||||
|
"""The model_list rows that came from config.yaml rather than the DB (``model_info.db_model`` unset)."""
|
||||||
|
for deployment in self.model_list:
|
||||||
|
if not isinstance(deployment, Mapping):
|
||||||
|
continue
|
||||||
|
model_info = deployment.get("model_info")
|
||||||
|
if not (isinstance(model_info, Mapping) and model_info.get("db_model")):
|
||||||
|
yield deployment
|
||||||
|
|
||||||
def get_deployment(self, model_id: str) -> Deployment | None:
|
def get_deployment(self, model_id: str) -> Deployment | None:
|
||||||
"""
|
"""
|
||||||
Returns -> Deployment or None
|
Returns -> Deployment or None
|
||||||
|
|
|
||||||
|
|
@ -922,7 +922,9 @@ def _is_classifier_timeout(exc: BaseException) -> bool:
|
||||||
# LiteLLM still supports 3.10, where they are distinct exception classes.
|
# LiteLLM still supports 3.10, where they are distinct exception classes.
|
||||||
if isinstance(exc, (TimeoutError, asyncio.TimeoutError)):
|
if isinstance(exc, (TimeoutError, asyncio.TimeoutError)):
|
||||||
return True
|
return True
|
||||||
return type(exc).__name__.endswith("TimeoutError")
|
from litellm.exceptions import Timeout as LiteLLMTimeout
|
||||||
|
|
||||||
|
return isinstance(exc, LiteLLMTimeout)
|
||||||
|
|
||||||
|
|
||||||
class _SessionAffinityPin(NamedTuple):
|
class _SessionAffinityPin(NamedTuple):
|
||||||
|
|
|
||||||
|
|
@ -160,6 +160,7 @@ class ModelInfo(MirroredPricingParams):
|
||||||
|
|
||||||
# the model_name that can be used by the team when making LLM calls
|
# the model_name that can be used by the team when making LLM calls
|
||||||
team_public_model_name: str | None = None
|
team_public_model_name: str | None = None
|
||||||
|
member_auto_router: bool = False
|
||||||
|
|
||||||
# admin-toggled pause flag; mirrors LiteLLM_ProxyModelTable.blocked
|
# admin-toggled pause flag; mirrors LiteLLM_ProxyModelTable.blocked
|
||||||
blocked: bool | None = None
|
blocked: bool | None = None
|
||||||
|
|
|
||||||
File diff suppressed because it is too large
Load diff
File diff suppressed because it is too large
Load diff
129
tests/test_litellm/proxy/auth/test_team_grants.py
Normal file
129
tests/test_litellm/proxy/auth/test_team_grants.py
Normal file
|
|
@ -0,0 +1,129 @@
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from litellm.proxy._types import (
|
||||||
|
LiteLLM_BudgetTable,
|
||||||
|
LiteLLM_ObjectPermissionTable,
|
||||||
|
LiteLLM_TeamMembership,
|
||||||
|
LiteLLM_TeamTable,
|
||||||
|
LiteLLM_VerificationTokenView,
|
||||||
|
Member,
|
||||||
|
UserAPIKeyAuth,
|
||||||
|
)
|
||||||
|
from litellm.models.team import LiteLLM_ModelTable
|
||||||
|
from litellm.proxy.auth.team_grants import team_grants, team_model_aliases
|
||||||
|
|
||||||
|
TEAM_ID = "team-grants"
|
||||||
|
USER_ID = "user-in-team"
|
||||||
|
ALIASES = {"fast": "gpt-4o-mini", "smart": "gpt-4o"}
|
||||||
|
|
||||||
|
|
||||||
|
def _alias_table(model_aliases) -> LiteLLM_ModelTable:
|
||||||
|
return LiteLLM_ModelTable(model_aliases=model_aliases, created_by="admin", updated_by="admin")
|
||||||
|
|
||||||
|
|
||||||
|
def _full_team(model_aliases=ALIASES) -> LiteLLM_TeamTable:
|
||||||
|
return LiteLLM_TeamTable(
|
||||||
|
team_id=TEAM_ID,
|
||||||
|
team_alias="grants-team",
|
||||||
|
tpm_limit=1000,
|
||||||
|
rpm_limit=10,
|
||||||
|
max_budget=50.0,
|
||||||
|
soft_budget=25.0,
|
||||||
|
spend=12.5,
|
||||||
|
models=["gpt-4o", "gpt-4o-mini"],
|
||||||
|
blocked=True,
|
||||||
|
metadata={"tier": "gold"},
|
||||||
|
litellm_model_table=_alias_table(model_aliases),
|
||||||
|
object_permission_id="op-1",
|
||||||
|
object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-1", mcp_servers=["mcp-a"]),
|
||||||
|
members_with_roles=[
|
||||||
|
Member(user_id="someone-else", role="user"),
|
||||||
|
Member(user_id=USER_ID, role="admin"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _membership() -> LiteLLM_TeamMembership:
|
||||||
|
return LiteLLM_TeamMembership(
|
||||||
|
user_id=USER_ID,
|
||||||
|
team_id=TEAM_ID,
|
||||||
|
spend=3.25,
|
||||||
|
litellm_budget_table=LiteLLM_BudgetTable(tpm_limit=500, rpm_limit=5),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_team_grants_cover_every_team_field_the_key_path_gets():
|
||||||
|
"""Class guard for LIT-5858 and its siblings: every ``team_*`` column the combined-view SQL hands the
|
||||||
|
virtual-key path must come out of the projection too, with the team's actual value, so adding a column
|
||||||
|
to ``LiteLLM_VerificationTokenView`` without teaching ``team_grants`` fails here instead of in prod."""
|
||||||
|
team = _full_team()
|
||||||
|
grants = team_grants(team_object=team, team_membership=_membership(), user_id=USER_ID)
|
||||||
|
token = UserAPIKeyAuth(team_id=TEAM_ID, **grants)
|
||||||
|
|
||||||
|
view_team_fields = {name for name in LiteLLM_VerificationTokenView.model_fields if name.startswith("team_")}
|
||||||
|
assert view_team_fields - {"team_id"} <= set(grants)
|
||||||
|
assert all(grants[name] is not None for name in view_team_fields - {"team_id"})
|
||||||
|
|
||||||
|
assert token.team_alias == "grants-team"
|
||||||
|
assert token.team_tpm_limit == 1000
|
||||||
|
assert token.team_rpm_limit == 10
|
||||||
|
assert token.team_max_budget == 50.0
|
||||||
|
assert token.team_soft_budget == 25.0
|
||||||
|
assert token.team_spend == 12.5
|
||||||
|
assert token.team_models == ["gpt-4o", "gpt-4o-mini"]
|
||||||
|
assert token.team_blocked is True
|
||||||
|
assert token.team_metadata == {"tier": "gold"}
|
||||||
|
assert token.team_model_aliases == ALIASES
|
||||||
|
assert token.team_object_permission_id == "op-1"
|
||||||
|
assert token.team_object_permission is not None
|
||||||
|
assert token.team_object_permission.mcp_servers == ["mcp-a"]
|
||||||
|
assert token.team_member == Member(user_id=USER_ID, role="admin")
|
||||||
|
assert token.team_member_spend == 3.25
|
||||||
|
assert token.team_member_tpm_limit == 500
|
||||||
|
assert token.team_member_rpm_limit == 5
|
||||||
|
|
||||||
|
|
||||||
|
def test_team_grants_without_team_leave_token_defaults():
|
||||||
|
token = UserAPIKeyAuth(**team_grants(team_object=None, team_membership=None, user_id=USER_ID))
|
||||||
|
assert token == UserAPIKeyAuth()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"stored_aliases",
|
||||||
|
[ALIASES, '{"fast": "gpt-4o-mini", "smart": "gpt-4o"}'],
|
||||||
|
ids=["json-object", "json-string-as-written-by-team-new"],
|
||||||
|
)
|
||||||
|
def test_team_model_aliases_decode_both_storage_shapes(stored_aliases):
|
||||||
|
team = _full_team(model_aliases=stored_aliases)
|
||||||
|
assert team_model_aliases(team) == ALIASES
|
||||||
|
assert team_grants(team_object=team, team_membership=None, user_id=None)["team_model_aliases"] == ALIASES
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("stored_aliases", [None, "not json", '["a", "b"]', {"fast": 3}], ids=str)
|
||||||
|
def test_team_model_aliases_treat_unusable_column_as_no_aliases(stored_aliases):
|
||||||
|
team = _full_team(model_aliases=stored_aliases)
|
||||||
|
assert team_model_aliases(team) is None
|
||||||
|
assert team_grants(team_object=team, team_membership=None, user_id=None)["team_model_aliases"] is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_team_model_aliases_none_without_relation_loaded():
|
||||||
|
team = _full_team()
|
||||||
|
team.litellm_model_table = None
|
||||||
|
assert team_model_aliases(team) is None
|
||||||
|
assert team_model_aliases(None) is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_team_member_is_the_callers_row_only():
|
||||||
|
team = _full_team()
|
||||||
|
assert team_grants(team_object=team, team_membership=None, user_id="someone-else")["team_member"] == Member(
|
||||||
|
user_id="someone-else", role="user"
|
||||||
|
)
|
||||||
|
assert team_grants(team_object=team, team_membership=None, user_id="stranger")["team_member"] is None
|
||||||
|
assert team_grants(team_object=team, team_membership=None, user_id=None)["team_member"] is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_membership_limits_absent_without_membership_row():
|
||||||
|
grants = team_grants(team_object=_full_team(), team_membership=None, user_id=USER_ID)
|
||||||
|
assert grants["team_member_spend"] is None
|
||||||
|
assert grants["team_member_tpm_limit"] is None
|
||||||
|
assert grants["team_member_rpm_limit"] is None
|
||||||
|
|
@ -187,7 +187,7 @@ async def test_budget_reservation_runs_when_not_disabled():
|
||||||
)
|
)
|
||||||
async def test_fail_closed_budget_enforcement_reaches_reservation(
|
async def test_fail_closed_budget_enforcement_reaches_reservation(
|
||||||
general_settings, expected_flag
|
general_settings, expected_flag
|
||||||
):
|
): # test-quality-ok: [TQ002] collaborator injected via its import site; there is no seam to patch otherwise
|
||||||
"""#33923: the strict flag must be threaded into reserve_budget_for_request so a
|
"""#33923: the strict flag must be threaded into reserve_budget_for_request so a
|
||||||
failed reservation write can reject instead of failing open."""
|
failed reservation write can reject instead of failing open."""
|
||||||
user_api_key_auth_obj = UserAPIKeyAuth(token="test_token")
|
user_api_key_auth_obj = UserAPIKeyAuth(token="test_token")
|
||||||
|
|
@ -210,10 +210,7 @@ async def test_fail_closed_budget_enforcement_reaches_reservation(
|
||||||
general_settings=general_settings,
|
general_settings=general_settings,
|
||||||
)
|
)
|
||||||
|
|
||||||
assert (
|
assert mock_reserve.await_args.kwargs["fail_closed_budget_enforcement"] is expected_flag
|
||||||
mock_reserve.await_args.kwargs["fail_closed_budget_enforcement"]
|
|
||||||
is expected_flag
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
@ -227,7 +224,7 @@ async def test_fail_closed_budget_enforcement_reaches_reservation(
|
||||||
)
|
)
|
||||||
async def test_apply_user_budget_to_team_keys_reaches_reservation(
|
async def test_apply_user_budget_to_team_keys_reaches_reservation(
|
||||||
general_settings, expected_flag
|
general_settings, expected_flag
|
||||||
):
|
): # test-quality-ok: [TQ002] collaborator injected via its import site; there is no seam to patch otherwise
|
||||||
"""The opt-in lives in general_settings but is consumed inside
|
"""The opt-in lives in general_settings but is consumed inside
|
||||||
_get_budget_counters, so it has to be threaded through reserve_budget_for_request
|
_get_budget_counters, so it has to be threaded through reserve_budget_for_request
|
||||||
or the reservation path keeps exempting team keys while the read path enforces."""
|
or the reservation path keeps exempting team keys while the read path enforces."""
|
||||||
|
|
@ -251,9 +248,7 @@ async def test_apply_user_budget_to_team_keys_reaches_reservation(
|
||||||
general_settings=general_settings,
|
general_settings=general_settings,
|
||||||
)
|
)
|
||||||
|
|
||||||
assert (
|
assert mock_reserve.await_args.kwargs["apply_user_budget_to_team_keys"] is expected_flag
|
||||||
mock_reserve.await_args.kwargs["apply_user_budget_to_team_keys"] is expected_flag
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
@ -353,9 +348,7 @@ async def test_custom_auth_honors_key_level_model_access_restriction_allowed_wit
|
||||||
"litellm.proxy.auth.user_api_key_auth.can_key_call_model",
|
"litellm.proxy.auth.user_api_key_auth.can_key_call_model",
|
||||||
new_callable=AsyncMock,
|
new_callable=AsyncMock,
|
||||||
) as mock_can_key,
|
) as mock_can_key,
|
||||||
patch(
|
patch("litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock),
|
||||||
"litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock
|
|
||||||
),
|
|
||||||
patch(
|
patch(
|
||||||
"litellm.proxy.proxy_server.general_settings",
|
"litellm.proxy.proxy_server.general_settings",
|
||||||
{"custom_auth_run_common_checks": True},
|
{"custom_auth_run_common_checks": True},
|
||||||
|
|
@ -386,9 +379,7 @@ async def test_custom_auth_enforces_key_model_access_from_file_route_header_with
|
||||||
"litellm.proxy.auth.user_api_key_auth.can_key_call_model",
|
"litellm.proxy.auth.user_api_key_auth.can_key_call_model",
|
||||||
new_callable=AsyncMock,
|
new_callable=AsyncMock,
|
||||||
) as mock_can_key,
|
) as mock_can_key,
|
||||||
patch(
|
patch("litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock),
|
||||||
"litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock
|
|
||||||
),
|
|
||||||
patch(
|
patch(
|
||||||
"litellm.proxy.proxy_server.general_settings",
|
"litellm.proxy.proxy_server.general_settings",
|
||||||
{"custom_auth_run_common_checks": True},
|
{"custom_auth_run_common_checks": True},
|
||||||
|
|
@ -419,9 +410,7 @@ async def test_custom_auth_honors_key_level_model_access_restriction_denied_with
|
||||||
"litellm.proxy.auth.user_api_key_auth.can_key_call_model",
|
"litellm.proxy.auth.user_api_key_auth.can_key_call_model",
|
||||||
new_callable=AsyncMock,
|
new_callable=AsyncMock,
|
||||||
) as mock_can_key,
|
) as mock_can_key,
|
||||||
patch(
|
patch("litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock),
|
||||||
"litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock
|
|
||||||
),
|
|
||||||
patch(
|
patch(
|
||||||
"litellm.proxy.proxy_server.general_settings",
|
"litellm.proxy.proxy_server.general_settings",
|
||||||
{"custom_auth_run_common_checks": True},
|
{"custom_auth_run_common_checks": True},
|
||||||
|
|
@ -457,9 +446,7 @@ def _proxy_server_attrs_for_custom_auth(*, user_custom_auth):
|
||||||
mock_proxy_logging_obj = MagicMock()
|
mock_proxy_logging_obj = MagicMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
|
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
|
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
|
||||||
AsyncMock()
|
|
||||||
)
|
|
||||||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
|
|
@ -721,9 +708,7 @@ async def test_enterprise_custom_auth_runs_post_custom_auth_checks_when_opt_in()
|
||||||
litellm.enable_post_custom_auth_checks = original_flag
|
litellm.enable_post_custom_auth_checks = original_flag
|
||||||
|
|
||||||
|
|
||||||
def _assert_get_api_key_with_custom_litellm_key_header(
|
def _assert_get_api_key_with_custom_litellm_key_header(custom_litellm_key_header, api_key, passed_in_key):
|
||||||
custom_litellm_key_header, api_key, passed_in_key
|
|
||||||
):
|
|
||||||
assert get_api_key(
|
assert get_api_key(
|
||||||
custom_litellm_key_header=custom_litellm_key_header,
|
custom_litellm_key_header=custom_litellm_key_header,
|
||||||
api_key=None,
|
api_key=None,
|
||||||
|
|
@ -780,9 +765,7 @@ def _assert_get_api_key_with_custom_litellm_key_header(
|
||||||
("App:LiteLLM", None, False, False),
|
("App:LiteLLM", None, False, False),
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
def test_routing_selector_matches_claim_parametrized(
|
def test_routing_selector_matches_claim_parametrized(selector_value, claim_value, expected, split_space_delimited):
|
||||||
selector_value, claim_value, expected, split_space_delimited
|
|
||||||
):
|
|
||||||
assert (
|
assert (
|
||||||
_routing_selector_matches_claim(
|
_routing_selector_matches_claim(
|
||||||
selector_value=selector_value,
|
selector_value=selector_value,
|
||||||
|
|
@ -876,10 +859,7 @@ def test_routing_selector_matches_claim_parametrized(
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
def test_matches_routing_override_parametrized(override, token_claims, expected):
|
def test_matches_routing_override_parametrized(override, token_claims, expected):
|
||||||
assert (
|
assert _matches_routing_override(token_claims=token_claims, override=override) is expected
|
||||||
_matches_routing_override(token_claims=token_claims, override=override)
|
|
||||||
is expected
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def test_get_api_key_with_custom_litellm_key_header_bearer_prefix():
|
def test_get_api_key_with_custom_litellm_key_header_bearer_prefix():
|
||||||
|
|
@ -958,12 +938,9 @@ def test_team_metadata_with_tags_flows_through_jwt_auth():
|
||||||
)
|
)
|
||||||
|
|
||||||
# Verify team_metadata is set
|
# Verify team_metadata is set
|
||||||
assert (
|
assert user_api_key_auth.team_metadata is not None, "team_metadata should be populated"
|
||||||
user_api_key_auth.team_metadata is not None
|
|
||||||
), "team_metadata should be populated"
|
|
||||||
assert user_api_key_auth.team_metadata == team_object.metadata, (
|
assert user_api_key_auth.team_metadata == team_object.metadata, (
|
||||||
f"team_metadata not correctly mapped. "
|
f"team_metadata not correctly mapped. Expected: {team_object.metadata}, Got: {user_api_key_auth.team_metadata}"
|
||||||
f"Expected: {team_object.metadata}, Got: {user_api_key_auth.team_metadata}"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# Specifically verify tags are present
|
# Specifically verify tags are present
|
||||||
|
|
@ -1002,9 +979,7 @@ def test_route_checks_is_llm_api_route():
|
||||||
]
|
]
|
||||||
|
|
||||||
for route in openai_routes:
|
for route in openai_routes:
|
||||||
assert RouteChecks.is_llm_api_route(
|
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
|
||||||
route=route
|
|
||||||
), f"Route {route} should be identified as LLM API route"
|
|
||||||
|
|
||||||
# Test Anthropic routes
|
# Test Anthropic routes
|
||||||
anthropic_routes = [
|
anthropic_routes = [
|
||||||
|
|
@ -1013,9 +988,7 @@ def test_route_checks_is_llm_api_route():
|
||||||
]
|
]
|
||||||
|
|
||||||
for route in anthropic_routes:
|
for route in anthropic_routes:
|
||||||
assert RouteChecks.is_llm_api_route(
|
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
|
||||||
route=route
|
|
||||||
), f"Route {route} should be identified as LLM API route"
|
|
||||||
|
|
||||||
# Test passthrough routes (this is the key improvement over the old route checking)
|
# Test passthrough routes (this is the key improvement over the old route checking)
|
||||||
passthrough_routes = [
|
passthrough_routes = [
|
||||||
|
|
@ -1035,9 +1008,7 @@ def test_route_checks_is_llm_api_route():
|
||||||
]
|
]
|
||||||
|
|
||||||
for route in passthrough_routes:
|
for route in passthrough_routes:
|
||||||
assert RouteChecks.is_llm_api_route(
|
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
|
||||||
route=route
|
|
||||||
), f"Route {route} should be identified as LLM API route"
|
|
||||||
|
|
||||||
# Test MCP routes
|
# Test MCP routes
|
||||||
mcp_routes = [
|
mcp_routes = [
|
||||||
|
|
@ -1047,9 +1018,7 @@ def test_route_checks_is_llm_api_route():
|
||||||
]
|
]
|
||||||
|
|
||||||
for route in mcp_routes:
|
for route in mcp_routes:
|
||||||
assert RouteChecks.is_llm_api_route(
|
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
|
||||||
route=route
|
|
||||||
), f"Route {route} should be identified as LLM API route"
|
|
||||||
|
|
||||||
# Test LiteLLM native RAG routes
|
# Test LiteLLM native RAG routes
|
||||||
rag_routes = [
|
rag_routes = [
|
||||||
|
|
@ -1059,9 +1028,7 @@ def test_route_checks_is_llm_api_route():
|
||||||
"/v1/rag/query",
|
"/v1/rag/query",
|
||||||
]
|
]
|
||||||
for route in rag_routes:
|
for route in rag_routes:
|
||||||
assert RouteChecks.is_llm_api_route(
|
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
|
||||||
route=route
|
|
||||||
), f"Route {route} should be identified as LLM API route"
|
|
||||||
|
|
||||||
# Test routes with placeholders
|
# Test routes with placeholders
|
||||||
placeholder_routes = [
|
placeholder_routes = [
|
||||||
|
|
@ -1076,9 +1043,7 @@ def test_route_checks_is_llm_api_route():
|
||||||
]
|
]
|
||||||
|
|
||||||
for route in placeholder_routes:
|
for route in placeholder_routes:
|
||||||
assert RouteChecks.is_llm_api_route(
|
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
|
||||||
route=route
|
|
||||||
), f"Route {route} should be identified as LLM API route"
|
|
||||||
|
|
||||||
# Test Azure OpenAI routes
|
# Test Azure OpenAI routes
|
||||||
azure_routes = [
|
azure_routes = [
|
||||||
|
|
@ -1089,9 +1054,7 @@ def test_route_checks_is_llm_api_route():
|
||||||
]
|
]
|
||||||
|
|
||||||
for route in azure_routes:
|
for route in azure_routes:
|
||||||
assert RouteChecks.is_llm_api_route(
|
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
|
||||||
route=route
|
|
||||||
), f"Route {route} should be identified as LLM API route"
|
|
||||||
|
|
||||||
# Test non-LLM routes (should return False)
|
# Test non-LLM routes (should return False)
|
||||||
non_llm_routes = [
|
non_llm_routes = [
|
||||||
|
|
@ -1110,9 +1073,7 @@ def test_route_checks_is_llm_api_route():
|
||||||
]
|
]
|
||||||
|
|
||||||
for route in non_llm_routes:
|
for route in non_llm_routes:
|
||||||
assert not RouteChecks.is_llm_api_route(
|
assert not RouteChecks.is_llm_api_route(route=route), f"Route {route} should NOT be identified as LLM API route"
|
||||||
route=route
|
|
||||||
), f"Route {route} should NOT be identified as LLM API route"
|
|
||||||
|
|
||||||
# Test invalid inputs
|
# Test invalid inputs
|
||||||
invalid_inputs = [
|
invalid_inputs = [
|
||||||
|
|
@ -1124,9 +1085,9 @@ def test_route_checks_is_llm_api_route():
|
||||||
]
|
]
|
||||||
|
|
||||||
for invalid_input in invalid_inputs:
|
for invalid_input in invalid_inputs:
|
||||||
assert not RouteChecks.is_llm_api_route(
|
assert not RouteChecks.is_llm_api_route(route=invalid_input), (
|
||||||
route=invalid_input
|
f"Invalid input {invalid_input} should return False"
|
||||||
), f"Invalid input {invalid_input} should return False"
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
@ -1173,9 +1134,7 @@ async def test_proxy_admin_expired_key_from_cache():
|
||||||
mock_proxy_logging_obj = MagicMock()
|
mock_proxy_logging_obj = MagicMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
|
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
|
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
|
||||||
AsyncMock()
|
|
||||||
)
|
|
||||||
# Mock post_call_failure_hook as async function returning None (no transformation)
|
# Mock post_call_failure_hook as async function returning None (no transformation)
|
||||||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
||||||
|
|
||||||
|
|
@ -1212,9 +1171,7 @@ async def test_proxy_admin_expired_key_from_cache():
|
||||||
"jwt_handler": None,
|
"jwt_handler": None,
|
||||||
"litellm_proxy_admin_name": "admin",
|
"litellm_proxy_admin_name": "admin",
|
||||||
}
|
}
|
||||||
_original_values = {
|
_original_values = {attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set}
|
||||||
attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set
|
|
||||||
}
|
|
||||||
try:
|
try:
|
||||||
for attr, val in _attrs_to_set.items():
|
for attr, val in _attrs_to_set.items():
|
||||||
setattr(_proxy_server_mod, attr, val)
|
setattr(_proxy_server_mod, attr, val)
|
||||||
|
|
@ -1238,36 +1195,30 @@ async def test_proxy_admin_expired_key_from_cache():
|
||||||
)
|
)
|
||||||
|
|
||||||
# Verify that ProxyException was raised with expired_key type
|
# Verify that ProxyException was raised with expired_key type
|
||||||
assert hasattr(
|
assert hasattr(exc_info.value, "type"), "Exception should have 'type' attribute"
|
||||||
exc_info.value, "type"
|
assert exc_info.value.type == ProxyErrorTypes.expired_key, (
|
||||||
), "Exception should have 'type' attribute"
|
f"Expected expired_key error type, got {exc_info.value.type}"
|
||||||
assert (
|
)
|
||||||
exc_info.value.type == ProxyErrorTypes.expired_key
|
|
||||||
), f"Expected expired_key error type, got {exc_info.value.type}"
|
|
||||||
assert int(exc_info.value.code) == status.HTTP_401_UNAUTHORIZED
|
assert int(exc_info.value.code) == status.HTTP_401_UNAUTHORIZED
|
||||||
assert "Expired Key" in str(
|
assert "Expired Key" in str(exc_info.value.message), (
|
||||||
exc_info.value.message
|
f"Exception message should mention 'Expired Key', got: {exc_info.value.message}"
|
||||||
), f"Exception message should mention 'Expired Key', got: {exc_info.value.message}"
|
)
|
||||||
|
|
||||||
# Verify that the param field does NOT leak the full API key (Issue #18731)
|
# Verify that the param field does NOT leak the full API key (Issue #18731)
|
||||||
# The param should be abbreviated like "sk-...XXXX" not the full plaintext key
|
# The param should be abbreviated like "sk-...XXXX" not the full plaintext key
|
||||||
assert (
|
assert exc_info.value.param is not None, "Exception should have 'param' attribute"
|
||||||
exc_info.value.param is not None
|
|
||||||
), "Exception should have 'param' attribute"
|
|
||||||
assert exc_info.value.param != api_key, (
|
assert exc_info.value.param != api_key, (
|
||||||
f"SECURITY: Full API key should NOT be in param field! "
|
f"SECURITY: Full API key should NOT be in param field! "
|
||||||
f"Got: {exc_info.value.param}, Expected abbreviated format like 'sk-...XXXX'"
|
f"Got: {exc_info.value.param}, Expected abbreviated format like 'sk-...XXXX'"
|
||||||
)
|
)
|
||||||
assert exc_info.value.param.startswith(
|
assert exc_info.value.param.startswith("sk-..."), (
|
||||||
"sk-..."
|
f"Param should be abbreviated to 'sk-...XXXX' format. Got: {exc_info.value.param}"
|
||||||
), f"Param should be abbreviated to 'sk-...XXXX' format. Got: {exc_info.value.param}"
|
)
|
||||||
|
|
||||||
# Verify that cache deletion was called
|
# Verify that cache deletion was called
|
||||||
mock_delete_cache.assert_called_once()
|
mock_delete_cache.assert_called_once()
|
||||||
call_args = mock_delete_cache.call_args
|
call_args = mock_delete_cache.call_args
|
||||||
assert (
|
assert call_args[1]["hashed_token"] == hashed_key, "Cache deletion should be called with the hashed key"
|
||||||
call_args[1]["hashed_token"] == hashed_key
|
|
||||||
), "Cache deletion should be called with the hashed key"
|
|
||||||
finally:
|
finally:
|
||||||
# Restore all module-level attributes so subsequent tests are not affected
|
# Restore all module-level attributes so subsequent tests are not affected
|
||||||
for attr, val in _original_values.items():
|
for attr, val in _original_values.items():
|
||||||
|
|
@ -1305,9 +1256,7 @@ async def test_scim_deactivated_user_key_is_rejected():
|
||||||
mock_proxy_logging_obj = MagicMock()
|
mock_proxy_logging_obj = MagicMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
|
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
|
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
|
||||||
AsyncMock()
|
|
||||||
)
|
|
||||||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
||||||
|
|
||||||
mock_prisma_client = MagicMock()
|
mock_prisma_client = MagicMock()
|
||||||
|
|
@ -1328,9 +1277,7 @@ async def test_scim_deactivated_user_key_is_rejected():
|
||||||
"jwt_handler": None,
|
"jwt_handler": None,
|
||||||
"litellm_proxy_admin_name": "admin",
|
"litellm_proxy_admin_name": "admin",
|
||||||
}
|
}
|
||||||
_original_values = {
|
_original_values = {attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set}
|
||||||
attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set
|
|
||||||
}
|
|
||||||
try:
|
try:
|
||||||
for attr, val in _attrs_to_set.items():
|
for attr, val in _attrs_to_set.items():
|
||||||
setattr(_proxy_server_mod, attr, val)
|
setattr(_proxy_server_mod, attr, val)
|
||||||
|
|
@ -1397,9 +1344,7 @@ async def test_cached_proxy_admin_key_sets_via_virtual_key_marker():
|
||||||
mock_proxy_logging_obj = MagicMock()
|
mock_proxy_logging_obj = MagicMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
|
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
|
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
|
||||||
AsyncMock()
|
|
||||||
)
|
|
||||||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
||||||
|
|
||||||
import litellm.proxy.proxy_server as _proxy_server_mod
|
import litellm.proxy.proxy_server as _proxy_server_mod
|
||||||
|
|
@ -1418,9 +1363,7 @@ async def test_cached_proxy_admin_key_sets_via_virtual_key_marker():
|
||||||
"jwt_handler": None,
|
"jwt_handler": None,
|
||||||
"litellm_proxy_admin_name": "admin",
|
"litellm_proxy_admin_name": "admin",
|
||||||
}
|
}
|
||||||
_original_values = {
|
_original_values = {attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set}
|
||||||
attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set
|
|
||||||
}
|
|
||||||
try:
|
try:
|
||||||
for attr, val in _attrs_to_set.items():
|
for attr, val in _attrs_to_set.items():
|
||||||
setattr(_proxy_server_mod, attr, val)
|
setattr(_proxy_server_mod, attr, val)
|
||||||
|
|
@ -1472,9 +1415,7 @@ async def test_master_key_auth_sets_via_virtual_key_marker():
|
||||||
mock_proxy_logging_obj = MagicMock()
|
mock_proxy_logging_obj = MagicMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
|
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
|
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
|
||||||
AsyncMock()
|
|
||||||
)
|
|
||||||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
||||||
|
|
||||||
import litellm.proxy.proxy_server as _proxy_server_mod
|
import litellm.proxy.proxy_server as _proxy_server_mod
|
||||||
|
|
@ -1493,9 +1434,7 @@ async def test_master_key_auth_sets_via_virtual_key_marker():
|
||||||
"jwt_handler": None,
|
"jwt_handler": None,
|
||||||
"litellm_proxy_admin_name": "admin",
|
"litellm_proxy_admin_name": "admin",
|
||||||
}
|
}
|
||||||
_original_values = {
|
_original_values = {attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set}
|
||||||
attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set
|
|
||||||
}
|
|
||||||
try:
|
try:
|
||||||
for attr, val in _attrs_to_set.items():
|
for attr, val in _attrs_to_set.items():
|
||||||
setattr(_proxy_server_mod, attr, val)
|
setattr(_proxy_server_mod, attr, val)
|
||||||
|
|
@ -1548,9 +1487,7 @@ async def test_db_virtual_key_auth_sets_via_virtual_key_marker():
|
||||||
mock_proxy_logging_obj = MagicMock()
|
mock_proxy_logging_obj = MagicMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
|
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
|
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
|
||||||
AsyncMock()
|
|
||||||
)
|
|
||||||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
||||||
|
|
||||||
mock_prisma_client = MagicMock()
|
mock_prisma_client = MagicMock()
|
||||||
|
|
@ -1571,9 +1508,7 @@ async def test_db_virtual_key_auth_sets_via_virtual_key_marker():
|
||||||
"jwt_handler": None,
|
"jwt_handler": None,
|
||||||
"litellm_proxy_admin_name": "admin",
|
"litellm_proxy_admin_name": "admin",
|
||||||
}
|
}
|
||||||
_original_values = {
|
_original_values = {attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set}
|
||||||
attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set
|
|
||||||
}
|
|
||||||
try:
|
try:
|
||||||
for attr, val in _attrs_to_set.items():
|
for attr, val in _attrs_to_set.items():
|
||||||
setattr(_proxy_server_mod, attr, val)
|
setattr(_proxy_server_mod, attr, val)
|
||||||
|
|
@ -1993,10 +1928,7 @@ class TestJWTOAuth2Coexistence:
|
||||||
def test_is_jwt_detects_jwt_tokens(self):
|
def test_is_jwt_detects_jwt_tokens(self):
|
||||||
"""JWT tokens have 3 dot-separated parts."""
|
"""JWT tokens have 3 dot-separated parts."""
|
||||||
assert JWTHandler.is_jwt("header.payload.signature") is True
|
assert JWTHandler.is_jwt("header.payload.signature") is True
|
||||||
assert (
|
assert JWTHandler.is_jwt("eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.sig123") is True
|
||||||
JWTHandler.is_jwt("eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.sig123")
|
|
||||||
is True
|
|
||||||
)
|
|
||||||
|
|
||||||
def test_is_jwt_rejects_opaque_tokens(self):
|
def test_is_jwt_rejects_opaque_tokens(self):
|
||||||
"""Opaque OAuth2 tokens do not have 3 dot-separated parts."""
|
"""Opaque OAuth2 tokens do not have 3 dot-separated parts."""
|
||||||
|
|
@ -2105,10 +2037,7 @@ class TestJWTOAuth2Coexistence:
|
||||||
|
|
||||||
assert exc_info.value.type == ProxyErrorTypes.auth_error
|
assert exc_info.value.type == ProxyErrorTypes.auth_error
|
||||||
assert exc_info.value.code == "403"
|
assert exc_info.value.code == "403"
|
||||||
assert (
|
assert "Oauth2 token validation is only available for premium users" in exc_info.value.message
|
||||||
"Oauth2 token validation is only available for premium users"
|
|
||||||
in exc_info.value.message
|
|
||||||
)
|
|
||||||
mock_oauth2.assert_not_called()
|
mock_oauth2.assert_not_called()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
@ -2300,9 +2229,7 @@ class TestJWTOAuth2Coexistence:
|
||||||
assert mock_auto_register.call_args.kwargs["team_id"] == "validated-team"
|
assert mock_auto_register.call_args.kwargs["team_id"] == "validated-team"
|
||||||
assert mock_auto_register.call_args.kwargs["user_id"] == "validated-user"
|
assert mock_auto_register.call_args.kwargs["user_id"] == "validated-user"
|
||||||
assert mock_auto_register.call_args.kwargs["org_id"] == "validated-org"
|
assert mock_auto_register.call_args.kwargs["org_id"] == "validated-org"
|
||||||
assert (
|
assert mock_auto_register.call_args.kwargs["end_user_id"] == "validated-end-user"
|
||||||
mock_auto_register.call_args.kwargs["end_user_id"] == "validated-end-user"
|
|
||||||
)
|
|
||||||
assert result.org_id == "validated-org"
|
assert result.org_id == "validated-org"
|
||||||
assert result.user_email == "validated@example.com"
|
assert result.user_email == "validated@example.com"
|
||||||
|
|
||||||
|
|
@ -2380,10 +2307,7 @@ class TestJWTOAuth2Coexistence:
|
||||||
|
|
||||||
assert result.user_id == "mapped-user"
|
assert result.user_id == "mapped-user"
|
||||||
assert result.user_email == "mapped@example.com"
|
assert result.user_email == "mapped@example.com"
|
||||||
assert (
|
assert mock_get_user_object.call_args_list[0].kwargs["user_email"] == "mapped@example.com"
|
||||||
mock_get_user_object.call_args_list[0].kwargs["user_email"]
|
|
||||||
== "mapped@example.com"
|
|
||||||
)
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_mapped_virtual_key_does_not_backfill_mismatched_owner(self):
|
async def test_mapped_virtual_key_does_not_backfill_mismatched_owner(self):
|
||||||
|
|
@ -2459,8 +2383,7 @@ class TestJWTOAuth2Coexistence:
|
||||||
assert result.user_id == "other-owner"
|
assert result.user_id == "other-owner"
|
||||||
assert result.user_email is None
|
assert result.user_email is None
|
||||||
assert all(
|
assert all(
|
||||||
call.kwargs.get("user_email") != "principal@example.com"
|
call.kwargs.get("user_email") != "principal@example.com" for call in mock_get_user_object.call_args_list
|
||||||
for call in mock_get_user_object.call_args_list
|
|
||||||
)
|
)
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
@ -3265,9 +3188,7 @@ async def test_user_api_key_auth_builder_no_blocking_calls():
|
||||||
mock_proxy_logging_obj = MagicMock()
|
mock_proxy_logging_obj = MagicMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
|
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
|
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
|
||||||
AsyncMock()
|
|
||||||
)
|
|
||||||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
||||||
|
|
||||||
import litellm.proxy.proxy_server as _proxy_server_mod
|
import litellm.proxy.proxy_server as _proxy_server_mod
|
||||||
|
|
@ -3399,9 +3320,7 @@ async def test_team_metadata_refreshed_from_team_object_during_auth():
|
||||||
mock_proxy_logging_obj = MagicMock()
|
mock_proxy_logging_obj = MagicMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
|
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
||||||
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
|
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
|
||||||
AsyncMock()
|
|
||||||
)
|
|
||||||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
||||||
|
|
||||||
import litellm.proxy.proxy_server as _proxy_server_mod
|
import litellm.proxy.proxy_server as _proxy_server_mod
|
||||||
|
|
@ -3451,9 +3370,9 @@ async def test_team_metadata_refreshed_from_team_object_during_auth():
|
||||||
request_data={},
|
request_data={},
|
||||||
)
|
)
|
||||||
|
|
||||||
assert result.team_metadata == {
|
assert result.team_metadata == {"guardrails": ["test-guardrail-333"]}, (
|
||||||
"guardrails": ["test-guardrail-333"]
|
f"team_metadata was not updated from fresh team object. Got: {result.team_metadata}"
|
||||||
}, f"team_metadata was not updated from fresh team object. Got: {result.team_metadata}"
|
)
|
||||||
|
|
||||||
finally:
|
finally:
|
||||||
for k, v in _originals.items():
|
for k, v in _originals.items():
|
||||||
|
|
@ -3778,9 +3697,7 @@ async def test_auth_flow_fallback_team_object_permission_none_when_unreadable():
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
def _proxy_attrs_for_centralized_checks(
|
def _proxy_attrs_for_centralized_checks(user_custom_auth=None, flag=False, master_key="sk-test-master"):
|
||||||
user_custom_auth=None, flag=False, master_key="sk-test-master"
|
|
||||||
):
|
|
||||||
"""Build the minimal proxy_server module attributes that
|
"""Build the minimal proxy_server module attributes that
|
||||||
_run_centralized_common_checks reads.
|
_run_centralized_common_checks reads.
|
||||||
|
|
||||||
|
|
@ -3909,9 +3826,7 @@ async def test_centralized_common_checks_skipped_for_custom_auth_without_flag():
|
||||||
request = Request(scope={"type": "http"})
|
request = Request(scope={"type": "http"})
|
||||||
request._url = URL(url="/chat/completions")
|
request._url = URL(url="/chat/completions")
|
||||||
|
|
||||||
attrs = _proxy_attrs_for_centralized_checks(
|
attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=AsyncMock(), flag=False)
|
||||||
user_custom_auth=AsyncMock(), flag=False
|
|
||||||
)
|
|
||||||
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
|
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
|
||||||
try:
|
try:
|
||||||
for k, v in attrs.items():
|
for k, v in attrs.items():
|
||||||
|
|
@ -4316,9 +4231,7 @@ async def test_centralized_common_checks_reserves_request_end_user_budget():
|
||||||
"applied_adjustment": 0.0,
|
"applied_adjustment": 0.0,
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
assert counter_cache.in_memory_cache.get_cache(
|
assert counter_cache.in_memory_cache.get_cache(key="spend:end_user:alice") == pytest.approx(0.6)
|
||||||
key="spend:end_user:alice"
|
|
||||||
) == pytest.approx(0.6)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
@ -4333,9 +4246,7 @@ async def test_centralized_common_checks_short_circuits_when_master_key_unset():
|
||||||
|
|
||||||
from litellm.proxy._types import LitellmUserRoles
|
from litellm.proxy._types import LitellmUserRoles
|
||||||
|
|
||||||
token = UserAPIKeyAuth(
|
token = UserAPIKeyAuth(api_key="sk-test", user_id="u", user_role=LitellmUserRoles.INTERNAL_USER)
|
||||||
api_key="sk-test", user_id="u", user_role=LitellmUserRoles.INTERNAL_USER
|
|
||||||
)
|
|
||||||
request = Request(scope={"type": "http"})
|
request = Request(scope={"type": "http"})
|
||||||
request._url = URL(url="/get/config/callbacks")
|
request._url = URL(url="/get/config/callbacks")
|
||||||
|
|
||||||
|
|
@ -5136,9 +5047,7 @@ async def test_centralized_common_checks_user_http_exception_isolates_to_user_on
|
||||||
request._url = URL(url="/chat/completions")
|
request._url = URL(url="/chat/completions")
|
||||||
request._body = json.dumps({"user": "alice", "model": "gpt-4o"}).encode()
|
request._body = json.dumps({"user": "alice", "model": "gpt-4o"}).encode()
|
||||||
|
|
||||||
fetched_team = LiteLLM_TeamTableCachedObj(
|
fetched_team = LiteLLM_TeamTableCachedObj(team_id="t1", max_budget=20.0, models=["gpt-4o"])
|
||||||
team_id="t1", max_budget=20.0, models=["gpt-4o"]
|
|
||||||
)
|
|
||||||
fetched_end_user = LiteLLM_EndUserTable(user_id="alice", blocked=False, spend=1.0)
|
fetched_end_user = LiteLLM_EndUserTable(user_id="alice", blocked=False, spend=1.0)
|
||||||
fetched_project = LiteLLM_ProjectTableCachedObj(
|
fetched_project = LiteLLM_ProjectTableCachedObj(
|
||||||
project_id="proj-1",
|
project_id="proj-1",
|
||||||
|
|
@ -5433,9 +5342,7 @@ async def test_user_api_key_auth_sets_end_user_id_when_builder_skips_it():
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
request._url = URL(url="/chat/completions")
|
request._url = URL(url="/chat/completions")
|
||||||
request._body = json.dumps(
|
request._body = json.dumps({"model": "gpt-4o", "user": "alice@example.com"}).encode()
|
||||||
{"model": "gpt-4o", "user": "alice@example.com"}
|
|
||||||
).encode()
|
|
||||||
|
|
||||||
attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None)
|
attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None)
|
||||||
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
|
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
|
||||||
|
|
@ -5479,9 +5386,7 @@ async def test_user_api_key_auth_does_not_overwrite_end_user_id_set_by_builder()
|
||||||
|
|
||||||
import litellm.proxy.proxy_server as _proxy_server_mod
|
import litellm.proxy.proxy_server as _proxy_server_mod
|
||||||
|
|
||||||
builder_token = UserAPIKeyAuth(
|
builder_token = UserAPIKeyAuth(api_key="sk-test", user_id="u1", end_user_id="builder-resolved-id")
|
||||||
api_key="sk-test", user_id="u1", end_user_id="builder-resolved-id"
|
|
||||||
)
|
|
||||||
|
|
||||||
request = Request(
|
request = Request(
|
||||||
scope={
|
scope={
|
||||||
|
|
@ -5491,9 +5396,7 @@ async def test_user_api_key_auth_does_not_overwrite_end_user_id_set_by_builder()
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
request._url = URL(url="/chat/completions")
|
request._url = URL(url="/chat/completions")
|
||||||
request._body = json.dumps(
|
request._body = json.dumps({"model": "gpt-4o", "user": "different-id-from-body"}).encode()
|
||||||
{"model": "gpt-4o", "user": "different-id-from-body"}
|
|
||||||
).encode()
|
|
||||||
|
|
||||||
attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None)
|
attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None)
|
||||||
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
|
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
|
||||||
|
|
@ -5891,9 +5794,7 @@ def _mint_cli_session_token(monkeypatch, *, user_id="cli-admin"):
|
||||||
models=["gpt-3.5-turbo"],
|
models=["gpt-3.5-turbo"],
|
||||||
max_budget=100.0,
|
max_budget=100.0,
|
||||||
)
|
)
|
||||||
return ExperimentalUIJWTToken.get_cli_jwt_auth_token(
|
return ExperimentalUIJWTToken.get_cli_jwt_auth_token(user_info, team_id="cli-team", team_alias="cli-team-alias")
|
||||||
user_info, team_id="cli-team", team_alias="cli-team-alias"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
@ -5943,7 +5844,7 @@ async def test_random_non_sk_token_is_rejected(monkeypatch):
|
||||||
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
|
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
|
||||||
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
|
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
|
||||||
):
|
):
|
||||||
with pytest.raises(Exception, match='LiteLLM Virtual Key expected\\.') as exc_info:
|
with pytest.raises(Exception, match="LiteLLM Virtual Key expected\\.") as exc_info:
|
||||||
await user_api_key_auth(
|
await user_api_key_auth(
|
||||||
request=mock_request,
|
request=mock_request,
|
||||||
api_key="Bearer not-a-real-token",
|
api_key="Bearer not-a-real-token",
|
||||||
|
|
@ -6022,9 +5923,7 @@ async def test_non_admin_cli_session_token_reaches_production_auth_path(monkeypa
|
||||||
user_role=LitellmUserRoles.INTERNAL_USER.value,
|
user_role=LitellmUserRoles.INTERNAL_USER.value,
|
||||||
models=[],
|
models=[],
|
||||||
)
|
)
|
||||||
cli_token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(
|
cli_token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(user_info, team_id="team-abc", team_alias="my-team")
|
||||||
user_info, team_id="team-abc", team_alias="my-team"
|
|
||||||
)
|
|
||||||
|
|
||||||
import litellm.proxy.proxy_server as _proxy_server_mod
|
import litellm.proxy.proxy_server as _proxy_server_mod
|
||||||
from fastapi import Request
|
from fastapi import Request
|
||||||
|
|
@ -6145,7 +6044,7 @@ async def test_real_jwt_still_requires_license_when_jwt_auth_enabled(monkeypatch
|
||||||
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
|
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
|
||||||
patch("litellm.proxy.proxy_server.prisma_client", None),
|
patch("litellm.proxy.proxy_server.prisma_client", None),
|
||||||
):
|
):
|
||||||
with pytest.raises(Exception, match='JWT Auth is an enterprise only feature\\. You must be a') as exc_info:
|
with pytest.raises(Exception, match="JWT Auth is an enterprise only feature\\. You must be a") as exc_info:
|
||||||
await user_api_key_auth(
|
await user_api_key_auth(
|
||||||
request=mock_request,
|
request=mock_request,
|
||||||
api_key=f"Bearer {jwt_token}",
|
api_key=f"Bearer {jwt_token}",
|
||||||
|
|
@ -6184,13 +6083,9 @@ async def test_auth_does_not_rewrite_cached_key_object_back_into_cache():
|
||||||
metadata={"model_rpm_limit": {"gpt-5.4-mini": 3}},
|
metadata={"model_rpm_limit": {"gpt-5.4-mini": 3}},
|
||||||
last_refreshed_at=1000.0,
|
last_refreshed_at=1000.0,
|
||||||
)
|
)
|
||||||
await key_cache.async_set_cache(
|
await key_cache.async_set_cache(key=hashed_key, value=stale_token, model_type=UserAPIKeyAuth)
|
||||||
key=hashed_key, value=stale_token, model_type=UserAPIKeyAuth
|
|
||||||
)
|
|
||||||
|
|
||||||
fetch_from_db = AsyncMock(
|
fetch_from_db = AsyncMock(side_effect=AssertionError("cache-hit auth must not touch the DB"))
|
||||||
side_effect=AssertionError("cache-hit auth must not touch the DB")
|
|
||||||
)
|
|
||||||
|
|
||||||
proxy_logging_obj = MagicMock()
|
proxy_logging_obj = MagicMock()
|
||||||
proxy_logging_obj.internal_usage_cache = MagicMock()
|
proxy_logging_obj.internal_usage_cache = MagicMock()
|
||||||
|
|
@ -6237,9 +6132,7 @@ async def test_auth_does_not_rewrite_cached_key_object_back_into_cache():
|
||||||
assert result.token == hashed_key
|
assert result.token == hashed_key
|
||||||
fetch_from_db.assert_not_called()
|
fetch_from_db.assert_not_called()
|
||||||
|
|
||||||
cached_after = await key_cache.async_get_cache(
|
cached_after = await key_cache.async_get_cache(key=hashed_key, model_type=UserAPIKeyAuth)
|
||||||
key=hashed_key, model_type=UserAPIKeyAuth
|
|
||||||
)
|
|
||||||
assert cached_after is not None
|
assert cached_after is not None
|
||||||
assert cached_after.last_refreshed_at == 1000.0
|
assert cached_after.last_refreshed_at == 1000.0
|
||||||
assert cached_after.metadata == {"model_rpm_limit": {"gpt-5.4-mini": 3}}
|
assert cached_after.metadata == {"model_rpm_limit": {"gpt-5.4-mini": 3}}
|
||||||
|
|
@ -6352,9 +6245,7 @@ class TestCheckKeyModelBudgetWithFallback:
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_within_budget_does_not_reroute(self):
|
async def test_within_budget_does_not_reroute(self):
|
||||||
valid_token = UserAPIKeyAuth(
|
valid_token = UserAPIKeyAuth(token="test-key", budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]})
|
||||||
token="test-key", budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]}
|
|
||||||
)
|
|
||||||
limiter = AsyncMock()
|
limiter = AsyncMock()
|
||||||
limiter.is_key_within_model_budget.return_value = True
|
limiter.is_key_within_model_budget.return_value = True
|
||||||
request_data = {"model": "gpt-4o"}
|
request_data = {"model": "gpt-4o"}
|
||||||
|
|
@ -6379,9 +6270,7 @@ class TestCheckKeyModelBudgetWithFallback:
|
||||||
budget_fallbacks={"gpt-4o": ["gpt-4o-mini", "claude-haiku"]},
|
budget_fallbacks={"gpt-4o": ["gpt-4o-mini", "claude-haiku"]},
|
||||||
)
|
)
|
||||||
limiter = AsyncMock()
|
limiter = AsyncMock()
|
||||||
limiter.is_key_within_model_budget.side_effect = litellm.BudgetExceededError(
|
limiter.is_key_within_model_budget.side_effect = litellm.BudgetExceededError(current_cost=10, max_budget=5)
|
||||||
current_cost=10, max_budget=5
|
|
||||||
)
|
|
||||||
limiter.get_fallback_model_within_budget.return_value = "gpt-4o-mini"
|
limiter.get_fallback_model_within_budget.return_value = "gpt-4o-mini"
|
||||||
request_data = {"model": "gpt-4o"}
|
request_data = {"model": "gpt-4o"}
|
||||||
request = self._make_request()
|
request = self._make_request()
|
||||||
|
|
@ -6395,9 +6284,7 @@ class TestCheckKeyModelBudgetWithFallback:
|
||||||
)
|
)
|
||||||
|
|
||||||
assert request_data["model"] == "gpt-4o-mini"
|
assert request_data["model"] == "gpt-4o-mini"
|
||||||
limiter.get_fallback_model_within_budget.assert_awaited_once_with(
|
limiter.get_fallback_model_within_budget.assert_awaited_once_with(user_api_key_dict=valid_token, model="gpt-4o")
|
||||||
user_api_key_dict=valid_token, model="gpt-4o"
|
|
||||||
)
|
|
||||||
# the rerouted model must be visible to a later, separate
|
# the rerouted model must be visible to a later, separate
|
||||||
# `_read_request_body` call on the same `request` (route handlers
|
# `_read_request_body` call on the same `request` (route handlers
|
||||||
# re-parse the body from this cache instead of reusing the dict).
|
# re-parse the body from this cache instead of reusing the dict).
|
||||||
|
|
@ -6406,9 +6293,7 @@ class TestCheckKeyModelBudgetWithFallback:
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_raises_when_every_fallback_also_exceeded(self):
|
async def test_raises_when_every_fallback_also_exceeded(self):
|
||||||
valid_token = UserAPIKeyAuth(
|
valid_token = UserAPIKeyAuth(token="test-key", budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]})
|
||||||
token="test-key", budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]}
|
|
||||||
)
|
|
||||||
limiter = AsyncMock()
|
limiter = AsyncMock()
|
||||||
original_error = litellm.BudgetExceededError(current_cost=10, max_budget=5)
|
original_error = litellm.BudgetExceededError(current_cost=10, max_budget=5)
|
||||||
limiter.is_key_within_model_budget.side_effect = original_error
|
limiter.is_key_within_model_budget.side_effect = original_error
|
||||||
|
|
@ -6478,9 +6363,7 @@ class TestCheckKeyModelBudgetWithFallback:
|
||||||
budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]},
|
budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]},
|
||||||
)
|
)
|
||||||
limiter = AsyncMock()
|
limiter = AsyncMock()
|
||||||
limiter.is_key_within_model_budget.side_effect = litellm.BudgetExceededError(
|
limiter.is_key_within_model_budget.side_effect = litellm.BudgetExceededError(current_cost=10, max_budget=5)
|
||||||
current_cost=10, max_budget=5
|
|
||||||
)
|
|
||||||
limiter.get_fallback_model_within_budget.return_value = "gpt-4o-mini"
|
limiter.get_fallback_model_within_budget.return_value = "gpt-4o-mini"
|
||||||
request_data = {"model": "gpt-4o"}
|
request_data = {"model": "gpt-4o"}
|
||||||
request = self._make_request()
|
request = self._make_request()
|
||||||
|
|
@ -6548,9 +6431,7 @@ class TestCheckKeyModelBudgetWithFallback:
|
||||||
budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]},
|
budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]},
|
||||||
)
|
)
|
||||||
limiter = AsyncMock()
|
limiter = AsyncMock()
|
||||||
limiter.is_key_within_model_budget.side_effect = litellm.BudgetExceededError(
|
limiter.is_key_within_model_budget.side_effect = litellm.BudgetExceededError(current_cost=10, max_budget=5)
|
||||||
current_cost=10, max_budget=5
|
|
||||||
)
|
|
||||||
limiter.get_fallback_model_within_budget.return_value = "gpt-4o-mini"
|
limiter.get_fallback_model_within_budget.return_value = "gpt-4o-mini"
|
||||||
request_data = {"model": "gpt-4o"}
|
request_data = {"model": "gpt-4o"}
|
||||||
request = self._make_request()
|
request = self._make_request()
|
||||||
|
|
@ -6630,9 +6511,7 @@ async def test_global_proxy_spend_reads_resettable_proxy_budget_row():
|
||||||
)
|
)
|
||||||
|
|
||||||
assert result == 42.5
|
assert result == 42.5
|
||||||
prisma_client.db.litellm_usertable.find_unique.assert_awaited_once_with(
|
prisma_client.db.litellm_usertable.find_unique.assert_awaited_once_with(where={"user_id": "litellm-proxy-budget"})
|
||||||
where={"user_id": "litellm-proxy-budget"}
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
@ -6798,9 +6677,7 @@ async def test_jwt_shaped_key_error_names_enable_jwt_auth_when_disabled():
|
||||||
Prometheus invalid-key filter and the admin UI both substring-match it.
|
Prometheus invalid-key filter and the admin UI both substring-match it.
|
||||||
Keys that are not JWT-shaped must not pick up the hint.
|
Keys that are not JWT-shaped must not pick up the hint.
|
||||||
"""
|
"""
|
||||||
jwt_error = await _proxy_exception_for_key(
|
jwt_error = await _proxy_exception_for_key("eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJzdmMtMSJ9.c2lnbmF0dXJl", {}, True)
|
||||||
"eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJzdmMtMSJ9.c2lnbmF0dXJl", {}, True
|
|
||||||
)
|
|
||||||
|
|
||||||
assert jwt_error.code == "401"
|
assert jwt_error.code == "401"
|
||||||
assert "enable_jwt_auth" in jwt_error.message
|
assert "enable_jwt_auth" in jwt_error.message
|
||||||
|
|
@ -6810,9 +6687,7 @@ async def test_jwt_shaped_key_error_names_enable_jwt_auth_when_disabled():
|
||||||
assert "is a JWT" not in jwt_error.message
|
assert "is a JWT" not in jwt_error.message
|
||||||
|
|
||||||
opaque_error = await _proxy_exception_for_key("not-a-jwt-at-all", {}, True)
|
opaque_error = await _proxy_exception_for_key("not-a-jwt-at-all", {}, True)
|
||||||
two_segment_error = await _proxy_exception_for_key(
|
two_segment_error = await _proxy_exception_for_key("eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJzdmMtMSJ9", {}, True)
|
||||||
"eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJzdmMtMSJ9", {}, True
|
|
||||||
)
|
|
||||||
|
|
||||||
assert "enable_jwt_auth" not in opaque_error.message
|
assert "enable_jwt_auth" not in opaque_error.message
|
||||||
assert "enable_jwt_auth" not in two_segment_error.message
|
assert "enable_jwt_auth" not in two_segment_error.message
|
||||||
|
|
@ -6841,9 +6716,7 @@ class TestLitellmReceivedAtStamping:
|
||||||
on OTEL being configured to see a true request-arrival timestamp."""
|
on OTEL being configured to see a true request-arrival timestamp."""
|
||||||
|
|
||||||
def test_stamped_even_when_otel_is_not_configured(self, monkeypatch):
|
def test_stamped_even_when_otel_is_not_configured(self, monkeypatch):
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr("litellm.proxy.proxy_server.open_telemetry_logger", None)
|
||||||
"litellm.proxy.proxy_server.open_telemetry_logger", None
|
|
||||||
)
|
|
||||||
request = MagicMock()
|
request = MagicMock()
|
||||||
request.state = SimpleNamespace()
|
request.state = SimpleNamespace()
|
||||||
|
|
||||||
|
|
@ -6872,3 +6745,119 @@ class TestLitellmReceivedAtStamping:
|
||||||
|
|
||||||
assert result == earlier
|
assert result == earlier
|
||||||
assert request.state.litellm_received_at == earlier
|
assert request.state.litellm_received_at == earlier
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.parametrize("is_proxy_admin", [False, True], ids=["standard-return", "proxy-admin-return"])
|
||||||
|
async def test_jwt_builder_returns_every_team_grant_the_key_path_gets(is_proxy_admin):
|
||||||
|
"""LIT-5858: the team-based JWT path hand-built ``UserAPIKeyAuth`` from a short list of team fields, so the
|
||||||
|
team's model aliases (and on the admin return, its object permission) never reached the token and alias
|
||||||
|
requests 403'd. Both returns now go through ``team_grants``; pin the fields that used to be dropped."""
|
||||||
|
import litellm.proxy.proxy_server as _proxy_server_mod
|
||||||
|
from fastapi import Request
|
||||||
|
from starlette.datastructures import URL
|
||||||
|
|
||||||
|
from litellm.models.team import LiteLLM_ModelTable
|
||||||
|
from litellm.proxy._types import (
|
||||||
|
LiteLLM_ObjectPermissionTable,
|
||||||
|
LiteLLM_TeamMembership,
|
||||||
|
LiteLLM_TeamTable,
|
||||||
|
Member,
|
||||||
|
)
|
||||||
|
|
||||||
|
class _AcceptEveryJwt(JWTHandler):
|
||||||
|
def is_jwt(self, token: str) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
jwt_handler = _AcceptEveryJwt()
|
||||||
|
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth()
|
||||||
|
|
||||||
|
team = LiteLLM_TeamTable(
|
||||||
|
team_id="team-jwt-aliases",
|
||||||
|
team_alias="jwt-aliases",
|
||||||
|
models=["gpt-4o"],
|
||||||
|
max_budget=40.0,
|
||||||
|
spend=4.0,
|
||||||
|
blocked=False,
|
||||||
|
metadata={"tier": "gold"},
|
||||||
|
litellm_model_table=LiteLLM_ModelTable(
|
||||||
|
model_aliases='{"fast": "gpt-4o"}', created_by="admin", updated_by="admin"
|
||||||
|
),
|
||||||
|
object_permission_id="op-jwt",
|
||||||
|
object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-jwt", mcp_servers=["mcp-a"]),
|
||||||
|
members_with_roles=[Member(user_id="jwt-user", role="admin")],
|
||||||
|
)
|
||||||
|
membership = LiteLLM_TeamMembership(user_id="jwt-user", team_id="team-jwt-aliases", spend=1.5)
|
||||||
|
builder_result = {
|
||||||
|
"is_proxy_admin": is_proxy_admin,
|
||||||
|
"team_object": team,
|
||||||
|
"user_object": None,
|
||||||
|
"end_user_object": None,
|
||||||
|
"org_object": None,
|
||||||
|
"token": "jwt",
|
||||||
|
"team_id": "team-jwt-aliases",
|
||||||
|
"user_id": "jwt-user",
|
||||||
|
"user_email": "jwt-user@example.com",
|
||||||
|
"end_user_id": None,
|
||||||
|
"org_id": None,
|
||||||
|
"team_membership": membership,
|
||||||
|
"jwt_claims": {"sub": "jwt-user"},
|
||||||
|
}
|
||||||
|
|
||||||
|
mock_proxy_logging_obj = MagicMock()
|
||||||
|
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
|
||||||
|
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
||||||
|
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
||||||
|
attrs = {
|
||||||
|
"prisma_client": MagicMock(),
|
||||||
|
"user_api_key_cache": DualCache(),
|
||||||
|
"proxy_logging_obj": mock_proxy_logging_obj,
|
||||||
|
"master_key": "sk-master-key",
|
||||||
|
"general_settings": {"enable_jwt_auth": True},
|
||||||
|
"llm_model_list": [],
|
||||||
|
"llm_router": None,
|
||||||
|
"open_telemetry_logger": None,
|
||||||
|
"model_max_budget_limiter": MagicMock(),
|
||||||
|
"user_custom_auth": None,
|
||||||
|
"jwt_handler": jwt_handler,
|
||||||
|
"premium_user": True,
|
||||||
|
"litellm_proxy_admin_name": "admin",
|
||||||
|
}
|
||||||
|
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
|
||||||
|
try:
|
||||||
|
for k, v in attrs.items():
|
||||||
|
setattr(_proxy_server_mod, k, v)
|
||||||
|
request = Request(scope={"type": "http", "headers": [], "method": "POST"})
|
||||||
|
request._url = URL(url="/chat/completions")
|
||||||
|
with patch( # test-quality-ok: auth_builder is the claim-resolution seam; the regression is how its result is projected onto the token
|
||||||
|
"litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder",
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
return_value=builder_result,
|
||||||
|
):
|
||||||
|
token = await _user_api_key_auth_builder(
|
||||||
|
request=request,
|
||||||
|
api_key="Bearer header.payload.signature",
|
||||||
|
azure_api_key_header="",
|
||||||
|
anthropic_api_key_header=None,
|
||||||
|
google_ai_studio_api_key_header=None,
|
||||||
|
azure_apim_header=None,
|
||||||
|
request_data={},
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
for k, v in originals.items():
|
||||||
|
setattr(_proxy_server_mod, k, v)
|
||||||
|
|
||||||
|
assert token.team_id == "team-jwt-aliases"
|
||||||
|
assert token.user_role == (LitellmUserRoles.PROXY_ADMIN if is_proxy_admin else LitellmUserRoles.INTERNAL_USER)
|
||||||
|
assert token.team_model_aliases == {"fast": "gpt-4o"}
|
||||||
|
assert token.team_object_permission is not None
|
||||||
|
assert token.team_object_permission.mcp_servers == ["mcp-a"]
|
||||||
|
assert token.team_object_permission_id == "op-jwt"
|
||||||
|
assert token.team_alias == "jwt-aliases"
|
||||||
|
assert token.team_models == ["gpt-4o"]
|
||||||
|
assert token.team_max_budget == 40.0
|
||||||
|
assert token.team_spend == 4.0
|
||||||
|
assert token.team_metadata == {"tier": "gold"}
|
||||||
|
assert token.team_member == Member(user_id="jwt-user", role="admin")
|
||||||
|
assert token.team_member_spend == 1.5
|
||||||
|
assert token.jwt_claims == {"sub": "jwt-user"}
|
||||||
|
|
|
||||||
|
|
@ -31,6 +31,10 @@ from litellm.router import Router
|
||||||
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo, updateDeployment, updateLiteLLMParams
|
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo, updateDeployment, updateLiteLLMParams
|
||||||
|
|
||||||
|
|
||||||
|
async def _passthrough_row(update_data):
|
||||||
|
return update_data
|
||||||
|
|
||||||
|
|
||||||
class MockPrismaClient:
|
class MockPrismaClient:
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
|
|
@ -1027,7 +1031,7 @@ class TestTeamModelSiblingRouting:
|
||||||
team_id = "team_no_alias"
|
team_id = "team_no_alias"
|
||||||
public_name = "gpt-4.1-mini"
|
public_name = "gpt-4.1-mini"
|
||||||
|
|
||||||
async def mock_add_model_to_db(model_params, user_api_key_dict, prisma_client):
|
async def mock_add_model_to_db(model_params, user_api_key_dict, prisma_client, slot=None):
|
||||||
return MagicMock(model_id=str(uuid.uuid4()))
|
return MagicMock(model_id=str(uuid.uuid4()))
|
||||||
|
|
||||||
mock_team_model_add = AsyncMock()
|
mock_team_model_add = AsyncMock()
|
||||||
|
|
@ -1207,6 +1211,7 @@ class TestTeamModelUpdate:
|
||||||
patch_data=patch_data,
|
patch_data=patch_data,
|
||||||
user_api_key_dict=user_api_key_dict,
|
user_api_key_dict=user_api_key_dict,
|
||||||
prisma_client=prisma_client, # type: ignore
|
prisma_client=prisma_client, # type: ignore
|
||||||
|
write_row=_passthrough_row,
|
||||||
)
|
)
|
||||||
|
|
||||||
assert result.get("model_name", "").startswith("model_name_test_team_123_")
|
assert result.get("model_name", "").startswith("model_name_test_team_123_")
|
||||||
|
|
@ -1437,6 +1442,7 @@ class TestTeamModelUpdate:
|
||||||
patch_data=patch_data,
|
patch_data=patch_data,
|
||||||
user_api_key_dict=user_api_key_dict,
|
user_api_key_dict=user_api_key_dict,
|
||||||
prisma_client=prisma_client, # type: ignore
|
prisma_client=prisma_client, # type: ignore
|
||||||
|
write_row=_passthrough_row,
|
||||||
)
|
)
|
||||||
assert "403" in str(exc_info.value)
|
assert "403" in str(exc_info.value)
|
||||||
|
|
||||||
|
|
@ -1697,6 +1703,7 @@ class TestTeamModelUpdate:
|
||||||
patch_data=patch_data,
|
patch_data=patch_data,
|
||||||
user_api_key_dict=user_api_key_dict,
|
user_api_key_dict=user_api_key_dict,
|
||||||
prisma_client=prisma_client, # type: ignore
|
prisma_client=prisma_client, # type: ignore
|
||||||
|
write_row=_passthrough_row,
|
||||||
)
|
)
|
||||||
|
|
||||||
# team ACL must not be touched on a no-op edit
|
# team ACL must not be touched on a no-op edit
|
||||||
|
|
@ -4480,3 +4487,122 @@ class TestTeamMemberAutoRouterWrites:
|
||||||
assert saved == expected
|
assert saved == expected
|
||||||
assert row.litellm_params["complexity_router_config"] == stored_config
|
assert row.litellm_params["complexity_router_config"] == stored_config
|
||||||
assert request.litellm_params.complexity_router_config == config
|
assert request.litellm_params.complexity_router_config == config
|
||||||
|
|
||||||
|
async def test_admin_router_changes_release_member_scope(self, endpoint: str, change: str) -> None:
|
||||||
|
from litellm.proxy.management_endpoints.model_management_endpoints import patch_model, update_model
|
||||||
|
|
||||||
|
original: Final = self._row()
|
||||||
|
row: Final = original.model_copy(update={"model_info": {**original.model_info, "member_auto_router": True}})
|
||||||
|
database: Final = self._database(self._team(), row)
|
||||||
|
params: Final = {
|
||||||
|
"config": {"complexity_router_config": {"tiers": {"SIMPLE": "allowed"}, "session_affinity": True}},
|
||||||
|
"strategy": {"model": "auto_router/quality_router", "quality_router_default_model": "allowed"},
|
||||||
|
"unrelated": {"model": "auto_router/complexity_router", "max_tokens": 100},
|
||||||
|
}
|
||||||
|
request: Final = updateDeployment(
|
||||||
|
litellm_params=updateLiteLLMParams.model_validate(params[change]),
|
||||||
|
model_info=ModelInfo(id=row.model_id) if endpoint == "legacy" or change == "unrelated" else None,
|
||||||
|
)
|
||||||
|
with self._environment(database, row):
|
||||||
|
actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||||
|
if endpoint == "patch":
|
||||||
|
await patch_model(row.model_id, request, actor)
|
||||||
|
else:
|
||||||
|
await update_model(request, actor)
|
||||||
|
written: Final = database.db.litellm_proxymodeltable.update.await_args.kwargs["data"]
|
||||||
|
saved_info: Final = json.loads(written["model_info"]) if "model_info" in written else row.model_info
|
||||||
|
assert saved_info["member_auto_router"] is (change == "unrelated")
|
||||||
|
assert saved_info["team_id"] == "member-team"
|
||||||
|
assert saved_info["access_groups"] == ["retained-admin-group"]
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.parametrize("endpoint", ["patch", "legacy"])
|
||||||
|
@pytest.mark.parametrize("access", ["owner", "peer", "limited-key"])
|
||||||
|
async def test_both_update_entries_enforce_creator_and_stamp_member_scope(self, endpoint: str, access: str) -> None:
|
||||||
|
from fastapi import HTTPException
|
||||||
|
|
||||||
|
from litellm.proxy._types import ProxyException
|
||||||
|
from litellm.proxy.management_endpoints.model_management_endpoints import patch_model, update_model
|
||||||
|
|
||||||
|
row: Final = self._row()
|
||||||
|
database: Final = self._database(self._team(), row)
|
||||||
|
request: Final = updateDeployment(
|
||||||
|
litellm_params=updateLiteLLMParams(
|
||||||
|
complexity_router_config={"tiers": {"SIMPLE": "allowed"}, "session_affinity": True}
|
||||||
|
),
|
||||||
|
model_info=ModelInfo(id=row.model_id, team_id="member-team"),
|
||||||
|
)
|
||||||
|
actor: Final = UserAPIKeyAuth(
|
||||||
|
user_id="peer" if access == "peer" else "owner",
|
||||||
|
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||||
|
models=["personal-router"] if access == "limited-key" else ["allowed"],
|
||||||
|
config={"timeout": 60},
|
||||||
|
)
|
||||||
|
with self._environment(database, row):
|
||||||
|
operation: Final = (
|
||||||
|
patch_model(row.model_id, request, actor) if endpoint == "patch" else update_model(request, actor)
|
||||||
|
)
|
||||||
|
if access != "owner":
|
||||||
|
with pytest.raises((HTTPException, ProxyException)):
|
||||||
|
await operation
|
||||||
|
database.transaction.litellm_proxymodeltable.update.assert_not_awaited()
|
||||||
|
return
|
||||||
|
await operation
|
||||||
|
written: Final = database.transaction.litellm_proxymodeltable.update.await_args.kwargs["data"]
|
||||||
|
saved_info: Final = json.loads(written["model_info"])
|
||||||
|
assert saved_info["member_auto_router"] is True
|
||||||
|
assert saved_info["team_id"] == "member-team"
|
||||||
|
assert saved_info["access_groups"] == ["retained-admin-group"]
|
||||||
|
assert "created_by" not in written
|
||||||
|
assert json.loads(written["litellm_params"])["complexity_router_config"]["session_affinity"] is True
|
||||||
|
assert written.get("model_name", row.model_name) == row.model_name
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.parametrize("access", ["allowed", "opt-out", "limited-key"])
|
||||||
|
async def test_create_entry_requires_opt_in_and_appends_only_its_router(self, access: str) -> None:
|
||||||
|
from litellm.proxy._types import ProxyException
|
||||||
|
from litellm.proxy.management_endpoints.model_management_endpoints import add_new_model
|
||||||
|
|
||||||
|
row: Final = self._row()
|
||||||
|
database: Final = self._database(self._team(enabled=access != "opt-out"), row)
|
||||||
|
actor: Final = UserAPIKeyAuth(
|
||||||
|
user_id="owner",
|
||||||
|
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||||
|
models=["personal-router"] if access == "limited-key" else ["allowed"],
|
||||||
|
config={"timeout": 60},
|
||||||
|
)
|
||||||
|
deployment: Final = Deployment(
|
||||||
|
model_name="new-personal-router",
|
||||||
|
litellm_params=LiteLLM_Params(
|
||||||
|
model="auto_router/complexity_router", complexity_router_config={"tiers": {"SIMPLE": "allowed"}}
|
||||||
|
),
|
||||||
|
model_info=ModelInfo(id=row.model_id, team_id="member-team"),
|
||||||
|
)
|
||||||
|
with (
|
||||||
|
self._environment(database, row),
|
||||||
|
patch(
|
||||||
|
"litellm.proxy.proxy_server.proxy_config.add_deployment",
|
||||||
|
new=AsyncMock(
|
||||||
|
return_value=ReconcileOutcome( # test-quality-ok: [TQ008] model reload I/O boundary
|
||||||
|
still_desired=frozenset((row.model_id, "allowed-id")),
|
||||||
|
live_after=frozenset((row.model_id, "allowed-id")),
|
||||||
|
)
|
||||||
|
),
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"litellm.proxy.management_endpoints.model_management_endpoints.team_model_add", new=AsyncMock()
|
||||||
|
) as appended, # test-quality-ok: [TQ008] persistence boundary; the appended scope is asserted
|
||||||
|
):
|
||||||
|
if access != "allowed":
|
||||||
|
with pytest.raises(ProxyException) as denied:
|
||||||
|
await add_new_model(deployment, actor)
|
||||||
|
assert denied.value.code == "403"
|
||||||
|
database.transaction.litellm_proxymodeltable.create.assert_not_awaited()
|
||||||
|
appended.assert_not_awaited()
|
||||||
|
return
|
||||||
|
await add_new_model(deployment, actor)
|
||||||
|
written: Final = database.transaction.litellm_proxymodeltable.create.await_args.kwargs["data"]
|
||||||
|
assert written["created_by"] == "owner"
|
||||||
|
assert json.loads(written["model_info"])["member_auto_router"] is True
|
||||||
|
assert appended.await_args.kwargs["data"].models == ["new-personal-router"]
|
||||||
|
assert appended.await_args.kwargs["data"].team_id == "member-team"
|
||||||
|
|
|
||||||
File diff suppressed because it is too large
Load diff
21
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
21
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -28154,6 +28154,8 @@ export interface components {
|
||||||
team_id: string;
|
team_id: string;
|
||||||
/** Team Member Permissions */
|
/** Team Member Permissions */
|
||||||
team_member_permissions?: string[] | null;
|
team_member_permissions?: string[] | null;
|
||||||
|
/** Tpd Limit */
|
||||||
|
tpd_limit?: number | null;
|
||||||
/** Tpm Limit */
|
/** Tpm Limit */
|
||||||
tpm_limit?: number | null;
|
tpm_limit?: number | null;
|
||||||
/** Updated At */
|
/** Updated At */
|
||||||
|
|
@ -29477,6 +29479,8 @@ export interface components {
|
||||||
team_id: string;
|
team_id: string;
|
||||||
/** Team Member Permissions */
|
/** Team Member Permissions */
|
||||||
team_member_permissions?: string[] | null;
|
team_member_permissions?: string[] | null;
|
||||||
|
/** Tpd Limit */
|
||||||
|
tpd_limit?: number | null;
|
||||||
/** Tpm Limit */
|
/** Tpm Limit */
|
||||||
tpm_limit?: number | null;
|
tpm_limit?: number | null;
|
||||||
/** Updated At */
|
/** Updated At */
|
||||||
|
|
@ -31959,6 +31963,8 @@ export interface components {
|
||||||
team_member_rpm_limit?: number | null;
|
team_member_rpm_limit?: number | null;
|
||||||
/** Team Member Tpm Limit */
|
/** Team Member Tpm Limit */
|
||||||
team_member_tpm_limit?: number | null;
|
team_member_tpm_limit?: number | null;
|
||||||
|
/** Tpd Limit */
|
||||||
|
tpd_limit?: number | null;
|
||||||
/** Tpm Limit */
|
/** Tpm Limit */
|
||||||
tpm_limit?: number | null;
|
tpm_limit?: number | null;
|
||||||
/** Tpm Limit Type */
|
/** Tpm Limit Type */
|
||||||
|
|
@ -35923,6 +35929,8 @@ export interface components {
|
||||||
team_id: string;
|
team_id: string;
|
||||||
/** Team Member Permissions */
|
/** Team Member Permissions */
|
||||||
team_member_permissions?: string[] | null;
|
team_member_permissions?: string[] | null;
|
||||||
|
/** Tpd Limit */
|
||||||
|
tpd_limit?: number | null;
|
||||||
/** Tpm Limit */
|
/** Tpm Limit */
|
||||||
tpm_limit?: number | null;
|
tpm_limit?: number | null;
|
||||||
/** Updated At */
|
/** Updated At */
|
||||||
|
|
@ -36063,6 +36071,8 @@ export interface components {
|
||||||
team_id: string;
|
team_id: string;
|
||||||
/** Team Member Permissions */
|
/** Team Member Permissions */
|
||||||
team_member_permissions?: string[] | null;
|
team_member_permissions?: string[] | null;
|
||||||
|
/** Tpd Limit */
|
||||||
|
tpd_limit?: number | null;
|
||||||
/** Tpm Limit */
|
/** Tpm Limit */
|
||||||
tpm_limit?: number | null;
|
tpm_limit?: number | null;
|
||||||
/** Updated At */
|
/** Updated At */
|
||||||
|
|
@ -38034,6 +38044,10 @@ export interface components {
|
||||||
team_model_aliases?: {
|
team_model_aliases?: {
|
||||||
[key: string]: unknown;
|
[key: string]: unknown;
|
||||||
} | null;
|
} | null;
|
||||||
|
/** Team Model Max Budget */
|
||||||
|
team_model_max_budget?: {
|
||||||
|
[key: string]: unknown;
|
||||||
|
} | null;
|
||||||
/**
|
/**
|
||||||
* Team Models
|
* Team Models
|
||||||
* @default []
|
* @default []
|
||||||
|
|
@ -38048,6 +38062,8 @@ export interface components {
|
||||||
team_soft_budget?: number | null;
|
team_soft_budget?: number | null;
|
||||||
/** Team Spend */
|
/** Team Spend */
|
||||||
team_spend?: number | null;
|
team_spend?: number | null;
|
||||||
|
/** Team Tpd Limit */
|
||||||
|
team_tpd_limit?: number | null;
|
||||||
/** Team Tpm Limit */
|
/** Team Tpm Limit */
|
||||||
team_tpm_limit?: number | null;
|
team_tpm_limit?: number | null;
|
||||||
/** Token */
|
/** Token */
|
||||||
|
|
@ -38537,6 +38553,11 @@ export interface components {
|
||||||
input_cost_per_character?: number | null;
|
input_cost_per_character?: number | null;
|
||||||
/** Input Cost Per Token */
|
/** Input Cost Per Token */
|
||||||
input_cost_per_token?: number | null;
|
input_cost_per_token?: number | null;
|
||||||
|
/**
|
||||||
|
* Member Auto Router
|
||||||
|
* @default false
|
||||||
|
*/
|
||||||
|
member_auto_router: boolean;
|
||||||
/** Output Cost Per Character */
|
/** Output Cost Per Character */
|
||||||
output_cost_per_character?: number | null;
|
output_cost_per_character?: number | null;
|
||||||
/** Output Cost Per Token */
|
/** Output Cost Per Token */
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue