From f2d7cb152adba50a561fb72b6dde4f7ce9c96913 Mon Sep 17 00:00:00 2001 From: yucheng-berriai Date: Fri, 26 Jun 2026 16:33:52 -0700 Subject: [PATCH] refactor(proxy): type object_permission dict with ObjectPermissionDict Replace bare Optional[dict] on the object_permission validator surfaces with a typed TypedDict mirror of LiteLLM_ObjectPermissionBase. The TypedDict shape matches the Pydantic model field-for-field and supports .get() and item assignment, so the mutation in _rewrite_object_permission_mcp_identifiers continues to work at runtime (TypedDict is a plain dict). Propagated through the surfaces this PR touches: _object_permission_to_dict, _validate_mcp_servers_for_key_update, validate_key_mcp_servers_against_team, validate_key_search_tools_against_team, validate_key_vector_stores_against _team, the five _extract_requested_* helpers, and the two _rewrite_object_permission_mcp_* mutators. attach_object_permission_to_dict, handle_update_object_permission_common, and _set_object_permission keep their wider dict typing because they handle the full key/team data_json, which is a superset of ObjectPermissionDict and pre-dates this PR. No behavior change. 373 tests pass; ruff strict + type discipline gates green. --- litellm/proxy/_types.py | 17 +++++++++++ .../key_management_endpoints.py | 26 ++++++++-------- .../object_permission_utils.py | 30 +++++++++++-------- .../test_object_permission_utils.py | 17 ++++++++++- 4 files changed, 63 insertions(+), 27 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index d84588a4c24..7d470fba08b 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1005,6 +1005,23 @@ class LiteLLM_ObjectPermissionBase(LiteLLMPydanticObjectBase): search_tools: Optional[List[str]] = None +class ObjectPermissionDict(TypedDict, total=False): + """Plain-dict mirror of LiteLLM_ObjectPermissionBase used by validators + that need to mutate the payload before persistence (e.g. MCP server + identifier normalization in object_permission_utils).""" + + 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]] + + 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 f19ea6da529..aeb664a5d1f 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -349,6 +349,14 @@ 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, @@ -852,9 +860,7 @@ 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 - ) + _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, @@ -2074,7 +2080,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 @@ -4472,9 +4478,7 @@ async def regenerate_key_fn( detail={"error": "You are not authorized to regenerate this key"}, ) - if data is not None and ( - data.access_group_ids or data.object_permission is not None - ): + 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( @@ -4483,17 +4487,13 @@ 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 - ) + _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 - ) + _regen_object_permission_dict = _object_permission_to_dict(data.object_permission) await validate_key_mcp_servers_against_team( object_permission=_regen_object_permission_dict, team_obj=regenerate_team_table, diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index d980a6f8cfd..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. @@ -609,7 +609,9 @@ def _validate_requested_toolsets( ) -def _extract_requested_vector_stores(object_permission: Optional[dict]) -> set[str]: +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() @@ -620,7 +622,7 @@ def _extract_requested_vector_stores(object_permission: Optional[dict]) -> set[s async def validate_key_vector_stores_against_team( - object_permission: Optional[dict], + object_permission: Optional[ObjectPermissionDict], team_obj: Optional["LiteLLM_TeamTableCachedObj"], is_proxy_admin: bool = False, ) -> None: @@ -647,7 +649,9 @@ async def validate_key_vector_stores_against_team( ) -def _extract_requested_search_tools(object_permission: Optional[dict]) -> List[str]: +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 [] @@ -658,7 +662,7 @@ 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: 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 d81511d2322..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, @@ -1010,3 +1010,18 @@ async def test_empty_object_permission_passes_for_personal_non_admin(): 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)}" + )