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.
This commit is contained in:
yucheng-berriai 2026-06-26 16:33:52 -07:00
parent 2a5790fe55
commit f2d7cb152a
4 changed files with 63 additions and 27 deletions

View file

@ -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

View file

@ -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,

View file

@ -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:

View file

@ -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)}"
)