diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260221183800_add_policy_versioning/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260221183800_add_policy_versioning/migration.sql new file mode 100644 index 00000000000..087c5ecc01a --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260221183800_add_policy_versioning/migration.sql @@ -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"); + diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 777e9c6b971..5d2cad6da5b 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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 diff --git a/litellm/proxy/_experimental/out/404.html b/litellm/proxy/_experimental/out/404/index.html similarity index 100% rename from litellm/proxy/_experimental/out/404.html rename to litellm/proxy/_experimental/out/404/index.html diff --git a/litellm/proxy/_experimental/out/_not-found.html b/litellm/proxy/_experimental/out/_not-found/index.html similarity index 100% rename from litellm/proxy/_experimental/out/_not-found.html rename to litellm/proxy/_experimental/out/_not-found/index.html diff --git a/litellm/proxy/_experimental/out/api-reference.html b/litellm/proxy/_experimental/out/api-reference/index.html similarity index 100% rename from litellm/proxy/_experimental/out/api-reference.html rename to litellm/proxy/_experimental/out/api-reference/index.html diff --git a/litellm/proxy/_experimental/out/experimental/api-playground.html b/litellm/proxy/_experimental/out/experimental/api-playground/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/api-playground.html rename to litellm/proxy/_experimental/out/experimental/api-playground/index.html diff --git a/litellm/proxy/_experimental/out/experimental/budgets.html b/litellm/proxy/_experimental/out/experimental/budgets/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/budgets.html rename to litellm/proxy/_experimental/out/experimental/budgets/index.html diff --git a/litellm/proxy/_experimental/out/experimental/caching.html b/litellm/proxy/_experimental/out/experimental/caching/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/caching.html rename to litellm/proxy/_experimental/out/experimental/caching/index.html diff --git a/litellm/proxy/_experimental/out/experimental/claude-code-plugins.html b/litellm/proxy/_experimental/out/experimental/claude-code-plugins/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/claude-code-plugins.html rename to litellm/proxy/_experimental/out/experimental/claude-code-plugins/index.html diff --git a/litellm/proxy/_experimental/out/experimental/old-usage.html b/litellm/proxy/_experimental/out/experimental/old-usage/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/old-usage.html rename to litellm/proxy/_experimental/out/experimental/old-usage/index.html diff --git a/litellm/proxy/_experimental/out/experimental/prompts.html b/litellm/proxy/_experimental/out/experimental/prompts/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/prompts.html rename to litellm/proxy/_experimental/out/experimental/prompts/index.html diff --git a/litellm/proxy/_experimental/out/experimental/tag-management.html b/litellm/proxy/_experimental/out/experimental/tag-management/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/tag-management.html rename to litellm/proxy/_experimental/out/experimental/tag-management/index.html diff --git a/litellm/proxy/_experimental/out/guardrails.html b/litellm/proxy/_experimental/out/guardrails/index.html similarity index 100% rename from litellm/proxy/_experimental/out/guardrails.html rename to litellm/proxy/_experimental/out/guardrails/index.html diff --git a/litellm/proxy/_experimental/out/login.html b/litellm/proxy/_experimental/out/login/index.html similarity index 100% rename from litellm/proxy/_experimental/out/login.html rename to litellm/proxy/_experimental/out/login/index.html diff --git a/litellm/proxy/_experimental/out/logs.html b/litellm/proxy/_experimental/out/logs/index.html similarity index 100% rename from litellm/proxy/_experimental/out/logs.html rename to litellm/proxy/_experimental/out/logs/index.html diff --git a/litellm/proxy/_experimental/out/mcp/oauth/callback.html b/litellm/proxy/_experimental/out/mcp/oauth/callback/index.html similarity index 100% rename from litellm/proxy/_experimental/out/mcp/oauth/callback.html rename to litellm/proxy/_experimental/out/mcp/oauth/callback/index.html diff --git a/litellm/proxy/_experimental/out/model-hub.html b/litellm/proxy/_experimental/out/model-hub/index.html similarity index 100% rename from litellm/proxy/_experimental/out/model-hub.html rename to litellm/proxy/_experimental/out/model-hub/index.html diff --git a/litellm/proxy/_experimental/out/model_hub.html b/litellm/proxy/_experimental/out/model_hub/index.html similarity index 100% rename from litellm/proxy/_experimental/out/model_hub.html rename to litellm/proxy/_experimental/out/model_hub/index.html diff --git a/litellm/proxy/_experimental/out/model_hub_table.html b/litellm/proxy/_experimental/out/model_hub_table/index.html similarity index 100% rename from litellm/proxy/_experimental/out/model_hub_table.html rename to litellm/proxy/_experimental/out/model_hub_table/index.html diff --git a/litellm/proxy/_experimental/out/models-and-endpoints.html b/litellm/proxy/_experimental/out/models-and-endpoints/index.html similarity index 100% rename from litellm/proxy/_experimental/out/models-and-endpoints.html rename to litellm/proxy/_experimental/out/models-and-endpoints/index.html diff --git a/litellm/proxy/_experimental/out/onboarding.html b/litellm/proxy/_experimental/out/onboarding/index.html similarity index 100% rename from litellm/proxy/_experimental/out/onboarding.html rename to litellm/proxy/_experimental/out/onboarding/index.html diff --git a/litellm/proxy/_experimental/out/organizations.html b/litellm/proxy/_experimental/out/organizations/index.html similarity index 100% rename from litellm/proxy/_experimental/out/organizations.html rename to litellm/proxy/_experimental/out/organizations/index.html diff --git a/litellm/proxy/_experimental/out/playground.html b/litellm/proxy/_experimental/out/playground/index.html similarity index 100% rename from litellm/proxy/_experimental/out/playground.html rename to litellm/proxy/_experimental/out/playground/index.html diff --git a/litellm/proxy/_experimental/out/policies.html b/litellm/proxy/_experimental/out/policies/index.html similarity index 100% rename from litellm/proxy/_experimental/out/policies.html rename to litellm/proxy/_experimental/out/policies/index.html diff --git a/litellm/proxy/_experimental/out/settings/admin-settings.html b/litellm/proxy/_experimental/out/settings/admin-settings/index.html similarity index 100% rename from litellm/proxy/_experimental/out/settings/admin-settings.html rename to litellm/proxy/_experimental/out/settings/admin-settings/index.html diff --git a/litellm/proxy/_experimental/out/settings/logging-and-alerts.html b/litellm/proxy/_experimental/out/settings/logging-and-alerts/index.html similarity index 100% rename from litellm/proxy/_experimental/out/settings/logging-and-alerts.html rename to litellm/proxy/_experimental/out/settings/logging-and-alerts/index.html diff --git a/litellm/proxy/_experimental/out/settings/router-settings.html b/litellm/proxy/_experimental/out/settings/router-settings/index.html similarity index 100% rename from litellm/proxy/_experimental/out/settings/router-settings.html rename to litellm/proxy/_experimental/out/settings/router-settings/index.html diff --git a/litellm/proxy/_experimental/out/settings/ui-theme.html b/litellm/proxy/_experimental/out/settings/ui-theme/index.html similarity index 100% rename from litellm/proxy/_experimental/out/settings/ui-theme.html rename to litellm/proxy/_experimental/out/settings/ui-theme/index.html diff --git a/litellm/proxy/_experimental/out/teams.html b/litellm/proxy/_experimental/out/teams/index.html similarity index 100% rename from litellm/proxy/_experimental/out/teams.html rename to litellm/proxy/_experimental/out/teams/index.html diff --git a/litellm/proxy/_experimental/out/test-key.html b/litellm/proxy/_experimental/out/test-key/index.html similarity index 100% rename from litellm/proxy/_experimental/out/test-key.html rename to litellm/proxy/_experimental/out/test-key/index.html diff --git a/litellm/proxy/_experimental/out/tools/mcp-servers.html b/litellm/proxy/_experimental/out/tools/mcp-servers/index.html similarity index 100% rename from litellm/proxy/_experimental/out/tools/mcp-servers.html rename to litellm/proxy/_experimental/out/tools/mcp-servers/index.html diff --git a/litellm/proxy/_experimental/out/tools/vector-stores.html b/litellm/proxy/_experimental/out/tools/vector-stores/index.html similarity index 100% rename from litellm/proxy/_experimental/out/tools/vector-stores.html rename to litellm/proxy/_experimental/out/tools/vector-stores/index.html diff --git a/litellm/proxy/_experimental/out/usage.html b/litellm/proxy/_experimental/out/usage/index.html similarity index 100% rename from litellm/proxy/_experimental/out/usage.html rename to litellm/proxy/_experimental/out/usage/index.html diff --git a/litellm/proxy/_experimental/out/users.html b/litellm/proxy/_experimental/out/users/index.html similarity index 100% rename from litellm/proxy/_experimental/out/users.html rename to litellm/proxy/_experimental/out/users/index.html diff --git a/litellm/proxy/_experimental/out/virtual-keys.html b/litellm/proxy/_experimental/out/virtual-keys/index.html similarity index 100% rename from litellm/proxy/_experimental/out/virtual-keys.html rename to litellm/proxy/_experimental/out/virtual-keys/index.html diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 089731c473d..b61dfa5b263 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -10,15 +10,10 @@ import litellm from litellm._logging import verbose_logger, verbose_proxy_logger from litellm._service_logger import ServiceLogging from litellm.litellm_core_utils.safe_json_loads import safe_json_loads -from litellm.proxy._types import ( - AddTeamCallback, - CommonProxyErrors, - LitellmDataForBackendLLMCall, - LitellmUserRoles, - SpecialHeaders, - TeamCallbackMetadata, - UserAPIKeyAuth, -) +from litellm.proxy._types import (AddTeamCallback, CommonProxyErrors, + LitellmDataForBackendLLMCall, + LitellmUserRoles, SpecialHeaders, + TeamCallbackMetadata, UserAPIKeyAuth) # Cache special headers as a frozenset for O(1) lookup performance _SPECIAL_HEADERS_CACHE = frozenset( @@ -27,12 +22,9 @@ _SPECIAL_HEADERS_CACHE = frozenset( from litellm.router import Router from litellm.types.llms.anthropic import ANTHROPIC_API_HEADERS from litellm.types.services import ServiceTypes -from litellm.types.utils import ( - LlmProviders, - ProviderSpecificHeader, - StandardLoggingUserAPIKeyMetadata, - SupportedCacheControls, -) +from litellm.types.utils import (LlmProviders, ProviderSpecificHeader, + StandardLoggingUserAPIKeyMetadata, + SupportedCacheControls) service_logger_obj = ServiceLogging() # used for tracking latency on OTEL @@ -661,8 +653,7 @@ class LiteLLMProxyRequestSetup: return data from litellm.proxy._types import ( LiteLLM_ManagementEndpoint_MetadataFields, - LiteLLM_ManagementEndpoint_MetadataFields_Premium, - ) + LiteLLM_ManagementEndpoint_MetadataFields_Premium) # ignore any special fields added_metadata = {} @@ -1125,7 +1116,7 @@ async def add_litellm_data_to_request( # noqa: PLR0915 data["litellm_disabled_callbacks"] = disabled_callbacks # Guardrails from key/team metadata and policy engine - move_guardrails_to_metadata( + await move_guardrails_to_metadata( data=data, _metadata_variable_name=_metadata_variable_name, user_api_key_dict=user_api_key_dict, @@ -1458,7 +1449,7 @@ def _add_guardrails_from_policies_in_metadata( ) -def move_guardrails_to_metadata( +async def move_guardrails_to_metadata( data: dict, _metadata_variable_name: str, user_api_key_dict: UserAPIKeyAuth, @@ -1487,7 +1478,8 @@ def move_guardrails_to_metadata( # Only check policy engine if no local config (avoid import + registry lookup) if not (has_key_config or has_team_config or has_request_config): - from litellm.proxy.policy_engine.policy_registry import get_policy_registry + from litellm.proxy.policy_engine.policy_registry import \ + get_policy_registry if not get_policy_registry().is_initialized(): # Nothing configured anywhere - clean up request body fields and return @@ -1515,7 +1507,7 @@ def move_guardrails_to_metadata( ######################################################################################### # Add guardrails from policy engine based on team/key/model context ######################################################################################### - add_guardrails_from_policy_engine( + await add_guardrails_from_policy_engine( data=data, metadata_variable_name=_metadata_variable_name, user_api_key_dict=user_api_key_dict, @@ -1549,10 +1541,29 @@ def move_guardrails_to_metadata( ] = request_body_guardrail_config +def _is_policy_version_id(s: str) -> bool: + """Return True if string is a policy version ID (starts with policy_ prefix).""" + from litellm.proxy.policy_engine.policy_registry import \ + POLICY_VERSION_ID_PREFIX + + return isinstance(s, str) and s.startswith(POLICY_VERSION_ID_PREFIX) + + +def _extract_policy_id(s: str) -> Optional[str]: + """Extract raw UUID from policy_ string, or None if not a valid version ID.""" + from litellm.proxy.policy_engine.policy_registry import \ + POLICY_VERSION_ID_PREFIX + + if not _is_policy_version_id(s): + return None + return s[len(POLICY_VERSION_ID_PREFIX) :].strip() or None + + def _match_and_track_policies( data: dict, context: "PolicyMatchContext", request_body_policies: Any, + policies_override: Optional[Dict[str, Any]] = None, ) -> tuple[list[str], dict[str, str]]: """ Match policies via attachments and request body, track them in metadata. @@ -1562,10 +1573,9 @@ def _match_and_track_policies( """ from litellm._logging import verbose_proxy_logger from litellm.proxy.common_utils.callback_utils import ( - add_policy_sources_to_metadata, - add_policy_to_applied_policies_header, - ) - from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry + add_policy_sources_to_metadata, add_policy_to_applied_policies_header) + from litellm.proxy.policy_engine.attachment_registry import \ + get_attachment_registry from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher # Get matching policies via attachments (with match reasons for attribution) @@ -1595,6 +1605,7 @@ def _match_and_track_policies( applied_policy_names = PolicyMatcher.get_policies_with_matching_conditions( policy_names=list(all_policy_names), context=context, + policies=policies_override, ) verbose_proxy_logger.debug( @@ -1622,20 +1633,30 @@ def _apply_resolved_guardrails_to_metadata( data: dict, metadata_variable_name: str, context: "PolicyMatchContext", + policy_names: Optional[List[str]] = None, + policies: Optional[Dict[str, Any]] = None, ) -> None: """Apply resolved guardrails and pipelines to request metadata.""" from litellm._logging import verbose_proxy_logger from litellm.proxy.policy_engine.policy_resolver import PolicyResolver # Resolve guardrails from matching policies - resolved_guardrails = PolicyResolver.resolve_guardrails_for_context(context=context) + resolved_guardrails = PolicyResolver.resolve_guardrails_for_context( + context=context, + policies=policies, + policy_names=policy_names, + ) verbose_proxy_logger.debug( f"Policy engine: resolved guardrails: {resolved_guardrails}" ) # Resolve pipelines from matching policies - pipelines = PolicyResolver.resolve_pipelines_for_context(context=context) + pipelines = PolicyResolver.resolve_pipelines_for_context( + context=context, + policies=policies, + policy_names=policy_names, + ) # Add resolved guardrails to request metadata if metadata_variable_name not in data: @@ -1675,7 +1696,7 @@ def _apply_resolved_guardrails_to_metadata( ) -def add_guardrails_from_policy_engine( +async def add_guardrails_from_policy_engine( data: dict, metadata_variable_name: str, user_api_key_dict: UserAPIKeyAuth, @@ -1685,12 +1706,13 @@ def add_guardrails_from_policy_engine( This function: 1. Extracts "policies" from request body (if present) for dynamic policy application - 2. Gets matching policies based on team_alias, key_alias, and model (via attachments) - 3. Combines dynamic policies with attachment-based policies - 4. Resolves guardrails from all policies (including inheritance) - 5. Adds guardrails to request metadata - 6. Tracks applied policies in metadata for response headers - 7. Removes "policies" from request body so it's not forwarded to LLM provider + 2. Supports policy_ in policies to execute a specific version (e.g. published) + 3. Gets matching policies based on team_alias, key_alias, and model (via attachments) + 4. Combines dynamic policies with attachment-based policies + 5. Resolves guardrails from all policies (including inheritance) + 6. Adds guardrails to request metadata + 7. Tracks applied policies in metadata for response headers + 8. Removes "policies" from request body so it's not forwarded to LLM provider Args: data: The request data to update @@ -1698,12 +1720,13 @@ def add_guardrails_from_policy_engine( user_api_key_dict: The user's API key authentication info """ from litellm._logging import verbose_proxy_logger - from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body + from litellm.proxy.common_utils.http_parsing_utils import \ + get_tags_from_request_body from litellm.proxy.policy_engine.policy_registry import get_policy_registry from litellm.types.proxy.policy_engine import PolicyMatchContext # Extract dynamic policies from request body (if present) - request_body_policies = data.pop("policies", None) + request_body_policies_raw = data.pop("policies", None) registry = get_policy_registry() verbose_proxy_logger.debug( @@ -1730,13 +1753,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_) + 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( diff --git a/litellm/proxy/policy_engine/policy_endpoints.py b/litellm/proxy/policy_engine/policy_endpoints.py index af12a8598f6..d8de028d6a0 100644 --- a/litellm/proxy/policy_engine/policy_endpoints.py +++ b/litellm/proxy/policy_engine/policy_endpoints.py @@ -4,25 +4,24 @@ CRUD ENDPOINTS FOR POLICIES Provides REST API endpoints for managing policies and policy attachments. """ +from typing import Optional + from fastapi import APIRouter, Depends, HTTPException from litellm._logging import verbose_proxy_logger from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry +from litellm.proxy.policy_engine.attachment_registry import \ + get_attachment_registry from litellm.proxy.policy_engine.pipeline_executor import PipelineExecutor from litellm.proxy.policy_engine.policy_registry import get_policy_registry from litellm.types.proxy.policy_engine import ( - GuardrailPipeline, - PipelineTestRequest, - PolicyAttachmentCreateRequest, - PolicyAttachmentDBResponse, - PolicyAttachmentListResponse, - PolicyCreateRequest, - PolicyDBResponse, - PolicyListDBResponse, - PolicyUpdateRequest, -) + GuardrailPipeline, PipelineTestRequest, PolicyAttachmentCreateRequest, + PolicyAttachmentDBResponse, PolicyAttachmentListResponse, + PolicyCreateRequest, PolicyDBResponse, PolicyListDBResponse, + PolicyUpdateRequest, PolicyVersionCompareResponse, + PolicyVersionCreateRequest, PolicyVersionListResponse, + PolicyVersionStatusUpdateRequest) router = APIRouter() @@ -38,14 +37,20 @@ router = APIRouter() dependencies=[Depends(user_api_key_auth)], response_model=PolicyListDBResponse, ) -async def list_policies(): +async def list_policies(version_status: Optional[str] = None): """ - List all policies from the database. + List all policies from the database. Optionally filter by version_status. + + Query params: + - version_status: Optional. One of "draft", "published", "production". + If omitted, all versions are returned. Example Request: ```bash curl -X GET "http://localhost:4000/policies/list" \\ -H "Authorization: Bearer " + curl -X GET "http://localhost:4000/policies/list?version_status=production" \\ + -H "Authorization: Bearer " ``` Example Response: @@ -55,6 +60,8 @@ async def list_policies(): { "policy_id": "123e4567-e89b-12d3-a456-426614174000", "policy_name": "global-baseline", + "version_number": 1, + "version_status": "production", "inherit": null, "description": "Base guardrails for all requests", "guardrails_add": ["pii_masking"], @@ -74,7 +81,9 @@ async def list_policies(): raise HTTPException(status_code=500, detail="Database not connected") try: - policies = await get_policy_registry().get_all_policies_from_db(prisma_client) + policies = await get_policy_registry().get_all_policies_from_db( + prisma_client, version_status=version_status + ) return PolicyListDBResponse(policies=policies, total_count=len(policies)) except Exception as e: verbose_proxy_logger.exception(f"Error listing policies: {e}") @@ -145,6 +154,172 @@ async def create_policy( raise HTTPException(status_code=500, detail=str(e)) +# ───────────────────────────────────────────────────────────────────────────── +# Policy Versioning Endpoints (must be before /policies/{policy_id} to avoid path conflicts) +# ───────────────────────────────────────────────────────────────────────────── + + +@router.get( + "/policies/name/{policy_name}/versions", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyVersionListResponse, +) +async def list_policy_versions(policy_name: str): + """ + List all versions of a policy by name, ordered by version_number descending. + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + return await get_policy_registry().get_versions_by_policy_name( + policy_name=policy_name, + prisma_client=prisma_client, + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error listing policy versions: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +@router.post( + "/policies/name/{policy_name}/versions", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyDBResponse, +) +async def create_policy_version( + policy_name: str, + request: PolicyVersionCreateRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Create a new draft version of a policy. Copies all fields from the source. + Source is current production if source_policy_id is not provided. + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + created_by = user_api_key_dict.user_id + return await get_policy_registry().create_new_version( + policy_name=policy_name, + prisma_client=prisma_client, + source_policy_id=request.source_policy_id, + created_by=created_by, + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error creating policy version: {e}") + if "not found" in str(e).lower() or "no production" in str(e).lower(): + raise HTTPException(status_code=404, detail=str(e)) + raise HTTPException(status_code=500, detail=str(e)) + + +@router.put( + "/policies/{policy_id}/status", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyDBResponse, +) +async def update_policy_version_status( + policy_id: str, + request: PolicyVersionStatusUpdateRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Update a policy version's status. Valid transitions: + - draft -> published + - published -> production (demotes current production to published) + - production -> published (demotes, policy becomes inactive) + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + updated_by = user_api_key_dict.user_id + return await get_policy_registry().update_version_status( + policy_id=policy_id, + new_status=request.version_status, + prisma_client=prisma_client, + updated_by=updated_by, + ) + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception(f"Error updating version status: {e}") + if "invalid status" in str(e).lower() or "only draft" in str(e).lower() or "cannot promote" in str(e).lower(): + raise HTTPException(status_code=400, detail=str(e)) + if "not found" in str(e).lower(): + raise HTTPException(status_code=404, detail=str(e)) + raise HTTPException(status_code=500, detail=str(e)) + + +@router.get( + "/policies/compare", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyVersionCompareResponse, +) +async def compare_policy_versions( + version_a: str, + version_b: str, +): + """ + Compare two policy versions. Query params: version_a, version_b (policy version IDs). + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + return await get_policy_registry().compare_versions( + policy_id_a=version_a, + policy_id_b=version_b, + prisma_client=prisma_client, + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error comparing versions: {e}") + if "not found" in str(e).lower(): + raise HTTPException(status_code=404, detail=str(e)) + raise HTTPException(status_code=500, detail=str(e)) + + +@router.delete( + "/policies/name/{policy_name}/all-versions", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], +) +async def delete_all_policy_versions(policy_name: str): + """ + Delete all versions of a policy. Also removes from in-memory registry. + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + return await get_policy_registry().delete_all_versions( + policy_name=policy_name, + prisma_client=prisma_client, + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error deleting all versions: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +# ───────────────────────────────────────────────────────────────────────────── +# Policy CRUD by ID +# ───────────────────────────────────────────────────────────────────────────── + + @router.get( "/policies/{policy_id}", tags=["Policies"], @@ -214,7 +389,7 @@ async def update_policy( raise HTTPException(status_code=500, detail="Database not connected") try: - # Check if policy exists + # Check if policy exists and is draft (only drafts can be updated) existing = await get_policy_registry().get_policy_by_id_from_db( policy_id=policy_id, prisma_client=prisma_client, @@ -223,6 +398,11 @@ async def update_policy( raise HTTPException( status_code=404, detail=f"Policy with ID {policy_id} not found" ) + if getattr(existing, "version_status", "production") != "draft": + raise HTTPException( + status_code=400, + detail="Only draft versions can be updated. Publish or create a new version to change published/production.", + ) updated_by = user_api_key_dict.user_id result = await get_policy_registry().update_policy_in_db( @@ -281,6 +461,7 @@ async def delete_policy(policy_id: str): policy_id=policy_id, prisma_client=prisma_client, ) + # Result may include "warning" if production was deleted return result except HTTPException: raise @@ -527,9 +708,11 @@ async def create_policy_attachment( raise HTTPException(status_code=500, detail="Database not connected") try: - # Verify the policy exists - policy = await get_policy_registry().get_all_policies_from_db(prisma_client) - policy_names = [p.policy_name for p in policy] + # Verify the policy has a production version (attachments resolve against production) + policies = await get_policy_registry().get_all_policies_from_db( + prisma_client, version_status="production" + ) + policy_names = {p.policy_name for p in policies} if request.policy_name not in policy_names: raise HTTPException( status_code=404, diff --git a/litellm/proxy/policy_engine/policy_registry.py b/litellm/proxy/policy_engine/policy_registry.py index 1ca6062e5ac..f8e1ebd7ba1 100644 --- a/litellm/proxy/policy_engine/policy_registry.py +++ b/litellm/proxy/policy_engine/policy_registry.py @@ -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_ 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_ 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_ 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 diff --git a/litellm/proxy/policy_engine/policy_resolver.py b/litellm/proxy/policy_engine/policy_resolver.py index a8ad78d6491..c802a970a80 100644 --- a/litellm/proxy/policy_engine/policy_resolver.py +++ b/litellm/proxy/policy_engine/policy_resolver.py @@ -11,12 +11,9 @@ Handles: from typing import Dict, List, Optional, Set, Tuple from litellm._logging import verbose_proxy_logger -from litellm.types.proxy.policy_engine import ( - GuardrailPipeline, - Policy, - PolicyMatchContext, - ResolvedPolicy, -) +from litellm.types.proxy.policy_engine import (GuardrailPipeline, Policy, + PolicyMatchContext, + ResolvedPolicy) class PolicyResolver: @@ -90,7 +87,8 @@ class PolicyResolver: Returns: ResolvedPolicy with final guardrails list """ - from litellm.proxy.policy_engine.condition_evaluator import ConditionEvaluator + from litellm.proxy.policy_engine.condition_evaluator import \ + ConditionEvaluator inheritance_chain = PolicyResolver.resolve_inheritance_chain( policy_name=policy_name, policies=policies @@ -134,12 +132,13 @@ class PolicyResolver: def resolve_guardrails_for_context( context: PolicyMatchContext, policies: Optional[Dict[str, Policy]] = None, + policy_names: Optional[List[str]] = None, ) -> List[str]: """ Resolve the final list of guardrails for a request context. This: - 1. Finds all policies that match the context via policy_attachments + 1. Finds all policies that match the context via policy_attachments (or policy_names if provided) 2. Resolves each policy's guardrails (including inheritance) 3. Evaluates model conditions 4. Combines all guardrails (union) @@ -147,12 +146,14 @@ class PolicyResolver: Args: context: The request context policies: Dictionary of all policies (if None, uses global registry) + policy_names: If provided, use this list instead of attachment matching Returns: List of guardrail names to apply """ from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher - from litellm.proxy.policy_engine.policy_registry import get_policy_registry + from litellm.proxy.policy_engine.policy_registry import \ + get_policy_registry if policies is None: registry = get_policy_registry() @@ -160,8 +161,12 @@ class PolicyResolver: return [] policies = registry.get_all_policies() - # Get matching policies via attachments - matching_policy_names = PolicyMatcher.get_matching_policies(context=context) + # Use provided policy names or get matching policies via attachments + matching_policy_names = ( + policy_names + if policy_names is not None + else PolicyMatcher.get_matching_policies(context=context) + ) if not matching_policy_names: verbose_proxy_logger.debug( @@ -195,6 +200,7 @@ class PolicyResolver: def resolve_pipelines_for_context( context: PolicyMatchContext, policies: Optional[Dict[str, Policy]] = None, + policy_names: Optional[List[str]] = None, ) -> List[Tuple[str, GuardrailPipeline]]: """ Resolve pipelines from matching policies for a request context. @@ -206,12 +212,14 @@ class PolicyResolver: Args: context: The request context policies: Dictionary of all policies (if None, uses global registry) + policy_names: If provided, use this list instead of attachment matching Returns: List of (policy_name, GuardrailPipeline) tuples """ from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher - from litellm.proxy.policy_engine.policy_registry import get_policy_registry + from litellm.proxy.policy_engine.policy_registry import \ + get_policy_registry if policies is None: registry = get_policy_registry() @@ -219,7 +227,11 @@ class PolicyResolver: return [] policies = registry.get_all_policies() - matching_policy_names = PolicyMatcher.get_matching_policies(context=context) + matching_policy_names = ( + policy_names + if policy_names is not None + else PolicyMatcher.get_matching_policies(context=context) + ) if not matching_policy_names: return [] @@ -269,7 +281,8 @@ class PolicyResolver: Returns: Dictionary mapping policy names to ResolvedPolicy objects """ - from litellm.proxy.policy_engine.policy_registry import get_policy_registry + from litellm.proxy.policy_engine.policy_registry import \ + get_policy_registry if policies is None: registry = get_policy_registry() diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 4128ab5f23e..5d2cad6da5b 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -963,20 +963,29 @@ model LiteLLM_SkillsTable { updated_by String? } -// Policy table for storing guardrail policies +// Policy table for storing guardrail policies (versioned) model LiteLLM_PolicyTable { - policy_id String @id @default(uuid()) - policy_name String @unique - inherit String? // Name of parent policy to inherit from - description String? - guardrails_add String[] @default([]) - guardrails_remove String[] @default([]) - condition Json? @default("{}") // Policy conditions (e.g., model matching) - pipeline Json? // Optional guardrail pipeline (mode + steps[]) - created_at DateTime @default(now()) - created_by String? - updated_at DateTime @default(now()) @updatedAt - updated_by String? + policy_id String @id @default(uuid()) + policy_name String // No longer @unique; use @@unique([policy_name, version_number]) + version_number Int @default(1) + version_status String @default("production") // "draft" | "published" | "production" + parent_version_id String? + is_latest Boolean @default(true) + published_at DateTime? + production_at DateTime? + inherit String? // Name of parent policy to inherit from + description String? + guardrails_add String[] @default([]) + guardrails_remove String[] @default([]) + condition Json? @default("{}") // Policy conditions (e.g., model matching) + pipeline Json? // Optional guardrail pipeline (mode + steps[]) + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? + + @@unique([policy_name, version_number]) + @@index([policy_name, version_status]) } // Policy attachment table for defining where policies apply diff --git a/litellm/types/proxy/policy_engine/__init__.py b/litellm/types/proxy/policy_engine/__init__.py index e0c1d6f30da..4df9f21e806 100644 --- a/litellm/types/proxy/policy_engine/__init__.py +++ b/litellm/types/proxy/policy_engine/__init__.py @@ -11,48 +11,28 @@ Configuration: """ from litellm.types.proxy.policy_engine.pipeline_types import ( - GuardrailPipeline, - PipelineExecutionResult, - PipelineStep, - PipelineStepResult, -) -from litellm.types.proxy.policy_engine.policy_types import ( - Policy, - PolicyAttachment, - PolicyCondition, - PolicyConfig, - PolicyGuardrails, - PolicyScope, -) + GuardrailPipeline, PipelineExecutionResult, PipelineStep, + PipelineStepResult) +from litellm.types.proxy.policy_engine.policy_types import (Policy, + PolicyAttachment, + PolicyCondition, + PolicyConfig, + PolicyGuardrails, + PolicyScope) from litellm.types.proxy.policy_engine.resolver_types import ( - AttachmentImpactResponse, - PipelineTestRequest, - PolicyAttachmentCreateRequest, - PolicyAttachmentDBResponse, - PolicyAttachmentListResponse, - PolicyConditionRequest, - PolicyCreateRequest, - PolicyDBResponse, - PolicyGuardrailsResponse, - PolicyInfoResponse, - PolicyListDBResponse, - PolicyListResponse, - PolicyMatchContext, - PolicyMatchDetail, - PolicyResolveRequest, - PolicyResolveResponse, - PolicyScopeResponse, - PolicySummaryItem, - PolicyTestResponse, - PolicyUpdateRequest, - ResolvedPolicy, -) + AttachmentImpactResponse, PipelineTestRequest, + PolicyAttachmentCreateRequest, PolicyAttachmentDBResponse, + PolicyAttachmentListResponse, PolicyConditionRequest, PolicyCreateRequest, + PolicyDBResponse, PolicyGuardrailsResponse, PolicyInfoResponse, + PolicyListDBResponse, PolicyListResponse, PolicyMatchContext, + PolicyMatchDetail, PolicyResolveRequest, PolicyResolveResponse, + PolicyScopeResponse, PolicySummaryItem, PolicyTestResponse, + PolicyUpdateRequest, PolicyVersionCompareResponse, + PolicyVersionCreateRequest, PolicyVersionListResponse, + PolicyVersionStatusUpdateRequest, ResolvedPolicy) from litellm.types.proxy.policy_engine.validation_types import ( - PolicyValidateRequest, - PolicyValidationError, - PolicyValidationErrorType, - PolicyValidationResponse, -) + PolicyValidateRequest, PolicyValidationError, PolicyValidationErrorType, + PolicyValidationResponse) __all__ = [ # Pipeline types @@ -98,4 +78,9 @@ __all__ = [ "PolicyResolveResponse", "PolicyMatchDetail", "AttachmentImpactResponse", + # Policy versioning + "PolicyVersionCreateRequest", + "PolicyVersionStatusUpdateRequest", + "PolicyVersionListResponse", + "PolicyVersionCompareResponse", ] diff --git a/litellm/types/proxy/policy_engine/resolver_types.py b/litellm/types/proxy/policy_engine/resolver_types.py index a5a2334ae4b..2df450dc2ba 100644 --- a/litellm/types/proxy/policy_engine/resolver_types.py +++ b/litellm/types/proxy/policy_engine/resolver_types.py @@ -198,6 +198,23 @@ class PolicyDBResponse(BaseModel): policy_id: str = Field(description="Unique ID of the policy.") policy_name: str = Field(description="Name of the policy.") + version_number: int = Field(default=1, description="Version number of this policy.") + version_status: str = Field( + default="production", + description="One of: draft, published, production.", + ) + parent_version_id: Optional[str] = Field( + default=None, description="Policy ID this version was cloned from." + ) + is_latest: bool = Field( + default=True, description="True if this is the latest version by version_number." + ) + published_at: Optional[datetime] = Field( + default=None, description="When this version was published." + ) + production_at: Optional[datetime] = Field( + default=None, description="When this version was promoted to production." + ) inherit: Optional[str] = Field(default=None, description="Parent policy name.") description: Optional[str] = Field(default=None, description="Policy description.") guardrails_add: List[str] = Field( @@ -233,6 +250,49 @@ class PolicyListDBResponse(BaseModel): total_count: int = Field(default=0, description="Total number of policies.") +# ───────────────────────────────────────────────────────────────────────────── +# Policy Versioning Types +# ───────────────────────────────────────────────────────────────────────────── + + +class PolicyVersionCreateRequest(BaseModel): + """Request body for creating a new policy version (draft).""" + + source_policy_id: Optional[str] = Field( + default=None, + description="Policy ID to clone from. If None, clone from current production version.", + ) + + +class PolicyVersionStatusUpdateRequest(BaseModel): + """Request body for updating a policy version's status.""" + + version_status: str = Field( + description="New status: 'published' or 'production'.", + ) + + +class PolicyVersionListResponse(BaseModel): + """Response for listing all versions of a policy.""" + + policy_name: str = Field(description="Name of the policy.") + versions: List[PolicyDBResponse] = Field( + default_factory=list, description="All versions ordered by version_number desc." + ) + total_count: int = Field(default=0, description="Total number of versions.") + + +class PolicyVersionCompareResponse(BaseModel): + """Response for comparing two policy versions.""" + + version_a: PolicyDBResponse = Field(description="First version.") + version_b: PolicyDBResponse = Field(description="Second version.") + field_diffs: Dict[str, Dict[str, Any]] = Field( + default_factory=dict, + description="Field name -> {version_a: val, version_b: val} for differing fields.", + ) + + # ───────────────────────────────────────────────────────────────────────────── # Policy Attachment CRUD Types # ───────────────────────────────────────────────────────────────────────────── diff --git a/schema.prisma b/schema.prisma index 4128ab5f23e..5d2cad6da5b 100644 --- a/schema.prisma +++ b/schema.prisma @@ -963,20 +963,29 @@ model LiteLLM_SkillsTable { updated_by String? } -// Policy table for storing guardrail policies +// Policy table for storing guardrail policies (versioned) model LiteLLM_PolicyTable { - policy_id String @id @default(uuid()) - policy_name String @unique - inherit String? // Name of parent policy to inherit from - description String? - guardrails_add String[] @default([]) - guardrails_remove String[] @default([]) - condition Json? @default("{}") // Policy conditions (e.g., model matching) - pipeline Json? // Optional guardrail pipeline (mode + steps[]) - created_at DateTime @default(now()) - created_by String? - updated_at DateTime @default(now()) @updatedAt - updated_by String? + policy_id String @id @default(uuid()) + policy_name String // No longer @unique; use @@unique([policy_name, version_number]) + version_number Int @default(1) + version_status String @default("production") // "draft" | "published" | "production" + parent_version_id String? + is_latest Boolean @default(true) + published_at DateTime? + production_at DateTime? + inherit String? // Name of parent policy to inherit from + description String? + guardrails_add String[] @default([]) + guardrails_remove String[] @default([]) + condition Json? @default("{}") // Policy conditions (e.g., model matching) + pipeline Json? // Optional guardrail pipeline (mode + steps[]) + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? + + @@unique([policy_name, version_number]) + @@index([policy_name, version_status]) } // Policy attachment table for defining where policies apply diff --git a/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py b/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py new file mode 100644 index 00000000000..738c611d928 --- /dev/null +++ b/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py @@ -0,0 +1,442 @@ +""" +Unit tests for policy versioning: registry behavior, status transitions, and version CRUD. +""" + +from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.policy_engine.policy_registry import ( + PolicyRegistry, + _row_to_policy_db_response, + get_policy_registry, +) +from litellm.types.proxy.policy_engine import ( + PolicyCreateRequest, + PolicyDBResponse, + PolicyUpdateRequest, +) + + +def _make_row( + policy_id="pid-1", + policy_name="test-policy", + version_number=1, + version_status="production", + parent_version_id=None, + is_latest=True, + published_at=None, + production_at=None, + inherit=None, + description="desc", + guardrails_add=None, + guardrails_remove=None, + condition=None, + pipeline=None, + created_at=None, + updated_at=None, + created_by=None, + updated_by=None, +): + row = MagicMock() + row.policy_id = policy_id + row.policy_name = policy_name + row.version_number = version_number + row.version_status = version_status + row.parent_version_id = parent_version_id + row.is_latest = is_latest + row.published_at = published_at + row.production_at = production_at + row.inherit = inherit + row.description = description + row.guardrails_add = guardrails_add or [] + row.guardrails_remove = guardrails_remove or [] + row.condition = condition + row.pipeline = pipeline + row.created_at = created_at or datetime.now(timezone.utc) + row.updated_at = updated_at or datetime.now(timezone.utc) + row.created_by = created_by + row.updated_by = updated_by + return row + + +class TestRowToPolicyDBResponse: + """Test _row_to_policy_db_response includes all version fields.""" + + def test_includes_version_fields(self): + row = _make_row( + version_number=2, + version_status="draft", + parent_version_id="pid-0", + is_latest=True, + published_at=None, + production_at=None, + ) + resp = _row_to_policy_db_response(row) + assert isinstance(resp, PolicyDBResponse) + assert resp.policy_id == "pid-1" + assert resp.policy_name == "test-policy" + assert resp.version_number == 2 + assert resp.version_status == "draft" + assert resp.parent_version_id == "pid-0" + assert resp.is_latest is True + assert resp.published_at is None + assert resp.production_at is None + + def test_backward_compat_missing_version_attrs(self): + row = _make_row() + del row.version_number + del row.version_status + del row.parent_version_id + del row.is_latest + del row.published_at + del row.production_at + resp = _row_to_policy_db_response(row) + assert resp.version_number == 1 + assert resp.version_status == "production" + assert resp.parent_version_id is None + assert resp.is_latest is True + + +class TestSyncPoliciesFromDbProductionOnly: + """Test that sync_policies_from_db only loads production versions.""" + + @pytest.mark.asyncio + async def test_get_all_policies_with_version_status_calls_find_many_with_where(self): + registry = PolicyRegistry() + prisma = MagicMock() + prod_row = _make_row(policy_id="prod-1", version_status="production") + prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[prod_row]) + + result = await registry.get_all_policies_from_db( + prisma, version_status="production" + ) + + assert len(result) == 1 + assert result[0].version_status == "production" + prisma.db.litellm_policytable.find_many.assert_called_once() + call_kw = prisma.db.litellm_policytable.find_many.call_args[1] + assert call_kw.get("where") == {"version_status": "production"} + + @pytest.mark.asyncio + async def test_sync_policies_from_db_only_loads_production(self): + registry = PolicyRegistry() + prisma = MagicMock() + prod_row = _make_row( + policy_id="prod-1", + policy_name="foo", + version_status="production", + guardrails_add=["g1"], + ) + prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[prod_row]) + + await registry.sync_policies_from_db(prisma) + + assert registry.has_policy("foo") + policy = registry.get_policy("foo") + assert policy is not None + assert policy.guardrails.add == ["g1"] + # find_many was called with version_status=production (via get_all_policies_from_db) + find_many_calls = prisma.db.litellm_policytable.find_many.call_args_list + assert len(find_many_calls) >= 1 + assert find_many_calls[0][1].get("where") == {"version_status": "production"} + + +class TestUpdatePolicyDraftOnly: + """Test that update_policy_in_db only allows draft versions.""" + + @pytest.mark.asyncio + async def test_update_production_raises(self): + registry = PolicyRegistry() + prisma = MagicMock() + prod_row = _make_row(policy_id="pid-1", version_status="production") + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=prod_row) + + with pytest.raises(Exception) as exc_info: + await registry.update_policy_in_db( + policy_id="pid-1", + policy_request=PolicyUpdateRequest(description="new"), + prisma_client=prisma, + ) + assert "Only draft" in str(exc_info.value) or "draft" in str(exc_info.value).lower() + prisma.db.litellm_policytable.update.assert_not_called() + + @pytest.mark.asyncio + async def test_update_draft_succeeds_and_does_not_update_registry(self): + registry = PolicyRegistry() + registry.add_policy("test-policy", MagicMock()) # in-memory state + prisma = MagicMock() + draft_row = _make_row( + policy_id="draft-1", + policy_name="test-policy", + version_status="draft", + description="old", + ) + updated_row = _make_row( + policy_id="draft-1", + policy_name="test-policy", + version_status="draft", + description="new", + ) + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=draft_row) + prisma.db.litellm_policytable.update = AsyncMock(return_value=updated_row) + + result = await registry.update_policy_in_db( + policy_id="draft-1", + policy_request=PolicyUpdateRequest(description="new"), + prisma_client=prisma, + ) + + assert result.description == "new" + prisma.db.litellm_policytable.update.assert_called_once() + # Registry still has old in-memory policy (drafts are not in registry; we don't add) + assert registry.has_policy("test-policy") + + +class TestDeletePolicyFromDb: + """Test delete_policy_from_db removes production from registry and returns warning.""" + + @pytest.mark.asyncio + async def test_delete_production_removes_from_registry_and_returns_warning(self): + registry = PolicyRegistry() + registry.add_policy("deleted-policy", MagicMock()) + prisma = MagicMock() + prod_row = _make_row( + policy_id="prod-1", + policy_name="deleted-policy", + version_status="production", + ) + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=prod_row) + prisma.db.litellm_policytable.delete = AsyncMock() + + result = await registry.delete_policy_from_db( + policy_id="prod-1", + prisma_client=prisma, + ) + + assert result["message"] + assert "warning" in result + assert "Production" in result["warning"] or "production" in result["warning"] + assert not registry.has_policy("deleted-policy") + + @pytest.mark.asyncio + async def test_delete_draft_does_not_remove_from_registry_no_warning(self): + registry = PolicyRegistry() + registry.add_policy("my-policy", MagicMock()) + prisma = MagicMock() + draft_row = _make_row( + policy_id="draft-1", + policy_name="my-policy", + version_status="draft", + ) + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=draft_row) + prisma.db.litellm_policytable.delete = AsyncMock() + + result = await registry.delete_policy_from_db( + policy_id="draft-1", + prisma_client=prisma, + ) + + assert "warning" not in result + assert registry.has_policy("my-policy") + + +class TestCreateNewVersion: + """Test create_new_version copies all fields and sets draft.""" + + @pytest.mark.asyncio + async def test_create_new_version_from_production_increments_version(self): + registry = PolicyRegistry() + prisma = MagicMock() + prod = _make_row( + policy_id="prod-1", + policy_name="foo", + version_number=1, + version_status="production", + guardrails_add=["g1"], + description="base", + inherit=None, + pipeline={"mode": "pre_call", "steps": []}, + ) + # find_first for production + prisma.db.litellm_policytable.find_first = AsyncMock(return_value=prod) + # find_first for latest version number + prisma.db.litellm_policytable.find_first.side_effect = [ + prod, # production lookup + prod, # latest version_number lookup + ] + # update_many for is_latest=False + prisma.db.litellm_policytable.update_many = AsyncMock() + new_row = _make_row( + policy_id="new-id", + policy_name="foo", + version_number=2, + version_status="draft", + parent_version_id="prod-1", + is_latest=True, + guardrails_add=["g1"], + description="base", + pipeline={"mode": "pre_call", "steps": []}, + ) + prisma.db.litellm_policytable.create = AsyncMock(return_value=new_row) + + result = await registry.create_new_version( + policy_name="foo", + prisma_client=prisma, + source_policy_id=None, + created_by="user", + ) + + assert result.version_number == 2 + assert result.version_status == "draft" + assert result.parent_version_id == "prod-1" + assert result.guardrails_add == ["g1"] + assert result.description == "base" + create_call = prisma.db.litellm_policytable.create.call_args[1]["data"] + assert create_call["version_number"] == 2 + assert create_call["version_status"] == "draft" + assert create_call["parent_version_id"] == "prod-1" + assert create_call["guardrails_add"] == ["g1"] + + +class TestUpdateVersionStatus: + """Test status transitions: valid succeed, invalid return error.""" + + @pytest.mark.asyncio + async def test_draft_to_published_sets_published_at(self): + registry = PolicyRegistry() + prisma = MagicMock() + draft = _make_row(policy_id="d-1", version_status="draft") + updated = _make_row( + policy_id="d-1", + version_status="published", + published_at=datetime.now(timezone.utc), + ) + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=draft) + prisma.db.litellm_policytable.update = AsyncMock(return_value=updated) + + result = await registry.update_version_status( + policy_id="d-1", + new_status="published", + prisma_client=prisma, + ) + + assert result.version_status == "published" + update_data = prisma.db.litellm_policytable.update.call_args[1]["data"] + assert update_data["version_status"] == "published" + assert "published_at" in update_data + + @pytest.mark.asyncio + async def test_draft_to_production_raises(self): + registry = PolicyRegistry() + prisma = MagicMock() + draft = _make_row(policy_id="d-1", version_status="draft") + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=draft) + + with pytest.raises(Exception) as exc_info: + await registry.update_version_status( + policy_id="d-1", + new_status="production", + prisma_client=prisma, + ) + assert "publish" in str(exc_info.value).lower() or "draft" in str(exc_info.value).lower() + + @pytest.mark.asyncio + async def test_published_to_production_demotes_old_and_updates_registry(self): + registry = PolicyRegistry() + prisma = MagicMock() + published_row = _make_row( + policy_id="pub-1", + policy_name="foo", + version_status="published", + ) + updated_row = _make_row( + policy_id="pub-1", + policy_name="foo", + version_status="production", + production_at=datetime.now(timezone.utc), + ) + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=published_row) + prisma.db.litellm_policytable.update_many = AsyncMock() + prisma.db.litellm_policytable.update = AsyncMock(return_value=updated_row) + + result = await registry.update_version_status( + policy_id="pub-1", + new_status="production", + prisma_client=prisma, + ) + + assert result.version_status == "production" + # update_many should have been called to demote current production + assert prisma.db.litellm_policytable.update_many.called + # Registry should have been updated with new production + assert registry.has_policy("foo") + + +class TestCompareVersions: + """Test compare_versions returns correct field diffs.""" + + @pytest.mark.asyncio + async def test_compare_versions_returns_diffs(self): + registry = PolicyRegistry() + prisma = MagicMock() + a = _make_row( + policy_id="a", + policy_name="p", + description="desc A", + guardrails_add=["g1"], + ) + b = _make_row( + policy_id="b", + policy_name="p", + description="desc B", + guardrails_add=["g1", "g2"], + ) + prisma.db.litellm_policytable.find_unique = AsyncMock(side_effect=[a, b]) + + result = await registry.compare_versions( + policy_id_a="a", + policy_id_b="b", + prisma_client=prisma, + ) + + assert result.version_a.policy_id == "a" + assert result.version_b.policy_id == "b" + assert "description" in result.field_diffs + assert result.field_diffs["description"]["version_a"] == "desc A" + assert result.field_diffs["description"]["version_b"] == "desc B" + assert "guardrails_add" in result.field_diffs + + +class TestResolveGuardrailsProductionOnly: + """Test that resolve_guardrails_from_db uses only production versions.""" + + @pytest.mark.asyncio + async def test_resolve_guardrails_calls_get_all_with_production_filter(self): + registry = PolicyRegistry() + prisma = MagicMock() + prod_row = _make_row( + policy_name="base", + version_status="production", + guardrails_add=["g1"], + ) + prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[prod_row]) + + result = await registry.resolve_guardrails_from_db( + policy_name="base", + prisma_client=prisma, + ) + + assert "g1" in result + call_kw = prisma.db.litellm_policytable.find_many.call_args[1] + assert call_kw.get("where") == {"version_status": "production"} + + +class TestGetPolicyRegistrySingleton: + """Test get_policy_registry returns same instance.""" + + def test_returns_singleton(self): + a = get_policy_registry() + b = get_policy_registry() + assert a is b diff --git a/tests/test_litellm/proxy/policy_engine/test_policy_versioning_e2e.py b/tests/test_litellm/proxy/policy_engine/test_policy_versioning_e2e.py new file mode 100644 index 00000000000..5d6f3a05ae6 --- /dev/null +++ b/tests/test_litellm/proxy/policy_engine/test_policy_versioning_e2e.py @@ -0,0 +1,217 @@ +""" +Integration-style tests for policy versioning: full lifecycle with mocked DB. +""" + +from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.policy_engine.policy_registry import PolicyRegistry +from litellm.types.proxy.policy_engine import (PolicyCreateRequest, + PolicyUpdateRequest) + + +def _make_row( + policy_id, + policy_name, + version_number=1, + version_status="production", + parent_version_id=None, + is_latest=True, + published_at=None, + production_at=None, + inherit=None, + description="", + guardrails_add=None, + guardrails_remove=None, + condition=None, + pipeline=None, +): + row = MagicMock() + row.policy_id = policy_id + row.policy_name = policy_name + row.version_number = version_number + row.version_status = version_status + row.parent_version_id = parent_version_id + row.is_latest = is_latest + row.published_at = published_at + row.production_at = production_at + row.inherit = inherit + row.description = description + row.guardrails_add = guardrails_add or [] + row.guardrails_remove = guardrails_remove or [] + row.condition = condition + row.pipeline = pipeline + row.created_at = datetime.now(timezone.utc) + row.updated_at = datetime.now(timezone.utc) + row.created_by = None + row.updated_by = None + return row + + +@pytest.mark.asyncio +async def test_full_lifecycle_create_draft_edit_publish_promote(): + """ + Full lifecycle: create policy -> create draft version -> edit draft -> + publish -> promote to production -> verify old version demoted -> + verify in-memory updated. + """ + registry = PolicyRegistry() + prisma = MagicMock() + now = datetime.now(timezone.utc) + + # 1) Create initial policy (v1 production) + create_data = {} + created_v1 = _make_row( + policy_id="v1-id", + policy_name="lifecycle-policy", + version_number=1, + version_status="production", + production_at=now, + guardrails_add=["g1"], + description="Initial", + ) + + async def create_impl(data=None, **kwargs): + create_data.update(kwargs.get("data", data or {})) + return created_v1 + + prisma.db.litellm_policytable.create = AsyncMock(side_effect=create_impl) + req = PolicyCreateRequest( + policy_name="lifecycle-policy", + description="Initial", + guardrails_add=["g1"], + ) + created = await registry.add_policy_to_db(req, prisma, created_by="user") + assert created.version_number == 1 + assert created.version_status == "production" + assert registry.has_policy("lifecycle-policy") + + # 2) Create new draft version (v2) + v2_row = _make_row( + policy_id="v2-id", + policy_name="lifecycle-policy", + version_number=2, + version_status="draft", + parent_version_id="v1-id", + is_latest=True, + guardrails_add=["g1", "g2"], + description="Draft v2", + ) + prisma.db.litellm_policytable.find_first = AsyncMock(return_value=created_v1) + prisma.db.litellm_policytable.update_many = AsyncMock() + prisma.db.litellm_policytable.create = AsyncMock(return_value=v2_row) + + draft_v2 = await registry.create_new_version( + policy_name="lifecycle-policy", + prisma_client=prisma, + source_policy_id=None, + created_by="user", + ) + assert draft_v2.version_number == 2 + assert draft_v2.version_status == "draft" + assert draft_v2.parent_version_id == "v1-id" + # In-memory still has v1 (only production is in registry) + assert registry.has_policy("lifecycle-policy") + policy = registry.get_policy("lifecycle-policy") + assert policy.guardrails.add == ["g1"] # still v1 + + # 3) Edit draft v2 + v2_updated_row = _make_row( + policy_id="v2-id", + policy_name="lifecycle-policy", + version_number=2, + version_status="draft", + guardrails_add=["g1", "g2", "g3"], + description="Draft v2 edited", + ) + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=v2_row) + prisma.db.litellm_policytable.update = AsyncMock(return_value=v2_updated_row) + + updated_draft = await registry.update_policy_in_db( + policy_id="v2-id", + policy_request=PolicyUpdateRequest( + description="Draft v2 edited", + guardrails_add=["g1", "g2", "g3"], + ), + prisma_client=prisma, + updated_by="user", + ) + assert updated_draft.description == "Draft v2 edited" + assert updated_draft.guardrails_add == ["g1", "g2", "g3"] + + # 4) Publish v2 (draft -> published) + v2_published = _make_row( + policy_id="v2-id", + policy_name="lifecycle-policy", + version_number=2, + version_status="published", + published_at=now, + guardrails_add=["g1", "g2", "g3"], + description="Draft v2 edited", + ) + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=v2_updated_row) + prisma.db.litellm_policytable.update = AsyncMock(return_value=v2_published) + + published = await registry.update_version_status( + policy_id="v2-id", + new_status="published", + prisma_client=prisma, + updated_by="user", + ) + assert published.version_status == "published" + + # 5) Promote v2 to production (demote v1 to published, update registry) + prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=v2_published) + prisma.db.litellm_policytable.update_many = AsyncMock() + v2_production = _make_row( + policy_id="v2-id", + policy_name="lifecycle-policy", + version_number=2, + version_status="production", + production_at=now, + guardrails_add=["g1", "g2", "g3"], + description="Draft v2 edited", + ) + prisma.db.litellm_policytable.update = AsyncMock(return_value=v2_production) + + prod = await registry.update_version_status( + policy_id="v2-id", + new_status="production", + prisma_client=prisma, + updated_by="user", + ) + assert prod.version_status == "production" + # In-memory registry should now have v2 content + assert registry.has_policy("lifecycle-policy") + policy = registry.get_policy("lifecycle-policy") + assert policy.guardrails.add == ["g1", "g2", "g3"] + + +@pytest.mark.asyncio +async def test_attachments_resolve_against_production_after_promotion(): + """ + After promoting a new version to production, resolve_guardrails_from_db + returns guardrails from the new production version (inheritance resolves + against production). + """ + registry = PolicyRegistry() + prisma = MagicMock() + # Simulate only production versions loaded for resolution + prod_row = _make_row( + policy_id="prod-1", + policy_name="att-policy", + version_status="production", + guardrails_add=["ga", "gb"], + ) + prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[prod_row]) + + resolved = await registry.resolve_guardrails_from_db( + policy_name="att-policy", + prisma_client=prisma, + ) + assert "ga" in resolved + assert "gb" in resolved + call_kw = prisma.db.litellm_policytable.find_many.call_args[1] + assert call_kw.get("where") == {"version_status": "production"} diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index e4b8613d204..ce79caeaf57 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -3,7 +3,7 @@ import copy import json import os import sys -from unittest.mock import MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import Request @@ -11,16 +11,11 @@ from fastapi import Request import litellm from litellm.proxy._types import TeamCallbackMetadata, UserAPIKeyAuth from litellm.proxy.litellm_pre_call_utils import ( - KeyAndTeamLoggingSettings, - LiteLLMProxyRequestSetup, - _get_dynamic_logging_metadata, - _get_enforced_params, - _get_metadata_variable_name, - _update_model_if_key_alias_exists, - add_guardrails_from_policy_engine, - add_litellm_data_to_request, - check_if_token_is_service_account, -) + KeyAndTeamLoggingSettings, LiteLLMProxyRequestSetup, + _get_dynamic_logging_metadata, _get_enforced_params, + _get_metadata_variable_name, _update_model_if_key_alias_exists, + add_guardrails_from_policy_engine, add_litellm_data_to_request, + check_if_token_is_service_account) sys.path.insert( 0, os.path.abspath("../../..") @@ -159,7 +154,8 @@ def test_get_enforced_params( @pytest.mark.asyncio async def test_add_litellm_data_to_request_parses_string_metadata(): - from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.litellm_pre_call_utils import \ + add_litellm_data_to_request # Setup request_mock = MagicMock(spec=Request) @@ -205,7 +201,8 @@ async def test_add_litellm_data_to_request_parses_string_metadata(): @pytest.mark.asyncio async def test_add_litellm_data_to_request_user_spend_and_budget(): - from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.litellm_pre_call_utils import \ + add_litellm_data_to_request request_mock = MagicMock(spec=Request) request_mock.url.path = "/v1/completions" @@ -243,7 +240,8 @@ async def test_add_litellm_data_to_request_user_spend_and_budget(): @pytest.mark.asyncio async def test_add_litellm_data_to_request_audio_transcription_multipart(): - from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.litellm_pre_call_utils import \ + add_litellm_data_to_request # Setup request mock for /v1/audio/transcriptions request_mock = MagicMock(spec=Request) @@ -308,7 +306,8 @@ async def test_add_litellm_data_to_request_disabled_callbacks(): """ Test that litellm_disabled_callbacks from key metadata is properly added to the request data. """ - from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.litellm_pre_call_utils import \ + add_litellm_data_to_request # Setup mock request request_mock = MagicMock(spec=Request) @@ -361,7 +360,8 @@ async def test_add_litellm_data_to_request_disabled_callbacks_empty(): """ Test that litellm_disabled_callbacks is not added when it's empty. """ - from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.litellm_pre_call_utils import \ + add_litellm_data_to_request # Setup mock request request_mock = MagicMock(spec=Request) @@ -413,7 +413,8 @@ async def test_add_litellm_data_to_request_disabled_callbacks_not_present(): """ Test that litellm_disabled_callbacks is not added when it's not present in metadata. """ - from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.litellm_pre_call_utils import \ + add_litellm_data_to_request # Setup mock request request_mock = MagicMock(spec=Request) @@ -465,7 +466,8 @@ async def test_add_litellm_data_to_request_disabled_callbacks_invalid_type(): """ Test that litellm_disabled_callbacks is not added when it's not a list. """ - from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.litellm_pre_call_utils import \ + add_litellm_data_to_request # Setup mock request request_mock = MagicMock(spec=Request) @@ -517,7 +519,8 @@ async def test_add_litellm_data_to_request_disabled_callbacks_with_logging_setti """ Test that litellm_disabled_callbacks works correctly alongside logging settings. """ - from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.litellm_pre_call_utils import \ + add_litellm_data_to_request # Setup mock request request_mock = MagicMock(spec=Request) @@ -1027,7 +1030,8 @@ from unittest.mock import AsyncMock from fastapi.responses import Response from litellm.integrations.custom_logger import CustomLogger -from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing +from litellm.proxy.common_request_processing import \ + ProxyBaseLLMRequestProcessing from litellm.proxy.utils import ProxyLogging from litellm.types.utils import StandardLoggingPayload @@ -1403,7 +1407,8 @@ async def test_embedding_header_forwarding_with_model_group(): importlib.reload(pre_call_utils_module) # Re-import the function after reload to get the fresh version - from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.litellm_pre_call_utils import \ + add_litellm_data_to_request # Setup mock request for embeddings request_mock = MagicMock(spec=Request) @@ -1531,18 +1536,17 @@ async def test_embedding_header_forwarding_without_model_group_config(): litellm.model_group_settings = original_model_group_settings -def test_add_guardrails_from_policy_engine(): +@pytest.mark.asyncio +async def test_add_guardrails_from_policy_engine(): """ Test that add_guardrails_from_policy_engine adds guardrails from matching policies and tracks applied policies in metadata. """ - from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry + from litellm.proxy.policy_engine.attachment_registry import \ + get_attachment_registry from litellm.proxy.policy_engine.policy_registry import get_policy_registry - from litellm.types.proxy.policy_engine import ( - Policy, - PolicyAttachment, - PolicyGuardrails, - ) + from litellm.types.proxy.policy_engine import (Policy, PolicyAttachment, + PolicyGuardrails) # Setup test data data = { @@ -1578,7 +1582,7 @@ def test_add_guardrails_from_policy_engine(): attachment_registry._initialized = True # Call the function - add_guardrails_from_policy_engine( + await add_guardrails_from_policy_engine( data=data, metadata_variable_name="metadata", user_api_key_dict=user_api_key_dict, @@ -1601,11 +1605,12 @@ def test_add_guardrails_from_policy_engine(): attachment_registry._initialized = False -def test_add_guardrails_from_policy_engine_accepts_dynamic_policies_and_pops_from_data(): +@pytest.mark.asyncio +async def test_add_guardrails_from_policy_engine_accepts_dynamic_policies_and_pops_from_data(): """ Test that add_guardrails_from_policy_engine accepts dynamic 'policies' from the request body and removes them to prevent forwarding to the LLM provider. - + This is critical because 'policies' is a LiteLLM proxy-specific parameter that should not be sent to the actual LLM API (e.g., OpenAI, Anthropic, etc.). """ @@ -1631,7 +1636,7 @@ def test_add_guardrails_from_policy_engine_accepts_dynamic_policies_and_pops_fro policy_registry._initialized = False # Call the function - should accept dynamic policies and not raise an error - add_guardrails_from_policy_engine( + await add_guardrails_from_policy_engine( data=data, metadata_variable_name="metadata", user_api_key_dict=user_api_key_dict, @@ -1646,3 +1651,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_ 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 diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index a16afe9c8a2..a618988fa1a 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -5812,6 +5812,102 @@ export const updatePolicyCall = async (accessToken: string, policyId: string, po } }; +export const listPolicyVersions = async ( + accessToken: string, + policyName: string +): Promise<{ policy_name: string; versions: any[]; total_count: number }> => { + try { + const encodedName = encodeURIComponent(policyName); + const url = proxyBaseUrl + ? `${proxyBaseUrl}/policies/name/${encodedName}/versions` + : `/policies/name/${encodedName}/versions`; + const response = await fetch(url, { + method: "GET", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + handleError(errorMessage); + throw new Error(errorMessage); + } + + return await response.json(); + } catch (error) { + console.error("Failed to list policy versions:", error); + throw error; + } +}; + +export const createPolicyVersion = async ( + accessToken: string, + policyName: string, + sourcePolicyId?: string | null +): Promise => { + try { + const encodedName = encodeURIComponent(policyName); + const url = proxyBaseUrl + ? `${proxyBaseUrl}/policies/name/${encodedName}/versions` + : `/policies/name/${encodedName}/versions`; + const response = await fetch(url, { + method: "POST", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + body: JSON.stringify({ source_policy_id: sourcePolicyId ?? undefined }), + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + handleError(errorMessage); + throw new Error(errorMessage); + } + + return await response.json(); + } catch (error) { + console.error("Failed to create policy version:", error); + throw error; + } +}; + +export const updatePolicyVersionStatus = async ( + accessToken: string, + policyId: string, + versionStatus: "published" | "production" +): Promise => { + try { + const url = proxyBaseUrl + ? `${proxyBaseUrl}/policies/${policyId}/status` + : `/policies/${policyId}/status`; + const response = await fetch(url, { + method: "PUT", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + body: JSON.stringify({ version_status: versionStatus }), + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + handleError(errorMessage); + throw new Error(errorMessage); + } + + return await response.json(); + } catch (error) { + console.error("Failed to update policy version status:", error); + throw error; + } +}; + export const deletePolicyCall = async (accessToken: string, policyId: string) => { try { const url = proxyBaseUrl ? `${proxyBaseUrl}/policies/${policyId}` : `/policies/${policyId}`; diff --git a/ui/litellm-dashboard/src/components/playground/complianceUI/ComplianceUI.tsx b/ui/litellm-dashboard/src/components/playground/complianceUI/ComplianceUI.tsx index 721fda31f2a..ec151036b2c 100644 --- a/ui/litellm-dashboard/src/components/playground/complianceUI/ComplianceUI.tsx +++ b/ui/litellm-dashboard/src/components/playground/complianceUI/ComplianceUI.tsx @@ -8,9 +8,10 @@ import { } from "@/data/compliancePrompts"; import { getGuardrailsList, - getPoliciesList, testPoliciesAndGuardrails, } from "@/components/networking"; +import PolicySelector, { getPolicyOptionEntries } from "@/components/policies/PolicySelector"; +import { Policy } from "@/components/policies/types"; import { makeOpenAIChatCompletionRequest } from "../llm_calls/chat_completion"; import { AlertTriangle, @@ -98,11 +99,6 @@ interface QuickTestMessage { type ResultFilter = "all" | "matches" | "mismatches" | "pending"; type RightPanelTab = "quick-test" | "batch-results"; -interface PolicyOption { - id: string; - name: string; -} - interface GuardrailOption { id: string; name: string; @@ -132,11 +128,10 @@ export default function ComplianceUI({ }: ComplianceUIProps) { const frameworks = getFrameworks(); - const [policyOptions, setPolicyOptions] = useState([]); + const [policyValueToLabel, setPolicyValueToLabel] = useState>(new Map()); const [guardrailOptions, setGuardrailOptions] = useState([]); const [selectedPolicies, setSelectedPolicies] = useState([]); const [selectedGuardrails, setSelectedGuardrails] = useState([]); - const [showPolicyDropdown, setShowPolicyDropdown] = useState(false); const [showGuardrailDropdown, setShowGuardrailDropdown] = useState(false); const [selectedPromptIds, setSelectedPromptIds] = useState>(new Set()); @@ -164,20 +159,16 @@ export default function ComplianceUI({ const [expandedResults, setExpandedResults] = useState>(new Set()); const batchAbortControllerRef = useRef(null); + const handlePoliciesLoaded = useCallback((policies: Policy[]) => { + const entries = getPolicyOptionEntries(policies); + setPolicyValueToLabel(new Map(entries.map((e) => [e.value, e.label]))); + }, []); + useEffect(() => { if (!accessToken) return; - const fetchOptions = async () => { + const fetchGuardrails = async () => { try { - const [policiesRes, guardrailsRes] = await Promise.all([ - getPoliciesList(accessToken).catch(() => ({ policies: [] })), - getGuardrailsList(accessToken).catch(() => ({ guardrails: [] })), - ]); - setPolicyOptions( - (policiesRes.policies || []).map((p: { policy_name: string; policy_id?: string }) => ({ - id: p.policy_id ?? p.policy_name, - name: p.policy_name, - })) - ); + const guardrailsRes = await getGuardrailsList(accessToken).catch(() => ({ guardrails: [] })); setGuardrailOptions( (guardrailsRes.guardrails || []).map((g: { guardrail_name: string }) => ({ id: g.guardrail_name, @@ -186,11 +177,10 @@ export default function ComplianceUI({ })) ); } catch { - setPolicyOptions([]); setGuardrailOptions([]); } }; - fetchOptions(); + fetchGuardrails(); }, [accessToken]); useEffect(() => { @@ -282,12 +272,6 @@ export default function ComplianceUI({ const deselectAll = () => setSelectedPromptIds(new Set()); - const togglePolicy = (id: string) => { - setSelectedPolicies((prev) => - prev.includes(id) ? prev.filter((p) => p !== id) : [...prev, id] - ); - }; - const toggleGuardrail = (id: string) => { setSelectedGuardrails((prev) => prev.includes(id) ? prev.filter((g) => g !== id) : [...prev, id] @@ -766,76 +750,13 @@ export default function ComplianceUI({ -
- - {showPolicyDropdown && ( -
- {policyOptions.length === 0 ? ( -
- No policies available. Create policies in the Policies page. -
- ) : ( - policyOptions.map((policy) => ( - - )) - )} -
- )} -
- {selectedPolicies.length > 0 && ( -
- {selectedPolicies.map((id) => { - const p = policyOptions.find((x) => x.id === id); - return ( - - {p?.name} - - - ); - })} -
+ {accessToken && ( + )} @@ -852,10 +773,7 @@ export default function ComplianceUI({