mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
feat(proxy): let personal keys inherit model access from the owner's current teams
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
04fa760bf2
commit
bcc41d60bc
6 changed files with 639 additions and 2 deletions
|
|
@ -474,6 +474,8 @@ max_end_user_budget_id: Optional[str] = None
|
|||
# pass through unchanged.
|
||||
validate_end_user_id_in_db: bool = False
|
||||
block_requests_for_models_without_pricing: bool = False
|
||||
personal_key_model_access_from_teams: bool = False
|
||||
personal_key_multi_team_access: Literal["union", "intersection"] = "union"
|
||||
disable_end_user_cost_tracking: Optional[bool] = None
|
||||
disable_end_user_cost_tracking_prometheus_only: Optional[bool] = None
|
||||
enable_end_user_cost_tracking_prometheus_only: Optional[bool] = None
|
||||
|
|
|
|||
|
|
@ -1893,6 +1893,8 @@ LITELLM_SETTINGS_SAFE_DB_OVERRIDES: Final = [
|
|||
"openai_system_messages_first",
|
||||
"max_ui_session_budget",
|
||||
"budget_rollover",
|
||||
"personal_key_model_access_from_teams",
|
||||
"personal_key_multi_team_access",
|
||||
"mcp_tool_search",
|
||||
"turn_off_message_logging",
|
||||
"datadog_params",
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ from litellm.constants import (
|
|||
DEFAULT_MAX_RECURSE_DEPTH,
|
||||
EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE,
|
||||
END_USER_RESTRICTED_REGISTRY_MAX_SIZE,
|
||||
LITELLM_PROXY_MASTER_KEY_ALIAS,
|
||||
MODEL_ACCESS_GROUP_REGISTRY_MAX_SIZE,
|
||||
REGISTRY_ERROR_NEGATIVE_CACHE_TTL,
|
||||
TAG_REGISTRY_MAX_SIZE,
|
||||
|
|
@ -1077,6 +1078,18 @@ async def common_checks(
|
|||
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
|
||||
)
|
||||
|
||||
if _model and valid_token is not None and personal_key_team_model_access_applies(valid_token):
|
||||
with tracer.trace("litellm.proxy.auth.common_checks.can_personal_key_call_model_via_teams"):
|
||||
await can_personal_key_call_model_via_teams(
|
||||
model=_model,
|
||||
user_object=user_object,
|
||||
valid_token=valid_token,
|
||||
llm_router=llm_router,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
# 1.1 - 2.2 - 3.0.2 - 3.0.3: Project checks (blocked, model access, budget)
|
||||
with tracer.trace("litellm.proxy.auth.common_checks.run_project_checks"):
|
||||
await _run_project_checks(
|
||||
|
|
@ -5017,6 +5030,25 @@ async def can_key_call_resolved_model(
|
|||
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
|
||||
)
|
||||
|
||||
if personal_key_team_model_access_applies(valid_token) and valid_token.user_id is not None:
|
||||
owner: Final = await get_user_object(
|
||||
user_id=valid_token.user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
parent_otel_span=valid_token.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
await can_personal_key_call_model_via_teams(
|
||||
model=model,
|
||||
user_object=owner,
|
||||
valid_token=valid_token,
|
||||
llm_router=llm_router,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
if valid_token.project_id is not None:
|
||||
project_object: Final = await get_project_object(
|
||||
project_id=valid_token.project_id,
|
||||
|
|
@ -5230,6 +5262,261 @@ async def can_user_call_model(
|
|||
)
|
||||
|
||||
|
||||
def personal_key_team_model_access_applies(valid_token: UserAPIKeyAuth | None) -> bool:
|
||||
return (
|
||||
bool(litellm.personal_key_model_access_from_teams)
|
||||
and valid_token is not None
|
||||
and valid_token.team_id is None
|
||||
and valid_token.user_id is not None
|
||||
and valid_token.api_key != LITELLM_PROXY_MASTER_KEY_ALIAS
|
||||
and valid_token.user_role != LitellmUserRoles.PROXY_ADMIN
|
||||
)
|
||||
|
||||
|
||||
async def _load_personal_key_team(
|
||||
team_id: str,
|
||||
valid_token: UserAPIKeyAuth,
|
||||
prisma_client: PrismaClient | None,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
) -> LiteLLM_TeamTableCachedObj | None:
|
||||
try:
|
||||
team_object: Final = await get_team_object(
|
||||
team_id=team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=valid_token.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # fail closed: a team that cannot be loaded grants nothing
|
||||
verbose_proxy_logger.warning(
|
||||
"Personal key team model access: team=%s lookup failed, granting nothing: %s", team_id, e
|
||||
)
|
||||
return None
|
||||
return None if team_object.blocked is True else team_object
|
||||
|
||||
|
||||
async def _load_personal_key_teams(
|
||||
user_object: LiteLLM_UserTable,
|
||||
valid_token: UserAPIKeyAuth,
|
||||
prisma_client: PrismaClient | None,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
) -> tuple[LiteLLM_TeamTableCachedObj | None, ...]:
|
||||
return tuple(
|
||||
await asyncio.gather(
|
||||
*(
|
||||
_load_personal_key_team(
|
||||
team_id=team_id,
|
||||
valid_token=valid_token,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
for team_id in dict.fromkeys(user_object.teams)
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
async def _team_grants_personal_key_model(
|
||||
model: str,
|
||||
team_object: LiteLLM_TeamTableCachedObj,
|
||||
valid_token: UserAPIKeyAuth,
|
||||
llm_router: Router | None,
|
||||
prisma_client: PrismaClient | None,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
) -> bool:
|
||||
key_model_aliases: Final = key_model_aliases_for_auth_check(valid_token)
|
||||
try:
|
||||
await can_team_access_model(
|
||||
model=model,
|
||||
team_object=team_object,
|
||||
llm_router=llm_router,
|
||||
key_model_aliases=key_model_aliases,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
await _check_team_member_model_access(
|
||||
model=model,
|
||||
team_object=team_object,
|
||||
valid_token=valid_token,
|
||||
llm_router=llm_router,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
key_model_aliases=key_model_aliases,
|
||||
)
|
||||
except ProxyException:
|
||||
return False
|
||||
except Exception as e: # noqa: BLE001 # fail closed: a team check that errors grants nothing
|
||||
verbose_proxy_logger.warning(
|
||||
"Personal key team model access: team=%s check failed, granting nothing: %s", team_object.team_id, e
|
||||
)
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
async def _loaded_teams_allow_model(
|
||||
model: str,
|
||||
teams: tuple[LiteLLM_TeamTableCachedObj | None, ...],
|
||||
valid_token: UserAPIKeyAuth,
|
||||
llm_router: Router | None,
|
||||
prisma_client: PrismaClient | None,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
) -> bool:
|
||||
is_union: Final = litellm.personal_key_multi_team_access == "union"
|
||||
loaded: Final = tuple(team for team in teams if team is not None)
|
||||
if not loaded or (not is_union and len(loaded) != len(teams)):
|
||||
return False
|
||||
grants: Final = await asyncio.gather(
|
||||
*(
|
||||
_team_grants_personal_key_model(
|
||||
model=model,
|
||||
team_object=team,
|
||||
valid_token=valid_token,
|
||||
llm_router=llm_router,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
for team in loaded
|
||||
)
|
||||
)
|
||||
return any(grants) if is_union else all(grants)
|
||||
|
||||
|
||||
async def personal_key_teams_allow_model(
|
||||
model: str,
|
||||
user_object: LiteLLM_UserTable,
|
||||
valid_token: UserAPIKeyAuth,
|
||||
llm_router: Router | None,
|
||||
prisma_client: PrismaClient | None,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
) -> bool:
|
||||
return await _loaded_teams_allow_model(
|
||||
model=model,
|
||||
teams=await _load_personal_key_teams(
|
||||
user_object=user_object,
|
||||
valid_token=valid_token,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
),
|
||||
valid_token=valid_token,
|
||||
llm_router=llm_router,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
|
||||
async def _personal_key_listing_teams(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
prisma_client: PrismaClient | None,
|
||||
user_api_key_cache: UserApiKeyCache | None,
|
||||
proxy_logging_obj: ProxyLogging | None,
|
||||
) -> tuple[LiteLLM_TeamTableCachedObj | None, ...]:
|
||||
if (
|
||||
user_api_key_dict.user_id is None
|
||||
or prisma_client is None
|
||||
or user_api_key_cache is None
|
||||
or proxy_logging_obj is None
|
||||
):
|
||||
return ()
|
||||
try:
|
||||
owner: Final = await get_user_object(
|
||||
user_id=user_api_key_dict.user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
parent_otel_span=user_api_key_dict.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # fail closed: an owner that cannot be loaded sees no models
|
||||
verbose_proxy_logger.warning("Personal key team model listing: owner lookup failed, listing nothing: %s", e)
|
||||
return ()
|
||||
if owner is None:
|
||||
return ()
|
||||
return await _load_personal_key_teams(
|
||||
user_object=owner,
|
||||
valid_token=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
|
||||
async def personal_key_team_visible_models(
|
||||
models: Sequence[str],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
llm_router: Router | None,
|
||||
prisma_client: PrismaClient | None,
|
||||
user_api_key_cache: UserApiKeyCache | None,
|
||||
proxy_logging_obj: ProxyLogging | None,
|
||||
) -> list[str]:
|
||||
teams: Final = await _personal_key_listing_teams(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
return [ # mutable-ok: callers expect a list
|
||||
model
|
||||
for model in (models if teams else ())
|
||||
if (
|
||||
prisma_client is not None
|
||||
and user_api_key_cache is not None
|
||||
and proxy_logging_obj is not None
|
||||
and await _loaded_teams_allow_model(
|
||||
model=model,
|
||||
teams=teams,
|
||||
valid_token=user_api_key_dict,
|
||||
llm_router=llm_router,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
async def can_personal_key_call_model_via_teams(
|
||||
model: str | list[str],
|
||||
user_object: LiteLLM_UserTable | None,
|
||||
valid_token: UserAPIKeyAuth,
|
||||
llm_router: Router | None,
|
||||
prisma_client: PrismaClient | None,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
) -> Literal[True]:
|
||||
for requested_model in (model,) if isinstance(model, str) else tuple(model):
|
||||
if user_object is None or not await personal_key_teams_allow_model(
|
||||
model=requested_model,
|
||||
user_object=user_object,
|
||||
valid_token=valid_token,
|
||||
llm_router=llm_router,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
):
|
||||
raise ModelAccessDeniedProxyException(
|
||||
message=model_access_denied_client_message(model=requested_model),
|
||||
internal_message=(
|
||||
f"Personal key model access is derived from the owner's teams. User={valid_token.user_id}, "
|
||||
f"teams={user_object.teams if user_object is not None else None}, "
|
||||
f"strategy={litellm.personal_key_multi_team_access}, "
|
||||
f"model={requested_model}"
|
||||
),
|
||||
type=ProxyErrorTypes.key_model_access_denied,
|
||||
param="model",
|
||||
code=status.HTTP_403_FORBIDDEN,
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
def _search_tool_names_from_object_permission(
|
||||
object_permission: LiteLLM_ObjectPermissionTable | None,
|
||||
) -> list[str]:
|
||||
|
|
|
|||
|
|
@ -18500,7 +18500,7 @@ class GeneralSettingsUILiteLLMFieldSpec(TypedDict):
|
|||
description: str
|
||||
options: NotRequired[tuple[str, ...]]
|
||||
tab: NotRequired[str] # Admin UI sub-tab this field renders under; None groups it with the rest
|
||||
default: NotRequired[float] # reset/clear restores this instead of None; fields whose None means fail-open set it
|
||||
default: NotRequired[float | str] # reset/clear restores this, not None; fields whose None means fail-open set it
|
||||
|
||||
|
||||
_GENERAL_SETTINGS_UI_LITELLM_FIELDS: Final[dict[str, GeneralSettingsUILiteLLMFieldSpec]] = {
|
||||
|
|
@ -18543,6 +18543,22 @@ _GENERAL_SETTINGS_UI_LITELLM_FIELDS: Final[dict[str, GeneralSettingsUILiteLLMFie
|
|||
"forgiving it. Applies to key, user, team, team member, org, tag and end-user budgets."
|
||||
),
|
||||
},
|
||||
"personal_key_model_access_from_teams": { # mutable-ok: registry literal, frozen with its siblings below
|
||||
"type": "Boolean",
|
||||
"description": (
|
||||
"Caps each personal key (a key with no team) to the models its owner's current teams can call, "
|
||||
"re-evaluated on every request. A user in no team gets no model access. Team keys are unaffected."
|
||||
),
|
||||
},
|
||||
"personal_key_multi_team_access": { # mutable-ok: registry literal, frozen with its siblings below
|
||||
"type": "Select",
|
||||
"options": ("union", "intersection"),
|
||||
"default": "union",
|
||||
"description": (
|
||||
"How a personal key's owner in several teams is resolved when team-derived model access is on. "
|
||||
"union allows a model any of the teams can call, intersection only a model every team can call."
|
||||
),
|
||||
},
|
||||
"max_ui_session_budget": {
|
||||
"type": "Dollar",
|
||||
"default": 1.0,
|
||||
|
|
|
|||
|
|
@ -8544,6 +8544,10 @@ async def get_available_models_for_user(
|
|||
Returns:
|
||||
List of model names available to the user
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
personal_key_team_model_access_applies,
|
||||
personal_key_team_visible_models,
|
||||
)
|
||||
from litellm.proxy.auth.model_checks import (
|
||||
get_complete_model_list,
|
||||
get_key_models,
|
||||
|
|
@ -8614,7 +8618,7 @@ async def get_available_models_for_user(
|
|||
granted_team_models: Final = (*team_models, *access_group_models) if team_models else team_models
|
||||
|
||||
# Get complete model list
|
||||
all_models: Final = get_complete_model_list(
|
||||
complete_models: Final = get_complete_model_list(
|
||||
key_models=granted_key_models,
|
||||
team_models=granted_team_models,
|
||||
proxy_model_list=proxy_model_list,
|
||||
|
|
@ -8627,6 +8631,18 @@ async def get_available_models_for_user(
|
|||
only_model_access_groups=only_model_access_groups,
|
||||
team_id=effective_team_id,
|
||||
)
|
||||
all_models: Final = (
|
||||
await personal_key_team_visible_models(
|
||||
models=complete_models,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
llm_router=llm_router,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if personal_key_team_model_access_applies(user_api_key_dict)
|
||||
else complete_models
|
||||
)
|
||||
|
||||
agent_visible: Final = await _agent_access_group_visible_models(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
|
|||
|
|
@ -0,0 +1,314 @@
|
|||
from collections.abc import Iterator, Mapping
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
import litellm
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_TeamTableCachedObj,
|
||||
LiteLLM_UserTable,
|
||||
LitellmUserRoles,
|
||||
ModelAccessDeniedProxyException,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth import auth_checks
|
||||
from litellm.proxy.auth.auth_checks import can_key_call_resolved_model, common_checks
|
||||
from litellm.proxy.proxy_server import _validate_general_settings_ui_litellm_value
|
||||
from litellm.proxy.utils import get_available_models_for_user
|
||||
|
||||
_TEAMS: Final[Mapping[str, LiteLLM_TeamTableCachedObj]] = {
|
||||
"team-a": LiteLLM_TeamTableCachedObj(team_id="team-a", models=["gpt-a", "shared"]),
|
||||
"team-b": LiteLLM_TeamTableCachedObj(team_id="team-b", models=["gpt-b", "shared"]),
|
||||
"team-open": LiteLLM_TeamTableCachedObj(team_id="team-open", models=[]),
|
||||
"team-blocked": LiteLLM_TeamTableCachedObj(team_id="team-blocked", models=["gpt-a"], blocked=True),
|
||||
}
|
||||
|
||||
_ROUTER: Final = litellm.Router(
|
||||
model_list=[
|
||||
{"model_name": name, "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-fake"}}
|
||||
for name in ("gpt-a", "gpt-b", "shared", "gpt-c")
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def teams(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
|
||||
async def _get_team_object(team_id: str, **_: object) -> LiteLLM_TeamTableCachedObj:
|
||||
if team_id not in _TEAMS:
|
||||
raise Exception(f"team {team_id} lookup failed")
|
||||
return _TEAMS[team_id]
|
||||
|
||||
async def _no_membership(**_: object) -> None:
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(auth_checks, "get_team_object", _get_team_object)
|
||||
monkeypatch.setattr(auth_checks, "get_team_membership", _no_membership)
|
||||
yield
|
||||
|
||||
|
||||
def _enable(monkeypatch: pytest.MonkeyPatch, strategy: str = "union") -> None:
|
||||
monkeypatch.setattr(litellm, "personal_key_model_access_from_teams", True)
|
||||
monkeypatch.setattr(litellm, "personal_key_multi_team_access", strategy)
|
||||
|
||||
|
||||
async def _common_checks(model: str, user: LiteLLM_UserTable | None, token: UserAPIKeyAuth) -> bool:
|
||||
return await common_checks(
|
||||
request_body={"model": model},
|
||||
team_object=None,
|
||||
user_object=user,
|
||||
end_user_object=None,
|
||||
global_proxy_spend=None,
|
||||
general_settings={},
|
||||
route="/chat/completions",
|
||||
llm_router=_ROUTER,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
valid_token=token,
|
||||
request=MagicMock(spec=Request),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"strategy, model, allowed",
|
||||
[
|
||||
("union", "gpt-a", True),
|
||||
("union", "gpt-b", True),
|
||||
("union", "shared", True),
|
||||
("union", "gpt-c", False),
|
||||
("intersection", "gpt-a", False),
|
||||
("intersection", "gpt-b", False),
|
||||
("intersection", "shared", True),
|
||||
("intersection", "gpt-c", False),
|
||||
],
|
||||
)
|
||||
async def test_personal_key_inherits_team_models_by_strategy(
|
||||
monkeypatch: pytest.MonkeyPatch, teams: None, strategy: str, model: str, allowed: bool
|
||||
) -> None:
|
||||
_enable(monkeypatch, strategy)
|
||||
user: Final = LiteLLM_UserTable(user_id="u1", teams=["team-a", "team-b"])
|
||||
token: Final = UserAPIKeyAuth(token="k1", user_id="u1")
|
||||
|
||||
if allowed:
|
||||
assert await _common_checks(model, user, token) is True
|
||||
else:
|
||||
with pytest.raises(ModelAccessDeniedProxyException):
|
||||
await _common_checks(model, user, token)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_personal_key_without_teams_gets_no_models(monkeypatch: pytest.MonkeyPatch, teams: None) -> None:
|
||||
_enable(monkeypatch)
|
||||
with pytest.raises(ModelAccessDeniedProxyException):
|
||||
await _common_checks(
|
||||
"gpt-a", LiteLLM_UserTable(user_id="u1", teams=[]), UserAPIKeyAuth(token="k1", user_id="u1")
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_personal_key_denied_when_owner_not_loaded(monkeypatch: pytest.MonkeyPatch, teams: None) -> None:
|
||||
_enable(monkeypatch)
|
||||
with pytest.raises(ModelAccessDeniedProxyException):
|
||||
await _common_checks("gpt-a", None, UserAPIKeyAuth(token="k1", user_id="u1"))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_admin_personal_key_is_not_capped(monkeypatch: pytest.MonkeyPatch, teams: None) -> None:
|
||||
_enable(monkeypatch, "intersection")
|
||||
user: Final = LiteLLM_UserTable(user_id="admin", teams=[], user_role=LitellmUserRoles.PROXY_ADMIN.value)
|
||||
token: Final = UserAPIKeyAuth(token="k1", user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
assert await _common_checks("gpt-c", user, token) is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_personal_key_unrestricted_when_flag_off(monkeypatch: pytest.MonkeyPatch, teams: None) -> None:
|
||||
monkeypatch.setattr(litellm, "personal_key_model_access_from_teams", False)
|
||||
user: Final = LiteLLM_UserTable(user_id="u1", teams=[])
|
||||
assert await _common_checks("gpt-c", user, UserAPIKeyAuth(token="k1", user_id="u1")) is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_key_is_not_capped_by_owner_teams(monkeypatch: pytest.MonkeyPatch, teams: None) -> None:
|
||||
_enable(monkeypatch, "intersection")
|
||||
user: Final = LiteLLM_UserTable(user_id="u1", teams=[])
|
||||
token: Final = UserAPIKeyAuth(token="k1", user_id="u1", team_id="team-open")
|
||||
assert await _common_checks("gpt-c", user, token) is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"strategy, teams_of_user, allowed",
|
||||
[
|
||||
("union", ["team-missing", "team-a"], True),
|
||||
("union", ["team-missing"], False),
|
||||
("intersection", ["team-missing", "team-a"], False),
|
||||
("union", ["team-blocked"], False),
|
||||
("intersection", ["team-a", "team-open"], True),
|
||||
],
|
||||
)
|
||||
async def test_unloadable_or_blocked_team_grants_nothing(
|
||||
monkeypatch: pytest.MonkeyPatch, teams: None, strategy: str, teams_of_user: list[str], allowed: bool
|
||||
) -> None:
|
||||
_enable(monkeypatch, strategy)
|
||||
user: Final = LiteLLM_UserTable(user_id="u1", teams=teams_of_user)
|
||||
token: Final = UserAPIKeyAuth(token="k1", user_id="u1")
|
||||
if allowed:
|
||||
assert await _common_checks("gpt-a", user, token) is True
|
||||
else:
|
||||
with pytest.raises(ModelAccessDeniedProxyException):
|
||||
await _common_checks("gpt-a", user, token)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_membership_change_applies_to_existing_key(monkeypatch: pytest.MonkeyPatch, teams: None) -> None:
|
||||
_enable(monkeypatch)
|
||||
token: Final = UserAPIKeyAuth(token="k1", user_id="u1")
|
||||
assert await _common_checks("gpt-a", LiteLLM_UserTable(user_id="u1", teams=["team-a"]), token) is True
|
||||
with pytest.raises(ModelAccessDeniedProxyException):
|
||||
await _common_checks("gpt-a", LiteLLM_UserTable(user_id="u1", teams=["team-b"]), token)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolved_model_check_caps_personal_key(monkeypatch: pytest.MonkeyPatch, teams: None) -> None:
|
||||
_enable(monkeypatch)
|
||||
owner: Final = LiteLLM_UserTable(user_id="u1", teams=["team-a"])
|
||||
|
||||
async def _get_user_object(**_: object) -> LiteLLM_UserTable:
|
||||
return owner
|
||||
|
||||
monkeypatch.setattr(auth_checks, "get_user_object", _get_user_object)
|
||||
token: Final = UserAPIKeyAuth(token="k1", user_id="u1")
|
||||
|
||||
await can_key_call_resolved_model(model="gpt-a", llm_model_list=None, valid_token=token, llm_router=_ROUTER)
|
||||
with pytest.raises(ModelAccessDeniedProxyException):
|
||||
await can_key_call_resolved_model(model="gpt-b", llm_model_list=None, valid_token=token, llm_router=_ROUTER)
|
||||
|
||||
|
||||
async def _list_models(monkeypatch: pytest.MonkeyPatch, owner: LiteLLM_UserTable, token: UserAPIKeyAuth) -> list[str]:
|
||||
async def _get_user_object(**_: object) -> LiteLLM_UserTable:
|
||||
return owner
|
||||
|
||||
monkeypatch.setattr(auth_checks, "get_user_object", _get_user_object)
|
||||
return await get_available_models_for_user(
|
||||
user_api_key_dict=token,
|
||||
llm_router=_ROUTER,
|
||||
general_settings={},
|
||||
user_model=None,
|
||||
prisma_client=MagicMock(),
|
||||
proxy_logging_obj=MagicMock(),
|
||||
user_api_key_cache=MagicMock(),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"strategy, key_models, expected",
|
||||
[
|
||||
("union", [], {"gpt-a", "gpt-b", "shared"}),
|
||||
("intersection", [], {"shared"}),
|
||||
("union", ["gpt-a", "gpt-c"], {"gpt-a"}),
|
||||
],
|
||||
)
|
||||
async def test_model_listing_matches_team_derived_access(
|
||||
monkeypatch: pytest.MonkeyPatch, teams: None, strategy: str, key_models: list[str], expected: set[str]
|
||||
) -> None:
|
||||
_enable(monkeypatch, strategy)
|
||||
owner: Final = LiteLLM_UserTable(user_id="u1", teams=["team-a", "team-b"])
|
||||
listed: Final = await _list_models(monkeypatch, owner, UserAPIKeyAuth(token="k1", user_id="u1", models=key_models))
|
||||
assert set(listed) == expected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_listing_empty_without_teams(monkeypatch: pytest.MonkeyPatch, teams: None) -> None:
|
||||
_enable(monkeypatch)
|
||||
owner: Final = LiteLLM_UserTable(user_id="u1", teams=[])
|
||||
assert await _list_models(monkeypatch, owner, UserAPIKeyAuth(token="k1", user_id="u1")) == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_listing_unchanged_when_flag_off(monkeypatch: pytest.MonkeyPatch, teams: None) -> None:
|
||||
monkeypatch.setattr(litellm, "personal_key_model_access_from_teams", False)
|
||||
owner: Final = LiteLLM_UserTable(user_id="u1", teams=[])
|
||||
listed: Final = await _list_models(monkeypatch, owner, UserAPIKeyAuth(token="k1", user_id="u1"))
|
||||
assert set(listed) == {"gpt-a", "gpt-b", "shared", "gpt-c"}
|
||||
|
||||
|
||||
def test_multi_team_strategy_resets_to_union_and_rejects_unknown() -> None:
|
||||
assert _validate_general_settings_ui_litellm_value("personal_key_multi_team_access", None) == "union"
|
||||
assert (
|
||||
_validate_general_settings_ui_litellm_value("personal_key_multi_team_access", "intersection") == "intersection"
|
||||
)
|
||||
with pytest.raises(HTTPException):
|
||||
_validate_general_settings_ui_litellm_value("personal_key_multi_team_access", "any")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_listing_loads_each_team_once(monkeypatch: pytest.MonkeyPatch, teams: None) -> None:
|
||||
_enable(monkeypatch)
|
||||
get_team: Final = AsyncMock(wraps=auth_checks.get_team_object)
|
||||
monkeypatch.setattr(auth_checks, "get_team_object", get_team)
|
||||
owner: Final = LiteLLM_UserTable(user_id="u1", teams=["team-a", "team-b", "team-a"])
|
||||
listed: Final = await _list_models(monkeypatch, owner, UserAPIKeyAuth(token="k1", user_id="u1"))
|
||||
assert set(listed) == {"gpt-a", "gpt-b", "shared"}
|
||||
assert sorted(call.kwargs["team_id"] for call in get_team.call_args_list) == ["team-a", "team-b"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_check_error_grants_nothing_for_that_team(monkeypatch: pytest.MonkeyPatch, teams: None) -> None:
|
||||
_enable(monkeypatch)
|
||||
check_team: Final = auth_checks.can_team_access_model
|
||||
|
||||
async def _failing_for_team_b(model: str, team_object: LiteLLM_TeamTableCachedObj, **_: object) -> bool:
|
||||
if team_object.team_id == "team-b":
|
||||
raise RuntimeError("access group lookup failed")
|
||||
return await check_team(model=model, team_object=team_object, llm_router=_ROUTER)
|
||||
|
||||
monkeypatch.setattr(auth_checks, "can_team_access_model", _failing_for_team_b)
|
||||
owner: Final = LiteLLM_UserTable(user_id="u1", teams=["team-a", "team-b"])
|
||||
token: Final = UserAPIKeyAuth(token="k1", user_id="u1")
|
||||
assert await _common_checks("gpt-a", owner, token) is True
|
||||
with pytest.raises(ModelAccessDeniedProxyException):
|
||||
await _common_checks("gpt-b", owner, token)
|
||||
|
||||
|
||||
async def _owner_lookup_fails(**_: object) -> LiteLLM_UserTable:
|
||||
raise RuntimeError("database unavailable")
|
||||
|
||||
|
||||
async def _owner_missing(**_: object) -> None:
|
||||
return None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("get_owner", [_owner_lookup_fails, _owner_missing])
|
||||
async def test_model_listing_empty_when_owner_not_loaded(
|
||||
monkeypatch: pytest.MonkeyPatch, teams: None, get_owner: object
|
||||
) -> None:
|
||||
_enable(monkeypatch)
|
||||
monkeypatch.setattr(auth_checks, "get_user_object", get_owner)
|
||||
listed: Final = await get_available_models_for_user(
|
||||
user_api_key_dict=UserAPIKeyAuth(token="k1", user_id="u1"),
|
||||
llm_router=_ROUTER,
|
||||
general_settings={},
|
||||
user_model=None,
|
||||
prisma_client=MagicMock(),
|
||||
proxy_logging_obj=MagicMock(),
|
||||
user_api_key_cache=MagicMock(),
|
||||
)
|
||||
assert listed == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_listing_empty_without_database(monkeypatch: pytest.MonkeyPatch, teams: None) -> None:
|
||||
_enable(monkeypatch)
|
||||
listed: Final = await get_available_models_for_user(
|
||||
user_api_key_dict=UserAPIKeyAuth(token="k1", user_id="u1"),
|
||||
llm_router=_ROUTER,
|
||||
general_settings={},
|
||||
user_model=None,
|
||||
prisma_client=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
user_api_key_cache=MagicMock(),
|
||||
)
|
||||
assert listed == []
|
||||
Loading…
Add table
Reference in a new issue