From fc181d425e6a546e87ed961cc38ae46301f3c254 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 21 Feb 2026 15:45:13 -0800 Subject: [PATCH 1/6] feat: initial commit, adding support for policy versioning on litellm --- .../proxy/policy_engine/policy_endpoints.py | 219 +++++++- .../proxy/policy_engine/policy_registry.py | 509 ++++++++++++++---- litellm/proxy/schema.prisma | 35 +- litellm/types/proxy/policy_engine/__init__.py | 65 +-- .../proxy/policy_engine/resolver_types.py | 60 +++ schema.prisma | 35 +- .../policy_engine/test_policy_versioning.py | 442 +++++++++++++++ .../test_policy_versioning_e2e.py | 217 ++++++++ 8 files changed, 1396 insertions(+), 186 deletions(-) create mode 100644 tests/test_litellm/proxy/policy_engine/test_policy_versioning.py create mode 100644 tests/test_litellm/proxy/policy_engine/test_policy_versioning_e2e.py diff --git a/litellm/proxy/policy_engine/policy_endpoints.py b/litellm/proxy/policy_engine/policy_endpoints.py index af12a8598f6..d8de028d6a0 100644 --- a/litellm/proxy/policy_engine/policy_endpoints.py +++ b/litellm/proxy/policy_engine/policy_endpoints.py @@ -4,25 +4,24 @@ CRUD ENDPOINTS FOR POLICIES Provides REST API endpoints for managing policies and policy attachments. """ +from typing import 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.attachment_registry import \ + get_attachment_registry from litellm.proxy.policy_engine.pipeline_executor import PipelineExecutor from litellm.proxy.policy_engine.policy_registry import get_policy_registry from litellm.types.proxy.policy_engine import ( - GuardrailPipeline, - PipelineTestRequest, - PolicyAttachmentCreateRequest, - PolicyAttachmentDBResponse, - PolicyAttachmentListResponse, - PolicyCreateRequest, - PolicyDBResponse, - PolicyListDBResponse, - PolicyUpdateRequest, -) + GuardrailPipeline, PipelineTestRequest, PolicyAttachmentCreateRequest, + PolicyAttachmentDBResponse, PolicyAttachmentListResponse, + PolicyCreateRequest, PolicyDBResponse, PolicyListDBResponse, + PolicyUpdateRequest, PolicyVersionCompareResponse, + PolicyVersionCreateRequest, PolicyVersionListResponse, + PolicyVersionStatusUpdateRequest) router = APIRouter() @@ -38,14 +37,20 @@ router = APIRouter() dependencies=[Depends(user_api_key_auth)], response_model=PolicyListDBResponse, ) -async def list_policies(): +async def list_policies(version_status: Optional[str] = None): """ - List all policies from the database. + List all policies from the database. Optionally filter by version_status. + + Query params: + - version_status: Optional. One of "draft", "published", "production". + If omitted, all versions are returned. Example Request: ```bash curl -X GET "http://localhost:4000/policies/list" \\ -H "Authorization: Bearer " + curl -X GET "http://localhost:4000/policies/list?version_status=production" \\ + -H "Authorization: Bearer " ``` Example Response: @@ -55,6 +60,8 @@ async def list_policies(): { "policy_id": "123e4567-e89b-12d3-a456-426614174000", "policy_name": "global-baseline", + "version_number": 1, + "version_status": "production", "inherit": null, "description": "Base guardrails for all requests", "guardrails_add": ["pii_masking"], @@ -74,7 +81,9 @@ async def list_policies(): raise HTTPException(status_code=500, detail="Database not connected") try: - policies = await get_policy_registry().get_all_policies_from_db(prisma_client) + policies = await get_policy_registry().get_all_policies_from_db( + prisma_client, version_status=version_status + ) return PolicyListDBResponse(policies=policies, total_count=len(policies)) except Exception as e: verbose_proxy_logger.exception(f"Error listing policies: {e}") @@ -145,6 +154,172 @@ async def create_policy( raise HTTPException(status_code=500, detail=str(e)) +# ───────────────────────────────────────────────────────────────────────────── +# Policy Versioning Endpoints (must be before /policies/{policy_id} to avoid path conflicts) +# ───────────────────────────────────────────────────────────────────────────── + + +@router.get( + "/policies/name/{policy_name}/versions", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyVersionListResponse, +) +async def list_policy_versions(policy_name: str): + """ + List all versions of a policy by name, ordered by version_number descending. + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + return await get_policy_registry().get_versions_by_policy_name( + policy_name=policy_name, + prisma_client=prisma_client, + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error listing policy versions: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +@router.post( + "/policies/name/{policy_name}/versions", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyDBResponse, +) +async def create_policy_version( + policy_name: str, + request: PolicyVersionCreateRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Create a new draft version of a policy. Copies all fields from the source. + Source is current production if source_policy_id is not provided. + """ + 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 + return await get_policy_registry().create_new_version( + policy_name=policy_name, + prisma_client=prisma_client, + source_policy_id=request.source_policy_id, + created_by=created_by, + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error creating policy version: {e}") + if "not found" in str(e).lower() or "no production" in str(e).lower(): + raise HTTPException(status_code=404, detail=str(e)) + raise HTTPException(status_code=500, detail=str(e)) + + +@router.put( + "/policies/{policy_id}/status", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyDBResponse, +) +async def update_policy_version_status( + policy_id: str, + request: PolicyVersionStatusUpdateRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Update a policy version's status. Valid transitions: + - draft -> published + - published -> production (demotes current production to published) + - production -> published (demotes, policy becomes inactive) + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + updated_by = user_api_key_dict.user_id + return await get_policy_registry().update_version_status( + policy_id=policy_id, + new_status=request.version_status, + prisma_client=prisma_client, + updated_by=updated_by, + ) + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception(f"Error updating version status: {e}") + if "invalid status" in str(e).lower() or "only draft" in str(e).lower() or "cannot promote" in str(e).lower(): + raise HTTPException(status_code=400, detail=str(e)) + if "not found" in str(e).lower(): + raise HTTPException(status_code=404, detail=str(e)) + raise HTTPException(status_code=500, detail=str(e)) + + +@router.get( + "/policies/compare", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyVersionCompareResponse, +) +async def compare_policy_versions( + version_a: str, + version_b: str, +): + """ + Compare two policy versions. Query params: version_a, version_b (policy version IDs). + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + return await get_policy_registry().compare_versions( + policy_id_a=version_a, + policy_id_b=version_b, + prisma_client=prisma_client, + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error comparing versions: {e}") + if "not found" in str(e).lower(): + raise HTTPException(status_code=404, detail=str(e)) + raise HTTPException(status_code=500, detail=str(e)) + + +@router.delete( + "/policies/name/{policy_name}/all-versions", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], +) +async def delete_all_policy_versions(policy_name: str): + """ + Delete all versions of a policy. Also removes from in-memory registry. + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + return await get_policy_registry().delete_all_versions( + policy_name=policy_name, + prisma_client=prisma_client, + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error deleting all versions: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +# ───────────────────────────────────────────────────────────────────────────── +# Policy CRUD by ID +# ───────────────────────────────────────────────────────────────────────────── + + @router.get( "/policies/{policy_id}", tags=["Policies"], @@ -214,7 +389,7 @@ async def update_policy( raise HTTPException(status_code=500, detail="Database not connected") try: - # Check if policy exists + # Check if policy exists and is draft (only drafts can be updated) existing = await get_policy_registry().get_policy_by_id_from_db( policy_id=policy_id, prisma_client=prisma_client, @@ -223,6 +398,11 @@ async def update_policy( raise HTTPException( status_code=404, detail=f"Policy with ID {policy_id} not found" ) + if getattr(existing, "version_status", "production") != "draft": + raise HTTPException( + status_code=400, + detail="Only draft versions can be updated. Publish or create a new version to change published/production.", + ) updated_by = user_api_key_dict.user_id result = await get_policy_registry().update_policy_in_db( @@ -281,6 +461,7 @@ async def delete_policy(policy_id: str): policy_id=policy_id, prisma_client=prisma_client, ) + # Result may include "warning" if production was deleted return result except HTTPException: raise @@ -527,9 +708,11 @@ async def create_policy_attachment( raise HTTPException(status_code=500, detail="Database not connected") try: - # Verify the policy exists - policy = await get_policy_registry().get_all_policies_from_db(prisma_client) - policy_names = [p.policy_name for p in policy] + # Verify the policy has a production version (attachments resolve against production) + policies = await get_policy_registry().get_all_policies_from_db( + prisma_client, version_status="production" + ) + policy_names = {p.policy_name for p in policies} if request.policy_name not in policy_names: raise HTTPException( status_code=404, diff --git a/litellm/proxy/policy_engine/policy_registry.py b/litellm/proxy/policy_engine/policy_registry.py index 1ca6062e5ac..f53f953c0e1 100644 --- a/litellm/proxy/policy_engine/policy_registry.py +++ b/litellm/proxy/policy_engine/policy_registry.py @@ -12,21 +12,43 @@ 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 ( - GuardrailPipeline, - PipelineStep, - Policy, - PolicyCondition, - PolicyCreateRequest, - PolicyDBResponse, - PolicyGuardrails, - PolicyUpdateRequest, -) +from litellm.types.proxy.policy_engine import (GuardrailPipeline, PipelineStep, + Policy, PolicyCondition, + PolicyCreateRequest, + PolicyDBResponse, + PolicyGuardrails, + PolicyUpdateRequest, + PolicyVersionCompareResponse, + PolicyVersionListResponse) if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient +def _row_to_policy_db_response(row: Any) -> PolicyDBResponse: + """Build PolicyDBResponse from a Prisma LiteLLM_PolicyTable row.""" + return PolicyDBResponse( + policy_id=row.policy_id, + policy_name=row.policy_name, + version_number=getattr(row, "version_number", 1), + version_status=getattr(row, "version_status", "production"), + parent_version_id=getattr(row, "parent_version_id", None), + is_latest=getattr(row, "is_latest", True), + published_at=getattr(row, "published_at", None), + production_at=getattr(row, "production_at", None), + inherit=row.inherit, + description=row.description, + guardrails_add=row.guardrails_add or [], + guardrails_remove=row.guardrails_remove or [], + condition=row.condition, + pipeline=row.pipeline, + created_at=row.created_at, + updated_at=row.updated_at, + created_by=row.created_by, + updated_by=row.updated_by, + ) + + class PolicyRegistry: """ In-memory registry for storing and managing policies. @@ -231,13 +253,18 @@ class PolicyRegistry: PolicyDBResponse with the created policy """ try: - # Build data dict, only include condition if it's set + now = datetime.now(timezone.utc) + # Build data dict; new policy is v1 production data: Dict[str, Any] = { "policy_name": policy_request.policy_name, + "version_number": 1, + "version_status": "production", + "is_latest": True, + "production_at": now, "guardrails_add": policy_request.guardrails_add or [], "guardrails_remove": policy_request.guardrails_remove or [], - "created_at": datetime.now(timezone.utc), - "updated_at": datetime.now(timezone.utc), + "created_at": now, + "updated_at": now, } # Only add optional fields if they have values @@ -276,20 +303,7 @@ class PolicyRegistry: ) 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, - pipeline=created_policy.pipeline, - created_at=created_policy.created_at, - updated_at=created_policy.updated_at, - created_by=created_policy.created_by, - updated_by=created_policy.updated_by, - ) + return _row_to_policy_db_response(created_policy) 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)}") @@ -302,7 +316,7 @@ class PolicyRegistry: updated_by: Optional[str] = None, ) -> PolicyDBResponse: """ - Update a policy in the database. + Update a policy in the database. Only draft versions can be updated. Args: policy_id: The ID of the policy to update @@ -312,8 +326,22 @@ class PolicyRegistry: Returns: PolicyDBResponse with the updated policy + + Raises: + Exception: If policy is not in draft status (only drafts are editable). """ try: + existing = await prisma_client.db.litellm_policytable.find_unique( + where={"policy_id": policy_id} + ) + if existing is None: + raise Exception(f"Policy with ID {policy_id} not found") + version_status = getattr(existing, "version_status", "production") + if version_status != "draft": + raise Exception( + f"Only draft versions can be updated. This policy has status '{version_status}'." + ) + # Build update data - only include fields that are set update_data: Dict[str, Any] = { "updated_at": datetime.now(timezone.utc), @@ -341,36 +369,9 @@ class PolicyRegistry: 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, - "pipeline": updated_policy.pipeline, - }, - ) - self.add_policy(updated_policy.policy_name, policy) + # Do NOT update in-memory registry: drafts are not loaded into memory. - 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, - pipeline=updated_policy.pipeline, - created_at=updated_policy.created_at, - updated_at=updated_policy.updated_at, - created_by=updated_policy.created_by, - updated_by=updated_policy.updated_by, - ) + return _row_to_policy_db_response(updated_policy) 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)}") @@ -379,19 +380,21 @@ class PolicyRegistry: self, policy_id: str, prisma_client: "PrismaClient", - ) -> Dict[str, str]: + ) -> Dict[str, Any]: """ - Delete a policy from the database. + Delete a policy version from the database. + + If the deleted version was production, it is removed from the in-memory + registry. No other version is auto-promoted; admin must explicitly promote. Args: - policy_id: The ID of the policy to delete + policy_id: The ID of the policy version to delete prisma_client: The Prisma client instance Returns: - Dict with success message + Dict with "message" and optional "warning" if production was deleted. """ try: - # Get policy name before deleting policy = await prisma_client.db.litellm_policytable.find_unique( where={"policy_id": policy_id} ) @@ -399,15 +402,25 @@ class PolicyRegistry: if policy is None: raise Exception(f"Policy with ID {policy_id} not found") + version_status = getattr(policy, "version_status", "production") + policy_name = policy.policy_name + # 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) + result: Dict[str, Any] = {"message": f"Policy {policy_id} deleted successfully"} - return {"message": f"Policy {policy_id} deleted successfully"} + # Remove from in-memory registry only if this was the production version + if version_status == "production": + self.remove_policy(policy_name) + result["warning"] = ( + "Production version was deleted. No other version was promoted. " + "Promote another version to production if this policy should remain active." + ) + + return result 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)}") @@ -435,20 +448,7 @@ class PolicyRegistry: 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, - pipeline=policy.pipeline, - created_at=policy.created_at, - updated_at=policy.updated_at, - created_by=policy.created_by, - updated_by=policy.updated_by, - ) + return _row_to_policy_db_response(policy) 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)}") @@ -456,38 +456,30 @@ class PolicyRegistry: async def get_all_policies_from_db( self, prisma_client: "PrismaClient", + version_status: Optional[str] = None, ) -> List[PolicyDBResponse]: """ - Get all policies from the database. + Get all policies from the database, optionally filtered by version_status. Args: prisma_client: The Prisma client instance + version_status: If set, only return policies with this status + ("draft", "published", "production"). Returns: List of PolicyDBResponse objects """ try: + where: Dict[str, Any] = {} + if version_status is not None: + where["version_status"] = version_status + policies = await prisma_client.db.litellm_policytable.find_many( + where=where if where else None, 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, - pipeline=p.pipeline, - created_at=p.created_at, - updated_at=p.updated_at, - created_by=p.created_by, - updated_by=p.updated_by, - ) - for p in policies - ] + return [_row_to_policy_db_response(p) 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)}") @@ -498,12 +490,15 @@ class PolicyRegistry: ) -> None: """ Sync policies from the database to in-memory registry. + Only production versions are loaded. Args: prisma_client: The Prisma client instance """ try: - policies = await self.get_all_policies_from_db(prisma_client) + policies = await self.get_all_policies_from_db( + prisma_client, version_status="production" + ) for policy_response in policies: policy = self._parse_policy( @@ -547,11 +542,13 @@ class PolicyRegistry: List of resolved guardrail names """ from litellm.proxy.policy_engine.policy_resolver import PolicyResolver - + try: - # Load all policies from DB to ensure we have the full inheritance chain - policies = await self.get_all_policies_from_db(prisma_client) - + # Load only production versions so inheritance resolves against production + policies = await self.get_all_policies_from_db( + prisma_client, version_status="production" + ) + # Build a temporary in-memory map for resolution temp_policies = {} for policy_response in policies: @@ -582,6 +579,314 @@ class PolicyRegistry: verbose_proxy_logger.exception(f"Error resolving guardrails from DB: {e}") raise Exception(f"Error resolving guardrails from DB: {str(e)}") + async def get_versions_by_policy_name( + self, + policy_name: str, + prisma_client: "PrismaClient", + ) -> PolicyVersionListResponse: + """ + Get all versions of a policy by name, ordered by version_number descending. + + Args: + policy_name: Name of the policy + prisma_client: The Prisma client instance + + Returns: + PolicyVersionListResponse with policy_name and list of versions + """ + try: + rows = await prisma_client.db.litellm_policytable.find_many( + where={"policy_name": policy_name}, + order={"version_number": "desc"}, + ) + versions = [_row_to_policy_db_response(r) for r in rows] + return PolicyVersionListResponse( + policy_name=policy_name, + versions=versions, + total_count=len(versions), + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error getting versions: {e}") + raise Exception(f"Error getting versions: {str(e)}") + + async def create_new_version( + self, + policy_name: str, + prisma_client: "PrismaClient", + source_policy_id: Optional[str] = None, + created_by: Optional[str] = None, + ) -> PolicyDBResponse: + """ + Create a new draft version of a policy. Copies all fields from the source. + Source is current production if source_policy_id is None. + + Args: + policy_name: Name of the policy + prisma_client: The Prisma client instance + source_policy_id: Policy ID to clone from; if None, use current production + created_by: User who created the version + + Returns: + PolicyDBResponse for the new draft version + """ + try: + if source_policy_id is not None: + source = await prisma_client.db.litellm_policytable.find_unique( + where={"policy_id": source_policy_id} + ) + if source is None: + raise Exception(f"Source policy {source_policy_id} not found") + if source.policy_name != policy_name: + raise Exception( + f"Source policy name '{source.policy_name}' does not match '{policy_name}'" + ) + else: + # Find current production version for this policy_name + prod = await prisma_client.db.litellm_policytable.find_first( + where={ + "policy_name": policy_name, + "version_status": "production", + } + ) + if prod is None: + raise Exception( + f"No production version found for policy '{policy_name}'" + ) + source = prod + + # Next version number + latest = await prisma_client.db.litellm_policytable.find_first( + where={"policy_name": policy_name}, + order={"version_number": "desc"}, + ) + next_num = (latest.version_number + 1) if latest else 1 + + now = datetime.now(timezone.utc) + # Set is_latest=False on all existing versions for this policy_name + await prisma_client.db.litellm_policytable.update_many( + where={"policy_name": policy_name}, + data={"is_latest": False}, + ) + + data: Dict[str, Any] = { + "policy_name": policy_name, + "version_number": next_num, + "version_status": "draft", + "parent_version_id": source.policy_id, + "is_latest": True, + "published_at": None, + "production_at": None, + "inherit": source.inherit, + "description": source.description, + "guardrails_add": source.guardrails_add or [], + "guardrails_remove": source.guardrails_remove or [], + "condition": source.condition, + "pipeline": source.pipeline, + "created_at": now, + "updated_at": now, + "created_by": created_by, + "updated_by": created_by, + } + + created = await prisma_client.db.litellm_policytable.create(data=data) + return _row_to_policy_db_response(created) + except Exception as e: + verbose_proxy_logger.exception(f"Error creating new version: {e}") + raise Exception(f"Error creating new version: {str(e)}") + + async def update_version_status( + self, + policy_id: str, + new_status: str, + prisma_client: "PrismaClient", + updated_by: Optional[str] = None, + ) -> PolicyDBResponse: + """ + Update a policy version's status. Valid transitions: + - draft -> published (sets published_at) + - published -> production (sets production_at, demotes current production to published, updates in-memory) + - production -> published (demotes, removes from in-memory) + - draft -> production: NOT allowed (must publish first) + - published -> draft: NOT allowed + + Args: + policy_id: The policy version ID + new_status: "published" or "production" + prisma_client: The Prisma client instance + updated_by: User who updated + + Returns: + PolicyDBResponse for the updated version + """ + try: + if new_status not in ("published", "production"): + raise Exception(f"Invalid status '{new_status}'. Use 'published' or 'production'.") + + row = await prisma_client.db.litellm_policytable.find_unique( + where={"policy_id": policy_id} + ) + if row is None: + raise Exception(f"Policy with ID {policy_id} not found") + + current = getattr(row, "version_status", "production") + policy_name = row.policy_name + now = datetime.now(timezone.utc) + + if new_status == "published": + if current != "draft": + raise Exception( + f"Only draft versions can be published. Current status: '{current}'." + ) + updated = await prisma_client.db.litellm_policytable.update( + where={"policy_id": policy_id}, + data={ + "version_status": "published", + "published_at": now, + "updated_at": now, + "updated_by": updated_by, + }, + ) + return _row_to_policy_db_response(updated) + + # new_status == "production" + if current not in ("draft", "published"): + raise Exception( + f"Only draft or published versions can be promoted to production. Current: '{current}'." + ) + # Plan: "draft -> production" NOT allowed + if current == "draft": + raise Exception( + "Cannot promote draft directly to production. Publish the version first." + ) + + # Demote current production to published + await prisma_client.db.litellm_policytable.update_many( + where={ + "policy_name": policy_name, + "version_status": "production", + }, + data={ + "version_status": "published", + "updated_at": now, + "updated_by": updated_by, + }, + ) + + # Promote this version to production + updated = await prisma_client.db.litellm_policytable.update( + where={"policy_id": policy_id}, + data={ + "version_status": "production", + "production_at": now, + "updated_at": now, + "updated_by": updated_by, + }, + ) + + # Update in-memory registry: remove old production (by name), add this one + self.remove_policy(policy_name) + policy = self._parse_policy( + policy_name, + { + "inherit": updated.inherit, + "description": updated.description, + "guardrails": { + "add": updated.guardrails_add or [], + "remove": updated.guardrails_remove or [], + }, + "condition": updated.condition, + "pipeline": updated.pipeline, + }, + ) + self.add_policy(policy_name, policy) + + return _row_to_policy_db_response(updated) + except Exception as e: + verbose_proxy_logger.exception(f"Error updating version status: {e}") + raise Exception(f"Error updating version status: {str(e)}") + + async def compare_versions( + self, + policy_id_a: str, + policy_id_b: str, + prisma_client: "PrismaClient", + ) -> PolicyVersionCompareResponse: + """ + Compare two policy versions and return field-by-field diffs. + + Args: + policy_id_a: First policy version ID + policy_id_b: Second policy version ID + prisma_client: The Prisma client instance + + Returns: + PolicyVersionCompareResponse with both versions and field_diffs + """ + try: + a = await prisma_client.db.litellm_policytable.find_unique( + where={"policy_id": policy_id_a} + ) + b = await prisma_client.db.litellm_policytable.find_unique( + where={"policy_id": policy_id_b} + ) + if a is None: + raise Exception(f"Policy {policy_id_a} not found") + if b is None: + raise Exception(f"Policy {policy_id_b} not found") + + resp_a = _row_to_policy_db_response(a) + resp_b = _row_to_policy_db_response(b) + + # Compare fields that are part of policy content (not metadata) + compare_fields = [ + "inherit", + "description", + "guardrails_add", + "guardrails_remove", + "condition", + "pipeline", + ] + field_diffs: Dict[str, Dict[str, Any]] = {} + for field in compare_fields: + val_a = getattr(resp_a, field) + val_b = getattr(resp_b, field) + if val_a != val_b: + field_diffs[field] = {"version_a": val_a, "version_b": val_b} + + return PolicyVersionCompareResponse( + version_a=resp_a, + version_b=resp_b, + field_diffs=field_diffs, + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error comparing versions: {e}") + raise Exception(f"Error comparing versions: {str(e)}") + + async def delete_all_versions( + self, + policy_name: str, + prisma_client: "PrismaClient", + ) -> Dict[str, str]: + """ + Delete all versions of a policy. Also removes from in-memory registry. + + Args: + policy_name: Name of the policy + prisma_client: The Prisma client instance + + Returns: + Dict with success message + """ + try: + await prisma_client.db.litellm_policytable.delete_many( + where={"policy_name": policy_name} + ) + self.remove_policy(policy_name) + return {"message": f"All versions of policy '{policy_name}' deleted successfully"} + except Exception as e: + verbose_proxy_logger.exception(f"Error deleting all versions: {e}") + raise Exception(f"Error deleting all versions: {str(e)}") + # Global singleton instance _policy_registry: Optional[PolicyRegistry] = None diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 4128ab5f23e..5d2cad6da5b 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -963,20 +963,29 @@ model LiteLLM_SkillsTable { updated_by String? } -// Policy table for storing guardrail policies +// Policy table for storing guardrail policies (versioned) model LiteLLM_PolicyTable { - policy_id String @id @default(uuid()) - policy_name String @unique - inherit String? // Name of parent policy to inherit from - description String? - guardrails_add String[] @default([]) - guardrails_remove String[] @default([]) - condition Json? @default("{}") // Policy conditions (e.g., model matching) - pipeline Json? // Optional guardrail pipeline (mode + steps[]) - created_at DateTime @default(now()) - created_by String? - updated_at DateTime @default(now()) @updatedAt - updated_by String? + policy_id String @id @default(uuid()) + policy_name String // No longer @unique; use @@unique([policy_name, version_number]) + version_number Int @default(1) + version_status String @default("production") // "draft" | "published" | "production" + parent_version_id String? + is_latest Boolean @default(true) + published_at DateTime? + production_at DateTime? + inherit String? // Name of parent policy to inherit from + description String? + guardrails_add String[] @default([]) + guardrails_remove String[] @default([]) + condition Json? @default("{}") // Policy conditions (e.g., model matching) + pipeline Json? // Optional guardrail pipeline (mode + steps[]) + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? + + @@unique([policy_name, version_number]) + @@index([policy_name, version_status]) } // Policy attachment table for defining where policies apply diff --git a/litellm/types/proxy/policy_engine/__init__.py b/litellm/types/proxy/policy_engine/__init__.py index e0c1d6f30da..4df9f21e806 100644 --- a/litellm/types/proxy/policy_engine/__init__.py +++ b/litellm/types/proxy/policy_engine/__init__.py @@ -11,48 +11,28 @@ Configuration: """ from litellm.types.proxy.policy_engine.pipeline_types import ( - GuardrailPipeline, - PipelineExecutionResult, - PipelineStep, - PipelineStepResult, -) -from litellm.types.proxy.policy_engine.policy_types import ( - Policy, - PolicyAttachment, - PolicyCondition, - PolicyConfig, - PolicyGuardrails, - PolicyScope, -) + GuardrailPipeline, PipelineExecutionResult, PipelineStep, + PipelineStepResult) +from litellm.types.proxy.policy_engine.policy_types import (Policy, + PolicyAttachment, + PolicyCondition, + PolicyConfig, + PolicyGuardrails, + PolicyScope) from litellm.types.proxy.policy_engine.resolver_types import ( - AttachmentImpactResponse, - PipelineTestRequest, - PolicyAttachmentCreateRequest, - PolicyAttachmentDBResponse, - PolicyAttachmentListResponse, - PolicyConditionRequest, - PolicyCreateRequest, - PolicyDBResponse, - PolicyGuardrailsResponse, - PolicyInfoResponse, - PolicyListDBResponse, - PolicyListResponse, - PolicyMatchContext, - PolicyMatchDetail, - PolicyResolveRequest, - PolicyResolveResponse, - PolicyScopeResponse, - PolicySummaryItem, - PolicyTestResponse, - PolicyUpdateRequest, - ResolvedPolicy, -) + AttachmentImpactResponse, PipelineTestRequest, + PolicyAttachmentCreateRequest, PolicyAttachmentDBResponse, + PolicyAttachmentListResponse, PolicyConditionRequest, PolicyCreateRequest, + PolicyDBResponse, PolicyGuardrailsResponse, PolicyInfoResponse, + PolicyListDBResponse, PolicyListResponse, PolicyMatchContext, + PolicyMatchDetail, PolicyResolveRequest, PolicyResolveResponse, + PolicyScopeResponse, PolicySummaryItem, PolicyTestResponse, + PolicyUpdateRequest, PolicyVersionCompareResponse, + PolicyVersionCreateRequest, PolicyVersionListResponse, + PolicyVersionStatusUpdateRequest, ResolvedPolicy) from litellm.types.proxy.policy_engine.validation_types import ( - PolicyValidateRequest, - PolicyValidationError, - PolicyValidationErrorType, - PolicyValidationResponse, -) + PolicyValidateRequest, PolicyValidationError, PolicyValidationErrorType, + PolicyValidationResponse) __all__ = [ # Pipeline types @@ -98,4 +78,9 @@ __all__ = [ "PolicyResolveResponse", "PolicyMatchDetail", "AttachmentImpactResponse", + # Policy versioning + "PolicyVersionCreateRequest", + "PolicyVersionStatusUpdateRequest", + "PolicyVersionListResponse", + "PolicyVersionCompareResponse", ] diff --git a/litellm/types/proxy/policy_engine/resolver_types.py b/litellm/types/proxy/policy_engine/resolver_types.py index a5a2334ae4b..2df450dc2ba 100644 --- a/litellm/types/proxy/policy_engine/resolver_types.py +++ b/litellm/types/proxy/policy_engine/resolver_types.py @@ -198,6 +198,23 @@ class PolicyDBResponse(BaseModel): policy_id: str = Field(description="Unique ID of the policy.") policy_name: str = Field(description="Name of the policy.") + version_number: int = Field(default=1, description="Version number of this policy.") + version_status: str = Field( + default="production", + description="One of: draft, published, production.", + ) + parent_version_id: Optional[str] = Field( + default=None, description="Policy ID this version was cloned from." + ) + is_latest: bool = Field( + default=True, description="True if this is the latest version by version_number." + ) + published_at: Optional[datetime] = Field( + default=None, description="When this version was published." + ) + production_at: Optional[datetime] = Field( + default=None, description="When this version was promoted to production." + ) inherit: Optional[str] = Field(default=None, description="Parent policy name.") description: Optional[str] = Field(default=None, description="Policy description.") guardrails_add: List[str] = Field( @@ -233,6 +250,49 @@ class PolicyListDBResponse(BaseModel): total_count: int = Field(default=0, description="Total number of policies.") +# ───────────────────────────────────────────────────────────────────────────── +# Policy Versioning Types +# ───────────────────────────────────────────────────────────────────────────── + + +class PolicyVersionCreateRequest(BaseModel): + """Request body for creating a new policy version (draft).""" + + source_policy_id: Optional[str] = Field( + default=None, + description="Policy ID to clone from. If None, clone from current production version.", + ) + + +class PolicyVersionStatusUpdateRequest(BaseModel): + """Request body for updating a policy version's status.""" + + version_status: str = Field( + description="New status: 'published' or 'production'.", + ) + + +class PolicyVersionListResponse(BaseModel): + """Response for listing all versions of a policy.""" + + policy_name: str = Field(description="Name of the policy.") + versions: List[PolicyDBResponse] = Field( + default_factory=list, description="All versions ordered by version_number desc." + ) + total_count: int = Field(default=0, description="Total number of versions.") + + +class PolicyVersionCompareResponse(BaseModel): + """Response for comparing two policy versions.""" + + version_a: PolicyDBResponse = Field(description="First version.") + version_b: PolicyDBResponse = Field(description="Second version.") + field_diffs: Dict[str, Dict[str, Any]] = Field( + default_factory=dict, + description="Field name -> {version_a: val, version_b: val} for differing fields.", + ) + + # ───────────────────────────────────────────────────────────────────────────── # Policy Attachment CRUD Types # ───────────────────────────────────────────────────────────────────────────── diff --git a/schema.prisma b/schema.prisma index 4128ab5f23e..5d2cad6da5b 100644 --- a/schema.prisma +++ b/schema.prisma @@ -963,20 +963,29 @@ model LiteLLM_SkillsTable { updated_by String? } -// Policy table for storing guardrail policies +// Policy table for storing guardrail policies (versioned) model LiteLLM_PolicyTable { - policy_id String @id @default(uuid()) - policy_name String @unique - inherit String? // Name of parent policy to inherit from - description String? - guardrails_add String[] @default([]) - guardrails_remove String[] @default([]) - condition Json? @default("{}") // Policy conditions (e.g., model matching) - pipeline Json? // Optional guardrail pipeline (mode + steps[]) - created_at DateTime @default(now()) - created_by String? - updated_at DateTime @default(now()) @updatedAt - updated_by String? + policy_id String @id @default(uuid()) + policy_name String // No longer @unique; use @@unique([policy_name, version_number]) + version_number Int @default(1) + version_status String @default("production") // "draft" | "published" | "production" + parent_version_id String? + is_latest Boolean @default(true) + published_at DateTime? + production_at DateTime? + inherit String? // Name of parent policy to inherit from + description String? + guardrails_add String[] @default([]) + guardrails_remove String[] @default([]) + condition Json? @default("{}") // Policy conditions (e.g., model matching) + pipeline Json? // Optional guardrail pipeline (mode + steps[]) + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? + + @@unique([policy_name, version_number]) + @@index([policy_name, version_status]) } // Policy attachment table for defining where policies apply diff --git a/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py b/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py new file mode 100644 index 00000000000..738c611d928 --- /dev/null +++ b/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py @@ -0,0 +1,442 @@ +""" +Unit tests for policy versioning: registry behavior, status transitions, and version CRUD. +""" + +from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.policy_engine.policy_registry import ( + PolicyRegistry, + _row_to_policy_db_response, + get_policy_registry, +) +from litellm.types.proxy.policy_engine import ( + PolicyCreateRequest, + PolicyDBResponse, + PolicyUpdateRequest, +) + + +def _make_row( + policy_id="pid-1", + policy_name="test-policy", + version_number=1, + version_status="production", + parent_version_id=None, + is_latest=True, + published_at=None, + production_at=None, + inherit=None, + description="desc", + guardrails_add=None, + guardrails_remove=None, + condition=None, + pipeline=None, + created_at=None, + updated_at=None, + created_by=None, + updated_by=None, +): + row = MagicMock() + row.policy_id = policy_id + row.policy_name = policy_name + row.version_number = version_number + row.version_status = version_status + row.parent_version_id = parent_version_id + row.is_latest = is_latest + row.published_at = published_at + row.production_at = production_at + row.inherit = inherit + row.description = description + row.guardrails_add = guardrails_add or [] + row.guardrails_remove = guardrails_remove or [] + row.condition = condition + row.pipeline = pipeline + row.created_at = created_at or datetime.now(timezone.utc) + row.updated_at = updated_at or datetime.now(timezone.utc) + row.created_by = created_by + row.updated_by = updated_by + return row + + +class TestRowToPolicyDBResponse: + """Test _row_to_policy_db_response includes all version fields.""" + + def test_includes_version_fields(self): + row = _make_row( + version_number=2, + version_status="draft", + parent_version_id="pid-0", + is_latest=True, + published_at=None, + production_at=None, + ) + resp = _row_to_policy_db_response(row) + assert isinstance(resp, PolicyDBResponse) + assert resp.policy_id == "pid-1" + assert resp.policy_name == "test-policy" + assert resp.version_number == 2 + assert resp.version_status == "draft" + assert resp.parent_version_id == "pid-0" + assert resp.is_latest is True + assert resp.published_at is None + assert resp.production_at is None + + def test_backward_compat_missing_version_attrs(self): + row = _make_row() + del row.version_number + del row.version_status + del row.parent_version_id + del row.is_latest + del row.published_at + del row.production_at + resp = _row_to_policy_db_response(row) + assert resp.version_number == 1 + assert resp.version_status == "production" + assert resp.parent_version_id is None + assert resp.is_latest is True + + +class TestSyncPoliciesFromDbProductionOnly: + """Test that sync_policies_from_db only loads production versions.""" + + @pytest.mark.asyncio + async def test_get_all_policies_with_version_status_calls_find_many_with_where(self): + registry = PolicyRegistry() + prisma = MagicMock() + prod_row = _make_row(policy_id="prod-1", version_status="production") + prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[prod_row]) + + result = await registry.get_all_policies_from_db( + prisma, version_status="production" + ) + + assert len(result) == 1 + assert result[0].version_status == "production" + prisma.db.litellm_policytable.find_many.assert_called_once() + call_kw = prisma.db.litellm_policytable.find_many.call_args[1] + assert call_kw.get("where") == {"version_status": "production"} + + @pytest.mark.asyncio + async def test_sync_policies_from_db_only_loads_production(self): + registry = PolicyRegistry() + prisma = MagicMock() + prod_row = _make_row( + policy_id="prod-1", + policy_name="foo", + version_status="production", + guardrails_add=["g1"], + ) + prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[prod_row]) + + await registry.sync_policies_from_db(prisma) + + assert registry.has_policy("foo") + policy = registry.get_policy("foo") + assert policy is not None + assert policy.guardrails.add == ["g1"] + # find_many was called with version_status=production (via get_all_policies_from_db) + find_many_calls = prisma.db.litellm_policytable.find_many.call_args_list + assert len(find_many_calls) >= 1 + assert find_many_calls[0][1].get("where") == {"version_status": "production"} + + +class TestUpdatePolicyDraftOnly: + """Test that update_policy_in_db only allows draft versions.""" + + @pytest.mark.asyncio + async def test_update_production_raises(self): + registry = PolicyRegistry() + prisma = MagicMock() + prod_row = _make_row(policy_id="pid-1", version_status="production") + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=prod_row) + + with pytest.raises(Exception) as exc_info: + await registry.update_policy_in_db( + policy_id="pid-1", + policy_request=PolicyUpdateRequest(description="new"), + prisma_client=prisma, + ) + assert "Only draft" in str(exc_info.value) or "draft" in str(exc_info.value).lower() + prisma.db.litellm_policytable.update.assert_not_called() + + @pytest.mark.asyncio + async def test_update_draft_succeeds_and_does_not_update_registry(self): + registry = PolicyRegistry() + registry.add_policy("test-policy", MagicMock()) # in-memory state + prisma = MagicMock() + draft_row = _make_row( + policy_id="draft-1", + policy_name="test-policy", + version_status="draft", + description="old", + ) + updated_row = _make_row( + policy_id="draft-1", + policy_name="test-policy", + version_status="draft", + description="new", + ) + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=draft_row) + prisma.db.litellm_policytable.update = AsyncMock(return_value=updated_row) + + result = await registry.update_policy_in_db( + policy_id="draft-1", + policy_request=PolicyUpdateRequest(description="new"), + prisma_client=prisma, + ) + + assert result.description == "new" + prisma.db.litellm_policytable.update.assert_called_once() + # Registry still has old in-memory policy (drafts are not in registry; we don't add) + assert registry.has_policy("test-policy") + + +class TestDeletePolicyFromDb: + """Test delete_policy_from_db removes production from registry and returns warning.""" + + @pytest.mark.asyncio + async def test_delete_production_removes_from_registry_and_returns_warning(self): + registry = PolicyRegistry() + registry.add_policy("deleted-policy", MagicMock()) + prisma = MagicMock() + prod_row = _make_row( + policy_id="prod-1", + policy_name="deleted-policy", + version_status="production", + ) + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=prod_row) + prisma.db.litellm_policytable.delete = AsyncMock() + + result = await registry.delete_policy_from_db( + policy_id="prod-1", + prisma_client=prisma, + ) + + assert result["message"] + assert "warning" in result + assert "Production" in result["warning"] or "production" in result["warning"] + assert not registry.has_policy("deleted-policy") + + @pytest.mark.asyncio + async def test_delete_draft_does_not_remove_from_registry_no_warning(self): + registry = PolicyRegistry() + registry.add_policy("my-policy", MagicMock()) + prisma = MagicMock() + draft_row = _make_row( + policy_id="draft-1", + policy_name="my-policy", + version_status="draft", + ) + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=draft_row) + prisma.db.litellm_policytable.delete = AsyncMock() + + result = await registry.delete_policy_from_db( + policy_id="draft-1", + prisma_client=prisma, + ) + + assert "warning" not in result + assert registry.has_policy("my-policy") + + +class TestCreateNewVersion: + """Test create_new_version copies all fields and sets draft.""" + + @pytest.mark.asyncio + async def test_create_new_version_from_production_increments_version(self): + registry = PolicyRegistry() + prisma = MagicMock() + prod = _make_row( + policy_id="prod-1", + policy_name="foo", + version_number=1, + version_status="production", + guardrails_add=["g1"], + description="base", + inherit=None, + pipeline={"mode": "pre_call", "steps": []}, + ) + # find_first for production + prisma.db.litellm_policytable.find_first = AsyncMock(return_value=prod) + # find_first for latest version number + prisma.db.litellm_policytable.find_first.side_effect = [ + prod, # production lookup + prod, # latest version_number lookup + ] + # update_many for is_latest=False + prisma.db.litellm_policytable.update_many = AsyncMock() + new_row = _make_row( + policy_id="new-id", + policy_name="foo", + version_number=2, + version_status="draft", + parent_version_id="prod-1", + is_latest=True, + guardrails_add=["g1"], + description="base", + pipeline={"mode": "pre_call", "steps": []}, + ) + prisma.db.litellm_policytable.create = AsyncMock(return_value=new_row) + + result = await registry.create_new_version( + policy_name="foo", + prisma_client=prisma, + source_policy_id=None, + created_by="user", + ) + + assert result.version_number == 2 + assert result.version_status == "draft" + assert result.parent_version_id == "prod-1" + assert result.guardrails_add == ["g1"] + assert result.description == "base" + create_call = prisma.db.litellm_policytable.create.call_args[1]["data"] + assert create_call["version_number"] == 2 + assert create_call["version_status"] == "draft" + assert create_call["parent_version_id"] == "prod-1" + assert create_call["guardrails_add"] == ["g1"] + + +class TestUpdateVersionStatus: + """Test status transitions: valid succeed, invalid return error.""" + + @pytest.mark.asyncio + async def test_draft_to_published_sets_published_at(self): + registry = PolicyRegistry() + prisma = MagicMock() + draft = _make_row(policy_id="d-1", version_status="draft") + updated = _make_row( + policy_id="d-1", + version_status="published", + published_at=datetime.now(timezone.utc), + ) + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=draft) + prisma.db.litellm_policytable.update = AsyncMock(return_value=updated) + + result = await registry.update_version_status( + policy_id="d-1", + new_status="published", + prisma_client=prisma, + ) + + assert result.version_status == "published" + update_data = prisma.db.litellm_policytable.update.call_args[1]["data"] + assert update_data["version_status"] == "published" + assert "published_at" in update_data + + @pytest.mark.asyncio + async def test_draft_to_production_raises(self): + registry = PolicyRegistry() + prisma = MagicMock() + draft = _make_row(policy_id="d-1", version_status="draft") + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=draft) + + with pytest.raises(Exception) as exc_info: + await registry.update_version_status( + policy_id="d-1", + new_status="production", + prisma_client=prisma, + ) + assert "publish" in str(exc_info.value).lower() or "draft" in str(exc_info.value).lower() + + @pytest.mark.asyncio + async def test_published_to_production_demotes_old_and_updates_registry(self): + registry = PolicyRegistry() + prisma = MagicMock() + published_row = _make_row( + policy_id="pub-1", + policy_name="foo", + version_status="published", + ) + updated_row = _make_row( + policy_id="pub-1", + policy_name="foo", + version_status="production", + production_at=datetime.now(timezone.utc), + ) + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=published_row) + prisma.db.litellm_policytable.update_many = AsyncMock() + prisma.db.litellm_policytable.update = AsyncMock(return_value=updated_row) + + result = await registry.update_version_status( + policy_id="pub-1", + new_status="production", + prisma_client=prisma, + ) + + assert result.version_status == "production" + # update_many should have been called to demote current production + assert prisma.db.litellm_policytable.update_many.called + # Registry should have been updated with new production + assert registry.has_policy("foo") + + +class TestCompareVersions: + """Test compare_versions returns correct field diffs.""" + + @pytest.mark.asyncio + async def test_compare_versions_returns_diffs(self): + registry = PolicyRegistry() + prisma = MagicMock() + a = _make_row( + policy_id="a", + policy_name="p", + description="desc A", + guardrails_add=["g1"], + ) + b = _make_row( + policy_id="b", + policy_name="p", + description="desc B", + guardrails_add=["g1", "g2"], + ) + prisma.db.litellm_policytable.find_unique = AsyncMock(side_effect=[a, b]) + + result = await registry.compare_versions( + policy_id_a="a", + policy_id_b="b", + prisma_client=prisma, + ) + + assert result.version_a.policy_id == "a" + assert result.version_b.policy_id == "b" + assert "description" in result.field_diffs + assert result.field_diffs["description"]["version_a"] == "desc A" + assert result.field_diffs["description"]["version_b"] == "desc B" + assert "guardrails_add" in result.field_diffs + + +class TestResolveGuardrailsProductionOnly: + """Test that resolve_guardrails_from_db uses only production versions.""" + + @pytest.mark.asyncio + async def test_resolve_guardrails_calls_get_all_with_production_filter(self): + registry = PolicyRegistry() + prisma = MagicMock() + prod_row = _make_row( + policy_name="base", + version_status="production", + guardrails_add=["g1"], + ) + prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[prod_row]) + + result = await registry.resolve_guardrails_from_db( + policy_name="base", + prisma_client=prisma, + ) + + assert "g1" in result + call_kw = prisma.db.litellm_policytable.find_many.call_args[1] + assert call_kw.get("where") == {"version_status": "production"} + + +class TestGetPolicyRegistrySingleton: + """Test get_policy_registry returns same instance.""" + + def test_returns_singleton(self): + a = get_policy_registry() + b = get_policy_registry() + assert a is b diff --git a/tests/test_litellm/proxy/policy_engine/test_policy_versioning_e2e.py b/tests/test_litellm/proxy/policy_engine/test_policy_versioning_e2e.py new file mode 100644 index 00000000000..5d6f3a05ae6 --- /dev/null +++ b/tests/test_litellm/proxy/policy_engine/test_policy_versioning_e2e.py @@ -0,0 +1,217 @@ +""" +Integration-style tests for policy versioning: full lifecycle with mocked DB. +""" + +from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.policy_engine.policy_registry import PolicyRegistry +from litellm.types.proxy.policy_engine import (PolicyCreateRequest, + PolicyUpdateRequest) + + +def _make_row( + policy_id, + policy_name, + version_number=1, + version_status="production", + parent_version_id=None, + is_latest=True, + published_at=None, + production_at=None, + inherit=None, + description="", + guardrails_add=None, + guardrails_remove=None, + condition=None, + pipeline=None, +): + row = MagicMock() + row.policy_id = policy_id + row.policy_name = policy_name + row.version_number = version_number + row.version_status = version_status + row.parent_version_id = parent_version_id + row.is_latest = is_latest + row.published_at = published_at + row.production_at = production_at + row.inherit = inherit + row.description = description + row.guardrails_add = guardrails_add or [] + row.guardrails_remove = guardrails_remove or [] + row.condition = condition + row.pipeline = pipeline + row.created_at = datetime.now(timezone.utc) + row.updated_at = datetime.now(timezone.utc) + row.created_by = None + row.updated_by = None + return row + + +@pytest.mark.asyncio +async def test_full_lifecycle_create_draft_edit_publish_promote(): + """ + Full lifecycle: create policy -> create draft version -> edit draft -> + publish -> promote to production -> verify old version demoted -> + verify in-memory updated. + """ + registry = PolicyRegistry() + prisma = MagicMock() + now = datetime.now(timezone.utc) + + # 1) Create initial policy (v1 production) + create_data = {} + created_v1 = _make_row( + policy_id="v1-id", + policy_name="lifecycle-policy", + version_number=1, + version_status="production", + production_at=now, + guardrails_add=["g1"], + description="Initial", + ) + + async def create_impl(data=None, **kwargs): + create_data.update(kwargs.get("data", data or {})) + return created_v1 + + prisma.db.litellm_policytable.create = AsyncMock(side_effect=create_impl) + req = PolicyCreateRequest( + policy_name="lifecycle-policy", + description="Initial", + guardrails_add=["g1"], + ) + created = await registry.add_policy_to_db(req, prisma, created_by="user") + assert created.version_number == 1 + assert created.version_status == "production" + assert registry.has_policy("lifecycle-policy") + + # 2) Create new draft version (v2) + v2_row = _make_row( + policy_id="v2-id", + policy_name="lifecycle-policy", + version_number=2, + version_status="draft", + parent_version_id="v1-id", + is_latest=True, + guardrails_add=["g1", "g2"], + description="Draft v2", + ) + prisma.db.litellm_policytable.find_first = AsyncMock(return_value=created_v1) + prisma.db.litellm_policytable.update_many = AsyncMock() + prisma.db.litellm_policytable.create = AsyncMock(return_value=v2_row) + + draft_v2 = await registry.create_new_version( + policy_name="lifecycle-policy", + prisma_client=prisma, + source_policy_id=None, + created_by="user", + ) + assert draft_v2.version_number == 2 + assert draft_v2.version_status == "draft" + assert draft_v2.parent_version_id == "v1-id" + # In-memory still has v1 (only production is in registry) + assert registry.has_policy("lifecycle-policy") + policy = registry.get_policy("lifecycle-policy") + assert policy.guardrails.add == ["g1"] # still v1 + + # 3) Edit draft v2 + v2_updated_row = _make_row( + policy_id="v2-id", + policy_name="lifecycle-policy", + version_number=2, + version_status="draft", + guardrails_add=["g1", "g2", "g3"], + description="Draft v2 edited", + ) + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=v2_row) + prisma.db.litellm_policytable.update = AsyncMock(return_value=v2_updated_row) + + updated_draft = await registry.update_policy_in_db( + policy_id="v2-id", + policy_request=PolicyUpdateRequest( + description="Draft v2 edited", + guardrails_add=["g1", "g2", "g3"], + ), + prisma_client=prisma, + updated_by="user", + ) + assert updated_draft.description == "Draft v2 edited" + assert updated_draft.guardrails_add == ["g1", "g2", "g3"] + + # 4) Publish v2 (draft -> published) + v2_published = _make_row( + policy_id="v2-id", + policy_name="lifecycle-policy", + version_number=2, + version_status="published", + published_at=now, + guardrails_add=["g1", "g2", "g3"], + description="Draft v2 edited", + ) + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=v2_updated_row) + prisma.db.litellm_policytable.update = AsyncMock(return_value=v2_published) + + published = await registry.update_version_status( + policy_id="v2-id", + new_status="published", + prisma_client=prisma, + updated_by="user", + ) + assert published.version_status == "published" + + # 5) Promote v2 to production (demote v1 to published, update registry) + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=v2_published) + prisma.db.litellm_policytable.update_many = AsyncMock() + v2_production = _make_row( + policy_id="v2-id", + policy_name="lifecycle-policy", + version_number=2, + version_status="production", + production_at=now, + guardrails_add=["g1", "g2", "g3"], + description="Draft v2 edited", + ) + prisma.db.litellm_policytable.update = AsyncMock(return_value=v2_production) + + prod = await registry.update_version_status( + policy_id="v2-id", + new_status="production", + prisma_client=prisma, + updated_by="user", + ) + assert prod.version_status == "production" + # In-memory registry should now have v2 content + assert registry.has_policy("lifecycle-policy") + policy = registry.get_policy("lifecycle-policy") + assert policy.guardrails.add == ["g1", "g2", "g3"] + + +@pytest.mark.asyncio +async def test_attachments_resolve_against_production_after_promotion(): + """ + After promoting a new version to production, resolve_guardrails_from_db + returns guardrails from the new production version (inheritance resolves + against production). + """ + registry = PolicyRegistry() + prisma = MagicMock() + # Simulate only production versions loaded for resolution + prod_row = _make_row( + policy_id="prod-1", + policy_name="att-policy", + version_status="production", + guardrails_add=["ga", "gb"], + ) + prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[prod_row]) + + resolved = await registry.resolve_guardrails_from_db( + policy_name="att-policy", + prisma_client=prisma, + ) + assert "ga" in resolved + assert "gb" in resolved + call_kw = prisma.db.litellm_policytable.find_many.call_args[1] + assert call_kw.get("where") == {"version_status": "production"} From 7f3b047a1c8bab2c7511ab274d24adc65dca5dfd Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 21 Feb 2026 16:00:01 -0800 Subject: [PATCH 2/6] fix(policy_registry): support policy versioning --- .../proxy/policy_engine/policy_registry.py | 57 +++++++++++++------ 1 file changed, 41 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/policy_engine/policy_registry.py b/litellm/proxy/policy_engine/policy_registry.py index f53f953c0e1..891c4e303a0 100644 --- a/litellm/proxy/policy_engine/policy_registry.py +++ b/litellm/proxy/policy_engine/policy_registry.py @@ -110,7 +110,9 @@ class PolicyRegistry: ) else: # Handle legacy format where guardrails might be a list - guardrails = PolicyGuardrails(add=guardrails_data if guardrails_data else None) + guardrails = PolicyGuardrails( + add=guardrails_data if guardrails_data else None + ) # Parse condition (simple model-based condition) condition = None @@ -130,7 +132,9 @@ class PolicyRegistry: ) @staticmethod - def _parse_pipeline(pipeline_data: Optional[Dict[str, Any]]) -> Optional[GuardrailPipeline]: + def _parse_pipeline( + pipeline_data: Optional[Dict[str, Any]], + ) -> Optional[GuardrailPipeline]: """Parse a pipeline configuration from raw data.""" if pipeline_data is None: return None @@ -295,9 +299,11 @@ class PolicyRegistry: "add": policy_request.guardrails_add, "remove": policy_request.guardrails_remove, }, - "condition": policy_request.condition.model_dump() - if policy_request.condition - else None, + "condition": ( + policy_request.condition.model_dump() + if policy_request.condition + else None + ), "pipeline": policy_request.pipeline, }, ) @@ -359,7 +365,9 @@ class PolicyRegistry: 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"] = json.dumps(policy_request.condition.model_dump()) + update_data["condition"] = json.dumps( + policy_request.condition.model_dump() + ) if policy_request.pipeline is not None: validated_pipeline = GuardrailPipeline(**policy_request.pipeline) update_data["pipeline"] = json.dumps(validated_pipeline.model_dump()) @@ -410,7 +418,9 @@ class PolicyRegistry: where={"policy_id": policy_id} ) - result: Dict[str, Any] = {"message": f"Policy {policy_id} deleted successfully"} + result: Dict[str, Any] = { + "message": f"Policy {policy_id} deleted successfully" + } # Remove from in-memory registry only if this was the production version if version_status == "production": @@ -531,13 +541,13 @@ class PolicyRegistry: ) -> List[str]: """ Resolve all guardrails for a policy from the database. - + Uses the existing PolicyResolver to handle inheritance chain resolution. - + Args: policy_name: Name of the policy to resolve prisma_client: The Prisma client instance - + Returns: List of resolved guardrail names """ @@ -566,14 +576,14 @@ class PolicyRegistry: }, ) temp_policies[policy_response.policy_name] = policy - + # Use the existing PolicyResolver to resolve guardrails resolved_policy = PolicyResolver.resolve_policy_guardrails( policy_name=policy_name, policies=temp_policies, context=None, # No context needed for simple resolution ) - + return sorted(resolved_policy.guardrails) except Exception as e: verbose_proxy_logger.exception(f"Error resolving guardrails from DB: {e}") @@ -680,13 +690,24 @@ class PolicyRegistry: "description": source.description, "guardrails_add": source.guardrails_add or [], "guardrails_remove": source.guardrails_remove or [], - "condition": source.condition, - "pipeline": source.pipeline, "created_at": now, "updated_at": now, "created_by": created_by, "updated_by": created_by, } + # Prisma expects Json fields as JSON strings on create (same as add_policy_to_db) + if source.condition is not None: + data["condition"] = ( + json.dumps(source.condition) + if isinstance(source.condition, dict) + else source.condition + ) + if source.pipeline is not None: + data["pipeline"] = ( + json.dumps(source.pipeline) + if isinstance(source.pipeline, dict) + else source.pipeline + ) created = await prisma_client.db.litellm_policytable.create(data=data) return _row_to_policy_db_response(created) @@ -720,7 +741,9 @@ class PolicyRegistry: """ try: if new_status not in ("published", "production"): - raise Exception(f"Invalid status '{new_status}'. Use 'published' or 'production'.") + raise Exception( + f"Invalid status '{new_status}'. Use 'published' or 'production'." + ) row = await prisma_client.db.litellm_policytable.find_unique( where={"policy_id": policy_id} @@ -882,7 +905,9 @@ class PolicyRegistry: where={"policy_name": policy_name} ) self.remove_policy(policy_name) - return {"message": f"All versions of policy '{policy_name}' deleted successfully"} + return { + "message": f"All versions of policy '{policy_name}' deleted successfully" + } except Exception as e: verbose_proxy_logger.exception(f"Error deleting all versions: {e}") raise Exception(f"Error deleting all versions: {str(e)}") From c68ee52a4d2cb9a5c6083c2d7f9c88a432b1629e Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 21 Feb 2026 16:49:28 -0800 Subject: [PATCH 3/6] fix: multiple QA fixes for policy flow builder with guardrail versioning on litellm --- .../out/{404.html => 404/index.html} | 0 .../index.html} | 0 .../index.html} | 0 .../index.html} | 0 .../{budgets.html => budgets/index.html} | 0 .../{caching.html => caching/index.html} | 0 .../index.html} | 0 .../{old-usage.html => old-usage/index.html} | 0 .../{prompts.html => prompts/index.html} | 0 .../index.html} | 0 .../index.html} | 0 .../out/{login.html => login/index.html} | 0 .../out/{logs.html => logs/index.html} | 0 .../{callback.html => callback/index.html} | 0 .../{model-hub.html => model-hub/index.html} | 0 .../{model_hub.html => model_hub/index.html} | 0 .../index.html} | 0 .../index.html} | 0 .../index.html} | 0 .../index.html} | 0 .../index.html} | 0 .../{policies.html => policies/index.html} | 0 .../index.html} | 0 .../index.html} | 0 .../index.html} | 0 .../{ui-theme.html => ui-theme/index.html} | 0 .../out/{teams.html => teams/index.html} | 0 .../{test-key.html => test-key/index.html} | 0 .../index.html} | 0 .../index.html} | 0 .../out/{usage.html => usage/index.html} | 0 .../out/{users.html => users/index.html} | 0 .../index.html} | 0 .../src/components/networking.tsx | 96 ++++ .../src/components/policies/index.tsx | 17 +- .../policies/pipeline_flow_builder.tsx | 533 +++++++++++++++++- .../src/components/policies/policy_table.tsx | 110 ++-- .../src/components/policies/types.ts | 9 + .../src/data/compliancePrompts.ts | 7 + ui/litellm-dashboard/tsconfig.json | 2 +- 40 files changed, 718 insertions(+), 56 deletions(-) rename litellm/proxy/_experimental/out/{404.html => 404/index.html} (100%) rename litellm/proxy/_experimental/out/{_not-found.html => _not-found/index.html} (100%) rename litellm/proxy/_experimental/out/{api-reference.html => api-reference/index.html} (100%) rename litellm/proxy/_experimental/out/experimental/{api-playground.html => api-playground/index.html} (100%) rename litellm/proxy/_experimental/out/experimental/{budgets.html => budgets/index.html} (100%) rename litellm/proxy/_experimental/out/experimental/{caching.html => caching/index.html} (100%) rename litellm/proxy/_experimental/out/experimental/{claude-code-plugins.html => claude-code-plugins/index.html} (100%) rename litellm/proxy/_experimental/out/experimental/{old-usage.html => old-usage/index.html} (100%) rename litellm/proxy/_experimental/out/experimental/{prompts.html => prompts/index.html} (100%) rename litellm/proxy/_experimental/out/experimental/{tag-management.html => tag-management/index.html} (100%) rename litellm/proxy/_experimental/out/{guardrails.html => guardrails/index.html} (100%) rename litellm/proxy/_experimental/out/{login.html => login/index.html} (100%) rename litellm/proxy/_experimental/out/{logs.html => logs/index.html} (100%) rename litellm/proxy/_experimental/out/mcp/oauth/{callback.html => callback/index.html} (100%) rename litellm/proxy/_experimental/out/{model-hub.html => model-hub/index.html} (100%) rename litellm/proxy/_experimental/out/{model_hub.html => model_hub/index.html} (100%) rename litellm/proxy/_experimental/out/{model_hub_table.html => model_hub_table/index.html} (100%) rename litellm/proxy/_experimental/out/{models-and-endpoints.html => models-and-endpoints/index.html} (100%) rename litellm/proxy/_experimental/out/{onboarding.html => onboarding/index.html} (100%) rename litellm/proxy/_experimental/out/{organizations.html => organizations/index.html} (100%) rename litellm/proxy/_experimental/out/{playground.html => playground/index.html} (100%) rename litellm/proxy/_experimental/out/{policies.html => policies/index.html} (100%) rename litellm/proxy/_experimental/out/settings/{admin-settings.html => admin-settings/index.html} (100%) rename litellm/proxy/_experimental/out/settings/{logging-and-alerts.html => logging-and-alerts/index.html} (100%) rename litellm/proxy/_experimental/out/settings/{router-settings.html => router-settings/index.html} (100%) rename litellm/proxy/_experimental/out/settings/{ui-theme.html => ui-theme/index.html} (100%) rename litellm/proxy/_experimental/out/{teams.html => teams/index.html} (100%) rename litellm/proxy/_experimental/out/{test-key.html => test-key/index.html} (100%) rename litellm/proxy/_experimental/out/tools/{mcp-servers.html => mcp-servers/index.html} (100%) rename litellm/proxy/_experimental/out/tools/{vector-stores.html => vector-stores/index.html} (100%) rename litellm/proxy/_experimental/out/{usage.html => usage/index.html} (100%) rename litellm/proxy/_experimental/out/{users.html => users/index.html} (100%) rename litellm/proxy/_experimental/out/{virtual-keys.html => virtual-keys/index.html} (100%) diff --git a/litellm/proxy/_experimental/out/404.html b/litellm/proxy/_experimental/out/404/index.html similarity index 100% rename from litellm/proxy/_experimental/out/404.html rename to litellm/proxy/_experimental/out/404/index.html diff --git a/litellm/proxy/_experimental/out/_not-found.html b/litellm/proxy/_experimental/out/_not-found/index.html similarity index 100% rename from litellm/proxy/_experimental/out/_not-found.html rename to litellm/proxy/_experimental/out/_not-found/index.html diff --git a/litellm/proxy/_experimental/out/api-reference.html b/litellm/proxy/_experimental/out/api-reference/index.html similarity index 100% rename from litellm/proxy/_experimental/out/api-reference.html rename to litellm/proxy/_experimental/out/api-reference/index.html diff --git a/litellm/proxy/_experimental/out/experimental/api-playground.html b/litellm/proxy/_experimental/out/experimental/api-playground/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/api-playground.html rename to litellm/proxy/_experimental/out/experimental/api-playground/index.html diff --git a/litellm/proxy/_experimental/out/experimental/budgets.html b/litellm/proxy/_experimental/out/experimental/budgets/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/budgets.html rename to litellm/proxy/_experimental/out/experimental/budgets/index.html diff --git a/litellm/proxy/_experimental/out/experimental/caching.html b/litellm/proxy/_experimental/out/experimental/caching/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/caching.html rename to litellm/proxy/_experimental/out/experimental/caching/index.html diff --git a/litellm/proxy/_experimental/out/experimental/claude-code-plugins.html b/litellm/proxy/_experimental/out/experimental/claude-code-plugins/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/claude-code-plugins.html rename to litellm/proxy/_experimental/out/experimental/claude-code-plugins/index.html diff --git a/litellm/proxy/_experimental/out/experimental/old-usage.html b/litellm/proxy/_experimental/out/experimental/old-usage/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/old-usage.html rename to litellm/proxy/_experimental/out/experimental/old-usage/index.html diff --git a/litellm/proxy/_experimental/out/experimental/prompts.html b/litellm/proxy/_experimental/out/experimental/prompts/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/prompts.html rename to litellm/proxy/_experimental/out/experimental/prompts/index.html diff --git a/litellm/proxy/_experimental/out/experimental/tag-management.html b/litellm/proxy/_experimental/out/experimental/tag-management/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/tag-management.html rename to litellm/proxy/_experimental/out/experimental/tag-management/index.html diff --git a/litellm/proxy/_experimental/out/guardrails.html b/litellm/proxy/_experimental/out/guardrails/index.html similarity index 100% rename from litellm/proxy/_experimental/out/guardrails.html rename to litellm/proxy/_experimental/out/guardrails/index.html diff --git a/litellm/proxy/_experimental/out/login.html b/litellm/proxy/_experimental/out/login/index.html similarity index 100% rename from litellm/proxy/_experimental/out/login.html rename to litellm/proxy/_experimental/out/login/index.html diff --git a/litellm/proxy/_experimental/out/logs.html b/litellm/proxy/_experimental/out/logs/index.html similarity index 100% rename from litellm/proxy/_experimental/out/logs.html rename to litellm/proxy/_experimental/out/logs/index.html diff --git a/litellm/proxy/_experimental/out/mcp/oauth/callback.html b/litellm/proxy/_experimental/out/mcp/oauth/callback/index.html similarity index 100% rename from litellm/proxy/_experimental/out/mcp/oauth/callback.html rename to litellm/proxy/_experimental/out/mcp/oauth/callback/index.html diff --git a/litellm/proxy/_experimental/out/model-hub.html b/litellm/proxy/_experimental/out/model-hub/index.html similarity index 100% rename from litellm/proxy/_experimental/out/model-hub.html rename to litellm/proxy/_experimental/out/model-hub/index.html diff --git a/litellm/proxy/_experimental/out/model_hub.html b/litellm/proxy/_experimental/out/model_hub/index.html similarity index 100% rename from litellm/proxy/_experimental/out/model_hub.html rename to litellm/proxy/_experimental/out/model_hub/index.html diff --git a/litellm/proxy/_experimental/out/model_hub_table.html b/litellm/proxy/_experimental/out/model_hub_table/index.html similarity index 100% rename from litellm/proxy/_experimental/out/model_hub_table.html rename to litellm/proxy/_experimental/out/model_hub_table/index.html diff --git a/litellm/proxy/_experimental/out/models-and-endpoints.html b/litellm/proxy/_experimental/out/models-and-endpoints/index.html similarity index 100% rename from litellm/proxy/_experimental/out/models-and-endpoints.html rename to litellm/proxy/_experimental/out/models-and-endpoints/index.html diff --git a/litellm/proxy/_experimental/out/onboarding.html b/litellm/proxy/_experimental/out/onboarding/index.html similarity index 100% rename from litellm/proxy/_experimental/out/onboarding.html rename to litellm/proxy/_experimental/out/onboarding/index.html diff --git a/litellm/proxy/_experimental/out/organizations.html b/litellm/proxy/_experimental/out/organizations/index.html similarity index 100% rename from litellm/proxy/_experimental/out/organizations.html rename to litellm/proxy/_experimental/out/organizations/index.html diff --git a/litellm/proxy/_experimental/out/playground.html b/litellm/proxy/_experimental/out/playground/index.html similarity index 100% rename from litellm/proxy/_experimental/out/playground.html rename to litellm/proxy/_experimental/out/playground/index.html diff --git a/litellm/proxy/_experimental/out/policies.html b/litellm/proxy/_experimental/out/policies/index.html similarity index 100% rename from litellm/proxy/_experimental/out/policies.html rename to litellm/proxy/_experimental/out/policies/index.html diff --git a/litellm/proxy/_experimental/out/settings/admin-settings.html b/litellm/proxy/_experimental/out/settings/admin-settings/index.html similarity index 100% rename from litellm/proxy/_experimental/out/settings/admin-settings.html rename to litellm/proxy/_experimental/out/settings/admin-settings/index.html diff --git a/litellm/proxy/_experimental/out/settings/logging-and-alerts.html b/litellm/proxy/_experimental/out/settings/logging-and-alerts/index.html similarity index 100% rename from litellm/proxy/_experimental/out/settings/logging-and-alerts.html rename to litellm/proxy/_experimental/out/settings/logging-and-alerts/index.html diff --git a/litellm/proxy/_experimental/out/settings/router-settings.html b/litellm/proxy/_experimental/out/settings/router-settings/index.html similarity index 100% rename from litellm/proxy/_experimental/out/settings/router-settings.html rename to litellm/proxy/_experimental/out/settings/router-settings/index.html diff --git a/litellm/proxy/_experimental/out/settings/ui-theme.html b/litellm/proxy/_experimental/out/settings/ui-theme/index.html similarity index 100% rename from litellm/proxy/_experimental/out/settings/ui-theme.html rename to litellm/proxy/_experimental/out/settings/ui-theme/index.html diff --git a/litellm/proxy/_experimental/out/teams.html b/litellm/proxy/_experimental/out/teams/index.html similarity index 100% rename from litellm/proxy/_experimental/out/teams.html rename to litellm/proxy/_experimental/out/teams/index.html diff --git a/litellm/proxy/_experimental/out/test-key.html b/litellm/proxy/_experimental/out/test-key/index.html similarity index 100% rename from litellm/proxy/_experimental/out/test-key.html rename to litellm/proxy/_experimental/out/test-key/index.html diff --git a/litellm/proxy/_experimental/out/tools/mcp-servers.html b/litellm/proxy/_experimental/out/tools/mcp-servers/index.html similarity index 100% rename from litellm/proxy/_experimental/out/tools/mcp-servers.html rename to litellm/proxy/_experimental/out/tools/mcp-servers/index.html diff --git a/litellm/proxy/_experimental/out/tools/vector-stores.html b/litellm/proxy/_experimental/out/tools/vector-stores/index.html similarity index 100% rename from litellm/proxy/_experimental/out/tools/vector-stores.html rename to litellm/proxy/_experimental/out/tools/vector-stores/index.html diff --git a/litellm/proxy/_experimental/out/usage.html b/litellm/proxy/_experimental/out/usage/index.html similarity index 100% rename from litellm/proxy/_experimental/out/usage.html rename to litellm/proxy/_experimental/out/usage/index.html diff --git a/litellm/proxy/_experimental/out/users.html b/litellm/proxy/_experimental/out/users/index.html similarity index 100% rename from litellm/proxy/_experimental/out/users.html rename to litellm/proxy/_experimental/out/users/index.html diff --git a/litellm/proxy/_experimental/out/virtual-keys.html b/litellm/proxy/_experimental/out/virtual-keys/index.html similarity index 100% rename from litellm/proxy/_experimental/out/virtual-keys.html rename to litellm/proxy/_experimental/out/virtual-keys/index.html diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index a16afe9c8a2..a618988fa1a 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -5812,6 +5812,102 @@ export const updatePolicyCall = async (accessToken: string, policyId: string, po } }; +export const listPolicyVersions = async ( + accessToken: string, + policyName: string +): Promise<{ policy_name: string; versions: any[]; total_count: number }> => { + try { + const encodedName = encodeURIComponent(policyName); + const url = proxyBaseUrl + ? `${proxyBaseUrl}/policies/name/${encodedName}/versions` + : `/policies/name/${encodedName}/versions`; + const response = await fetch(url, { + method: "GET", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + handleError(errorMessage); + throw new Error(errorMessage); + } + + return await response.json(); + } catch (error) { + console.error("Failed to list policy versions:", error); + throw error; + } +}; + +export const createPolicyVersion = async ( + accessToken: string, + policyName: string, + sourcePolicyId?: string | null +): Promise => { + try { + const encodedName = encodeURIComponent(policyName); + const url = proxyBaseUrl + ? `${proxyBaseUrl}/policies/name/${encodedName}/versions` + : `/policies/name/${encodedName}/versions`; + const response = await fetch(url, { + method: "POST", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + body: JSON.stringify({ source_policy_id: sourcePolicyId ?? undefined }), + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + handleError(errorMessage); + throw new Error(errorMessage); + } + + return await response.json(); + } catch (error) { + console.error("Failed to create policy version:", error); + throw error; + } +}; + +export const updatePolicyVersionStatus = async ( + accessToken: string, + policyId: string, + versionStatus: "published" | "production" +): Promise => { + try { + const url = proxyBaseUrl + ? `${proxyBaseUrl}/policies/${policyId}/status` + : `/policies/${policyId}/status`; + const response = await fetch(url, { + method: "PUT", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + body: JSON.stringify({ version_status: versionStatus }), + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + handleError(errorMessage); + throw new Error(errorMessage); + } + + return await response.json(); + } catch (error) { + console.error("Failed to update policy version status:", error); + throw error; + } +}; + export const deletePolicyCall = async (accessToken: string, policyId: string) => { try { const url = proxyBaseUrl ? `${proxyBaseUrl}/policies/${policyId}` : `/policies/${policyId}`; diff --git a/ui/litellm-dashboard/src/components/policies/index.tsx b/ui/litellm-dashboard/src/components/policies/index.tsx index 9ca382d3432..9c636d432da 100644 --- a/ui/litellm-dashboard/src/components/policies/index.tsx +++ b/ui/litellm-dashboard/src/components/policies/index.tsx @@ -469,11 +469,7 @@ const PoliciesPanel: React.FC = ({ onDeleteClick={handleDeleteClick} onEditClick={(policy) => { setEditingPolicy(policy); - if (policy.pipeline) { - setShowFlowBuilder(true); - } else { - setIsAddPolicyModalVisible(true); - } + setShowFlowBuilder(true); }} onViewClick={(policyId) => setSelectedPolicyId(policyId)} isAdmin={isAdmin} @@ -643,6 +639,17 @@ const PoliciesPanel: React.FC = ({ availableGuardrails={guardrailsList} createPolicy={createPolicyCall} updatePolicy={updatePolicyCall} + onVersionCreated={(newPolicy) => { + setEditingPolicy(newPolicy); + fetchPolicies(); + }} + onSelectVersion={(policy) => { + setEditingPolicy(policy); + }} + onVersionStatusUpdated={(updatedPolicy) => { + setEditingPolicy(updatedPolicy); + fetchPolicies(); + }} /> )} diff --git a/ui/litellm-dashboard/src/components/policies/pipeline_flow_builder.tsx b/ui/litellm-dashboard/src/components/policies/pipeline_flow_builder.tsx index 259378f2e19..22c45b0fb6a 100644 --- a/ui/litellm-dashboard/src/components/policies/pipeline_flow_builder.tsx +++ b/ui/litellm-dashboard/src/components/policies/pipeline_flow_builder.tsx @@ -1,12 +1,14 @@ import React, { useState } from "react"; -import { Select, Typography, message } from "antd"; +import { Select, Typography, message, Spin } from "antd"; import { Button, TextInput } from "@tremor/react"; import { ArrowLeftIcon, PlusIcon } from "@heroicons/react/outline"; import { DotsVerticalIcon } from "@heroicons/react/solid"; import { GuardrailPipeline, PipelineStep, PipelineTestResult, PolicyCreateRequest, PolicyUpdateRequest, Policy } from "./types"; import { Guardrail } from "../guardrails/types"; -import { testPipelineCall } from "../networking"; +import { testPipelineCall, listPolicyVersions, createPolicyVersion, updatePolicyVersionStatus } from "../networking"; import NotificationsManager from "../molecules/notifications_manager"; +import { getComplianceDatasetPrompts } from "../../data/compliancePrompts"; +import type { CompliancePrompt } from "../../data/compliancePrompts"; const { Text } = Typography; @@ -55,6 +57,34 @@ function updateStepAtIndex( return steps.map((s, i) => (i === index ? { ...s, ...updated } : s)); } +/** + * Derives a pipeline from a policy. When the policy has a pipeline, use it. + * When it only has guardrails_add (legacy/simple form), convert those guardrails + * into pipeline steps in order. + */ +function derivePipelineFromPolicy(policy: Policy | null | undefined): GuardrailPipeline { + if (!policy) { + return { mode: "pre_call", steps: [createDefaultStep()] }; + } + if (policy.pipeline?.steps?.length) { + return policy.pipeline; + } + const guardrails = policy.guardrails_add || []; + if (guardrails.length > 0) { + return { + mode: policy.pipeline?.mode ?? "pre_call", + steps: guardrails.map((g) => ({ + guardrail: g, + on_pass: "next" as const, + on_fail: "block" as const, + pass_data: false, + modify_response_message: null, + })), + }; + } + return { mode: "pre_call", steps: [createDefaultStep()] }; +} + // ───────────────────────────────────────────────────────────────────────────── // Icons (matching the reference image) // ───────────────────────────────────────────────────────────────────────────── @@ -559,6 +589,20 @@ const TERMINAL_STYLES: Record = { modify_response: { bg: "#eff6ff", color: "#2563eb" }, }; +interface ComplianceRunEntry { + prompt: CompliancePrompt; + result: PipelineTestResult | null; + error?: string; + matched: boolean; +} + +function complianceMatchExpected(expected: "pass" | "fail", terminalAction: string): boolean { + if (expected === "pass") { + return terminalAction === "allow" || terminalAction === "modify_response"; + } + return terminalAction === "block"; +} + const PipelineTestPanel: React.FC = ({ pipeline, accessToken, @@ -568,6 +612,8 @@ const PipelineTestPanel: React.FC = ({ const [isRunning, setIsRunning] = useState(false); const [result, setResult] = useState(null); const [error, setError] = useState(null); + const [complianceRunning, setComplianceRunning] = useState(false); + const [complianceResults, setComplianceResults] = useState([]); const handleRunTest = async () => { if (!accessToken) return; @@ -596,6 +642,43 @@ const PipelineTestPanel: React.FC = ({ } }; + const handleRunComplianceDataset = async () => { + if (!accessToken) return; + + const emptySteps = pipeline.steps.filter((s) => !s.guardrail); + if (emptySteps.length > 0) { + setError("All steps must have a guardrail selected"); + return; + } + + setError(null); + setResult(null); + setComplianceRunning(true); + const prompts = getComplianceDatasetPrompts(); + const entries: ComplianceRunEntry[] = []; + + for (const prompt of prompts) { + try { + const data = await testPipelineCall(accessToken, pipeline, [ + { role: "user", content: prompt.prompt }, + ]); + const matched = complianceMatchExpected(prompt.expectedResult, data.terminal_action); + entries.push({ prompt, result: data, matched }); + } catch (e) { + const errMsg = e instanceof Error ? e.message : String(e); + entries.push({ + prompt, + result: null, + error: errMsg, + matched: false, + }); + } + } + + setComplianceResults(entries); + setComplianceRunning(false); + }; + return (
= ({ +
{/* Results section */} @@ -773,11 +866,318 @@ const PipelineTestPanel: React.FC = ({ )} - {!result && !error && ( -
- Enter a test message and click "Run Test" to execute the pipeline + {complianceResults.length > 0 && ( +
+
+ Compliance dataset +
+
+ {complianceResults.filter((e) => e.matched).length} / {complianceResults.length} matched + expected +
+
+ {complianceResults.map((entry, i) => { + const actual = + entry.result?.terminal_action ?? (entry.error ? "error" : "—"); + const matchStyle = entry.matched + ? { bg: "#f0fdf4", color: "#16a34a" } + : { bg: "#fef2f2", color: "#dc2626" }; + return ( +
+
+ {entry.prompt.prompt} +
+
+ + expected: {entry.prompt.expectedResult} + + → + + actual: {actual} + + + {entry.matched ? "✓" : "✗"} + +
+ {entry.error && ( +
+ {entry.error} +
+ )} +
+ ); + })} +
)} + + {!result && !error && complianceResults.length === 0 && ( +
+ Enter a test message and click "Run Test" or "Test pipeline (compliance dataset)" to + execute the pipeline +
+ )} +
+ + ); +}; + +// ───────────────────────────────────────────────────────────────────────────── +// Policy Versions Sidebar (left sidebar when editing a policy) +// ───────────────────────────────────────────────────────────────────────────── + +const VERSION_STATUS_STYLES: Record< + string, + { bg: string; color: string } +> = { + draft: { bg: "#f3f4f6", color: "#6b7280" }, + published: { bg: "#eff6ff", color: "#2563eb" }, + production: { bg: "#f0fdf4", color: "#16a34a" }, +}; + +interface PolicyVersionsSidebarProps { + policyName: string; + editingPolicyId: string | null; + editingVersionStatus?: "draft" | "published" | "production"; + accessToken: string | null; + versions: Policy[]; + isLoading: boolean; + isCreatingVersion?: boolean; + isUpdatingStatus?: boolean; + onNewVersion: () => void; + onSelectVersion: (policy: Policy) => void; + onPublish?: () => void; + onPromoteToProduction?: () => void; +} + +const PolicyVersionsSidebar: React.FC = ({ + policyName, + editingPolicyId, + editingVersionStatus, + accessToken, + versions, + isLoading, + isCreatingVersion = false, + isUpdatingStatus = false, + onNewVersion, + onSelectVersion, + onPublish, + onPromoteToProduction, +}) => { + const canPublish = editingVersionStatus === "draft" && onPublish; + const canPromote = editingVersionStatus === "published" && onPromoteToProduction; + + return ( +
+
+ {/* Versions section */} +
+ + Versions + + + {isLoading ? ( +
+ +
+ ) : versions.length === 0 ? ( + + No versions found + + ) : ( +
+ {versions.map((v) => { + const statusStyle = + VERSION_STATUS_STYLES[v.version_status ?? "draft"] ?? + VERSION_STATUS_STYLES.draft; + const isActive = v.policy_id === editingPolicyId; + return ( + + ); + })} +
+ )} + + {/* Publish / Promote to production for selected version */} + {(canPublish || canPromote) && ( +
+ {canPublish && ( + + )} + {canPromote && ( + + )} +
+ )} +
+ + {/* Silent Mirroring section */} +
+
+ + Silent Mirroring + + + COMING SOON + +
+ + Test policy versions on production traffic without blocking requests. + Shadow testing helps validate changes before full rollout. + +
); @@ -795,6 +1195,9 @@ interface FlowBuilderPageProps { availableGuardrails: Guardrail[]; createPolicy: (accessToken: string, policyData: any) => Promise; updatePolicy: (accessToken: string, policyId: string, policyData: any) => Promise; + onVersionCreated?: (newPolicy: Policy) => void; + onSelectVersion?: (policy: Policy) => void; + onVersionStatusUpdated?: (updatedPolicy: Policy) => void; } export const FlowBuilderPage: React.FC = ({ @@ -805,16 +1208,110 @@ export const FlowBuilderPage: React.FC = ({ availableGuardrails, createPolicy, updatePolicy, + onVersionCreated, + onSelectVersion, + onVersionStatusUpdated, }) => { const isEditing = !!editingPolicy?.policy_id; + const showVersionsSidebar = !!editingPolicy?.policy_name; const [policyName, setPolicyName] = useState(editingPolicy?.policy_name || ""); const [description, setDescription] = useState(editingPolicy?.description || ""); const [isSubmitting, setIsSubmitting] = useState(false); const [showTestPanel, setShowTestPanel] = useState(false); const [pipeline, setPipeline] = useState( - editingPolicy?.pipeline || { mode: "pre_call", steps: [createDefaultStep()] } + () => derivePipelineFromPolicy(editingPolicy) ); + const [versions, setVersions] = useState([]); + const [isVersionsLoading, setIsVersionsLoading] = useState(false); + const [isCreatingVersion, setIsCreatingVersion] = useState(false); + const [isUpdatingStatus, setIsUpdatingStatus] = useState(false); + + // Sync local state when editingPolicy changes (e.g. user switched version) + React.useEffect(() => { + setPolicyName(editingPolicy?.policy_name || ""); + setDescription(editingPolicy?.description || ""); + setPipeline(derivePipelineFromPolicy(editingPolicy)); + }, [editingPolicy?.policy_id, editingPolicy?.policy_name, editingPolicy?.description, editingPolicy?.pipeline, editingPolicy?.guardrails_add]); + + // Fetch versions when editing an existing policy by name + React.useEffect(() => { + if (!showVersionsSidebar || !editingPolicy?.policy_name || !accessToken) { + setVersions([]); + return; + } + let cancelled = false; + setIsVersionsLoading(true); + listPolicyVersions(accessToken, editingPolicy.policy_name) + .then((res) => { + if (!cancelled) setVersions(res.versions || []); + }) + .catch(() => { + if (!cancelled) setVersions([]); + }) + .finally(() => { + if (!cancelled) setIsVersionsLoading(false); + }); + return () => { + cancelled = true; + }; + }, [showVersionsSidebar, editingPolicy?.policy_name, accessToken]); + + const handleNewVersion = async () => { + if (!accessToken || !editingPolicy?.policy_name) return; + setIsCreatingVersion(true); + try { + const newPolicy = await createPolicyVersion(accessToken, editingPolicy.policy_name); + NotificationsManager.success("New draft version created"); + onVersionCreated?.(newPolicy); + } catch (error) { + NotificationsManager.fromBackend( + "Failed to create version: " + (error instanceof Error ? error.message : String(error)) + ); + } finally { + setIsCreatingVersion(false); + } + }; + + const handleSelectVersion = (policy: Policy) => { + onSelectVersion?.(policy); + }; + + const handlePublishVersion = async () => { + if (!accessToken || !editingPolicy?.policy_id) return; + setIsUpdatingStatus(true); + try { + const updated = await updatePolicyVersionStatus(accessToken, editingPolicy.policy_id, "published"); + NotificationsManager.success("Version published"); + const list = await listPolicyVersions(accessToken, editingPolicy.policy_name ?? ""); + setVersions(list.versions ?? []); + onVersionStatusUpdated?.(updated); + } catch (error) { + NotificationsManager.fromBackend( + "Failed to publish: " + (error instanceof Error ? error.message : String(error)) + ); + } finally { + setIsUpdatingStatus(false); + } + }; + + const handlePromoteToProduction = async () => { + if (!accessToken || !editingPolicy?.policy_id) return; + setIsUpdatingStatus(true); + try { + const updated = await updatePolicyVersionStatus(accessToken, editingPolicy.policy_id, "production"); + NotificationsManager.success("Version promoted to production"); + const list = await listPolicyVersions(accessToken, editingPolicy.policy_name ?? ""); + setVersions(list.versions ?? []); + onVersionStatusUpdated?.(updated); + } catch (error) { + NotificationsManager.fromBackend( + "Failed to promote to production: " + (error instanceof Error ? error.message : String(error)) + ); + } finally { + setIsUpdatingStatus(false); + } + }; const handleSave = async () => { if (!policyName.trim()) { @@ -849,13 +1346,13 @@ export const FlowBuilderPage: React.FC = ({ if (isEditing && editingPolicy) { await updatePolicy(accessToken, editingPolicy.policy_id, data as PolicyUpdateRequest); NotificationsManager.success("Policy updated successfully"); + onSuccess(); } else { await createPolicy(accessToken, data as PolicyCreateRequest); NotificationsManager.success("Policy created successfully"); + onSuccess(); + onBack(); } - - onSuccess(); - onBack(); } catch (error) { console.error("Failed to save policy:", error); NotificationsManager.fromBackend( @@ -963,8 +1460,24 @@ export const FlowBuilderPage: React.FC = ({ /> - {/* Flow builder canvas + test panel */} + {/* Sidebar (when editing) + Flow builder canvas + test panel */}
+ {showVersionsSidebar && ( + + )}
(); + for (const p of policies) { + const name = p.policy_name || "(unnamed)"; + if (!byName.has(name)) byName.set(name, []); + byName.get(name)!.push(p); + } + const rows: PolicyRow[] = []; + for (const [policyName, versions] of byName) { + // Prefer production, then highest version_number + const primary = + versions.find((v) => v.version_status === "production") ?? + [...versions].sort((a, b) => (b.version_number ?? 0) - (a.version_number ?? 0))[0] ?? + versions[0]; + rows.push({ policy_name: policyName, primaryPolicy: primary, versionCount: versions.length }); + } + return rows.sort((a, b) => a.policy_name.localeCompare(b.policy_name)); +} + interface PolicyTableProps { policies: Policy[]; isLoading: boolean; @@ -29,49 +55,48 @@ const PolicyTable: React.FC = ({ onViewClick, isAdmin = false, }) => { - const [sorting, setSorting] = useState([{ id: "created_at", desc: true }]); + const [sorting, setSorting] = useState([{ id: "policy_name", desc: false }]); + + const rows = useMemo(() => groupPoliciesByName(policies), [policies]); - // Format date helper function const formatDate = (dateString?: string) => { if (!dateString) return "-"; const date = new Date(dateString); return date.toLocaleString(); }; - const columns: ColumnDef[] = [ - { - header: "Policy ID", - accessorKey: "policy_id", - cell: (info: any) => ( - - - - ), - }, + const columns: ColumnDef[] = [ { header: "Name", accessorKey: "policy_name", cell: ({ row }) => { - const policy = row.original; + const { primaryPolicy, versionCount } = row.original; return ( - - {policy.policy_name || "-"} - +
+ 1 ? ` (${versionCount} versions)` : ""}`}> + + + {versionCount > 1 && ( + + {versionCount} version{versionCount !== 1 ? "s" : ""} + + )} +
); }, }, { header: "Description", - accessorKey: "description", + accessorFn: (row) => row.primaryPolicy.description ?? "", cell: ({ row }) => { - const policy = row.original; + const policy = row.original.primaryPolicy; return ( @@ -83,9 +108,9 @@ const PolicyTable: React.FC = ({ }, { header: "Inherits From", - accessorKey: "inherit", + accessorFn: (row) => row.primaryPolicy.inherit ?? "", cell: ({ row }) => { - const policy = row.original; + const policy = row.original.primaryPolicy; return policy.inherit ? ( {policy.inherit} @@ -97,9 +122,9 @@ const PolicyTable: React.FC = ({ }, { header: "Guardrails (Add)", - accessorKey: "guardrails_add", + accessorFn: (row) => (row.primaryPolicy.guardrails_add ?? []).join(", "), cell: ({ row }) => { - const policy = row.original; + const policy = row.original.primaryPolicy; const guardrails = policy.guardrails_add || []; if (guardrails.length === 0) { return -; @@ -122,9 +147,9 @@ const PolicyTable: React.FC = ({ }, { header: "Guardrails (Remove)", - accessorKey: "guardrails_remove", + accessorFn: (row) => (row.primaryPolicy.guardrails_remove ?? []).join(", "), cell: ({ row }) => { - const policy = row.original; + const policy = row.original.primaryPolicy; const guardrails = policy.guardrails_remove || []; if (guardrails.length === 0) { return -; @@ -147,9 +172,12 @@ const PolicyTable: React.FC = ({ }, { header: "Model Condition", - accessorKey: "condition", + accessorFn: (row) => { + const m = row.primaryPolicy.condition?.model; + return typeof m === "string" ? m : JSON.stringify(m ?? ""); + }, cell: ({ row }) => { - const policy = row.original; + const policy = row.original.primaryPolicy; const modelCondition = policy.condition?.model; if (!modelCondition) { return -; @@ -169,9 +197,10 @@ const PolicyTable: React.FC = ({ }, { header: "Created At", - accessorKey: "created_at", + id: "created_at", + accessorFn: (row) => row.primaryPolicy.created_at ?? "", cell: ({ row }) => { - const policy = row.original; + const policy = row.original.primaryPolicy; return ( {formatDate(policy.created_at)} @@ -183,7 +212,8 @@ const PolicyTable: React.FC = ({ id: "actions", header: "Actions", cell: ({ row }) => { - const policy = row.original; + const { primaryPolicy } = row.original; + const policy = primaryPolicy; return (
{isAdmin && ( @@ -216,7 +246,7 @@ const PolicyTable: React.FC = ({ ]; const table = useReactTable({ - data: policies, + data: rows, columns, state: { sorting, @@ -273,9 +303,9 @@ const PolicyTable: React.FC = ({
- ) : policies.length > 0 ? ( + ) : rows.length > 0 ? ( table.getRowModel().rows.map((row) => ( - + {row.getVisibleCells().map((cell) => ( = { }, }; +/** Flat list of all compliance prompts for pipeline testing (EU AI Act, GDPR, topic blocking, airline, etc.). */ +export function getComplianceDatasetPrompts(): CompliancePrompt[] { + return getFrameworks().flatMap((fw) => + fw.categories.flatMap((cat) => cat.prompts) + ); +} + export function getFrameworks(): ComplianceFramework[] { const frameworkMap = new Map< string, diff --git a/ui/litellm-dashboard/tsconfig.json b/ui/litellm-dashboard/tsconfig.json index d24bdd340f7..5b0352feb98 100644 --- a/ui/litellm-dashboard/tsconfig.json +++ b/ui/litellm-dashboard/tsconfig.json @@ -14,7 +14,7 @@ "moduleResolution": "bundler", "resolveJsonModule": true, "isolatedModules": true, - "jsx": "react-jsx", + "jsx": "preserve", "incremental": true, "plugins": [ { From 9bd4ae3df4eb33acb24188f00493abac69b7dff8 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 21 Feb 2026 17:35:44 -0800 Subject: [PATCH 4/6] feat: ui improvements --- litellm/proxy/litellm_pre_call_utils.py | 163 +++++++++---- .../proxy/policy_engine/policy_registry.py | 43 +++- .../proxy/policy_engine/policy_resolver.py | 41 ++-- .../proxy/test_litellm_pre_call_utils.py | 135 ++++++++--- .../playground/complianceUI/ComplianceUI.tsx | 139 +++-------- .../components/policies/PolicySelector.tsx | 65 +++-- .../src/components/policies/index.tsx | 6 +- .../policies/pipeline_flow_builder.tsx | 224 ++++++++++++------ 8 files changed, 515 insertions(+), 301 deletions(-) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 089731c473d..810ad19a9dd 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -10,15 +10,10 @@ import litellm from litellm._logging import verbose_logger, verbose_proxy_logger from litellm._service_logger import ServiceLogging from litellm.litellm_core_utils.safe_json_loads import safe_json_loads -from litellm.proxy._types import ( - AddTeamCallback, - CommonProxyErrors, - LitellmDataForBackendLLMCall, - LitellmUserRoles, - SpecialHeaders, - TeamCallbackMetadata, - UserAPIKeyAuth, -) +from litellm.proxy._types import (AddTeamCallback, CommonProxyErrors, + LitellmDataForBackendLLMCall, + LitellmUserRoles, SpecialHeaders, + TeamCallbackMetadata, UserAPIKeyAuth) # Cache special headers as a frozenset for O(1) lookup performance _SPECIAL_HEADERS_CACHE = frozenset( @@ -27,12 +22,9 @@ _SPECIAL_HEADERS_CACHE = frozenset( from litellm.router import Router from litellm.types.llms.anthropic import ANTHROPIC_API_HEADERS from litellm.types.services import ServiceTypes -from litellm.types.utils import ( - LlmProviders, - ProviderSpecificHeader, - StandardLoggingUserAPIKeyMetadata, - SupportedCacheControls, -) +from litellm.types.utils import (LlmProviders, ProviderSpecificHeader, + StandardLoggingUserAPIKeyMetadata, + SupportedCacheControls) service_logger_obj = ServiceLogging() # used for tracking latency on OTEL @@ -661,8 +653,7 @@ class LiteLLMProxyRequestSetup: return data from litellm.proxy._types import ( LiteLLM_ManagementEndpoint_MetadataFields, - LiteLLM_ManagementEndpoint_MetadataFields_Premium, - ) + LiteLLM_ManagementEndpoint_MetadataFields_Premium) # ignore any special fields added_metadata = {} @@ -1125,7 +1116,7 @@ async def add_litellm_data_to_request( # noqa: PLR0915 data["litellm_disabled_callbacks"] = disabled_callbacks # Guardrails from key/team metadata and policy engine - move_guardrails_to_metadata( + await move_guardrails_to_metadata( data=data, _metadata_variable_name=_metadata_variable_name, user_api_key_dict=user_api_key_dict, @@ -1458,7 +1449,7 @@ def _add_guardrails_from_policies_in_metadata( ) -def move_guardrails_to_metadata( +async def move_guardrails_to_metadata( data: dict, _metadata_variable_name: str, user_api_key_dict: UserAPIKeyAuth, @@ -1487,7 +1478,8 @@ def move_guardrails_to_metadata( # Only check policy engine if no local config (avoid import + registry lookup) if not (has_key_config or has_team_config or has_request_config): - from litellm.proxy.policy_engine.policy_registry import get_policy_registry + from litellm.proxy.policy_engine.policy_registry import \ + get_policy_registry if not get_policy_registry().is_initialized(): # Nothing configured anywhere - clean up request body fields and return @@ -1515,7 +1507,7 @@ def move_guardrails_to_metadata( ######################################################################################### # Add guardrails from policy engine based on team/key/model context ######################################################################################### - add_guardrails_from_policy_engine( + await add_guardrails_from_policy_engine( data=data, metadata_variable_name=_metadata_variable_name, user_api_key_dict=user_api_key_dict, @@ -1549,10 +1541,29 @@ def move_guardrails_to_metadata( ] = request_body_guardrail_config +def _is_policy_version_id(s: str) -> bool: + """Return True if string is a policy version ID (starts with policy_ prefix).""" + from litellm.proxy.policy_engine.policy_registry import \ + POLICY_VERSION_ID_PREFIX + + return isinstance(s, str) and s.startswith(POLICY_VERSION_ID_PREFIX) + + +def _extract_policy_id(s: str) -> Optional[str]: + """Extract raw UUID from policy_ string, or None if not a valid version ID.""" + from litellm.proxy.policy_engine.policy_registry import \ + POLICY_VERSION_ID_PREFIX + + if not _is_policy_version_id(s): + return None + return s[len(POLICY_VERSION_ID_PREFIX) :].strip() or None + + def _match_and_track_policies( data: dict, context: "PolicyMatchContext", request_body_policies: Any, + policies_override: Optional[Dict[str, Any]] = None, ) -> tuple[list[str], dict[str, str]]: """ Match policies via attachments and request body, track them in metadata. @@ -1562,10 +1573,9 @@ def _match_and_track_policies( """ from litellm._logging import verbose_proxy_logger from litellm.proxy.common_utils.callback_utils import ( - add_policy_sources_to_metadata, - add_policy_to_applied_policies_header, - ) - from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry + add_policy_sources_to_metadata, add_policy_to_applied_policies_header) + from litellm.proxy.policy_engine.attachment_registry import \ + get_attachment_registry from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher # Get matching policies via attachments (with match reasons for attribution) @@ -1595,6 +1605,7 @@ def _match_and_track_policies( applied_policy_names = PolicyMatcher.get_policies_with_matching_conditions( policy_names=list(all_policy_names), context=context, + policies=policies_override, ) verbose_proxy_logger.debug( @@ -1622,20 +1633,30 @@ def _apply_resolved_guardrails_to_metadata( data: dict, metadata_variable_name: str, context: "PolicyMatchContext", + policy_names: Optional[List[str]] = None, + policies: Optional[Dict[str, Any]] = None, ) -> None: """Apply resolved guardrails and pipelines to request metadata.""" from litellm._logging import verbose_proxy_logger from litellm.proxy.policy_engine.policy_resolver import PolicyResolver # Resolve guardrails from matching policies - resolved_guardrails = PolicyResolver.resolve_guardrails_for_context(context=context) + resolved_guardrails = PolicyResolver.resolve_guardrails_for_context( + context=context, + policies=policies, + policy_names=policy_names, + ) verbose_proxy_logger.debug( f"Policy engine: resolved guardrails: {resolved_guardrails}" ) # Resolve pipelines from matching policies - pipelines = PolicyResolver.resolve_pipelines_for_context(context=context) + pipelines = PolicyResolver.resolve_pipelines_for_context( + context=context, + policies=policies, + policy_names=policy_names, + ) # Add resolved guardrails to request metadata if metadata_variable_name not in data: @@ -1675,7 +1696,7 @@ def _apply_resolved_guardrails_to_metadata( ) -def add_guardrails_from_policy_engine( +async def add_guardrails_from_policy_engine( data: dict, metadata_variable_name: str, user_api_key_dict: UserAPIKeyAuth, @@ -1685,12 +1706,13 @@ def add_guardrails_from_policy_engine( This function: 1. Extracts "policies" from request body (if present) for dynamic policy application - 2. Gets matching policies based on team_alias, key_alias, and model (via attachments) - 3. Combines dynamic policies with attachment-based policies - 4. Resolves guardrails from all policies (including inheritance) - 5. Adds guardrails to request metadata - 6. Tracks applied policies in metadata for response headers - 7. Removes "policies" from request body so it's not forwarded to LLM provider + 2. Supports policy_ in policies to execute a specific version (e.g. published) + 3. Gets matching policies based on team_alias, key_alias, and model (via attachments) + 4. Combines dynamic policies with attachment-based policies + 5. Resolves guardrails from all policies (including inheritance) + 6. Adds guardrails to request metadata + 7. Tracks applied policies in metadata for response headers + 8. Removes "policies" from request body so it's not forwarded to LLM provider Args: data: The request data to update @@ -1698,12 +1720,13 @@ def add_guardrails_from_policy_engine( user_api_key_dict: The user's API key authentication info """ from litellm._logging import verbose_proxy_logger - from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body + from litellm.proxy.common_utils.http_parsing_utils import \ + get_tags_from_request_body from litellm.proxy.policy_engine.policy_registry import get_policy_registry from litellm.types.proxy.policy_engine import PolicyMatchContext # Extract dynamic policies from request body (if present) - request_body_policies = data.pop("policies", None) + request_body_policies_raw = data.pop("policies", None) registry = get_policy_registry() verbose_proxy_logger.debug( @@ -1730,13 +1753,69 @@ def add_guardrails_from_policy_engine( f"key_alias={context.key_alias}, model={context.model}, tags={context.tags}" ) - # Match and track policies based on attachments and request body - _match_and_track_policies(data, context, request_body_policies) + # Separate policy names from policy version IDs (policy_) + request_body_names: List[str] = [] + request_body_version_ids: List[str] = [] + if request_body_policies_raw and isinstance(request_body_policies_raw, list): + for item in request_body_policies_raw: + if not isinstance(item, str): + continue + if _is_policy_version_id(item): + policy_id = _extract_policy_id(item) + if policy_id: + request_body_version_ids.append(policy_id) + else: + request_body_names.append(item) - # Always resolve and apply guardrails, even if no policies matched above. - # PolicyResolver does its own independent matching and inheritance resolution, - # so guardrails can still be applied via inherited parent policies. - _apply_resolved_guardrails_to_metadata(data, metadata_variable_name, context) + # Fetch policy versions by ID from DB + merged_policies: Dict[str, Any] = dict(registry.get_all_policies()) + fetched_policy_names: List[str] = [] + if request_body_version_ids: + try: + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is not None: + for policy_id in request_body_version_ids: + result = await registry.get_policy_by_id_for_request( + policy_id=policy_id, + prisma_client=prisma_client, + ) + if result is not None: + pname, policy = result + merged_policies[pname] = policy + fetched_policy_names.append(pname) + verbose_proxy_logger.debug( + f"Policy engine: loaded version by ID policy_{policy_id} -> {pname}" + ) + else: + verbose_proxy_logger.debug( + f"Policy engine: policy version {policy_id} not found, skipping" + ) + except Exception as e: + verbose_proxy_logger.warning( + f"Policy engine: failed to fetch policy versions by ID: {e}" + ) + + # Build request body list: names + policy names from fetched versions + request_body_policies = request_body_names + fetched_policy_names + + # Match and track policies (with merged_policies when we have version overrides) + applied_policy_names, _ = _match_and_track_policies( + data, + context, + request_body_policies, + policies_override=merged_policies if request_body_version_ids else None, + ) + + # Resolve and apply guardrails. Use applied_policy_names so request-body policies + # (names + version IDs) are included. Use merged_policies when we have version overrides. + _apply_resolved_guardrails_to_metadata( + data, + metadata_variable_name, + context, + policy_names=applied_policy_names if applied_policy_names else None, + policies=merged_policies if request_body_version_ids else None, + ) def add_provider_specific_headers_to_request( diff --git a/litellm/proxy/policy_engine/policy_registry.py b/litellm/proxy/policy_engine/policy_registry.py index 891c4e303a0..ceee38897bd 100644 --- a/litellm/proxy/policy_engine/policy_registry.py +++ b/litellm/proxy/policy_engine/policy_registry.py @@ -9,7 +9,7 @@ by policy_attachments (see AttachmentRegistry). import json from datetime import datetime, timezone -from typing import TYPE_CHECKING, Any, Dict, List, Optional +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple from litellm._logging import verbose_proxy_logger from litellm.types.proxy.policy_engine import (GuardrailPipeline, PipelineStep, @@ -24,6 +24,9 @@ from litellm.types.proxy.policy_engine import (GuardrailPipeline, PipelineStep, if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient +# Prefix for policy version IDs in request body. Use policy_ to execute a specific version. +POLICY_VERSION_ID_PREFIX = "policy_" + def _row_to_policy_db_response(row: Any) -> PolicyDBResponse: """Build PolicyDBResponse from a Prisma LiteLLM_PolicyTable row.""" @@ -463,6 +466,44 @@ class PolicyRegistry: verbose_proxy_logger.exception(f"Error getting policy from DB: {e}") raise Exception(f"Error getting policy from DB: {str(e)}") + async def get_policy_by_id_for_request( + self, + policy_id: str, + prisma_client: "PrismaClient", + ) -> Optional[Tuple[str, Policy]]: + """ + Fetch a policy version by ID from the DB and convert to Policy for resolution. + + Used when the request body specifies policy_ to execute a specific version + (e.g. published or draft) instead of production. + + Args: + policy_id: The policy version ID (raw UUID, no prefix) + prisma_client: The Prisma client instance + + Returns: + (policy_name, Policy) if found, None otherwise + """ + response = await self.get_policy_by_id_from_db( + policy_id=policy_id, prisma_client=prisma_client + ) + if response is None: + return None + policy = self._parse_policy( + response.policy_name, + { + "inherit": response.inherit, + "description": response.description, + "guardrails": { + "add": response.guardrails_add, + "remove": response.guardrails_remove, + }, + "condition": response.condition, + "pipeline": response.pipeline, + }, + ) + return (response.policy_name, policy) + async def get_all_policies_from_db( self, prisma_client: "PrismaClient", diff --git a/litellm/proxy/policy_engine/policy_resolver.py b/litellm/proxy/policy_engine/policy_resolver.py index a8ad78d6491..c802a970a80 100644 --- a/litellm/proxy/policy_engine/policy_resolver.py +++ b/litellm/proxy/policy_engine/policy_resolver.py @@ -11,12 +11,9 @@ Handles: from typing import Dict, List, Optional, Set, Tuple from litellm._logging import verbose_proxy_logger -from litellm.types.proxy.policy_engine import ( - GuardrailPipeline, - Policy, - PolicyMatchContext, - ResolvedPolicy, -) +from litellm.types.proxy.policy_engine import (GuardrailPipeline, Policy, + PolicyMatchContext, + ResolvedPolicy) class PolicyResolver: @@ -90,7 +87,8 @@ class PolicyResolver: Returns: ResolvedPolicy with final guardrails list """ - from litellm.proxy.policy_engine.condition_evaluator import ConditionEvaluator + from litellm.proxy.policy_engine.condition_evaluator import \ + ConditionEvaluator inheritance_chain = PolicyResolver.resolve_inheritance_chain( policy_name=policy_name, policies=policies @@ -134,12 +132,13 @@ class PolicyResolver: def resolve_guardrails_for_context( context: PolicyMatchContext, policies: Optional[Dict[str, Policy]] = None, + policy_names: Optional[List[str]] = None, ) -> List[str]: """ Resolve the final list of guardrails for a request context. This: - 1. Finds all policies that match the context via policy_attachments + 1. Finds all policies that match the context via policy_attachments (or policy_names if provided) 2. Resolves each policy's guardrails (including inheritance) 3. Evaluates model conditions 4. Combines all guardrails (union) @@ -147,12 +146,14 @@ class PolicyResolver: Args: context: The request context policies: Dictionary of all policies (if None, uses global registry) + policy_names: If provided, use this list instead of attachment matching Returns: List of guardrail names to apply """ from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher - from litellm.proxy.policy_engine.policy_registry import get_policy_registry + from litellm.proxy.policy_engine.policy_registry import \ + get_policy_registry if policies is None: registry = get_policy_registry() @@ -160,8 +161,12 @@ class PolicyResolver: return [] policies = registry.get_all_policies() - # Get matching policies via attachments - matching_policy_names = PolicyMatcher.get_matching_policies(context=context) + # Use provided policy names or get matching policies via attachments + matching_policy_names = ( + policy_names + if policy_names is not None + else PolicyMatcher.get_matching_policies(context=context) + ) if not matching_policy_names: verbose_proxy_logger.debug( @@ -195,6 +200,7 @@ class PolicyResolver: def resolve_pipelines_for_context( context: PolicyMatchContext, policies: Optional[Dict[str, Policy]] = None, + policy_names: Optional[List[str]] = None, ) -> List[Tuple[str, GuardrailPipeline]]: """ Resolve pipelines from matching policies for a request context. @@ -206,12 +212,14 @@ class PolicyResolver: Args: context: The request context policies: Dictionary of all policies (if None, uses global registry) + policy_names: If provided, use this list instead of attachment matching Returns: List of (policy_name, GuardrailPipeline) tuples """ from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher - from litellm.proxy.policy_engine.policy_registry import get_policy_registry + from litellm.proxy.policy_engine.policy_registry import \ + get_policy_registry if policies is None: registry = get_policy_registry() @@ -219,7 +227,11 @@ class PolicyResolver: return [] policies = registry.get_all_policies() - matching_policy_names = PolicyMatcher.get_matching_policies(context=context) + matching_policy_names = ( + policy_names + if policy_names is not None + else PolicyMatcher.get_matching_policies(context=context) + ) if not matching_policy_names: return [] @@ -269,7 +281,8 @@ class PolicyResolver: Returns: Dictionary mapping policy names to ResolvedPolicy objects """ - from litellm.proxy.policy_engine.policy_registry import get_policy_registry + from litellm.proxy.policy_engine.policy_registry import \ + get_policy_registry if policies is None: registry = get_policy_registry() diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index e4b8613d204..9d5039a3e91 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -3,7 +3,7 @@ import copy import json import os import sys -from unittest.mock import MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import Request @@ -11,16 +11,11 @@ from fastapi import Request import litellm from litellm.proxy._types import TeamCallbackMetadata, UserAPIKeyAuth from litellm.proxy.litellm_pre_call_utils import ( - KeyAndTeamLoggingSettings, - LiteLLMProxyRequestSetup, - _get_dynamic_logging_metadata, - _get_enforced_params, - _get_metadata_variable_name, - _update_model_if_key_alias_exists, - add_guardrails_from_policy_engine, - add_litellm_data_to_request, - check_if_token_is_service_account, -) + KeyAndTeamLoggingSettings, LiteLLMProxyRequestSetup, + _get_dynamic_logging_metadata, _get_enforced_params, + _get_metadata_variable_name, _update_model_if_key_alias_exists, + add_guardrails_from_policy_engine, add_litellm_data_to_request, + check_if_token_is_service_account) sys.path.insert( 0, os.path.abspath("../../..") @@ -159,7 +154,8 @@ def test_get_enforced_params( @pytest.mark.asyncio async def test_add_litellm_data_to_request_parses_string_metadata(): - from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.litellm_pre_call_utils import \ + add_litellm_data_to_request # Setup request_mock = MagicMock(spec=Request) @@ -205,7 +201,8 @@ async def test_add_litellm_data_to_request_parses_string_metadata(): @pytest.mark.asyncio async def test_add_litellm_data_to_request_user_spend_and_budget(): - from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.litellm_pre_call_utils import \ + add_litellm_data_to_request request_mock = MagicMock(spec=Request) request_mock.url.path = "/v1/completions" @@ -243,7 +240,8 @@ async def test_add_litellm_data_to_request_user_spend_and_budget(): @pytest.mark.asyncio async def test_add_litellm_data_to_request_audio_transcription_multipart(): - from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.litellm_pre_call_utils import \ + add_litellm_data_to_request # Setup request mock for /v1/audio/transcriptions request_mock = MagicMock(spec=Request) @@ -308,7 +306,8 @@ async def test_add_litellm_data_to_request_disabled_callbacks(): """ Test that litellm_disabled_callbacks from key metadata is properly added to the request data. """ - from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.litellm_pre_call_utils import \ + add_litellm_data_to_request # Setup mock request request_mock = MagicMock(spec=Request) @@ -361,7 +360,8 @@ async def test_add_litellm_data_to_request_disabled_callbacks_empty(): """ Test that litellm_disabled_callbacks is not added when it's empty. """ - from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.litellm_pre_call_utils import \ + add_litellm_data_to_request # Setup mock request request_mock = MagicMock(spec=Request) @@ -413,7 +413,8 @@ async def test_add_litellm_data_to_request_disabled_callbacks_not_present(): """ Test that litellm_disabled_callbacks is not added when it's not present in metadata. """ - from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.litellm_pre_call_utils import \ + add_litellm_data_to_request # Setup mock request request_mock = MagicMock(spec=Request) @@ -465,7 +466,8 @@ async def test_add_litellm_data_to_request_disabled_callbacks_invalid_type(): """ Test that litellm_disabled_callbacks is not added when it's not a list. """ - from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.litellm_pre_call_utils import \ + add_litellm_data_to_request # Setup mock request request_mock = MagicMock(spec=Request) @@ -517,7 +519,8 @@ async def test_add_litellm_data_to_request_disabled_callbacks_with_logging_setti """ Test that litellm_disabled_callbacks works correctly alongside logging settings. """ - from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.litellm_pre_call_utils import \ + add_litellm_data_to_request # Setup mock request request_mock = MagicMock(spec=Request) @@ -1027,7 +1030,8 @@ from unittest.mock import AsyncMock from fastapi.responses import Response from litellm.integrations.custom_logger import CustomLogger -from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing +from litellm.proxy.common_request_processing import \ + ProxyBaseLLMRequestProcessing from litellm.proxy.utils import ProxyLogging from litellm.types.utils import StandardLoggingPayload @@ -1403,7 +1407,8 @@ async def test_embedding_header_forwarding_with_model_group(): importlib.reload(pre_call_utils_module) # Re-import the function after reload to get the fresh version - from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.litellm_pre_call_utils import \ + add_litellm_data_to_request # Setup mock request for embeddings request_mock = MagicMock(spec=Request) @@ -1531,18 +1536,17 @@ async def test_embedding_header_forwarding_without_model_group_config(): litellm.model_group_settings = original_model_group_settings -def test_add_guardrails_from_policy_engine(): +@pytest.mark.asyncio +async def test_add_guardrails_from_policy_engine(): """ Test that add_guardrails_from_policy_engine adds guardrails from matching policies and tracks applied policies in metadata. """ - from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry + 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 ( - Policy, - PolicyAttachment, - PolicyGuardrails, - ) + from litellm.types.proxy.policy_engine import (Policy, PolicyAttachment, + PolicyGuardrails) # Setup test data data = { @@ -1578,7 +1582,7 @@ def test_add_guardrails_from_policy_engine(): attachment_registry._initialized = True # Call the function - add_guardrails_from_policy_engine( + await add_guardrails_from_policy_engine( data=data, metadata_variable_name="metadata", user_api_key_dict=user_api_key_dict, @@ -1601,11 +1605,12 @@ def test_add_guardrails_from_policy_engine(): attachment_registry._initialized = False -def test_add_guardrails_from_policy_engine_accepts_dynamic_policies_and_pops_from_data(): +@pytest.mark.asyncio +async def test_add_guardrails_from_policy_engine_accepts_dynamic_policies_and_pops_from_data(): """ Test that add_guardrails_from_policy_engine accepts dynamic 'policies' from the request body and removes them to prevent forwarding to the LLM provider. - + This is critical because 'policies' is a LiteLLM proxy-specific parameter that should not be sent to the actual LLM API (e.g., OpenAI, Anthropic, etc.). """ @@ -1631,7 +1636,7 @@ def test_add_guardrails_from_policy_engine_accepts_dynamic_policies_and_pops_fro policy_registry._initialized = False # Call the function - should accept dynamic policies and not raise an error - add_guardrails_from_policy_engine( + await add_guardrails_from_policy_engine( data=data, metadata_variable_name="metadata", user_api_key_dict=user_api_key_dict, @@ -1646,3 +1651,69 @@ def test_add_guardrails_from_policy_engine_accepts_dynamic_policies_and_pops_fro assert "messages" in data assert data["messages"] == [{"role": "user", "content": "Hello"}] assert "metadata" in data + + +@pytest.mark.asyncio +async def test_add_guardrails_from_policy_engine_policy_version_by_id(): + """ + Test that add_guardrails_from_policy_engine executes a specific policy version + when policy_ is passed in the request body. + """ + 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 Policy, PolicyGuardrails + + policy_version_uuid = "12345678-1234-5678-1234-567812345678" + policy_version_ref = f"policy_{policy_version_uuid}" + + # Policy from the specific version (e.g. published) - different guardrail than production + published_version_policy = Policy( + guardrails=PolicyGuardrails(add=["published_version_guardrail"]), + ) + + data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + "policies": [policy_version_ref], + "metadata": {}, + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + team_alias="test-team", + key_alias="test-key", + ) + + policy_registry = get_policy_registry() + policy_registry._policies = {} + policy_registry._initialized = True + + attachment_registry = get_attachment_registry() + attachment_registry._attachments = [] + attachment_registry._initialized = True + + mock_prisma = MagicMock() + + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): + with patch.object( + policy_registry, + "get_policy_by_id_for_request", + new_callable=AsyncMock, + return_value=("test-policy-from-version", published_version_policy), + ): + await add_guardrails_from_policy_engine( + data=data, + metadata_variable_name="metadata", + user_api_key_dict=user_api_key_dict, + ) + + # Verify guardrails from the specific version were applied + assert "metadata" in data + assert "guardrails" in data["metadata"] + assert "published_version_guardrail" in data["metadata"]["guardrails"] + assert "policies" not in data + + # Clean up + policy_registry._policies = {} + policy_registry._initialized = False diff --git a/ui/litellm-dashboard/src/components/playground/complianceUI/ComplianceUI.tsx b/ui/litellm-dashboard/src/components/playground/complianceUI/ComplianceUI.tsx index 721fda31f2a..ec151036b2c 100644 --- a/ui/litellm-dashboard/src/components/playground/complianceUI/ComplianceUI.tsx +++ b/ui/litellm-dashboard/src/components/playground/complianceUI/ComplianceUI.tsx @@ -8,9 +8,10 @@ import { } from "@/data/compliancePrompts"; import { getGuardrailsList, - getPoliciesList, testPoliciesAndGuardrails, } from "@/components/networking"; +import PolicySelector, { getPolicyOptionEntries } from "@/components/policies/PolicySelector"; +import { Policy } from "@/components/policies/types"; import { makeOpenAIChatCompletionRequest } from "../llm_calls/chat_completion"; import { AlertTriangle, @@ -98,11 +99,6 @@ interface QuickTestMessage { type ResultFilter = "all" | "matches" | "mismatches" | "pending"; type RightPanelTab = "quick-test" | "batch-results"; -interface PolicyOption { - id: string; - name: string; -} - interface GuardrailOption { id: string; name: string; @@ -132,11 +128,10 @@ export default function ComplianceUI({ }: ComplianceUIProps) { const frameworks = getFrameworks(); - const [policyOptions, setPolicyOptions] = useState([]); + const [policyValueToLabel, setPolicyValueToLabel] = useState>(new Map()); const [guardrailOptions, setGuardrailOptions] = useState([]); const [selectedPolicies, setSelectedPolicies] = useState([]); const [selectedGuardrails, setSelectedGuardrails] = useState([]); - const [showPolicyDropdown, setShowPolicyDropdown] = useState(false); const [showGuardrailDropdown, setShowGuardrailDropdown] = useState(false); const [selectedPromptIds, setSelectedPromptIds] = useState>(new Set()); @@ -164,20 +159,16 @@ export default function ComplianceUI({ const [expandedResults, setExpandedResults] = useState>(new Set()); const batchAbortControllerRef = useRef(null); + const handlePoliciesLoaded = useCallback((policies: Policy[]) => { + const entries = getPolicyOptionEntries(policies); + setPolicyValueToLabel(new Map(entries.map((e) => [e.value, e.label]))); + }, []); + useEffect(() => { if (!accessToken) return; - const fetchOptions = async () => { + const fetchGuardrails = async () => { try { - const [policiesRes, guardrailsRes] = await Promise.all([ - getPoliciesList(accessToken).catch(() => ({ policies: [] })), - getGuardrailsList(accessToken).catch(() => ({ guardrails: [] })), - ]); - setPolicyOptions( - (policiesRes.policies || []).map((p: { policy_name: string; policy_id?: string }) => ({ - id: p.policy_id ?? p.policy_name, - name: p.policy_name, - })) - ); + const guardrailsRes = await getGuardrailsList(accessToken).catch(() => ({ guardrails: [] })); setGuardrailOptions( (guardrailsRes.guardrails || []).map((g: { guardrail_name: string }) => ({ id: g.guardrail_name, @@ -186,11 +177,10 @@ export default function ComplianceUI({ })) ); } catch { - setPolicyOptions([]); setGuardrailOptions([]); } }; - fetchOptions(); + fetchGuardrails(); }, [accessToken]); useEffect(() => { @@ -282,12 +272,6 @@ export default function ComplianceUI({ const deselectAll = () => setSelectedPromptIds(new Set()); - const togglePolicy = (id: string) => { - setSelectedPolicies((prev) => - prev.includes(id) ? prev.filter((p) => p !== id) : [...prev, id] - ); - }; - const toggleGuardrail = (id: string) => { setSelectedGuardrails((prev) => prev.includes(id) ? prev.filter((g) => g !== id) : [...prev, id] @@ -766,76 +750,13 @@ export default function ComplianceUI({ -
- - {showPolicyDropdown && ( -
- {policyOptions.length === 0 ? ( -
- No policies available. Create policies in the Policies page. -
- ) : ( - policyOptions.map((policy) => ( - - )) - )} -
- )} -
- {selectedPolicies.length > 0 && ( -
- {selectedPolicies.map((id) => { - const p = policyOptions.find((x) => x.id === id); - return ( - - {p?.name} - - - ); - })} -
+ {accessToken && ( + )}
@@ -852,10 +773,7 @@ export default function ComplianceUI({