diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index d84588a4c24..5fe17d79ab5 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1005,6 +1005,9 @@ class LiteLLM_ObjectPermissionBase(LiteLLMPydanticObjectBase): search_tools: Optional[List[str]] = None +from litellm.types.object_permission import ( # noqa: E402 + ObjectPermissionDict as ObjectPermissionDict, +) from litellm.models.team import BudgetLimitEntry as BudgetLimitEntry # noqa: E402 diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 3fd6b886811..4106eae606c 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -80,6 +80,7 @@ from litellm.proxy.management_helpers.object_permission_utils import ( handle_update_object_permission_common, validate_key_mcp_servers_against_team, validate_key_search_tools_against_team, + validate_key_vector_stores_against_team, ) from litellm.proxy.management_helpers.team_member_permission_checks import ( TeamMemberPermissionChecks, @@ -348,7 +349,21 @@ def _personal_key_membership_check( return True +def _object_permission_to_dict( + object_permission: Optional[LiteLLM_ObjectPermissionBase], +) -> Optional[ObjectPermissionDict]: + if object_permission is None: + return None + return cast(ObjectPermissionDict, object_permission.model_dump(exclude_unset=True)) + + def _personal_key_generation_check(user_api_key_dict: UserAPIKeyAuth, data: GenerateKeyRequest): + TeamMemberPermissionChecks.enforce_member_can_assign_access_groups( + user_api_key_dict=user_api_key_dict, + team_table=None, + access_group_ids=data.access_group_ids, + ) + if ( litellm.key_generation_settings is None or litellm.key_generation_settings.get("personal_key_generation") is None @@ -930,17 +945,24 @@ async def _common_key_generation_helper( data_json.pop("tags") # Validate MCP servers in object_permission are within team scope + _is_proxy_admin_caller = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value normalized_object_permission = await validate_key_mcp_servers_against_team( object_permission=data_json.get("object_permission"), team_obj=team_table, prisma_client=prisma_client, - is_proxy_admin=user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value, + is_proxy_admin=_is_proxy_admin_caller, ) if normalized_object_permission is not None: data_json["object_permission"] = normalized_object_permission await validate_key_search_tools_against_team( object_permission=data_json.get("object_permission"), team_obj=team_table, + is_proxy_admin=_is_proxy_admin_caller, + ) + await validate_key_vector_stores_against_team( + object_permission=data_json.get("object_permission"), + team_obj=team_table, + is_proxy_admin=_is_proxy_admin_caller, ) data_json = await _set_object_permission( @@ -2143,7 +2165,7 @@ async def _validate_mcp_servers_for_key_update( prisma_client: Any, user_api_key_cache: Any, is_proxy_admin: bool, -) -> Optional[dict]: +) -> Optional[ObjectPermissionDict]: """Validate MCP servers in object_permission against the effective team.""" effective_team_obj = team_obj # If team_id isn't being changed, resolve the existing key's team @@ -2154,13 +2176,7 @@ async def _validate_mcp_servers_for_key_update( user_api_key_cache=user_api_key_cache, check_db_only=True, ) - object_permission_dict: Optional[dict] = None - if data.object_permission is not None: - object_permission_dict = ( - data.object_permission.model_dump(exclude_unset=True) - if hasattr(data.object_permission, "model_dump") - else dict(data.object_permission) # type: ignore[arg-type] - ) + object_permission_dict = _object_permission_to_dict(data.object_permission) normalized_object_permission = await validate_key_mcp_servers_against_team( object_permission=object_permission_dict, team_obj=effective_team_obj, @@ -2170,6 +2186,12 @@ async def _validate_mcp_servers_for_key_update( await validate_key_search_tools_against_team( object_permission=object_permission_dict, team_obj=effective_team_obj, + is_proxy_admin=is_proxy_admin, + ) + await validate_key_vector_stores_against_team( + object_permission=object_permission_dict, + team_obj=effective_team_obj, + is_proxy_admin=is_proxy_admin, ) return normalized_object_permission @@ -2301,14 +2323,6 @@ async def _validate_update_key_data( detail=f"Team not found for team_id={data.team_id}. Non-admin users cannot set keys to non-existent teams.", ) - # Field-level opt-in: non-admin members may only assign access groups when - # the team has enabled KEY_ACCESS_GROUP_ASSIGNMENT. - TeamMemberPermissionChecks.enforce_member_can_assign_access_groups( - user_api_key_dict=user_api_key_dict, - team_table=team_obj, - access_group_ids=data.access_group_ids, - ) - if team_obj is not None: await _check_team_key_limits( team_table=team_obj, @@ -2316,6 +2330,12 @@ async def _validate_update_key_data( prisma_client=prisma_client, ) + TeamMemberPermissionChecks.enforce_member_can_assign_access_groups( + user_api_key_dict=user_api_key_dict, + team_table=team_obj, + access_group_ids=data.access_group_ids, + ) + # Validate key against project limits if project_id is being set _project_id_to_check = getattr(data, "project_id", None) or getattr(existing_key_row, "project_id", None) if _project_id_to_check is not None and (data.models is not None or data.max_budget is not None): @@ -4623,9 +4643,7 @@ async def regenerate_key_fn( detail={"error": "You are not authorized to regenerate this key"}, ) - # Gate access_group_ids on regenerate, same as /key/generate and - # /key/update. Use the existing key's team since the body may omit it. - if data is not None and data.access_group_ids: + if data is not None and (data.access_group_ids or data.object_permission is not None): regenerate_team_table: Optional[LiteLLM_TeamTableCachedObj] = None if _key_in_db.team_id is not None: regenerate_team_table = await get_team_object( @@ -4634,11 +4652,32 @@ async def regenerate_key_fn( user_api_key_cache=user_api_key_cache, check_db_only=True, ) + _regen_is_proxy_admin = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value TeamMemberPermissionChecks.enforce_member_can_assign_access_groups( user_api_key_dict=user_api_key_dict, team_table=regenerate_team_table, access_group_ids=data.access_group_ids, ) + _regen_object_permission_dict = _object_permission_to_dict(data.object_permission) + normalized_object_permission = await validate_key_mcp_servers_against_team( + object_permission=_regen_object_permission_dict, + team_obj=regenerate_team_table, + prisma_client=prisma_client, + is_proxy_admin=_regen_is_proxy_admin, + ) + if normalized_object_permission is not None: + data.object_permission = LiteLLM_ObjectPermissionBase(**normalized_object_permission) + _regen_object_permission_dict = normalized_object_permission + await validate_key_search_tools_against_team( + object_permission=_regen_object_permission_dict, + team_obj=regenerate_team_table, + is_proxy_admin=_regen_is_proxy_admin, + ) + await validate_key_vector_stores_against_team( + object_permission=_regen_object_permission_dict, + team_obj=regenerate_team_table, + is_proxy_admin=_regen_is_proxy_admin, + ) verbose_proxy_logger.info( "Key regeneration requested: key_alias=%s", diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index 9d5f716033f..fe96d9c260a 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -11,7 +11,7 @@ from fastapi import HTTPException, status from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid from litellm.litellm_core_utils.safe_json_dumps import safe_dumps -from litellm.proxy._types import SpecialMCPServerNames +from litellm.proxy._types import ObjectPermissionDict, SpecialMCPServerNames from litellm.proxy.utils import PrismaClient from litellm.repositories.object_permission_repository import ObjectPermissionRepository from litellm.repositories.table_repositories import MCPServerRepository @@ -257,7 +257,7 @@ async def _resolve_mcp_server_identifiers_to_ids( def _rewrite_object_permission_mcp_servers( - object_permission: dict, + object_permission: ObjectPermissionDict, identifier_to_server_ids: Dict[str, Set[str]], ) -> None: mcp_servers = object_permission.get("mcp_servers") @@ -274,7 +274,7 @@ def _rewrite_object_permission_mcp_servers( def _rewrite_object_permission_mcp_tool_permissions( - object_permission: dict, + object_permission: ObjectPermissionDict, identifier_to_server_ids: Dict[str, Set[str]], ) -> None: mcp_tool_permissions = object_permission.get("mcp_tool_permissions") @@ -295,7 +295,7 @@ def _rewrite_object_permission_mcp_tool_permissions( def _rewrite_object_permission_mcp_identifiers( - object_permission: Optional[dict], + object_permission: Optional[ObjectPermissionDict], identifier_to_server_ids: Dict[str, Set[str]], ) -> None: if not object_permission or not isinstance(object_permission, dict): @@ -383,7 +383,7 @@ async def _get_team_allowed_mcp_servers( def _extract_requested_mcp_server_ids( - object_permission: Optional[dict], + object_permission: Optional[ObjectPermissionDict], ) -> Set[str]: """ Extract all MCP server IDs referenced in a key's object_permission dict. @@ -409,7 +409,7 @@ def _extract_requested_mcp_server_ids( def _extract_requested_mcp_access_groups( - object_permission: Optional[dict], + object_permission: Optional[ObjectPermissionDict], ) -> Set[str]: """Extract MCP access groups from a key's object_permission dict.""" if not object_permission or not isinstance(object_permission, dict): @@ -422,7 +422,7 @@ def _extract_requested_mcp_access_groups( def _extract_requested_mcp_toolsets( - object_permission: Optional[dict], + object_permission: Optional[ObjectPermissionDict], ) -> Set[str]: """Extract MCP toolset IDs from a key's object_permission dict.""" if not object_permission or not isinstance(object_permission, dict): @@ -435,11 +435,11 @@ def _extract_requested_mcp_toolsets( async def validate_key_mcp_servers_against_team( - object_permission: Optional[dict], + object_permission: Optional[ObjectPermissionDict], team_obj: Optional["LiteLLM_TeamTableCachedObj"], prisma_client: Optional[PrismaClient] = None, is_proxy_admin: bool = False, -) -> Optional[dict]: +) -> Optional[ObjectPermissionDict]: """ Validate that MCP servers requested on a key are within the allowed scope. @@ -555,32 +555,103 @@ async def validate_key_mcp_servers_against_team( detail={"error": detail}, ) - # Validate requested toolsets against team's allowed toolsets. - # Only enforce the team-based restriction when a team is present — standalone - # keys (no team) can freely be granted any toolset by an admin. - if requested_toolsets and team_obj is not None: - team_op = team_obj.object_permission - team_mcp_toolsets = team_op.mcp_toolsets if team_op is not None else None - # None or [] means the team has no toolset restriction — allow any toolsets. - if team_mcp_toolsets: - disallowed_toolsets = requested_toolsets - set(team_mcp_toolsets) - if disallowed_toolsets: - team_id = team_obj.team_id - raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail={ - "error": ( - f"Key requests MCP toolsets not allowed by team '{team_id}': " - f"{sorted(disallowed_toolsets)}. " - f"Team allows: {sorted(team_mcp_toolsets)}." - ) - }, - ) + _validate_requested_toolsets( + requested_toolsets=requested_toolsets, + team_obj=team_obj, + is_proxy_admin=is_proxy_admin, + ) return object_permission -def _extract_requested_search_tools(object_permission: Optional[dict]) -> List[str]: +def _validate_requested_toolsets( + requested_toolsets: set[str], + team_obj: Optional["LiteLLM_TeamTableCachedObj"], + is_proxy_admin: bool, +) -> None: + """ + Validate mcp_toolsets requested on a key. + + Non-admin callers cannot assign toolsets to a personal (no team) key. Team + keys must request a subset of the team's own toolset allowlist. + """ + if not requested_toolsets: + return + if team_obj is None: + if is_proxy_admin: + return + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={ + "error": ( + "Key is not in a team. MCP toolsets cannot be assigned to " + "personal keys by non-admin callers. Disallowed toolsets: " + f"{sorted(requested_toolsets)}." + ) + }, + ) + team_op = team_obj.object_permission + team_mcp_toolsets = team_op.mcp_toolsets if team_op is not None else None + if not team_mcp_toolsets: + return + disallowed_toolsets = requested_toolsets - set(team_mcp_toolsets) + if not disallowed_toolsets: + return + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={ + "error": ( + f"Key requests MCP toolsets not allowed by team '{team_obj.team_id}': " + f"{sorted(disallowed_toolsets)}. " + f"Team allows: {sorted(team_mcp_toolsets)}." + ) + }, + ) + + +def _extract_requested_vector_stores( + object_permission: Optional[ObjectPermissionDict], +) -> set[str]: + """Return vector_store IDs from a key's object_permission dict.""" + if not object_permission or not isinstance(object_permission, dict): + return set() + raw = object_permission.get("vector_stores") + if isinstance(raw, list): + return {str(x) for x in raw if x} + return set() + + +async def validate_key_vector_stores_against_team( + object_permission: Optional[ObjectPermissionDict], + team_obj: Optional["LiteLLM_TeamTableCachedObj"], + is_proxy_admin: bool = False, +) -> None: + """ + Reject vector_stores requested on a personal (no team) key by a non-admin + caller. Vector store access is granted at use-time from the key's + object_permission.vector_stores list, so the assignment is the authorization + boundary. Team keys and proxy admins are unaffected. + """ + requested = _extract_requested_vector_stores(object_permission) + if not requested: + return + if team_obj is not None or is_proxy_admin: + return + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={ + "error": ( + "Key is not in a team. Vector stores cannot be assigned to " + "personal keys by non-admin callers. Disallowed vector stores: " + f"{sorted(requested)}." + ) + }, + ) + + +def _extract_requested_search_tools( + object_permission: Optional[ObjectPermissionDict], +) -> list[str]: """Return search_tool_name values from a key's object_permission dict.""" if not object_permission or not isinstance(object_permission, dict): return [] @@ -591,18 +662,32 @@ def _extract_requested_search_tools(object_permission: Optional[dict]) -> List[s async def validate_key_search_tools_against_team( - object_permission: Optional[dict], + object_permission: Optional[ObjectPermissionDict], team_obj: Optional["LiteLLM_TeamTableCachedObj"], + is_proxy_admin: bool = False, ) -> None: """ Validate key object_permission.search_tools is a subset of the team's allowlist. Empty team allowlist means no restriction at team layer (skip). + Non-admin callers cannot assign search_tools to a personal (no team) key. """ requested = _extract_requested_search_tools(object_permission) if not requested: return + if team_obj is None and not is_proxy_admin: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={ + "error": ( + "Key is not in a team. search_tools cannot be assigned to " + "personal keys by non-admin callers. Disallowed search tools: " + f"{sorted(requested)}." + ) + }, + ) + team_tools: List[str] = [] if team_obj is not None and team_obj.object_permission is not None: st = team_obj.object_permission.search_tools diff --git a/litellm/proxy/management_helpers/team_member_permission_checks.py b/litellm/proxy/management_helpers/team_member_permission_checks.py index 1353b9ed651..1532668ed19 100644 --- a/litellm/proxy/management_helpers/team_member_permission_checks.py +++ b/litellm/proxy/management_helpers/team_member_permission_checks.py @@ -167,9 +167,15 @@ class TeamMemberPermissionChecks: if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value: return - # Personal (non-team) keys are out of scope for team-member gating. if team_table is None: - return + raise HTTPException( + status_code=403, + detail=( + "Key is not in a team. Access groups cannot be assigned to " + "personal keys by non-admin callers. Disallowed access groups: " + f"{sorted(access_group_ids)}." + ), + ) team_member_object = _get_user_in_team(team_table=team_table, user_id=user_api_key_dict.user_id) diff --git a/litellm/types/object_permission.py b/litellm/types/object_permission.py new file mode 100644 index 00000000000..ff932dccd5d --- /dev/null +++ b/litellm/types/object_permission.py @@ -0,0 +1,26 @@ +""" +TypedDict mirror of ``LiteLLM_ObjectPermissionBase`` for in-memory dict +payloads passed through validators that mutate before persistence (e.g. +MCP server identifier normalization in object_permission_utils). + +Lives in ``litellm/types/`` so SDK-side modules (``litellm.types.agents``) +can adopt the type without violating the SDK-must-not-import-from-proxy +layering rule. +""" + +from typing import Optional + +from typing_extensions import TypedDict + + +class ObjectPermissionDict(TypedDict, total=False): + mcp_servers: Optional[list[str]] + mcp_access_groups: Optional[list[str]] + mcp_tool_permissions: Optional[dict[str, list[str]]] + mcp_toolsets: Optional[list[str]] + blocked_tools: Optional[list[str]] + vector_stores: Optional[list[str]] + agents: Optional[list[str]] + agent_access_groups: Optional[list[str]] + models: Optional[list[str]] + search_tools: Optional[list[str]] diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 8e4f6148dfe..04048020e18 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -526,6 +526,197 @@ async def test_key_generation_with_object_permission(monkeypatch): assert key_insert_calls[0]["data"].get("object_permission_id") == "objperm123" +@pytest.mark.asyncio +@pytest.mark.parametrize( + "field,request_kwargs,expected_in_error", + [ + ( + "access_group_ids", + {"access_group_ids": ["acme_private"]}, + "Access groups", + ), + ( + "mcp_toolsets", + {"object_permission": {"mcp_toolsets": ["acme_toolset"]}}, + "MCP toolsets", + ), + ( + "vector_stores", + {"object_permission": {"vector_stores": ["acme_vs"]}}, + "Vector stores", + ), + ( + "search_tools", + {"object_permission": {"search_tools": ["acme_search"]}}, + "search_tools", + ), + ], +) +async def test_generate_key_personal_non_admin_denied_for_team_scoped_fields( + monkeypatch, field, request_kwargs, expected_in_error +): + """generate_key_fn must reject access_group_ids and + object_permission.{mcp_toolsets, vector_stores, search_tools} when the + caller is a non-admin and the request has no team_id. Mutating any of the + three validator calls in _common_key_generation_helper or unmoving the + enforce_member_can_assign_access_groups call in _personal_key_generation_check + must break this test.""" + mock_prisma_client = AsyncMock() + mock_prisma_client.jsonify_object = lambda data: data # type: ignore + mock_prisma_client.db = MagicMock() + mock_prisma_client.db.litellm_objectpermissiontable = MagicMock() + mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock( + return_value=MagicMock(object_permission_id="should-not-create") + ) + mock_prisma_client.insert_data = AsyncMock(return_value=MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + from litellm.proxy._types import ( + GenerateKeyRequest, + LiteLLM_ObjectPermissionBase, + LitellmUserRoles, + ) + from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth + from litellm.proxy.management_endpoints.key_management_endpoints import ( + generate_key_fn, + ) + + if "object_permission" in request_kwargs: + request_kwargs = { + **request_kwargs, + "object_permission": LiteLLM_ObjectPermissionBase( + **request_kwargs["object_permission"] + ), + } + request_data = GenerateKeyRequest(**request_kwargs) + + from litellm.proxy._types import ProxyException + + with pytest.raises((HTTPException, ProxyException)) as exc: + await generate_key_fn( + data=request_data, + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-alice", + user_id="alice", + ), + ) + code = getattr(exc.value, "status_code", None) or getattr(exc.value, "code", None) + assert int(code) == 403 + body = str( + getattr(exc.value, "detail", None) or getattr(exc.value, "message", exc.value) + ) + assert expected_in_error in body + mock_prisma_client.db.litellm_objectpermissiontable.create.assert_not_called() + + +@pytest.mark.asyncio +async def test_update_key_personal_non_admin_denied_vector_stores(monkeypatch): + """/key/update must reject vector_stores on a personal key by a non-admin. + Reverting the enforce_member_can_assign_access_groups move (i.e. putting + it back inside `if _team_id_to_check is not None`) does NOT cover + object_permission fields; this test exercises _validate_update_key_data + which calls _validate_mcp_servers_for_key_update.""" + mock_prisma_client = AsyncMock() + mock_prisma_client.jsonify_object = lambda data: data # type: ignore + mock_prisma_client.db = MagicMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr( + "litellm.proxy.proxy_server.user_api_key_cache", + MagicMock(), + ) + + from litellm.proxy._types import ( + LiteLLM_ObjectPermissionBase, + LitellmUserRoles, + UpdateKeyRequest, + ) + from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _validate_update_key_data, + ) + + existing_key_row = MagicMock( + token="hashed_alice_personal_key", + user_id="alice", + team_id=None, + created_by="alice", + max_budget=None, + organization_id=None, + project_id=None, + ) + data = UpdateKeyRequest( + key="sk-alice-personal", + object_permission=LiteLLM_ObjectPermissionBase(vector_stores=["acme_vs"]), + ) + + with pytest.raises(HTTPException) as exc: + await _validate_update_key_data( + data=data, + existing_key_row=existing_key_row, + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-alice", + user_id="alice", + ), + llm_router=None, + premium_user=True, + prisma_client=mock_prisma_client, + user_api_key_cache=MagicMock(), + ) + assert exc.value.status_code == 403 + assert "Vector stores" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_update_key_personal_non_admin_denied_access_groups( + monkeypatch, +): + """/key/update on a personal key must also gate access_group_ids for + non-admins. Reverting the enforce move (putting it back inside + `if _team_id_to_check is not None`) breaks this test.""" + mock_prisma_client = AsyncMock() + mock_prisma_client.jsonify_object = lambda data: data # type: ignore + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + from litellm.proxy._types import LitellmUserRoles, UpdateKeyRequest + from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _validate_update_key_data, + ) + + existing_key_row = MagicMock( + token="hashed_alice_personal_key", + user_id="alice", + team_id=None, + created_by="alice", + max_budget=None, + organization_id=None, + project_id=None, + ) + data = UpdateKeyRequest( + key="sk-alice-personal", + access_group_ids=["ag-private"], + ) + + with pytest.raises(HTTPException) as exc: + await _validate_update_key_data( + data=data, + existing_key_row=existing_key_row, + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-alice", + user_id="alice", + ), + llm_router=None, + premium_user=True, + prisma_client=mock_prisma_client, + user_api_key_cache=MagicMock(), + ) + assert exc.value.status_code == 403 + assert "Access groups" in str(exc.value.detail) + + @pytest.mark.asyncio async def test_generate_key_helper_fn_with_access_group_ids(monkeypatch): """Ensure generate_key_helper_fn passes access_group_ids into the key insert payload.""" @@ -11643,6 +11834,85 @@ async def test_regenerate_premium_gate_allows_actual_master_key_holder(): assert result.token == "sk-new-master" +@pytest.mark.asyncio +async def test_regenerate_applies_normalized_mcp_object_permission(): + from litellm.proxy._types import ( + LiteLLM_ObjectPermissionBase, + RegenerateKeyRequest, + ) + from litellm.proxy.management_endpoints.key_management_endpoints import ( + regenerate_key_fn, + ) + + data = RegenerateKeyRequest( + key="sk-old", + object_permission=LiteLLM_ObjectPermissionBase(mcp_servers=["server-alias"]), + ) + existing_key = _make_regenerate_existing_key() + mock_prisma_client = AsyncMock() + mock_repo = MagicMock() + mock_repo.table.find_unique = AsyncMock(return_value=existing_key) + execute_mock = AsyncMock(return_value=MagicMock()) + + with ( + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.proxy_server.master_key", None), + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + patch("litellm.proxy.proxy_server.hash_token", lambda token: "hashed-old"), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.VerificationTokenRepository", + return_value=mock_repo, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.enforce_member_can_assign_access_groups", + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.can_modify_verification_token", + new_callable=AsyncMock, + return_value=True, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.validate_key_mcp_servers_against_team", + new_callable=AsyncMock, + return_value={"mcp_servers": ["server-id"]}, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.validate_key_search_tools_against_team", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.validate_key_vector_stores_against_team", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._persist_deleted_verification_tokens", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._execute_virtual_key_regeneration", + execute_mock, + ), + ): + await regenerate_key_fn( + key="sk-old", + data=data, + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN.value, + api_key="sk-admin", + user_id="admin", + ), + ) + + regenerated_data = execute_mock.await_args.kwargs["data"] + assert regenerated_data.object_permission.mcp_servers == ["server-id"] + + # --------------------------------------------------------------------------- # Regression tests for GHSA-q775-qw9r-2r4g: budget escalation via key/generate # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py index 2b38d732e9d..26c8c774812 100644 --- a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py @@ -9,7 +9,7 @@ sys.path.insert(0, os.path.abspath("../../../..")) from unittest.mock import AsyncMock, MagicMock, patch -from litellm.proxy._types import LiteLLM_ObjectPermissionTable +from litellm.proxy._types import LiteLLM_ObjectPermissionBase, LiteLLM_ObjectPermissionTable, ObjectPermissionDict from litellm.proxy.management_helpers.object_permission_utils import ( _extract_requested_mcp_access_groups, _extract_requested_mcp_server_ids, @@ -18,6 +18,7 @@ from litellm.proxy.management_helpers.object_permission_utils import ( _set_object_permission, validate_key_mcp_servers_against_team, validate_key_search_tools_against_team, + validate_key_vector_stores_against_team, ) @@ -890,3 +891,137 @@ async def test_validate_search_tools_raises_when_not_subset(): team_obj=_make_team_obj_search(search_tools=["t1"]), ) assert exc.value.status_code == 403 + + +# ---- Personal-key non-admin gates on toolsets / vector_stores / search_tools ---- + + +@pytest.mark.asyncio +@patch( + "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + return_value=set(), +) +@patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + return_value=[], +) +async def test_personal_non_admin_cannot_assign_mcp_toolsets( + mock_access_groups, mock_allow_all +): + with pytest.raises(HTTPException) as exc: + await validate_key_mcp_servers_against_team( + object_permission={"mcp_toolsets": ["ts-private"]}, + team_obj=None, + is_proxy_admin=False, + ) + assert exc.value.status_code == 403 + assert "ts-private" in str(exc.value.detail) + + +@pytest.mark.asyncio +@patch( + "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + return_value=set(), +) +@patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + return_value=[], +) +async def test_personal_admin_can_assign_mcp_toolsets( + mock_access_groups, mock_allow_all +): + await validate_key_mcp_servers_against_team( + object_permission={"mcp_toolsets": ["ts-private"]}, + team_obj=None, + is_proxy_admin=True, + ) + + +@pytest.mark.asyncio +async def test_personal_non_admin_cannot_assign_vector_stores(): + with pytest.raises(HTTPException) as exc: + await validate_key_vector_stores_against_team( + object_permission={"vector_stores": ["vs-private"]}, + team_obj=None, + is_proxy_admin=False, + ) + assert exc.value.status_code == 403 + assert "vs-private" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_personal_admin_can_assign_vector_stores(): + await validate_key_vector_stores_against_team( + object_permission={"vector_stores": ["vs-private"]}, + team_obj=None, + is_proxy_admin=True, + ) + + +@pytest.mark.asyncio +async def test_team_key_vector_stores_unrestricted_at_create(): + """Team-scoped keys retain their existing trust model at create time.""" + team_obj = _make_team_obj_search() + await validate_key_vector_stores_against_team( + object_permission={"vector_stores": ["vs-anything"]}, + team_obj=team_obj, + is_proxy_admin=False, + ) + + +@pytest.mark.asyncio +async def test_personal_non_admin_cannot_assign_search_tools(): + with pytest.raises(HTTPException) as exc: + await validate_key_search_tools_against_team( + object_permission={"search_tools": ["st-private"]}, + team_obj=None, + is_proxy_admin=False, + ) + assert exc.value.status_code == 403 + assert "st-private" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_personal_admin_can_assign_search_tools(): + await validate_key_search_tools_against_team( + object_permission={"search_tools": ["st-private"]}, + team_obj=None, + is_proxy_admin=True, + ) + + +@pytest.mark.asyncio +async def test_empty_object_permission_passes_for_personal_non_admin(): + """An empty / absent object_permission must not be blocked.""" + await validate_key_vector_stores_against_team( + object_permission=None, + team_obj=None, + is_proxy_admin=False, + ) + await validate_key_vector_stores_against_team( + object_permission={"vector_stores": []}, + team_obj=None, + is_proxy_admin=False, + ) + await validate_key_search_tools_against_team( + object_permission=None, + team_obj=None, + is_proxy_admin=False, + ) + + +def test_object_permission_dict_mirrors_pydantic_model(): + """ObjectPermissionDict must stay field-for-field aligned with + LiteLLM_ObjectPermissionBase. If a new field is added to the Pydantic + model, this test fails until the TypedDict is updated to match.""" + from typing import get_type_hints + + pydantic_fields = set(LiteLLM_ObjectPermissionBase.model_fields.keys()) + typeddict_fields = set(get_type_hints(ObjectPermissionDict).keys()) + assert pydantic_fields == typeddict_fields, ( + f"ObjectPermissionDict drifted from LiteLLM_ObjectPermissionBase.\n" + f"Only in Pydantic model: {sorted(pydantic_fields - typeddict_fields)}\n" + f"Only in TypedDict: {sorted(typeddict_fields - pydantic_fields)}" + ) diff --git a/tests/test_litellm/proxy/management_helpers/test_team_member_permission_checks.py b/tests/test_litellm/proxy/management_helpers/test_team_member_permission_checks.py index 29aa75a0f0a..71999e29f96 100644 --- a/tests/test_litellm/proxy/management_helpers/test_team_member_permission_checks.py +++ b/tests/test_litellm/proxy/management_helpers/test_team_member_permission_checks.py @@ -314,12 +314,45 @@ class TestEnforceMemberCanAssignAccessGroups: access_group_ids=["ag-1"], ) - def test_personal_key_out_of_scope(self): - """Personal (non-team) keys are not gated by team-member permissions.""" + def test_personal_key_non_admin_denied(self): + """A non-admin cannot self-grant access_group_ids on a personal (no + team) key. The access_group_id grants model access at use-time + without any team-membership cross-check, so the assignment is the + authorization boundary.""" + from fastapi import HTTPException + + with pytest.raises(HTTPException) as exc: + TeamMemberPermissionChecks.enforce_member_can_assign_access_groups( + user_api_key_dict=self._user(), + team_table=None, + access_group_ids=["ag-private"], + ) + assert exc.value.status_code == 403 + assert "ag-private" in str(exc.value.detail) + + def test_personal_key_proxy_admin_can_assign(self): + """Proxy admins bypass the personal-key gate and may assign access + groups on personal keys.""" + from litellm.proxy._types import LitellmUserRoles + + TeamMemberPermissionChecks.enforce_member_can_assign_access_groups( + user_api_key_dict=self._user(role=LitellmUserRoles.PROXY_ADMIN.value), + team_table=None, + access_group_ids=["ag-private"], + ) + + def test_personal_key_empty_access_groups_passes(self): + """An empty / absent access_group_ids list must not be rejected even + on a personal key — the gate only fires when the field is non-empty.""" TeamMemberPermissionChecks.enforce_member_can_assign_access_groups( user_api_key_dict=self._user(), team_table=None, - access_group_ids=["ag-1"], + access_group_ids=None, + ) + TeamMemberPermissionChecks.enforce_member_can_assign_access_groups( + user_api_key_dict=self._user(), + team_table=None, + access_group_ids=[], ) def test_team_admin_bypasses(self, monkeypatch):