diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py index 930cb7c73ae..c90917d5341 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py @@ -12,7 +12,7 @@ Endpoints for /project operations import json from collections.abc import Mapping, Sequence -from typing import TYPE_CHECKING, Final +from typing import TYPE_CHECKING, Final, cast from fastapi import APIRouter, Depends, HTTPException, Request from pydantic import TypeAdapter @@ -56,6 +56,15 @@ def _project_table(prisma_client: PrismaClient) -> TableActions["prisma_models.L return ProjectRepository(prisma_client).table +def _writer_project_table( + prisma_client: PrismaClient, +) -> TableActions["prisma_models.LiteLLM_ProjectTable"]: + return cast( # cast-ok: writer_db exposes generated Prisma tables through a dynamic wrapper + "TableActions[prisma_models.LiteLLM_ProjectTable]", + prisma_client.writer_db.litellm_projecttable, + ) + + def _verification_token_table( prisma_client: PrismaClient, ) -> TableActions["prisma_models.LiteLLM_VerificationToken"]: @@ -793,7 +802,7 @@ async def update_project( } if data.team_id is not None: - current_project_record: Final = await prisma_client.writer_db.litellm_projecttable.find_unique( + current_project_record: Final = await _writer_project_table(prisma_client).find_unique( where={"project_id": data.project_id} ) current_project: Final = ( diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 579d6848152..75c79fd77e1 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -262,6 +262,15 @@ def _prisma_table( ) +def _writer_project_table( + prisma_client: PrismaClient, +) -> "TableActions[prisma_models.LiteLLM_ProjectTable]": + return cast( # cast-ok: writer_db exposes generated Prisma tables through a dynamic wrapper + "TableActions[prisma_models.LiteLLM_ProjectTable]", + prisma_client.writer_db.litellm_projecttable, + ) + + def _deleted_verification_token_table( prisma_client: PrismaClient, ) -> "TableActions[prisma_models.LiteLLM_DeletedVerificationToken]": @@ -1332,7 +1341,8 @@ async def _common_key_generation_helper( apply_enterprise_key_management_params, ) - data = apply_enterprise_key_management_params(data, team_table) + enterprise_data: Final[object] = apply_enterprise_key_management_params(data, team_table) + data = GenerateKeyRequest.model_validate(enterprise_data) except Exception as e: verbose_proxy_logger.debug( "litellm.proxy.proxy_server.generate_key_fn(): Enterprise key management params not applied - %s", e @@ -1817,9 +1827,7 @@ async def _check_key_project_team( key_team_id: str | None, prisma_client: PrismaClient, ) -> None: - project_record: Final = await prisma_client.writer_db.litellm_projecttable.find_unique( - where={"project_id": project_id} - ) + project_record: Final = await _writer_project_table(prisma_client).find_unique(where={"project_id": project_id}) project_obj: Final = ( LiteLLM_ProjectTable.model_validate(record_to_dict(project_record)) if project_record is not None else None )