From 1a3cb2a1d3d43f9f9cbda6d71bc57b0a7dc33ec1 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Fri, 23 Jan 2026 09:26:16 -0800 Subject: [PATCH] init schema.prisma --- litellm/proxy/_types.py | 6 + .../key_management_endpoints.py | 9 +- .../management_endpoints/team_endpoints.py | 7 +- .../policy_engine/attachment_registry.py | 240 +++++++- .../proxy/policy_engine/policy_endpoints.py | 517 ++++++++++++++++++ .../proxy/policy_engine/policy_registry.py | 313 ++++++++++- litellm/proxy/schema.prisma | 33 ++ litellm/types/proxy/policy_engine/__init__.py | 17 + .../proxy/policy_engine/resolver_types.py | 168 +++++- 9 files changed, 1303 insertions(+), 7 deletions(-) create mode 100644 litellm/proxy/policy_engine/policy_endpoints.py diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 2fd03ac2128..0516a4aaa66 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -73,6 +73,7 @@ class SupportedDBObjectType(str, enum.Enum): MODELS = "models" MCP = "mcp" GUARDRAILS = "guardrails" + POLICIES = "policies" VECTOR_STORES = "vector_stores" PASS_THROUGH_ENDPOINTS = "pass_through_endpoints" PROMPTS = "prompts" @@ -844,6 +845,7 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase): model_rpm_limit: Optional[dict] = None model_tpm_limit: Optional[dict] = None guardrails: Optional[List[str]] = None + policies: Optional[List[str]] = None prompts: Optional[List[str]] = None blocked: Optional[bool] = None aliases: Optional[dict] = {} @@ -1477,6 +1479,7 @@ class NewTeamRequest(TeamBase): model_aliases: Optional[dict] = None tags: Optional[list] = None guardrails: Optional[List[str]] = None + policies: Optional[List[str]] = None prompts: Optional[List[str]] = None object_permission: Optional[LiteLLM_ObjectPermissionBase] = None allowed_passthrough_routes: Optional[list] = None @@ -1526,6 +1529,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase): blocked: Optional[bool] = None budget_duration: Optional[str] = None guardrails: Optional[List[str]] = None + policies: Optional[List[str]] = None """ team_id: str # required @@ -1541,6 +1545,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase): tags: Optional[list] = None model_aliases: Optional[dict] = None guardrails: Optional[List[str]] = None + policies: Optional[List[str]] = None object_permission: Optional[LiteLLM_ObjectPermissionBase] = None team_member_budget: Optional[float] = None team_member_budget_duration: Optional[str] = None @@ -3499,6 +3504,7 @@ LiteLLM_ManagementEndpoint_MetadataFields = [ LiteLLM_ManagementEndpoint_MetadataFields_Premium = [ "guardrails", + "policies", "tags", "team_member_key_duration", "prompts", diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index e40a44edf5c..51524947254 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -14,11 +14,11 @@ import copy import json import secrets import traceback -import yaml from datetime import datetime, timedelta, timezone from typing import Any, Dict, List, Literal, Optional, Tuple, cast -from litellm.litellm_core_utils.safe_json_dumps import safe_dumps + import fastapi +import yaml from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, status import litellm @@ -31,6 +31,7 @@ from litellm.constants import ( UI_SESSION_TOKEN_TEAM_ID, ) from litellm.litellm_core_utils.duration_parser import duration_in_seconds +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.proxy._experimental.mcp_server.db import ( rotate_mcp_server_credentials_master_key, ) @@ -2077,6 +2078,7 @@ async def generate_key_helper_fn( # noqa: PLR0915 model_rpm_limit: Optional[dict] = None, model_tpm_limit: Optional[dict] = None, guardrails: Optional[list] = None, + policies: Optional[list] = None, prompts: Optional[list] = None, teams: Optional[list] = None, organization_id: Optional[str] = None, @@ -2139,6 +2141,9 @@ async def generate_key_helper_fn( # noqa: PLR0915 if guardrails is not None: metadata = metadata or {} metadata["guardrails"] = guardrails + if policies is not None: + metadata = metadata or {} + metadata["policies"] = policies if prompts is not None: metadata = metadata or {} metadata["prompts"] = prompts diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 4d313fb1235..c77b60649ab 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -22,11 +22,13 @@ from pydantic import BaseModel import litellm 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 ( BlockTeamRequest, CommonProxyErrors, DeleteTeamRequest, LiteLLM_AuditLogs, + LiteLLM_DeletedTeamTable, LiteLLM_ManagementEndpoint_MetadataFields, LiteLLM_ManagementEndpoint_MetadataFields_Premium, LiteLLM_ModelTable, @@ -34,7 +36,6 @@ from litellm.proxy._types import ( LiteLLM_OrganizationTableWithMembers, LiteLLM_TeamMembership, LiteLLM_TeamTable, - LiteLLM_DeletedTeamTable, LiteLLM_TeamTableCachedObj, LiteLLM_UserTable, LiteLLM_VerificationToken, @@ -102,7 +103,7 @@ from litellm.types.proxy.management_endpoints.team_endpoints import ( TeamMemberAddResult, UpdateTeamMemberPermissionsRequest, ) -from litellm.litellm_core_utils.safe_json_dumps import safe_dumps + router = APIRouter() @@ -689,6 +690,7 @@ async def new_team( # noqa: PLR0915 - organization_id: Optional[str] - The organization id of the team. Default is None. Create via `/organization/new`. - model_aliases: Optional[dict] - Model aliases for the team. [Docs](https://docs.litellm.ai/docs/proxy/team_based_routing#create-team-with-model-alias) - guardrails: Optional[List[str]] - Guardrails for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails) + - policies: Optional[List[str]] - Policies for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails/guardrail_policies) - disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key. - object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "agents": ["agent_1", "agent_2"], "agent_access_groups": ["dev_group"]}. IF null or {} then no object permission. - team_member_budget: Optional[float] - The maximum budget allocated to an individual team member. @@ -1228,6 +1230,7 @@ async def update_team( # noqa: PLR0915 - organization_id: Optional[str] - The organization id of the team. Default is None. Create via `/organization/new`. - model_aliases: Optional[dict] - Model aliases for the team. [Docs](https://docs.litellm.ai/docs/proxy/team_based_routing#create-team-with-model-alias) - guardrails: Optional[List[str]] - Guardrails for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails) + - policies: Optional[List[str]] - Policies for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails/guardrail_policies) - disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key. - object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "agents": ["agent_1", "agent_2"], "agent_access_groups": ["dev_group"]}. IF null or {} then no object permission. - team_member_budget: Optional[float] - The maximum budget allocated to an individual team member. diff --git a/litellm/proxy/policy_engine/attachment_registry.py b/litellm/proxy/policy_engine/attachment_registry.py index b5d6f2fb745..4a335b54747 100644 --- a/litellm/proxy/policy_engine/attachment_registry.py +++ b/litellm/proxy/policy_engine/attachment_registry.py @@ -5,14 +5,20 @@ Attachments define WHERE policies apply, separate from the policy definitions. This allows the same policy to be attached to multiple scopes. """ -from typing import Any, Dict, List, Optional +from datetime import datetime, timezone +from typing import TYPE_CHECKING, Any, Dict, List, Optional from litellm._logging import verbose_proxy_logger from litellm.types.proxy.policy_engine import ( PolicyAttachment, + PolicyAttachmentCreateRequest, + PolicyAttachmentDBResponse, PolicyMatchContext, ) +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + class AttachmentRegistry: """ @@ -188,6 +194,238 @@ class AttachmentRegistry: ) return removed_count + def remove_attachment_by_id(self, attachment_id: str) -> bool: + """ + Remove an attachment by its ID (for DB-synced attachments). + + Args: + attachment_id: The ID of the attachment to remove + + Returns: + True if removed, False if not found + """ + # Note: In-memory attachments don't have IDs, so this is primarily + # for consistency after DB operations + return False + + # ───────────────────────────────────────────────────────────────────────── + # Database CRUD Methods + # ───────────────────────────────────────────────────────────────────────── + + async def add_attachment_to_db( + self, + attachment_request: PolicyAttachmentCreateRequest, + prisma_client: "PrismaClient", + created_by: Optional[str] = None, + ) -> PolicyAttachmentDBResponse: + """ + Add a policy attachment to the database. + + Args: + attachment_request: The attachment creation request + prisma_client: The Prisma client instance + created_by: User who created the attachment + + Returns: + PolicyAttachmentDBResponse with the created attachment + """ + try: + created_attachment = ( + await prisma_client.db.litellm_policyattachmenttable.create( + data={ + "policy_name": attachment_request.policy_name, + "scope": attachment_request.scope, + "teams": attachment_request.teams or [], + "keys": attachment_request.keys or [], + "models": attachment_request.models or [], + "created_at": datetime.now(timezone.utc), + "updated_at": datetime.now(timezone.utc), + "created_by": created_by, + "updated_by": created_by, + } + ) + ) + + # Also add to in-memory registry + attachment = PolicyAttachment( + policy=attachment_request.policy_name, + scope=attachment_request.scope, + teams=attachment_request.teams, + keys=attachment_request.keys, + models=attachment_request.models, + ) + self.add_attachment(attachment) + + return PolicyAttachmentDBResponse( + attachment_id=created_attachment.attachment_id, + policy_name=created_attachment.policy_name, + scope=created_attachment.scope, + teams=created_attachment.teams or [], + keys=created_attachment.keys or [], + models=created_attachment.models or [], + created_at=created_attachment.created_at, + updated_at=created_attachment.updated_at, + created_by=created_attachment.created_by, + updated_by=created_attachment.updated_by, + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error adding attachment to DB: {e}") + raise Exception(f"Error adding attachment to DB: {str(e)}") + + async def delete_attachment_from_db( + self, + attachment_id: str, + prisma_client: "PrismaClient", + ) -> Dict[str, str]: + """ + Delete a policy attachment from the database. + + Args: + attachment_id: The ID of the attachment to delete + prisma_client: The Prisma client instance + + Returns: + Dict with success message + """ + try: + # Get attachment before deleting + attachment = ( + await prisma_client.db.litellm_policyattachmenttable.find_unique( + where={"attachment_id": attachment_id} + ) + ) + + if attachment is None: + raise Exception(f"Attachment with ID {attachment_id} not found") + + # Delete from DB + await prisma_client.db.litellm_policyattachmenttable.delete( + where={"attachment_id": attachment_id} + ) + + # Note: In-memory attachments don't have IDs, so we need to sync from DB + # to properly update in-memory state + await self.sync_attachments_from_db(prisma_client) + + return {"message": f"Attachment {attachment_id} deleted successfully"} + except Exception as e: + verbose_proxy_logger.exception(f"Error deleting attachment from DB: {e}") + raise Exception(f"Error deleting attachment from DB: {str(e)}") + + async def get_attachment_by_id_from_db( + self, + attachment_id: str, + prisma_client: "PrismaClient", + ) -> Optional[PolicyAttachmentDBResponse]: + """ + Get a policy attachment by ID from the database. + + Args: + attachment_id: The ID of the attachment to retrieve + prisma_client: The Prisma client instance + + Returns: + PolicyAttachmentDBResponse if found, None otherwise + """ + try: + attachment = ( + await prisma_client.db.litellm_policyattachmenttable.find_unique( + where={"attachment_id": attachment_id} + ) + ) + + if attachment is None: + return None + + return PolicyAttachmentDBResponse( + attachment_id=attachment.attachment_id, + policy_name=attachment.policy_name, + scope=attachment.scope, + teams=attachment.teams or [], + keys=attachment.keys or [], + models=attachment.models or [], + created_at=attachment.created_at, + updated_at=attachment.updated_at, + created_by=attachment.created_by, + updated_by=attachment.updated_by, + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error getting attachment from DB: {e}") + raise Exception(f"Error getting attachment from DB: {str(e)}") + + async def get_all_attachments_from_db( + self, + prisma_client: "PrismaClient", + ) -> List[PolicyAttachmentDBResponse]: + """ + Get all policy attachments from the database. + + Args: + prisma_client: The Prisma client instance + + Returns: + List of PolicyAttachmentDBResponse objects + """ + try: + attachments = ( + await prisma_client.db.litellm_policyattachmenttable.find_many( + order={"created_at": "desc"}, + ) + ) + + return [ + PolicyAttachmentDBResponse( + attachment_id=a.attachment_id, + policy_name=a.policy_name, + scope=a.scope, + teams=a.teams or [], + keys=a.keys or [], + models=a.models or [], + created_at=a.created_at, + updated_at=a.updated_at, + created_by=a.created_by, + updated_by=a.updated_by, + ) + for a in attachments + ] + except Exception as e: + verbose_proxy_logger.exception(f"Error getting attachments from DB: {e}") + raise Exception(f"Error getting attachments from DB: {str(e)}") + + async def sync_attachments_from_db( + self, + prisma_client: "PrismaClient", + ) -> None: + """ + Sync policy attachments from the database to in-memory registry. + + Args: + prisma_client: The Prisma client instance + """ + try: + attachments = await self.get_all_attachments_from_db(prisma_client) + + # Clear existing attachments and reload from DB + self._attachments = [] + + for attachment_response in attachments: + attachment = PolicyAttachment( + policy=attachment_response.policy_name, + scope=attachment_response.scope, + teams=attachment_response.teams if attachment_response.teams else None, + keys=attachment_response.keys if attachment_response.keys else None, + models=attachment_response.models if attachment_response.models else None, + ) + self._attachments.append(attachment) + + self._initialized = True + verbose_proxy_logger.info( + f"Synced {len(attachments)} attachments from DB to in-memory registry" + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error syncing attachments from DB: {e}") + raise Exception(f"Error syncing attachments from DB: {str(e)}") + # Global singleton instance _attachment_registry: Optional[AttachmentRegistry] = None diff --git a/litellm/proxy/policy_engine/policy_endpoints.py b/litellm/proxy/policy_engine/policy_endpoints.py new file mode 100644 index 00000000000..a4aac7caf2c --- /dev/null +++ b/litellm/proxy/policy_engine/policy_endpoints.py @@ -0,0 +1,517 @@ +""" +CRUD ENDPOINTS FOR POLICIES + +Provides REST API endpoints for managing policies and policy attachments. +""" + +from typing import List, Optional + +from fastapi import APIRouter, Depends, HTTPException + +from litellm._logging import verbose_proxy_logger +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry +from litellm.proxy.policy_engine.policy_registry import get_policy_registry +from litellm.types.proxy.policy_engine import ( + PolicyAttachmentCreateRequest, + PolicyAttachmentDBResponse, + PolicyAttachmentListResponse, + PolicyCreateRequest, + PolicyDBResponse, + PolicyListDBResponse, + PolicyUpdateRequest, +) + +router = APIRouter() + +# Get singleton instances +POLICY_REGISTRY = get_policy_registry() +ATTACHMENT_REGISTRY = get_attachment_registry() + + +# ───────────────────────────────────────────────────────────────────────────── +# Policy CRUD Endpoints +# ───────────────────────────────────────────────────────────────────────────── + + +@router.get( + "/policies/list", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyListDBResponse, +) +async def list_policies(): + """ + List all policies from the database. + + Example Request: + ```bash + curl -X GET "http://localhost:4000/policies/list" \\ + -H "Authorization: Bearer " + ``` + + Example Response: + ```json + { + "policies": [ + { + "policy_id": "123e4567-e89b-12d3-a456-426614174000", + "policy_name": "global-baseline", + "inherit": null, + "description": "Base guardrails for all requests", + "guardrails_add": ["pii_masking"], + "guardrails_remove": [], + "condition": null, + "created_at": "2024-01-01T00:00:00Z", + "updated_at": "2024-01-01T00:00:00Z" + } + ], + "total_count": 1 + } + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + policies = await POLICY_REGISTRY.get_all_policies_from_db(prisma_client) + return PolicyListDBResponse(policies=policies, total_count=len(policies)) + except Exception as e: + verbose_proxy_logger.exception(f"Error listing policies: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +@router.post( + "/policies", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyDBResponse, +) +async def create_policy( + request: PolicyCreateRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Create a new policy. + + Example Request: + ```bash + curl -X POST "http://localhost:4000/policies" \\ + -H "Authorization: Bearer " \\ + -H "Content-Type: application/json" \\ + -d '{ + "policy_name": "global-baseline", + "description": "Base guardrails for all requests", + "guardrails_add": ["pii_masking", "prompt_injection"], + "guardrails_remove": [] + }' + ``` + + Example Response: + ```json + { + "policy_id": "123e4567-e89b-12d3-a456-426614174000", + "policy_name": "global-baseline", + "inherit": null, + "description": "Base guardrails for all requests", + "guardrails_add": ["pii_masking", "prompt_injection"], + "guardrails_remove": [], + "condition": null, + "created_at": "2024-01-01T00:00:00Z", + "updated_at": "2024-01-01T00:00:00Z" + } + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + created_by = user_api_key_dict.user_id + result = await POLICY_REGISTRY.add_policy_to_db( + policy_request=request, + prisma_client=prisma_client, + created_by=created_by, + ) + return result + except Exception as e: + verbose_proxy_logger.exception(f"Error creating policy: {e}") + if "unique constraint" in str(e).lower(): + raise HTTPException( + status_code=400, + detail=f"Policy with name '{request.policy_name}' already exists", + ) + raise HTTPException(status_code=500, detail=str(e)) + + +@router.get( + "/policies/{policy_id}", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyDBResponse, +) +async def get_policy(policy_id: str): + """ + Get a policy by ID. + + Example Request: + ```bash + curl -X GET "http://localhost:4000/policies/123e4567-e89b-12d3-a456-426614174000" \\ + -H "Authorization: Bearer " + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + result = await POLICY_REGISTRY.get_policy_by_id_from_db( + policy_id=policy_id, + prisma_client=prisma_client, + ) + if result is None: + raise HTTPException( + status_code=404, detail=f"Policy with ID {policy_id} not found" + ) + return result + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception(f"Error getting policy: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +@router.put( + "/policies/{policy_id}", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyDBResponse, +) +async def update_policy( + policy_id: str, + request: PolicyUpdateRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Update an existing policy. + + Example Request: + ```bash + curl -X PUT "http://localhost:4000/policies/123e4567-e89b-12d3-a456-426614174000" \\ + -H "Authorization: Bearer " \\ + -H "Content-Type: application/json" \\ + -d '{ + "description": "Updated description", + "guardrails_add": ["pii_masking", "toxicity_filter"] + }' + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + # Check if policy exists + existing = await POLICY_REGISTRY.get_policy_by_id_from_db( + policy_id=policy_id, + prisma_client=prisma_client, + ) + if existing is None: + raise HTTPException( + status_code=404, detail=f"Policy with ID {policy_id} not found" + ) + + updated_by = user_api_key_dict.user_id + result = await POLICY_REGISTRY.update_policy_in_db( + policy_id=policy_id, + policy_request=request, + prisma_client=prisma_client, + updated_by=updated_by, + ) + return result + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception(f"Error updating policy: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +@router.delete( + "/policies/{policy_id}", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], +) +async def delete_policy(policy_id: str): + """ + Delete a policy. + + Example Request: + ```bash + curl -X DELETE "http://localhost:4000/policies/123e4567-e89b-12d3-a456-426614174000" \\ + -H "Authorization: Bearer " + ``` + + Example Response: + ```json + { + "message": "Policy 123e4567-e89b-12d3-a456-426614174000 deleted successfully" + } + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + # Check if policy exists + existing = await POLICY_REGISTRY.get_policy_by_id_from_db( + policy_id=policy_id, + prisma_client=prisma_client, + ) + if existing is None: + raise HTTPException( + status_code=404, detail=f"Policy with ID {policy_id} not found" + ) + + result = await POLICY_REGISTRY.delete_policy_from_db( + policy_id=policy_id, + prisma_client=prisma_client, + ) + return result + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception(f"Error deleting policy: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +# ───────────────────────────────────────────────────────────────────────────── +# Policy Attachment CRUD Endpoints +# ───────────────────────────────────────────────────────────────────────────── + + +@router.get( + "/policies/attachments/list", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyAttachmentListResponse, +) +async def list_policy_attachments(): + """ + List all policy attachments from the database. + + Example Request: + ```bash + curl -X GET "http://localhost:4000/policies/attachments/list" \\ + -H "Authorization: Bearer " + ``` + + Example Response: + ```json + { + "attachments": [ + { + "attachment_id": "123e4567-e89b-12d3-a456-426614174000", + "policy_name": "global-baseline", + "scope": "*", + "teams": [], + "keys": [], + "models": [], + "created_at": "2024-01-01T00:00:00Z", + "updated_at": "2024-01-01T00:00:00Z" + } + ], + "total_count": 1 + } + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + attachments = await ATTACHMENT_REGISTRY.get_all_attachments_from_db( + prisma_client + ) + return PolicyAttachmentListResponse( + attachments=attachments, total_count=len(attachments) + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error listing policy attachments: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +@router.post( + "/policies/attachments", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyAttachmentDBResponse, +) +async def create_policy_attachment( + request: PolicyAttachmentCreateRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Create a new policy attachment. + + Example Request: + ```bash + curl -X POST "http://localhost:4000/policies/attachments" \\ + -H "Authorization: Bearer " \\ + -H "Content-Type: application/json" \\ + -d '{ + "policy_name": "global-baseline", + "scope": "*" + }' + ``` + + Example with team-specific attachment: + ```bash + curl -X POST "http://localhost:4000/policies/attachments" \\ + -H "Authorization: Bearer " \\ + -H "Content-Type: application/json" \\ + -d '{ + "policy_name": "healthcare-compliance", + "teams": ["healthcare-team", "medical-research"] + }' + ``` + + Example Response: + ```json + { + "attachment_id": "123e4567-e89b-12d3-a456-426614174000", + "policy_name": "global-baseline", + "scope": "*", + "teams": [], + "keys": [], + "models": [], + "created_at": "2024-01-01T00:00:00Z", + "updated_at": "2024-01-01T00:00:00Z" + } + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + # Verify the policy exists + policy = await POLICY_REGISTRY.get_all_policies_from_db(prisma_client) + policy_names = [p.policy_name for p in policy] + if request.policy_name not in policy_names: + raise HTTPException( + status_code=404, + detail=f"Policy '{request.policy_name}' not found. Create the policy first.", + ) + + created_by = user_api_key_dict.user_id + result = await ATTACHMENT_REGISTRY.add_attachment_to_db( + attachment_request=request, + prisma_client=prisma_client, + created_by=created_by, + ) + return result + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception(f"Error creating policy attachment: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +@router.get( + "/policies/attachments/{attachment_id}", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyAttachmentDBResponse, +) +async def get_policy_attachment(attachment_id: str): + """ + Get a policy attachment by ID. + + Example Request: + ```bash + curl -X GET "http://localhost:4000/policies/attachments/123e4567-e89b-12d3-a456-426614174000" \\ + -H "Authorization: Bearer " + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + result = await ATTACHMENT_REGISTRY.get_attachment_by_id_from_db( + attachment_id=attachment_id, + prisma_client=prisma_client, + ) + if result is None: + raise HTTPException( + status_code=404, + detail=f"Attachment with ID {attachment_id} not found", + ) + return result + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception(f"Error getting policy attachment: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +@router.delete( + "/policies/attachments/{attachment_id}", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], +) +async def delete_policy_attachment(attachment_id: str): + """ + Delete a policy attachment. + + Example Request: + ```bash + curl -X DELETE "http://localhost:4000/policies/attachments/123e4567-e89b-12d3-a456-426614174000" \\ + -H "Authorization: Bearer " + ``` + + Example Response: + ```json + { + "message": "Attachment 123e4567-e89b-12d3-a456-426614174000 deleted successfully" + } + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + # Check if attachment exists + existing = await ATTACHMENT_REGISTRY.get_attachment_by_id_from_db( + attachment_id=attachment_id, + prisma_client=prisma_client, + ) + if existing is None: + raise HTTPException( + status_code=404, + detail=f"Attachment with ID {attachment_id} not found", + ) + + result = await ATTACHMENT_REGISTRY.delete_attachment_from_db( + attachment_id=attachment_id, + prisma_client=prisma_client, + ) + return result + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception(f"Error deleting policy attachment: {e}") + raise HTTPException(status_code=500, detail=str(e)) diff --git a/litellm/proxy/policy_engine/policy_registry.py b/litellm/proxy/policy_engine/policy_registry.py index 68485f92489..b83c4d3e246 100644 --- a/litellm/proxy/policy_engine/policy_registry.py +++ b/litellm/proxy/policy_engine/policy_registry.py @@ -7,15 +7,23 @@ Policies define WHAT guardrails to apply. WHERE they apply is defined by policy_attachments (see AttachmentRegistry). """ -from typing import Any, Dict, List, Optional +from datetime import datetime, timezone +from typing import TYPE_CHECKING, Any, Dict, List, Optional from litellm._logging import verbose_proxy_logger +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.types.proxy.policy_engine import ( Policy, PolicyCondition, + PolicyCreateRequest, + PolicyDBResponse, PolicyGuardrails, + PolicyUpdateRequest, ) +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + class PolicyRegistry: """ @@ -178,6 +186,309 @@ class PolicyRegistry: return True return False + # ───────────────────────────────────────────────────────────────────────── + # Database CRUD Methods + # ───────────────────────────────────────────────────────────────────────── + + async def add_policy_to_db( + self, + policy_request: PolicyCreateRequest, + prisma_client: "PrismaClient", + created_by: Optional[str] = None, + ) -> PolicyDBResponse: + """ + Add a policy to the database. + + Args: + policy_request: The policy creation request + prisma_client: The Prisma client instance + created_by: User who created the policy + + Returns: + PolicyDBResponse with the created policy + """ + try: + condition_json = None + if policy_request.condition: + condition_json = safe_dumps(policy_request.condition.model_dump()) + + created_policy = await prisma_client.db.litellm_policytable.create( + data={ + "policy_name": policy_request.policy_name, + "inherit": policy_request.inherit, + "description": policy_request.description, + "guardrails_add": policy_request.guardrails_add or [], + "guardrails_remove": policy_request.guardrails_remove or [], + "condition": condition_json, + "created_at": datetime.now(timezone.utc), + "updated_at": datetime.now(timezone.utc), + "created_by": created_by, + "updated_by": created_by, + } + ) + + # Also add to in-memory registry + policy = self._parse_policy( + policy_request.policy_name, + { + "inherit": policy_request.inherit, + "description": policy_request.description, + "guardrails": { + "add": policy_request.guardrails_add, + "remove": policy_request.guardrails_remove, + }, + "condition": policy_request.condition.model_dump() + if policy_request.condition + else None, + }, + ) + self.add_policy(policy_request.policy_name, policy) + + return PolicyDBResponse( + policy_id=created_policy.policy_id, + policy_name=created_policy.policy_name, + inherit=created_policy.inherit, + description=created_policy.description, + guardrails_add=created_policy.guardrails_add or [], + guardrails_remove=created_policy.guardrails_remove or [], + condition=created_policy.condition, + created_at=created_policy.created_at, + updated_at=created_policy.updated_at, + created_by=created_policy.created_by, + updated_by=created_policy.updated_by, + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error adding policy to DB: {e}") + raise Exception(f"Error adding policy to DB: {str(e)}") + + async def update_policy_in_db( + self, + policy_id: str, + policy_request: PolicyUpdateRequest, + prisma_client: "PrismaClient", + updated_by: Optional[str] = None, + ) -> PolicyDBResponse: + """ + Update a policy in the database. + + Args: + policy_id: The ID of the policy to update + policy_request: The policy update request + prisma_client: The Prisma client instance + updated_by: User who updated the policy + + Returns: + PolicyDBResponse with the updated policy + """ + try: + # Build update data - only include fields that are set + update_data: Dict[str, Any] = { + "updated_at": datetime.now(timezone.utc), + "updated_by": updated_by, + } + + if policy_request.policy_name is not None: + update_data["policy_name"] = policy_request.policy_name + if policy_request.inherit is not None: + update_data["inherit"] = policy_request.inherit + if policy_request.description is not None: + update_data["description"] = policy_request.description + if policy_request.guardrails_add is not None: + update_data["guardrails_add"] = policy_request.guardrails_add + if policy_request.guardrails_remove is not None: + update_data["guardrails_remove"] = policy_request.guardrails_remove + if policy_request.condition is not None: + update_data["condition"] = safe_dumps( + policy_request.condition.model_dump() + ) + + updated_policy = await prisma_client.db.litellm_policytable.update( + where={"policy_id": policy_id}, + data=update_data, + ) + + # Update in-memory registry + policy = self._parse_policy( + updated_policy.policy_name, + { + "inherit": updated_policy.inherit, + "description": updated_policy.description, + "guardrails": { + "add": updated_policy.guardrails_add, + "remove": updated_policy.guardrails_remove, + }, + "condition": updated_policy.condition, + }, + ) + self.add_policy(updated_policy.policy_name, policy) + + return PolicyDBResponse( + policy_id=updated_policy.policy_id, + policy_name=updated_policy.policy_name, + inherit=updated_policy.inherit, + description=updated_policy.description, + guardrails_add=updated_policy.guardrails_add or [], + guardrails_remove=updated_policy.guardrails_remove or [], + condition=updated_policy.condition, + created_at=updated_policy.created_at, + updated_at=updated_policy.updated_at, + created_by=updated_policy.created_by, + updated_by=updated_policy.updated_by, + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error updating policy in DB: {e}") + raise Exception(f"Error updating policy in DB: {str(e)}") + + async def delete_policy_from_db( + self, + policy_id: str, + prisma_client: "PrismaClient", + ) -> Dict[str, str]: + """ + Delete a policy from the database. + + Args: + policy_id: The ID of the policy to delete + prisma_client: The Prisma client instance + + Returns: + Dict with success message + """ + try: + # Get policy name before deleting + policy = await prisma_client.db.litellm_policytable.find_unique( + where={"policy_id": policy_id} + ) + + if policy is None: + raise Exception(f"Policy with ID {policy_id} not found") + + # Delete from DB + await prisma_client.db.litellm_policytable.delete( + where={"policy_id": policy_id} + ) + + # Remove from in-memory registry + self.remove_policy(policy.policy_name) + + return {"message": f"Policy {policy_id} deleted successfully"} + except Exception as e: + verbose_proxy_logger.exception(f"Error deleting policy from DB: {e}") + raise Exception(f"Error deleting policy from DB: {str(e)}") + + async def get_policy_by_id_from_db( + self, + policy_id: str, + prisma_client: "PrismaClient", + ) -> Optional[PolicyDBResponse]: + """ + Get a policy by ID from the database. + + Args: + policy_id: The ID of the policy to retrieve + prisma_client: The Prisma client instance + + Returns: + PolicyDBResponse if found, None otherwise + """ + try: + policy = await prisma_client.db.litellm_policytable.find_unique( + where={"policy_id": policy_id} + ) + + if policy is None: + return None + + return PolicyDBResponse( + policy_id=policy.policy_id, + policy_name=policy.policy_name, + inherit=policy.inherit, + description=policy.description, + guardrails_add=policy.guardrails_add or [], + guardrails_remove=policy.guardrails_remove or [], + condition=policy.condition, + created_at=policy.created_at, + updated_at=policy.updated_at, + created_by=policy.created_by, + updated_by=policy.updated_by, + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error getting policy from DB: {e}") + raise Exception(f"Error getting policy from DB: {str(e)}") + + async def get_all_policies_from_db( + self, + prisma_client: "PrismaClient", + ) -> List[PolicyDBResponse]: + """ + Get all policies from the database. + + Args: + prisma_client: The Prisma client instance + + Returns: + List of PolicyDBResponse objects + """ + try: + policies = await prisma_client.db.litellm_policytable.find_many( + order={"created_at": "desc"}, + ) + + return [ + PolicyDBResponse( + policy_id=p.policy_id, + policy_name=p.policy_name, + inherit=p.inherit, + description=p.description, + guardrails_add=p.guardrails_add or [], + guardrails_remove=p.guardrails_remove or [], + condition=p.condition, + created_at=p.created_at, + updated_at=p.updated_at, + created_by=p.created_by, + updated_by=p.updated_by, + ) + for p in policies + ] + except Exception as e: + verbose_proxy_logger.exception(f"Error getting policies from DB: {e}") + raise Exception(f"Error getting policies from DB: {str(e)}") + + async def sync_policies_from_db( + self, + prisma_client: "PrismaClient", + ) -> None: + """ + Sync policies from the database to in-memory registry. + + Args: + prisma_client: The Prisma client instance + """ + try: + policies = await self.get_all_policies_from_db(prisma_client) + + for policy_response in policies: + policy = self._parse_policy( + policy_response.policy_name, + { + "inherit": policy_response.inherit, + "description": policy_response.description, + "guardrails": { + "add": policy_response.guardrails_add, + "remove": policy_response.guardrails_remove, + }, + "condition": policy_response.condition, + }, + ) + self.add_policy(policy_response.policy_name, policy) + + verbose_proxy_logger.info( + f"Synced {len(policies)} policies from DB to in-memory registry" + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error syncing policies from DB: {e}") + raise Exception(f"Error syncing policies from DB: {str(e)}") + # Global singleton instance _policy_registry: Optional[PolicyRegistry] = None diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 22888f6d3af..ebbd63f3b10 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -880,3 +880,36 @@ model LiteLLM_ClaudeCodePluginTable { @@index([name]) @@map("litellm_claudecodeplugin") } + +// Policy table for storing policy configurations +model LiteLLM_PolicyTable { + policy_id String @id @default(uuid()) + policy_name String @unique + inherit String? // Parent policy name for inheritance + description String? + guardrails_add String[] @default([]) // Guardrails to add + guardrails_remove String[] @default([]) // Guardrails to remove (from inherited) + condition Json? // Condition for when policy applies, e.g. {model: "gpt-4.*"} + created_at DateTime @default(now()) + updated_at DateTime @updatedAt + created_by String? + updated_by String? + + @@index([policy_name]) +} + +// Policy attachment table for storing where policies apply +model LiteLLM_PolicyAttachmentTable { + attachment_id String @id @default(uuid()) + policy_name String // Reference to policy + scope String? // "*" for global scope + teams String[] @default([]) // Team aliases or patterns + keys String[] @default([]) // Key aliases or patterns + models String[] @default([]) // Model names or patterns + created_at DateTime @default(now()) + updated_at DateTime @updatedAt + created_by String? + updated_by String? + + @@index([policy_name]) +} diff --git a/litellm/types/proxy/policy_engine/__init__.py b/litellm/types/proxy/policy_engine/__init__.py index 50ed4581013..bc54c3eb36b 100644 --- a/litellm/types/proxy/policy_engine/__init__.py +++ b/litellm/types/proxy/policy_engine/__init__.py @@ -19,13 +19,21 @@ from litellm.types.proxy.policy_engine.policy_types import ( PolicyScope, ) from litellm.types.proxy.policy_engine.resolver_types import ( + PolicyAttachmentCreateRequest, + PolicyAttachmentDBResponse, + PolicyAttachmentListResponse, + PolicyConditionRequest, + PolicyCreateRequest, + PolicyDBResponse, PolicyGuardrailsResponse, PolicyInfoResponse, + PolicyListDBResponse, PolicyListResponse, PolicyMatchContext, PolicyScopeResponse, PolicySummaryItem, PolicyTestResponse, + PolicyUpdateRequest, ResolvedPolicy, ) from litellm.types.proxy.policy_engine.validation_types import ( @@ -58,4 +66,13 @@ __all__ = [ "PolicyScopeResponse", "PolicySummaryItem", "PolicyTestResponse", + # CRUD Request/Response types + "PolicyConditionRequest", + "PolicyCreateRequest", + "PolicyUpdateRequest", + "PolicyDBResponse", + "PolicyListDBResponse", + "PolicyAttachmentCreateRequest", + "PolicyAttachmentDBResponse", + "PolicyAttachmentListResponse", ] diff --git a/litellm/types/proxy/policy_engine/resolver_types.py b/litellm/types/proxy/policy_engine/resolver_types.py index 81ae248d436..9488b8b0841 100644 --- a/litellm/types/proxy/policy_engine/resolver_types.py +++ b/litellm/types/proxy/policy_engine/resolver_types.py @@ -5,7 +5,8 @@ These types are used for matching requests to policies and resolving the final guardrails list. """ -from typing import Dict, List, Optional +from datetime import datetime +from typing import Any, Dict, List, Optional from pydantic import BaseModel, ConfigDict, Field @@ -108,3 +109,168 @@ class PolicyTestResponse(BaseModel): matching_policies: List[str] resolved_guardrails: List[str] message: Optional[str] = None + + +# ───────────────────────────────────────────────────────────────────────────── +# CRUD Request/Response Types for Policy Endpoints +# ───────────────────────────────────────────────────────────────────────────── + + +class PolicyConditionRequest(BaseModel): + """Condition for when a policy applies.""" + + model: Optional[str] = Field( + default=None, + description="Model name pattern (exact match or regex) for when policy applies.", + ) + + +class PolicyCreateRequest(BaseModel): + """Request body for creating a new policy.""" + + policy_name: str = Field(description="Unique name for the policy.") + inherit: Optional[str] = Field( + default=None, + description="Name of parent policy to inherit from.", + ) + description: Optional[str] = Field( + default=None, + description="Human-readable description of the policy.", + ) + guardrails_add: Optional[List[str]] = Field( + default=None, + description="List of guardrail names to add.", + ) + guardrails_remove: Optional[List[str]] = Field( + default=None, + description="List of guardrail names to remove (from inherited).", + ) + condition: Optional[PolicyConditionRequest] = Field( + default=None, + description="Condition for when this policy applies.", + ) + + +class PolicyUpdateRequest(BaseModel): + """Request body for updating a policy.""" + + policy_name: Optional[str] = Field( + default=None, + description="New name for the policy.", + ) + inherit: Optional[str] = Field( + default=None, + description="Name of parent policy to inherit from.", + ) + description: Optional[str] = Field( + default=None, + description="Human-readable description of the policy.", + ) + guardrails_add: Optional[List[str]] = Field( + default=None, + description="List of guardrail names to add.", + ) + guardrails_remove: Optional[List[str]] = Field( + default=None, + description="List of guardrail names to remove (from inherited).", + ) + condition: Optional[PolicyConditionRequest] = Field( + default=None, + description="Condition for when this policy applies.", + ) + + +class PolicyDBResponse(BaseModel): + """Response for a policy from the database.""" + + policy_id: str = Field(description="Unique ID of the policy.") + policy_name: str = Field(description="Name of the policy.") + inherit: Optional[str] = Field(default=None, description="Parent policy name.") + description: Optional[str] = Field(default=None, description="Policy description.") + guardrails_add: List[str] = Field( + default_factory=list, description="Guardrails to add." + ) + guardrails_remove: List[str] = Field( + default_factory=list, description="Guardrails to remove." + ) + condition: Optional[Dict[str, Any]] = Field( + default=None, description="Policy condition." + ) + created_at: Optional[datetime] = Field( + default=None, description="When the policy was created." + ) + updated_at: Optional[datetime] = Field( + default=None, description="When the policy was last updated." + ) + created_by: Optional[str] = Field(default=None, description="Who created the policy.") + updated_by: Optional[str] = Field( + default=None, description="Who last updated the policy." + ) + + +class PolicyListDBResponse(BaseModel): + """Response for listing policies from the database.""" + + policies: List[PolicyDBResponse] = Field( + default_factory=list, description="List of policies." + ) + total_count: int = Field(default=0, description="Total number of policies.") + + +# ───────────────────────────────────────────────────────────────────────────── +# Policy Attachment CRUD Types +# ───────────────────────────────────────────────────────────────────────────── + + +class PolicyAttachmentCreateRequest(BaseModel): + """Request body for creating a policy attachment.""" + + policy_name: str = Field(description="Name of the policy to attach.") + scope: Optional[str] = Field( + default=None, + description="Use '*' for global scope (applies to all requests).", + ) + teams: Optional[List[str]] = Field( + default=None, + description="Team aliases or patterns this attachment applies to.", + ) + keys: Optional[List[str]] = Field( + default=None, + description="Key aliases or patterns this attachment applies to.", + ) + models: Optional[List[str]] = Field( + default=None, + description="Model names or patterns this attachment applies to.", + ) + + +class PolicyAttachmentDBResponse(BaseModel): + """Response for a policy attachment from the database.""" + + attachment_id: str = Field(description="Unique ID of the attachment.") + policy_name: str = Field(description="Name of the attached policy.") + scope: Optional[str] = Field(default=None, description="Scope of the attachment.") + teams: List[str] = Field(default_factory=list, description="Team patterns.") + keys: List[str] = Field(default_factory=list, description="Key patterns.") + models: List[str] = Field(default_factory=list, description="Model patterns.") + created_at: Optional[datetime] = Field( + default=None, description="When the attachment was created." + ) + updated_at: Optional[datetime] = Field( + default=None, description="When the attachment was last updated." + ) + created_by: Optional[str] = Field( + default=None, description="Who created the attachment." + ) + updated_by: Optional[str] = Field( + default=None, description="Who last updated the attachment." + ) + + +class PolicyAttachmentListResponse(BaseModel): + """Response for listing policy attachments.""" + + attachments: List[PolicyAttachmentDBResponse] = Field( + default_factory=list, description="List of policy attachments." + ) + total_count: int = Field(default=0, description="Total number of attachments.")