From 6b8be4f6a06161fa167cc098a89813bc024349db Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Sun, 22 Feb 2026 03:37:13 +0000 Subject: [PATCH] fix(policy-engine): address policy versioning review issues Co-authored-by: Krish Dholakia --- .../proxy/policy_engine/policy_endpoints.py | 8 +- .../proxy/policy_engine/policy_registry.py | 177 +++++++++++------- .../proxy/policy_engine/resolver_types.py | 4 +- .../policy_engine/test_policy_versioning.py | 92 ++++++++- .../test_policy_versioning_e2e.py | 22 ++- 5 files changed, 221 insertions(+), 82 deletions(-) diff --git a/litellm/proxy/policy_engine/policy_endpoints.py b/litellm/proxy/policy_engine/policy_endpoints.py index d8de028d6a0..020b7741117 100644 --- a/litellm/proxy/policy_engine/policy_endpoints.py +++ b/litellm/proxy/policy_engine/policy_endpoints.py @@ -4,7 +4,7 @@ CRUD ENDPOINTS FOR POLICIES Provides REST API endpoints for managing policies and policy attachments. """ -from typing import Optional +from typing import Literal, Optional from fastapi import APIRouter, Depends, HTTPException @@ -37,7 +37,9 @@ router = APIRouter() dependencies=[Depends(user_api_key_auth)], response_model=PolicyListDBResponse, ) -async def list_policies(version_status: Optional[str] = None): +async def list_policies( + version_status: Optional[Literal["draft", "published", "production"]] = None, +): """ List all policies from the database. Optionally filter by version_status. @@ -187,7 +189,6 @@ async def list_policy_versions(policy_name: str): @router.post( "/policies/name/{policy_name}/versions", tags=["Policies"], - dependencies=[Depends(user_api_key_auth)], response_model=PolicyDBResponse, ) async def create_policy_version( @@ -222,7 +223,6 @@ async def create_policy_version( @router.put( "/policies/{policy_id}/status", tags=["Policies"], - dependencies=[Depends(user_api_key_auth)], response_model=PolicyDBResponse, ) async def update_policy_version_status( diff --git a/litellm/proxy/policy_engine/policy_registry.py b/litellm/proxy/policy_engine/policy_registry.py index f8e1ebd7ba1..a02e0eeca93 100644 --- a/litellm/proxy/policy_engine/policy_registry.py +++ b/litellm/proxy/policy_engine/policy_registry.py @@ -682,77 +682,78 @@ class PolicyRegistry: 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}'" + async with prisma_client.db.tx() as tx: + if source_policy_id is not None: + source = await tx.litellm_policytable.find_unique( + where={"policy_id": source_policy_id} ) - 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}'" + 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 tx.litellm_policytable.find_first( + where={ + "policy_name": policy_name, + "version_status": "production", + } ) - source = prod + 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 + # Next version number + latest = await tx.litellm_policytable.find_first( + where={"policy_name": policy_name}, + order={"version_number": "desc"}, ) - if source.pipeline is not None: - data["pipeline"] = ( - json.dumps(source.pipeline) - if isinstance(source.pipeline, dict) - else source.pipeline + 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 tx.litellm_policytable.update_many( + where={"policy_name": policy_name}, + data={"is_latest": False}, ) - created = await prisma_client.db.litellm_policytable.create(data=data) - return _row_to_policy_db_response(created) + 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 tx.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)}") @@ -811,6 +812,21 @@ class PolicyRegistry: "updated_by": updated_by, }, ) + # Keep policy_ cache in sync for non-production versions. + published_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._policies_by_id[policy_id] = (policy_name, published_policy) return _row_to_policy_db_response(updated) # new_status == "production" @@ -824,6 +840,14 @@ class PolicyRegistry: "Cannot promote draft directly to production. Publish the version first." ) + # Capture current production before demotion so cache can be updated. + current_production_rows = await prisma_client.db.litellm_policytable.find_many( + where={ + "policy_name": policy_name, + "version_status": "production", + }, + ) + # Demote current production to published await prisma_client.db.litellm_policytable.update_many( where={ @@ -837,6 +861,26 @@ class PolicyRegistry: }, ) + # Demoted production versions are now published, so they should resolve by policy_. + for demoted in current_production_rows: + demoted_policy = self._parse_policy( + demoted.policy_name, + { + "inherit": demoted.inherit, + "description": demoted.description, + "guardrails": { + "add": demoted.guardrails_add or [], + "remove": demoted.guardrails_remove or [], + }, + "condition": demoted.condition, + "pipeline": demoted.pipeline, + }, + ) + self._policies_by_id[demoted.policy_id] = ( + demoted.policy_name, + demoted_policy, + ) + # Promote this version to production updated = await prisma_client.db.litellm_policytable.update( where={"policy_id": policy_id}, @@ -848,6 +892,9 @@ class PolicyRegistry: }, ) + # Production versions should not resolve by policy_. + self._policies_by_id.pop(policy_id, None) + # Update in-memory registry: remove old production (by name), add this one self.remove_policy(policy_name) policy = self._parse_policy( diff --git a/litellm/types/proxy/policy_engine/resolver_types.py b/litellm/types/proxy/policy_engine/resolver_types.py index 2df450dc2ba..759c08ab813 100644 --- a/litellm/types/proxy/policy_engine/resolver_types.py +++ b/litellm/types/proxy/policy_engine/resolver_types.py @@ -6,7 +6,7 @@ the final guardrails list. """ from datetime import datetime -from typing import Any, Dict, List, Optional +from typing import Any, Dict, List, Literal, Optional from pydantic import BaseModel, ConfigDict, Field @@ -267,7 +267,7 @@ class PolicyVersionCreateRequest(BaseModel): class PolicyVersionStatusUpdateRequest(BaseModel): """Request body for updating a policy version's status.""" - version_status: str = Field( + version_status: Literal["published", "production"] = Field( description="New status: 'published' or 'production'.", ) diff --git a/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py b/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py index 738c611d928..27aeaf8b6e1 100644 --- a/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py +++ b/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py @@ -3,19 +3,23 @@ Unit tests for policy versioning: registry behavior, status transitions, and ver """ from datetime import datetime, timezone +from typing import Literal, get_args, get_origin from unittest.mock import AsyncMock, MagicMock import pytest +from pydantic import ValidationError from litellm.proxy.policy_engine.policy_registry import ( PolicyRegistry, _row_to_policy_db_response, get_policy_registry, ) +from litellm.proxy.policy_engine.policy_endpoints import list_policies from litellm.types.proxy.policy_engine import ( PolicyCreateRequest, PolicyDBResponse, PolicyUpdateRequest, + PolicyVersionStatusUpdateRequest, ) @@ -61,6 +65,21 @@ def _make_row( return row +def _mock_prisma_transaction(prisma: MagicMock) -> MagicMock: + """Attach an async tx() context manager and return tx mock.""" + tx = MagicMock() + + class _TxContextManager: + async def __aenter__(self): + return tx + + async def __aexit__(self, exc_type, exc, tb): + return False + + prisma.db.tx = MagicMock(return_value=_TxContextManager()) + return tx + + class TestRowToPolicyDBResponse: """Test _row_to_policy_db_response includes all version fields.""" @@ -249,6 +268,7 @@ class TestCreateNewVersion: async def test_create_new_version_from_production_increments_version(self): registry = PolicyRegistry() prisma = MagicMock() + tx = _mock_prisma_transaction(prisma) prod = _make_row( policy_id="prod-1", policy_name="foo", @@ -259,15 +279,13 @@ class TestCreateNewVersion: 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 = [ + tx.litellm_policytable.find_first = AsyncMock(side_effect=[ prod, # production lookup prod, # latest version_number lookup - ] + ]) # update_many for is_latest=False - prisma.db.litellm_policytable.update_many = AsyncMock() + tx.litellm_policytable.update_many = AsyncMock() new_row = _make_row( policy_id="new-id", policy_name="foo", @@ -279,7 +297,7 @@ class TestCreateNewVersion: description="base", pipeline={"mode": "pre_call", "steps": []}, ) - prisma.db.litellm_policytable.create = AsyncMock(return_value=new_row) + tx.litellm_policytable.create = AsyncMock(return_value=new_row) result = await registry.create_new_version( policy_name="foo", @@ -293,7 +311,7 @@ class TestCreateNewVersion: 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"] + create_call = tx.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" @@ -307,12 +325,25 @@ class TestUpdateVersionStatus: async def test_draft_to_published_sets_published_at(self): registry = PolicyRegistry() prisma = MagicMock() - draft = _make_row(policy_id="d-1", version_status="draft") + draft = _make_row( + policy_id="d-1", + policy_name="foo", + version_status="draft", + guardrails_add=["g1"], + ) updated = _make_row( policy_id="d-1", + policy_name="foo", version_status="published", published_at=datetime.now(timezone.utc), + guardrails_add=["g1", "g2"], ) + # Existing cache entry for this non-production policy should be refreshed. + stale_policy = registry._parse_policy( + "foo", + {"guardrails": {"add": ["old"], "remove": []}}, + ) + registry._policies_by_id["d-1"] = ("foo", stale_policy) prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=draft) prisma.db.litellm_policytable.update = AsyncMock(return_value=updated) @@ -326,6 +357,11 @@ class TestUpdateVersionStatus: update_data = prisma.db.litellm_policytable.update.call_args[1]["data"] assert update_data["version_status"] == "published" assert "published_at" in update_data + cached = registry.get_policy_by_id_for_request("d-1") + assert cached is not None + cached_name, cached_policy = cached + assert cached_name == "foo" + assert cached_policy.guardrails.add == ["g1", "g2"] @pytest.mark.asyncio async def test_draft_to_production_raises(self): @@ -351,13 +387,25 @@ class TestUpdateVersionStatus: policy_name="foo", version_status="published", ) + old_production = _make_row( + policy_id="prod-1", + policy_name="foo", + version_status="production", + guardrails_add=["legacy"], + ) updated_row = _make_row( policy_id="pub-1", policy_name="foo", version_status="production", production_at=datetime.now(timezone.utc), ) + # Promoted version should be removed from id cache. + registry._policies_by_id["pub-1"] = ( + "foo", + registry._parse_policy("foo", {"guardrails": {"add": ["stale"], "remove": []}}), + ) prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=published_row) + prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[old_production]) prisma.db.litellm_policytable.update_many = AsyncMock() prisma.db.litellm_policytable.update = AsyncMock(return_value=updated_row) @@ -372,6 +420,34 @@ class TestUpdateVersionStatus: assert prisma.db.litellm_policytable.update_many.called # Registry should have been updated with new production assert registry.has_policy("foo") + # Promoted production should no longer resolve via policy_ + assert registry.get_policy_by_id_for_request("pub-1") is None + # Demoted production should now be available in id cache + demoted = registry.get_policy_by_id_for_request("prod-1") + assert demoted is not None + assert demoted[0] == "foo" + + +class TestVersionStatusRequestValidation: + """Test request model validates allowed status transitions.""" + + def test_only_allows_published_or_production(self): + PolicyVersionStatusUpdateRequest(version_status="published") + PolicyVersionStatusUpdateRequest(version_status="production") + with pytest.raises(ValidationError): + PolicyVersionStatusUpdateRequest(version_status="active") + + +class TestListPoliciesValidation: + """Test query parameter validation shape for list endpoint.""" + + def test_list_policies_version_status_is_literal(self): + annotation = list_policies.__annotations__["version_status"] + annotation_args = get_args(annotation) + literal_arg = next(arg for arg in annotation_args if arg is not type(None)) + + assert get_origin(literal_arg) is Literal + assert set(get_args(literal_arg)) == {"draft", "published", "production"} class TestCompareVersions: 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 index 5d6f3a05ae6..59923a4ee03 100644 --- a/tests/test_litellm/proxy/policy_engine/test_policy_versioning_e2e.py +++ b/tests/test_litellm/proxy/policy_engine/test_policy_versioning_e2e.py @@ -50,6 +50,21 @@ def _make_row( return row +def _mock_prisma_transaction(prisma: MagicMock) -> MagicMock: + """Attach an async tx() context manager and return tx mock.""" + tx = MagicMock() + + class _TxContextManager: + async def __aenter__(self): + return tx + + async def __aexit__(self, exc_type, exc, tb): + return False + + prisma.db.tx = MagicMock(return_value=_TxContextManager()) + return tx + + @pytest.mark.asyncio async def test_full_lifecycle_create_draft_edit_publish_promote(): """ @@ -99,9 +114,10 @@ async def test_full_lifecycle_create_draft_edit_publish_promote(): 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) + tx = _mock_prisma_transaction(prisma) + tx.litellm_policytable.find_first = AsyncMock(side_effect=[created_v1, created_v1]) + tx.litellm_policytable.update_many = AsyncMock() + tx.litellm_policytable.create = AsyncMock(return_value=v2_row) draft_v2 = await registry.create_new_version( policy_name="lifecycle-policy",