diff --git a/litellm/__init__.py b/litellm/__init__.py index e1da202b9ee..5070157a33f 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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 diff --git a/litellm/constants.py b/litellm/constants.py index 530d678457d..9af5ea576e2 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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", diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index e8f335cb348..dfbf76d345e 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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]: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index a994293a4d8..3b0c3f4ac27 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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, diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 43ad433c19b..33566514c02 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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, diff --git a/tests/test_litellm/proxy/auth/test_personal_key_team_model_access.py b/tests/test_litellm/proxy/auth/test_personal_key_team_model_access.py new file mode 100644 index 00000000000..4e1f17215b0 --- /dev/null +++ b/tests/test_litellm/proxy/auth/test_personal_key_team_model_access.py @@ -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 == []