mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
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:
parent
2a5790fe55
commit
f2d7cb152a
4 changed files with 63 additions and 27 deletions
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)}"
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue