Merge PR #21862 branch for policy versioning

Co-authored-by: Krish Dholakia <krrishdholakia@gmail.com>
This commit is contained in:
Cursor Agent 2026-02-22 03:32:44 +00:00
commit b1c7ff9f58
55 changed files with 2646 additions and 570 deletions

View file

@ -0,0 +1,17 @@
-- DropIndex
DROP INDEX "LiteLLM_PolicyTable_policy_name_key";
-- AlterTable
ALTER TABLE "LiteLLM_PolicyTable" ADD COLUMN "is_latest" BOOLEAN NOT NULL DEFAULT true,
ADD COLUMN "parent_version_id" TEXT,
ADD COLUMN "production_at" TIMESTAMP(3),
ADD COLUMN "published_at" TIMESTAMP(3),
ADD COLUMN "version_number" INTEGER NOT NULL DEFAULT 1,
ADD COLUMN "version_status" TEXT NOT NULL DEFAULT 'production';
-- CreateIndex
CREATE INDEX "LiteLLM_PolicyTable_policy_name_version_status_idx" ON "LiteLLM_PolicyTable"("policy_name", "version_status");
-- CreateIndex
CREATE UNIQUE INDEX "LiteLLM_PolicyTable_policy_name_version_number_key" ON "LiteLLM_PolicyTable"("policy_name", "version_number");

View file

@ -213,53 +213,6 @@ model LiteLLM_DeletedTeamTable {
@@index([created_at])
}
// Audit table for deleted teams - preserves spend and team information for historical tracking
model LiteLLM_DeletedTeamTable {
id String @id @default(uuid())
team_id String // Original team_id
team_alias String?
organization_id String?
object_permission_id String?
admins String[]
members String[]
members_with_roles Json @default("{}")
metadata Json @default("{}")
max_budget Float?
soft_budget Float?
spend Float @default(0.0)
models String[]
max_parallel_requests Int?
tpm_limit BigInt?
rpm_limit BigInt?
budget_duration String?
budget_reset_at DateTime?
blocked Boolean @default(false)
model_spend Json @default("{}")
model_max_budget Json @default("{}")
router_settings Json? @default("{}")
team_member_permissions String[] @default([])
access_group_ids String[] @default([])
policies String[] @default([])
model_id Int? // id for LiteLLM_ModelTable -> stores team-level model aliases
allow_team_guardrail_config Boolean @default(false)
// Original timestamps from team creation/updates
created_at DateTime? @map("created_at")
updated_at DateTime? @map("updated_at")
// Deletion metadata
deleted_at DateTime @default(now()) @map("deleted_at")
deleted_by String? @map("deleted_by") // User who deleted the team
deleted_by_api_key String? @map("deleted_by_api_key") // API key hash that performed the deletion
litellm_changed_by String? @map("litellm_changed_by") // From litellm-changed-by header if provided
@@index([team_id])
@@index([deleted_at])
@@index([organization_id])
@@index([team_alias])
@@index([created_at])
}
// Track spend, rate limit, budget Users
model LiteLLM_UserTable {
user_id String @id
@ -320,6 +273,7 @@ model LiteLLM_MCPServerTable {
alias String?
description String?
url String?
spec_path String?
transport String @default("sse")
auth_type String?
credentials Json? @default("{}")
@ -1009,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

View file

@ -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_<uuid> 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_<uuid> 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_<uuid> 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,57 @@ 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_<uuid>)
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)
# Resolve policy versions by ID from in-memory cache (populated by sync job; no DB in hot path)
merged_policies: Dict[str, Any] = dict(registry.get_all_policies())
fetched_policy_names: List[str] = []
for policy_id in request_body_version_ids:
result = registry.get_policy_by_id_for_request(policy_id=policy_id)
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 in cache, skipping"
)
# 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(

View file

@ -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 <your_api_key>"
curl -X GET "http://localhost:4000/policies/list?version_status=production" \\
-H "Authorization: Bearer <your_api_key>"
```
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,

View file

@ -9,23 +9,48 @@ 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,
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
# Prefix for policy version IDs in request body. Use policy_<uuid> 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."""
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:
"""
@ -42,6 +67,7 @@ class PolicyRegistry:
def __init__(self):
self._policies: Dict[str, Policy] = {}
self._policies_by_id: Dict[str, Tuple[str, Policy]] = {}
self._initialized: bool = False
def load_policies(self, policies_config: Dict[str, Any]) -> None:
@ -53,6 +79,7 @@ class PolicyRegistry:
This is the raw config from the YAML file.
"""
self._policies = {}
self._policies_by_id = {}
for policy_name, policy_data in policies_config.items():
try:
@ -88,7 +115,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
@ -108,7 +137,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
@ -231,13 +262,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
@ -268,28 +304,17 @@ 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,
},
)
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 +327,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 +337,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),
@ -331,7 +370,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())
@ -341,36 +382,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 +393,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 +415,27 @@ 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,59 +463,54 @@ 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)}")
def get_policy_by_id_for_request(self, policy_id: str) -> Optional[Tuple[str, Policy]]:
"""
Return a policy version by ID from in-memory cache (no DB access).
Used when the request body specifies policy_<uuid> to execute a specific version
(e.g. published or draft). The cache is populated by sync_policies_from_db,
which loads draft and published versions keyed by policy_id.
Args:
policy_id: The policy version ID (raw UUID, no prefix)
Returns:
(policy_name, Policy) if found, None otherwise
"""
return self._policies_by_id.get(policy_id)
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,14 +521,16 @@ class PolicyRegistry:
) -> None:
"""
Sync policies from the database to in-memory registry.
Args:
prisma_client: The Prisma client instance
- Production versions are loaded into _policies (by policy name) for resolution.
- Draft and published versions are loaded into _policies_by_id so request-body
policy_<uuid> overrides can be resolved without DB access in the hot path.
"""
try:
policies = await self.get_all_policies_from_db(prisma_client)
for policy_response in policies:
self._policies = {}
production = await self.get_all_policies_from_db(
prisma_client, version_status="production"
)
for policy_response in production:
policy = self._parse_policy(
policy_response.policy_name,
{
@ -521,9 +546,31 @@ class PolicyRegistry:
)
self.add_policy(policy_response.policy_name, policy)
self._policies_by_id = {}
non_production = await prisma_client.db.litellm_policytable.find_many(
where={"version_status": {"in": ["draft", "published"]}},
order={"created_at": "desc"},
)
for row in non_production:
policy = self._parse_policy(
row.policy_name,
{
"inherit": row.inherit,
"description": row.description,
"guardrails": {
"add": row.guardrails_add or [],
"remove": row.guardrails_remove or [],
},
"condition": row.condition,
"pipeline": row.pipeline,
},
)
self._policies_by_id[row.policy_id] = (row.policy_name, policy)
self._initialized = True
verbose_proxy_logger.info(
f"Synced {len(policies)} policies from DB to in-memory registry"
f"Synced {len(production)} production policies and {len(non_production)} "
"draft/published (by ID) from DB to in-memory registry"
)
except Exception as e:
verbose_proxy_logger.exception(f"Error syncing policies from DB: {e}")
@ -536,22 +583,24 @@ 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
"""
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:
@ -569,19 +618,342 @@ 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}")
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 [],
"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)
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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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,65 @@ 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_<uuid> 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
with patch.object(
policy_registry,
"get_policy_by_id_for_request",
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

View file

@ -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<any> => {
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<any> => {
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}`;

View file

@ -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<PolicyOption[]>([]);
const [policyValueToLabel, setPolicyValueToLabel] = useState<Map<string, string>>(new Map());
const [guardrailOptions, setGuardrailOptions] = useState<GuardrailOption[]>([]);
const [selectedPolicies, setSelectedPolicies] = useState<string[]>([]);
const [selectedGuardrails, setSelectedGuardrails] = useState<string[]>([]);
const [showPolicyDropdown, setShowPolicyDropdown] = useState(false);
const [showGuardrailDropdown, setShowGuardrailDropdown] = useState(false);
const [selectedPromptIds, setSelectedPromptIds] = useState<Set<string>>(new Set());
@ -164,20 +159,16 @@ export default function ComplianceUI({
const [expandedResults, setExpandedResults] = useState<Set<string>>(new Set());
const batchAbortControllerRef = useRef<AbortController | null>(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({
<label className="text-[11px] font-medium text-gray-500 uppercase tracking-wide mb-1.5 block">
Policies
</label>
<div className="relative">
<button
type="button"
onClick={() => {
setShowPolicyDropdown(!showPolicyDropdown);
setShowGuardrailDropdown(false);
}}
className="w-full flex items-center justify-between border border-gray-200 rounded-lg px-3 py-2 text-sm text-left hover:border-gray-300 transition-colors"
>
<span
className={
selectedPolicies.length > 0 ? "text-gray-700" : "text-gray-400"
}
>
{selectedPolicies.length > 0
? `${selectedPolicies.length} selected`
: "None selected"}
</span>
<ChevronDown className="w-4 h-4 text-gray-400" />
</button>
{showPolicyDropdown && (
<div className="absolute z-30 top-full left-0 right-0 mt-1 bg-white border border-gray-200 rounded-lg shadow-lg py-1 max-h-52 overflow-y-auto">
{policyOptions.length === 0 ? (
<div className="px-3 py-2 text-xs text-gray-500">
No policies available. Create policies in the Policies page.
</div>
) : (
policyOptions.map((policy) => (
<button
key={policy.id}
type="button"
onClick={() => togglePolicy(policy.id)}
className="w-full flex items-center gap-2.5 px-3 py-2 text-sm text-left hover:bg-gray-50"
>
<div
className={`w-4 h-4 rounded border flex items-center justify-center flex-shrink-0 ${selectedPolicies.includes(policy.id) ? "bg-blue-500 border-blue-500" : "border-gray-300"}`}
>
{selectedPolicies.includes(policy.id) && (
<Check className="w-3 h-3 text-white" />
)}
</div>
<span className="text-gray-700">{policy.name}</span>
</button>
))
)}
</div>
)}
</div>
{selectedPolicies.length > 0 && (
<div className="flex flex-wrap gap-1 mt-1.5">
{selectedPolicies.map((id) => {
const p = policyOptions.find((x) => x.id === id);
return (
<span
key={id}
className="inline-flex items-center gap-1 text-[11px] bg-blue-50 text-blue-700 px-1.5 py-0.5 rounded font-medium"
>
{p?.name}
<button
type="button"
onClick={() => togglePolicy(id)}
className="hover:text-blue-900"
aria-label="Remove"
>
<X className="w-2.5 h-2.5" />
</button>
</span>
);
})}
</div>
{accessToken && (
<PolicySelector
value={selectedPolicies}
onChange={setSelectedPolicies}
accessToken={accessToken}
onPoliciesLoaded={handlePoliciesLoaded}
/>
)}
</div>
@ -852,10 +773,7 @@ export default function ComplianceUI({
<div className="relative">
<button
type="button"
onClick={() => {
setShowGuardrailDropdown(!showGuardrailDropdown);
setShowPolicyDropdown(false);
}}
onClick={() => setShowGuardrailDropdown(!showGuardrailDropdown)}
className="w-full flex items-center justify-between border border-gray-200 rounded-lg px-3 py-2 text-sm text-left hover:border-gray-300 transition-colors"
>
<span
@ -1342,17 +1260,14 @@ export default function ComplianceUI({
<span className="text-[11px] font-medium text-gray-500">
Testing against:
</span>
{selectedPolicies.map((id) => {
const p = policyOptions.find((x) => x.id === id);
return (
<span
key={id}
className="text-[11px] bg-blue-50 text-blue-700 px-2 py-0.5 rounded font-medium"
>
{p?.name}
</span>
);
})}
{selectedPolicies.map((id) => (
<span
key={id}
className="text-[11px] bg-blue-50 text-blue-700 px-2 py-0.5 rounded font-medium"
>
{policyValueToLabel.get(id) ?? id}
</span>
))}
{selectedGuardrails.map((id) => {
const g = guardrailOptions.find((x) => x.id === id);
return (

View file

@ -3,20 +3,53 @@ import { Select } from "antd";
import { Policy } from "./types";
import { getPoliciesList } from "../networking";
/** Prefix for policy version IDs in request body; must match backend POLICY_VERSION_ID_PREFIX. */
export const POLICY_VERSION_ID_PREFIX = "policy_";
/** Build the value sent in the request body: policy_<uuid> so backend executes this exact version. */
export function policyVersionRef(policyId: string): string {
return `${POLICY_VERSION_ID_PREFIX}${policyId}`;
}
/** Build select options from policies (filter non-draft, label with name/version/status). */
export function getPolicyOptionEntries(policies: Policy[]): { value: string; label: string }[] {
return policies
.filter((policy) => (policy.version_status ?? "draft") !== "draft")
.map((policy) => {
const versionNum = policy.version_number ?? 1;
const status = policy.version_status ?? "draft";
const label = `${policy.policy_name} — v${versionNum} (${status})${
policy.description ? ` — ${policy.description}` : ""
}`;
const isProduction = status === "production";
return {
label,
value: isProduction
? policy.policy_name
: policy.policy_id
? policyVersionRef(policy.policy_id)
: policy.policy_name,
};
});
}
interface PolicySelectorProps {
onChange: (selectedPolicies: string[]) => void;
value?: string[];
className?: string;
accessToken: string;
disabled?: boolean;
/** Called after policies are loaded; use to build value→label map for display elsewhere. */
onPoliciesLoaded?: (policies: Policy[]) => void;
}
const PolicySelector: React.FC<PolicySelectorProps> = ({
onChange,
value,
className,
accessToken,
disabled
const PolicySelector: React.FC<PolicySelectorProps> = ({
onChange,
value,
className,
accessToken,
disabled,
onPoliciesLoaded,
}) => {
const [policies, setPolicies] = useState<Policy[]>([]);
const [loading, setLoading] = useState(false);
@ -28,10 +61,9 @@ const PolicySelector: React.FC<PolicySelectorProps> = ({
setLoading(true);
try {
const response = await getPoliciesList(accessToken);
console.log("Policies response:", response);
if (response.policies) {
console.log("Policies data:", response.policies);
setPolicies(response.policies);
onPoliciesLoaded?.(response.policies);
}
} catch (error) {
console.error("Error fetching policies:", error);
@ -41,10 +73,9 @@ const PolicySelector: React.FC<PolicySelectorProps> = ({
};
fetchPolicies();
}, [accessToken]);
}, [accessToken, onPoliciesLoaded]);
const handlePolicyChange = (selectedValues: string[]) => {
console.log("Selected policies:", selectedValues);
onChange(selectedValues);
};
@ -53,19 +84,17 @@ const PolicySelector: React.FC<PolicySelectorProps> = ({
<Select
mode="multiple"
disabled={disabled}
placeholder={disabled ? "Setting policies is a premium feature." : "Select policies"}
placeholder={
disabled
? "Setting policies is a premium feature."
: "Select policies (production or published versions)"
}
onChange={handlePolicyChange}
value={value}
loading={loading}
className={className}
allowClear
options={policies.map((policy) => {
console.log("Mapping policy:", policy);
return {
label: `${policy.policy_name}${policy.description ? ` - ${policy.description}` : ""}`,
value: policy.policy_name,
};
})}
options={getPolicyOptionEntries(policies)}
optionFilterProp="label"
showSearch
style={{ width: "100%" }}

View file

@ -452,11 +452,7 @@ const PoliciesPanel: React.FC<PoliciesPanelProps> = ({
onEdit={(policy) => {
setEditingPolicy(policy);
setSelectedPolicyId(null);
if (policy.pipeline) {
setShowFlowBuilder(true);
} else {
setIsAddPolicyModalVisible(true);
}
setShowFlowBuilder(true);
}}
accessToken={accessToken}
isAdmin={isAdmin}
@ -469,11 +465,7 @@ const PoliciesPanel: React.FC<PoliciesPanelProps> = ({
onDeleteClick={handleDeleteClick}
onEditClick={(policy) => {
setEditingPolicy(policy);
if (policy.pipeline) {
setShowFlowBuilder(true);
} else {
setIsAddPolicyModalVisible(true);
}
setShowFlowBuilder(true);
}}
onViewClick={(policyId) => setSelectedPolicyId(policyId)}
isAdmin={isAdmin}
@ -643,6 +635,17 @@ const PoliciesPanel: React.FC<PoliciesPanelProps> = ({
availableGuardrails={guardrailsList}
createPolicy={createPolicyCall}
updatePolicy={updatePolicyCall}
onVersionCreated={(newPolicy) => {
setEditingPolicy(newPolicy);
fetchPolicies();
}}
onSelectVersion={(policy) => {
setEditingPolicy(policy);
}}
onVersionStatusUpdated={(updatedPolicy) => {
setEditingPolicy(updatedPolicy);
fetchPolicies();
}}
/>
)}
</div>

View file

@ -1,12 +1,27 @@
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,
getFrameworks,
} from "../../data/compliancePrompts";
import type { CompliancePrompt } from "../../data/compliancePrompts";
const TEST_SOURCE_QUICK = "quick_chat";
const TEST_SOURCE_ALL = "__all__";
function getPromptsForTestSource(source: string): CompliancePrompt[] {
if (source === TEST_SOURCE_QUICK) return [];
if (source === TEST_SOURCE_ALL) return getComplianceDatasetPrompts();
const fw = getFrameworks().find((f) => f.name === source);
return fw ? fw.categories.flatMap((c) => c.prompts) : [];
}
const { Text } = Typography;
@ -55,6 +70,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,15 +602,41 @@ const TERMINAL_STYLES: Record<string, { bg: string; color: string }> = {
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 testSourceOptions = [
{ value: TEST_SOURCE_QUICK, label: "Quick chat (custom message)" },
...getFrameworks().map((f) => ({ value: f.name, label: f.name })),
{ value: TEST_SOURCE_ALL, label: "All compliance datasets" },
];
const PipelineTestPanel: React.FC<PipelineTestPanelProps> = ({
pipeline,
accessToken,
onClose,
}) => {
const [testSource, setTestSource] = useState<string>(TEST_SOURCE_QUICK);
const [testMessage, setTestMessage] = useState("Hello, can you help me?");
const [isRunning, setIsRunning] = useState(false);
const [result, setResult] = useState<PipelineTestResult | null>(null);
const [error, setError] = useState<string | null>(null);
const [complianceResults, setComplianceResults] = useState<ComplianceRunEntry[]>([]);
const isQuickChat = testSource === TEST_SOURCE_QUICK;
const promptsForSource = getPromptsForTestSource(testSource);
const isDataset = promptsForSource.length > 0;
const handleRunTest = async () => {
if (!accessToken) return;
@ -578,22 +647,47 @@ const PipelineTestPanel: React.FC<PipelineTestPanelProps> = ({
return;
}
setError(null);
setIsRunning(true);
setResult(null);
setError(null);
setComplianceResults([]);
try {
const data = await testPipelineCall(
accessToken,
pipeline,
[{ role: "user", content: testMessage }]
);
setResult(data);
} catch (e) {
setError(e instanceof Error ? e.message : String(e));
} finally {
setIsRunning(false);
if (isQuickChat) {
try {
const data = await testPipelineCall(
accessToken,
pipeline,
[{ role: "user", content: testMessage }]
);
setResult(data);
} catch (e) {
setError(e instanceof Error ? e.message : String(e));
} finally {
setIsRunning(false);
}
return;
}
const entries: ComplianceRunEntry[] = [];
for (const prompt of promptsForSource) {
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);
setIsRunning(false);
};
return (
@ -637,23 +731,53 @@ const PipelineTestPanel: React.FC<PipelineTestPanelProps> = ({
{/* Input section */}
<div style={{ padding: 16, borderBottom: "1px solid #e5e7eb" }}>
<label style={{ fontSize: 12, fontWeight: 500, color: "#6b7280", display: "block", marginBottom: 6 }}>
Test Message
Test with
</label>
<textarea
value={testMessage}
onChange={(e) => setTestMessage(e.target.value)}
placeholder="Enter a test message..."
rows={3}
style={{
width: "100%",
border: "1px solid #d1d5db",
borderRadius: 6,
padding: "8px 10px",
fontSize: 13,
resize: "vertical",
fontFamily: "inherit",
}}
<Select
value={testSource}
onChange={setTestSource}
options={testSourceOptions}
style={{ width: "100%", marginBottom: 12 }}
size="middle"
/>
{isQuickChat && (
<>
<label style={{ fontSize: 12, fontWeight: 500, color: "#6b7280", display: "block", marginBottom: 6 }}>
Message
</label>
<textarea
value={testMessage}
onChange={(e) => setTestMessage(e.target.value)}
placeholder="Enter a test message..."
rows={3}
style={{
width: "100%",
border: "1px solid #d1d5db",
borderRadius: 6,
padding: "8px 10px",
fontSize: 13,
resize: "vertical",
fontFamily: "inherit",
}}
/>
</>
)}
{isDataset && (
<div
style={{
fontSize: 12,
color: "#6b7280",
padding: "8px 10px",
backgroundColor: "#f9fafb",
borderRadius: 6,
marginBottom: 8,
}}
>
{testSource === TEST_SOURCE_ALL
? "Run pipeline against all compliance prompts (EU AI Act, GDPR, Topic Blocking, Airline, etc.)."
: `Run pipeline against ${promptsForSource.length} prompts from "${testSource}".`}
</div>
)}
<Button
onClick={handleRunTest}
loading={isRunning}
@ -773,11 +897,353 @@ const PipelineTestPanel: React.FC<PipelineTestPanelProps> = ({
</div>
)}
{!result && !error && (
<div style={{ textAlign: "center", color: "#9ca3af", fontSize: 13, marginTop: 24 }}>
Enter a test message and click "Run Test" to execute the pipeline
{complianceResults.length > 0 && (
<div style={{ marginTop: 16 }}>
<div
style={{
fontSize: 13,
fontWeight: 600,
color: "#111827",
marginBottom: 8,
}}
>
Compliance dataset
</div>
<div
style={{
fontSize: 12,
color: "#6b7280",
marginBottom: 10,
}}
>
{complianceResults.filter((e) => e.matched).length} / {complianceResults.length} matched
expected
</div>
<div
style={{
maxHeight: 320,
overflowY: "auto",
border: "1px solid #e5e7eb",
borderRadius: 8,
}}
>
{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 (
<div
key={entry.prompt.id ?? i}
style={{
padding: "8px 10px",
borderBottom:
i < complianceResults.length - 1
? "1px solid #e5e7eb"
: "none",
fontSize: 12,
}}
>
<div
style={{
color: "#374151",
marginBottom: 4,
overflow: "hidden",
textOverflow: "ellipsis",
whiteSpace: "nowrap",
}}
title={entry.prompt.prompt}
>
{entry.prompt.prompt}
</div>
<div
style={{
display: "flex",
alignItems: "center",
gap: 8,
flexWrap: "wrap",
}}
>
<span style={{ color: "#6b7280" }}>
expected: {entry.prompt.expectedResult}
</span>
<span style={{ color: "#9ca3af" }}>→</span>
<span style={{ color: "#6b7280" }}>
actual: {actual}
</span>
<span
style={{
backgroundColor: matchStyle.bg,
color: matchStyle.color,
padding: "1px 6px",
borderRadius: 4,
fontWeight: 600,
}}
>
{entry.matched ? "✓" : "✗"}
</span>
</div>
{entry.error && (
<div style={{ color: "#dc2626", marginTop: 4 }}>
{entry.error}
</div>
)}
</div>
);
})}
</div>
</div>
)}
{!result && !error && complianceResults.length === 0 && (
<div style={{ textAlign: "center", color: "#9ca3af", fontSize: 13, marginTop: 24 }}>
Choose a test source above (quick chat or a compliance dataset) and click "Run Test"
</div>
)}
</div>
</div>
);
};
// ─────────────────────────────────────────────────────────────────────────────
// 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<PolicyVersionsSidebarProps> = ({
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 (
<div
style={{
width: 260,
flexShrink: 0,
backgroundColor: "#fff",
borderRight: "1px solid #e5e7eb",
display: "flex",
flexDirection: "column",
overflow: "hidden",
}}
>
<div style={{ padding: 16, overflowY: "auto", flex: 1 }}>
{/* Versions section */}
<div style={{ marginBottom: 24 }}>
<span
style={{
fontSize: 11,
fontWeight: 700,
textTransform: "uppercase",
color: "#6b7280",
letterSpacing: "0.06em",
display: "block",
marginBottom: 4,
}}
>
Versions
</span>
<span
style={{
fontSize: 11,
color: "#6b7280",
lineHeight: 1.4,
display: "block",
marginBottom: 12,
}}
>
Production = the version used when anyone calls this policy by name.
</span>
<Button
onClick={onNewVersion}
disabled={!accessToken || isCreatingVersion}
loading={isCreatingVersion}
style={{ width: "100%", marginBottom: 12 }}
>
+ New Version
</Button>
{isLoading ? (
<div style={{ display: "flex", justifyContent: "center", padding: 16 }}>
<Spin size="small" />
</div>
) : versions.length === 0 ? (
<span style={{ fontSize: 13, color: "#9ca3af" }}>
No versions found
</span>
) : (
<div className="flex flex-col gap-1">
{versions.map((v) => {
const statusStyle =
VERSION_STATUS_STYLES[v.version_status ?? "draft"] ??
VERSION_STATUS_STYLES.draft;
const isActive = v.policy_id === editingPolicyId;
return (
<button
key={v.policy_id}
type="button"
onClick={() => onSelectVersion(v)}
style={{
width: "100%",
textAlign: "left",
padding: "10px 12px",
borderRadius: 8,
border: isActive ? "1px solid #6366f1" : "1px solid #e5e7eb",
backgroundColor: isActive ? "#eef2ff" : "#fff",
cursor: "pointer",
}}
>
<div className="flex items-center justify-between" style={{ marginBottom: 4 }}>
<span style={{ fontSize: 13, fontWeight: 600, color: "#111827" }}>
v{v.version_number ?? 1}
</span>
<span
style={{
fontSize: 10,
fontWeight: 600,
textTransform: "uppercase",
backgroundColor: statusStyle.bg,
color: statusStyle.color,
padding: "2px 6px",
borderRadius: 4,
}}
>
{v.version_status ?? "draft"}
</span>
</div>
</button>
);
})}
</div>
)}
{/* Publish / Promote to production for selected version */}
{(canPublish || canPromote) && (
<div style={{ marginTop: 12, paddingTop: 12, borderTop: "1px solid #e5e7eb" }}>
{canPublish && (
<>
<Button
variant="secondary"
onClick={onPublish}
disabled={!accessToken || isUpdatingStatus}
loading={isUpdatingStatus}
style={{ width: "100%", marginBottom: 8 }}
>
Publish
</Button>
<span
style={{
fontSize: 11,
color: "#6b7280",
lineHeight: 1.4,
display: "block",
marginBottom: canPromote ? 8 : 0,
}}
>
Published versions can be tested in the Playground before promoting to production.
</span>
</>
)}
{canPromote && (
<>
<Button
onClick={onPromoteToProduction}
disabled={!accessToken || isUpdatingStatus}
loading={isUpdatingStatus}
style={{ width: "100%", marginBottom: 8 }}
>
Promote to production
</Button>
<span
style={{
fontSize: 11,
color: "#6b7280",
lineHeight: 1.4,
display: "block",
}}
>
This version will be used when anyone calls this policy by name.
</span>
</>
)}
</div>
)}
</div>
{/* Silent Mirroring section */}
<div>
<div className="flex items-center gap-2" style={{ marginBottom: 8 }}>
<span
style={{
fontSize: 11,
fontWeight: 700,
textTransform: "uppercase",
color: "#6b7280",
letterSpacing: "0.06em",
}}
>
Silent Mirroring
</span>
<span
style={{
fontSize: 10,
fontWeight: 600,
backgroundColor: "#eef2ff",
color: "#6366f1",
padding: "2px 6px",
borderRadius: 4,
}}
>
COMING SOON
</span>
</div>
<span
style={{
fontSize: 12,
color: "#6b7280",
lineHeight: 1.5,
display: "block",
}}
>
Test policy versions on production traffic without blocking requests.
Shadow testing helps validate changes before full rollout.
</span>
</div>
</div>
</div>
);
@ -795,6 +1261,9 @@ interface FlowBuilderPageProps {
availableGuardrails: Guardrail[];
createPolicy: (accessToken: string, policyData: any) => Promise<any>;
updatePolicy: (accessToken: string, policyId: string, policyData: any) => Promise<any>;
onVersionCreated?: (newPolicy: Policy) => void;
onSelectVersion?: (policy: Policy) => void;
onVersionStatusUpdated?: (updatedPolicy: Policy) => void;
}
export const FlowBuilderPage: React.FC<FlowBuilderPageProps> = ({
@ -805,16 +1274,114 @@ export const FlowBuilderPage: React.FC<FlowBuilderPageProps> = ({
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<GuardrailPipeline>(
editingPolicy?.pipeline || { mode: "pre_call", steps: [createDefaultStep()] }
() => derivePipelineFromPolicy(editingPolicy)
);
const [versions, setVersions] = useState<Policy[]>([]);
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);
const list = await listPolicyVersions(accessToken, editingPolicy.policy_name);
setVersions(list.versions ?? []);
} 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. You can test it in the Playground by selecting this version in the Policies dropdown."
);
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 +1416,13 @@ export const FlowBuilderPage: React.FC<FlowBuilderPageProps> = ({
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 +1530,24 @@ export const FlowBuilderPage: React.FC<FlowBuilderPageProps> = ({
/>
</div>
{/* Flow builder canvas + test panel */}
{/* Sidebar (when editing) + Flow builder canvas + test panel */}
<div style={{ flex: 1, display: "flex", overflow: "hidden" }}>
{showVersionsSidebar && (
<PolicyVersionsSidebar
policyName={policyName}
editingPolicyId={editingPolicy?.policy_id ?? null}
editingVersionStatus={editingPolicy?.version_status}
accessToken={accessToken}
versions={versions}
isLoading={isVersionsLoading}
isCreatingVersion={isCreatingVersion}
isUpdatingStatus={isUpdatingStatus}
onNewVersion={handleNewVersion}
onSelectVersion={handleSelectVersion}
onPublish={handlePublishVersion}
onPromoteToProduction={handlePromoteToProduction}
/>
)}
<div
style={{
flex: 1,

View file

@ -1,4 +1,4 @@
import React, { useState } from "react";
import React, { useMemo, useState } from "react";
import { Table, TableBody, TableCell, TableHead, TableHeaderCell, TableRow, Icon, Button, Badge } from "@tremor/react";
import { TrashIcon, PencilIcon, SwitchVerticalIcon, ChevronUpIcon, ChevronDownIcon } from "@heroicons/react/outline";
import { Tooltip, Tag } from "antd";
@ -12,6 +12,32 @@ import {
} from "@tanstack/react-table";
import { Policy } from "./types";
/** One row per policy name; primaryPolicy is used for display and for Edit (FlowBuilder loads all versions) */
interface PolicyRow {
policy_name: string;
primaryPolicy: Policy;
versionCount: number;
}
function groupPoliciesByName(policies: Policy[]): PolicyRow[] {
const byName = new Map<string, Policy[]>();
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<PolicyTableProps> = ({
onViewClick,
isAdmin = false,
}) => {
const [sorting, setSorting] = useState<SortingState>([{ id: "created_at", desc: true }]);
const [sorting, setSorting] = useState<SortingState>([{ 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<Policy>[] = [
{
header: "Policy ID",
accessorKey: "policy_id",
cell: (info: any) => (
<Tooltip title={String(info.getValue() || "")}>
<Button
size="xs"
variant="light"
className="font-mono text-blue-500 bg-blue-50 hover:bg-blue-100 text-xs font-normal px-2 py-0.5 text-left overflow-hidden truncate max-w-[200px]"
onClick={() => info.getValue() && onViewClick(info.getValue())}
>
{info.getValue() ? `${String(info.getValue()).slice(0, 7)}...` : ""}
</Button>
</Tooltip>
),
},
const columns: ColumnDef<PolicyRow>[] = [
{
header: "Name",
accessorKey: "policy_name",
cell: ({ row }) => {
const policy = row.original;
const { primaryPolicy, versionCount } = row.original;
return (
<Tooltip title={policy.policy_name}>
<span className="text-xs font-medium">{policy.policy_name || "-"}</span>
</Tooltip>
<div className="flex items-center gap-2">
<Tooltip title={`${primaryPolicy.policy_name || "-"}${versionCount > 1 ? ` (${versionCount} versions)` : ""}`}>
<Button
size="xs"
variant="light"
className="font-medium text-blue-500 bg-blue-50 hover:bg-blue-100 text-xs font-normal px-2 py-0.5 text-left"
onClick={() => primaryPolicy.policy_id && onViewClick(primaryPolicy.policy_id)}
>
{primaryPolicy.policy_name || "-"}
</Button>
</Tooltip>
{versionCount > 1 && (
<Badge color="gray" size="xs">
{versionCount} version{versionCount !== 1 ? "s" : ""}
</Badge>
)}
</div>
);
},
},
{
header: "Description",
accessorKey: "description",
accessorFn: (row) => row.primaryPolicy.description ?? "",
cell: ({ row }) => {
const policy = row.original;
const policy = row.original.primaryPolicy;
return (
<Tooltip title={policy.description}>
<span className="text-xs truncate max-w-[200px] block">
@ -83,9 +108,9 @@ const PolicyTable: React.FC<PolicyTableProps> = ({
},
{
header: "Inherits From",
accessorKey: "inherit",
accessorFn: (row) => row.primaryPolicy.inherit ?? "",
cell: ({ row }) => {
const policy = row.original;
const policy = row.original.primaryPolicy;
return policy.inherit ? (
<Badge color="blue" size="xs">
{policy.inherit}
@ -97,9 +122,9 @@ const PolicyTable: React.FC<PolicyTableProps> = ({
},
{
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 <span className="text-xs text-gray-400">-</span>;
@ -122,9 +147,9 @@ const PolicyTable: React.FC<PolicyTableProps> = ({
},
{
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 <span className="text-xs text-gray-400">-</span>;
@ -147,9 +172,12 @@ const PolicyTable: React.FC<PolicyTableProps> = ({
},
{
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 <span className="text-xs text-gray-400">-</span>;
@ -169,9 +197,10 @@ const PolicyTable: React.FC<PolicyTableProps> = ({
},
{
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 (
<Tooltip title={policy.created_at}>
<span className="text-xs">{formatDate(policy.created_at)}</span>
@ -183,7 +212,8 @@ const PolicyTable: React.FC<PolicyTableProps> = ({
id: "actions",
header: "Actions",
cell: ({ row }) => {
const policy = row.original;
const { primaryPolicy } = row.original;
const policy = primaryPolicy;
return (
<div className="flex space-x-2">
{isAdmin && (
@ -216,7 +246,7 @@ const PolicyTable: React.FC<PolicyTableProps> = ({
];
const table = useReactTable({
data: policies,
data: rows,
columns,
state: {
sorting,
@ -273,9 +303,9 @@ const PolicyTable: React.FC<PolicyTableProps> = ({
</div>
</TableCell>
</TableRow>
) : policies.length > 0 ? (
) : rows.length > 0 ? (
table.getRowModel().rows.map((row) => (
<TableRow key={row.id} className="h-8">
<TableRow key={row.original.policy_name} className="h-8">
{row.getVisibleCells().map((cell) => (
<TableCell
key={cell.id}

View file

@ -7,6 +7,9 @@ export interface Policy {
guardrails_remove: string[];
condition: PolicyCondition | null;
pipeline?: GuardrailPipeline | null;
version_number?: number;
version_status?: "draft" | "published" | "production";
parent_version_id?: string | null;
created_at?: string;
updated_at?: string;
created_by?: string;
@ -78,6 +81,12 @@ export interface PolicyListResponse {
total_count: number;
}
export interface PolicyVersionListResponse {
policy_name: string;
versions: Policy[];
total_count: number;
}
export interface PolicyAttachmentListResponse {
attachments: PolicyAttachment[];
total_count: number;

View file

@ -538,6 +538,13 @@ const frameworkMeta: Record<string, { icon: string; description: string }> = {
},
};
/** 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,

View file

@ -14,7 +14,7 @@
"moduleResolution": "bundler",
"resolveJsonModule": true,
"isolatedModules": true,
"jsx": "react-jsx",
"jsx": "preserve",
"incremental": true,
"plugins": [
{