fix(policy_engine): keep config-defined policies enforced and listed when a DB is connected

Config policies and attachments were wiped from the in-memory registries by the periodic DB sync, and the list endpoints only ever read the DB, so a DB-connected proxy silently stopped enforcing them and never showed them in the UI.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-07-30 19:21:36 +00:00
parent 71b825a7f0
commit 004b134a0a
7 changed files with 390 additions and 14 deletions

View file

@ -42,6 +42,7 @@ class AttachmentRegistry:
def __init__(self):
self._attachments: List[PolicyAttachment] = []
self._config_attachments: List[PolicyAttachment] = []
self._initialized: bool = False
def load_attachments(self, attachments_config: List[Dict[str, Any]]) -> None:
@ -52,11 +53,13 @@ class AttachmentRegistry:
attachments_config: List of attachment dictionaries from YAML.
"""
self._attachments = []
self._config_attachments = []
for attachment_data in attachments_config:
try:
attachment = self._parse_attachment(attachment_data)
self._attachments.append(attachment)
self._config_attachments.append(attachment)
verbose_proxy_logger.debug(f"Loaded attachment for policy: {attachment.policy}")
except Exception as e:
verbose_proxy_logger.error(f"Error loading attachment: {str(e)}")
@ -173,6 +176,12 @@ class AttachmentRegistry:
"""
return self._attachments.copy()
def get_config_attachments(self) -> List[PolicyAttachment]:
"""
Get the attachments that came from config.yaml.
"""
return self._config_attachments.copy()
def get_attachments_for_policy(self, policy_name: str) -> List[PolicyAttachment]:
"""
Get all attachments for a specific policy.
@ -199,6 +208,7 @@ class AttachmentRegistry:
Clear all attachments from the registry.
"""
self._attachments = []
self._config_attachments = []
self._initialized = False
def add_attachment(self, attachment: PolicyAttachment) -> None:
@ -435,8 +445,7 @@ class AttachmentRegistry:
try:
attachments = await self.get_all_attachments_from_db(prisma_client)
# Clear existing attachments and reload from DB
self._attachments = []
self._attachments = list(self._config_attachments)
for attachment_response in attachments:
attachment = PolicyAttachment(

View file

@ -32,6 +32,72 @@ from litellm.types.proxy.policy_engine import (
router = APIRouter()
CONFIG_ATTACHMENT_ID_PREFIX = "config-attachment-"
def _config_policies_as_responses() -> list[PolicyDBResponse]:
"""
Render config.yaml policies in the same shape as DB-backed ones.
Config policies are not versioned, so they are reported as the production version
and keyed by their name; they have no DB row to take a UUID from.
"""
return [
PolicyDBResponse(
policy_id=policy_name,
policy_name=policy_name,
inherit=policy.inherit,
description=policy.description,
guardrails_add=policy.guardrails.get_add(),
guardrails_remove=policy.guardrails.get_remove(),
condition=policy.condition.model_dump() if policy.condition else None,
pipeline=policy.pipeline.model_dump() if policy.pipeline else None,
policy_definition_location="config",
)
for policy_name, policy in get_policy_registry().get_config_policies().items()
]
def _config_policy_response(policy_id: str) -> PolicyDBResponse | None:
"""
Look up a config.yaml policy by the ID the list endpoint reports for it.
"""
return next(
(policy for policy in _config_policies_as_responses() if policy.policy_id == policy_id),
None,
)
def _reject_if_config_policy(policy_id: str) -> None:
"""
Config-defined policies have no DB row to mutate; say so instead of 404ing.
"""
if _config_policy_response(policy_id) is None:
return
raise HTTPException(
status_code=400,
detail=f"Policy '{policy_id}' is defined in config.yaml and can only be changed there",
)
def _config_attachments_as_responses() -> list[PolicyAttachmentDBResponse]:
"""
Render config.yaml policy attachments in the same shape as DB-backed ones.
"""
return [
PolicyAttachmentDBResponse(
attachment_id=f"{CONFIG_ATTACHMENT_ID_PREFIX}{index}",
policy_name=attachment.policy,
scope=attachment.scope,
teams=attachment.teams or [],
keys=attachment.keys or [],
models=attachment.models or [],
tags=attachment.tags or [],
policy_definition_location="config",
)
for index, attachment in enumerate(get_attachment_registry().get_config_attachments())
]
# ─────────────────────────────────────────────────────────────────────────────
# Policy CRUD Endpoints
@ -46,7 +112,10 @@ router = APIRouter()
)
async def list_policies(version_status: Optional[str] = None):
"""
List all policies from the database. Optionally filter by version_status.
List all policies from the database and from config.yaml. Optionally filter by version_status.
Config-defined policies are reported with policy_definition_location="config"; they have no
versions, so they are only included when version_status is unset or "production".
Query params:
- version_status: Optional. One of "draft", "published", "production".
@ -84,11 +153,14 @@ async def list_policies(version_status: Optional[str] = None):
"""
from litellm.proxy.proxy_server import prisma_client
config_policies = _config_policies_as_responses() if version_status in (None, "production") else []
if prisma_client is None:
raise HTTPException(status_code=500, detail="Database not connected")
return PolicyListDBResponse(policies=config_policies, total_count=len(config_policies))
try:
policies = await get_policy_registry().get_all_policies_from_db(prisma_client, version_status=version_status)
db_policies = await get_policy_registry().get_all_policies_from_db(prisma_client, version_status=version_status)
policies = db_policies + config_policies
return PolicyListDBResponse(policies=policies, total_count=len(policies))
except Exception as e:
verbose_proxy_logger.exception(f"Error listing policies: {e}")
@ -333,7 +405,7 @@ async def delete_all_policy_versions(policy_name: str):
)
async def get_policy(policy_id: str):
"""
Get a policy by ID.
Get a policy by ID. Config-defined policies are looked up by their name.
Example Request:
```bash
@ -343,14 +415,20 @@ async def get_policy(policy_id: str):
"""
from litellm.proxy.proxy_server import prisma_client
config_policy = _config_policy_response(policy_id)
if prisma_client is None:
raise HTTPException(status_code=500, detail="Database not connected")
if config_policy is None:
raise HTTPException(status_code=404, detail=f"Policy with ID {policy_id} not found")
return config_policy
try:
result = await get_policy_registry().get_policy_by_id_from_db(
policy_id=policy_id,
prisma_client=prisma_client,
)
if result is None:
result = config_policy
if result is None:
raise HTTPException(status_code=404, detail=f"Policy with ID {policy_id} not found")
return result
@ -388,6 +466,8 @@ async def update_policy(
"""
from litellm.proxy.proxy_server import prisma_client
_reject_if_config_policy(policy_id)
if prisma_client is None:
raise HTTPException(status_code=500, detail="Database not connected")
@ -444,6 +524,8 @@ async def delete_policy(policy_id: str):
"""
from litellm.proxy.proxy_server import prisma_client
_reject_if_config_policy(policy_id)
if prisma_client is None:
raise HTTPException(status_code=500, detail="Database not connected")
@ -606,7 +688,9 @@ async def test_pipeline(
)
async def list_policy_attachments():
"""
List all policy attachments from the database.
List all policy attachments from the database and from config.yaml.
Config-defined attachments are reported with policy_definition_location="config".
Example Request:
```bash
@ -635,11 +719,14 @@ async def list_policy_attachments():
"""
from litellm.proxy.proxy_server import prisma_client
config_attachments = _config_attachments_as_responses()
if prisma_client is None:
raise HTTPException(status_code=500, detail="Database not connected")
return PolicyAttachmentListResponse(attachments=config_attachments, total_count=len(config_attachments))
try:
attachments = await get_attachment_registry().get_all_attachments_from_db(prisma_client)
db_attachments = await get_attachment_registry().get_all_attachments_from_db(prisma_client)
attachments = db_attachments + config_attachments
return PolicyAttachmentListResponse(attachments=attachments, total_count=len(attachments))
except Exception as e:
verbose_proxy_logger.exception(f"Error listing policy attachments: {e}")
@ -756,8 +843,18 @@ async def get_policy_attachment(attachment_id: str):
"""
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(status_code=500, detail="Database not connected")
config_attachment = next(
(a for a in _config_attachments_as_responses() if a.attachment_id == attachment_id),
None,
)
if prisma_client is None or config_attachment is not None:
if config_attachment is None:
raise HTTPException(
status_code=404,
detail=f"Attachment with ID {attachment_id} not found",
)
return config_attachment
try:
result = await get_attachment_registry().get_attachment_by_id_from_db(
@ -801,6 +898,12 @@ async def delete_policy_attachment(attachment_id: str):
"""
from litellm.proxy.proxy_server import prisma_client
if attachment_id.startswith(CONFIG_ATTACHMENT_ID_PREFIX):
raise HTTPException(
status_code=400,
detail=f"Attachment '{attachment_id}' is defined in config.yaml and can only be removed there",
)
if prisma_client is None:
raise HTTPException(status_code=500, detail="Database not connected")

View file

@ -13,6 +13,7 @@ from datetime import datetime, timezone
from typing import (
TYPE_CHECKING,
Any,
Literal,
Optional,
Protocol,
TypedDict,
@ -163,6 +164,7 @@ class PolicyRegistry:
def __init__(self):
self._policies: dict[str, Policy] = {}
self._policies_by_id: dict[str, tuple[str, Policy]] = {}
self._config_policies: dict[str, Policy] = {}
self._initialized: bool = False
def load_policies(self, policies_config: Mapping[str, dict[str, object]]) -> None:
@ -175,11 +177,13 @@ class PolicyRegistry:
"""
self._policies = {}
self._policies_by_id = {}
self._config_policies = {}
for policy_name, policy_data in policies_config.items():
try:
policy = self._parse_policy(policy_name, policy_data)
self._policies[policy_name] = policy
self._config_policies[policy_name] = policy
verbose_proxy_logger.debug(f"Loaded policy: {policy_name}")
except Exception as e:
verbose_proxy_logger.error(f"Error loading policy '{policy_name}': {str(e)}")
@ -299,8 +303,27 @@ class PolicyRegistry:
Clear all policies from the registry.
"""
self._policies = {}
self._config_policies = {}
self._initialized = False
def get_config_policies(self) -> dict[str, Policy]:
"""
Get the policies that came from config.yaml, excluding names a DB row shadows.
"""
return {
name: policy for name, policy in self._config_policies.items() if self.get_policy_source(name) == "config"
}
def get_policy_source(self, policy_name: str) -> Literal["config", "db"] | None:
"""
Return the provenance of an in-memory policy, or None if it isn't loaded.
"""
if policy_name not in self._policies:
return None
if self._policies.get(policy_name) is self._config_policies.get(policy_name):
return "config"
return "db"
def add_policy(self, policy_name: str, policy: Policy) -> None:
"""
Add or update a single policy.
@ -595,7 +618,7 @@ class PolicyRegistry:
policy_<uuid> overrides can be resolved without DB access in the hot path.
"""
try:
self._policies = {}
self._policies = dict(self._config_policies)
production = await self.get_all_policies_from_db(prisma_client, version_status="production")
for policy_response in production:
policy = self._parse_policy(

View file

@ -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
@ -220,6 +220,10 @@ class PolicyDBResponse(BaseModel):
updated_at: Optional[datetime] = Field(default=None, description="When the policy was last updated.")
created_by: Optional[str] = Field(default=None, description="Who created the policy.")
updated_by: Optional[str] = Field(default=None, description="Who last updated the policy.")
policy_definition_location: Literal["config", "db"] = Field(
default="db",
description="Where the policy is defined: 'config' for config.yaml, 'db' for the database.",
)
class PolicyListDBResponse(BaseModel):
@ -317,6 +321,10 @@ class PolicyAttachmentDBResponse(BaseModel):
updated_at: Optional[datetime] = Field(default=None, description="When the attachment was last updated.")
created_by: Optional[str] = Field(default=None, description="Who created the attachment.")
updated_by: Optional[str] = Field(default=None, description="Who last updated the attachment.")
policy_definition_location: Literal["config", "db"] = Field(
default="db",
description="Where the attachment is defined: 'config' for config.yaml, 'db' for the database.",
)
class PolicyAttachmentListResponse(BaseModel):

View file

@ -389,3 +389,37 @@ class TestAttachmentRegistrySingleton:
registry1 = get_attachment_registry()
registry2 = get_attachment_registry()
assert registry1 is registry2
class TestConfigAttachmentsSurviveDbSync:
"""A DB-connected proxy must keep enforcing attachments defined in config.yaml."""
@pytest.mark.asyncio
async def test_config_attachment_survives_sync(self):
from datetime import datetime, timezone
from unittest.mock import AsyncMock, MagicMock
registry = AttachmentRegistry()
registry.load_attachments([{"policy": "config-policy", "scope": "*"}])
db_row = MagicMock()
db_row.attachment_id = "aid-1"
db_row.policy_name = "db-policy"
db_row.scope = "*"
db_row.teams = []
db_row.keys = []
db_row.models = []
db_row.tags = []
db_row.created_at = datetime.now(timezone.utc)
db_row.updated_at = datetime.now(timezone.utc)
db_row.created_by = None
db_row.updated_by = None
prisma = MagicMock()
prisma.db.litellm_policyattachmenttable.find_many = AsyncMock(return_value=[db_row])
await registry.sync_attachments_from_db(prisma)
await registry.sync_attachments_from_db(prisma)
context = PolicyMatchContext(team_alias="t", key_alias="k", model="gpt-4")
attached = registry.get_attached_policies(context)
assert attached == ["config-policy", "db-policy"]
assert [a.policy for a in registry.get_config_attachments()] == ["config-policy"]

View file

@ -0,0 +1,86 @@
"""
Unit tests for the policy list endpoints: config-defined policies and attachments
must be listed alongside DB-backed ones.
"""
import pytest
from fastapi import HTTPException
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry
from litellm.proxy.policy_engine.policy_endpoints import (
delete_policy,
delete_policy_attachment,
get_policy,
list_policies,
list_policy_attachments,
)
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
@pytest.fixture
def config_policy_engine(monkeypatch):
"""Load one config policy + attachment into the registries, with no DB attached."""
monkeypatch.setattr(proxy_server, "prisma_client", None)
policy_registry = get_policy_registry()
attachment_registry = get_attachment_registry()
policy_registry.load_policies({"config-policy": {"description": "from config", "guardrails": {"add": ["tooling"]}}})
attachment_registry.load_attachments([{"policy": "config-policy", "scope": "*"}])
yield
policy_registry.clear()
attachment_registry.clear()
@pytest.mark.asyncio
async def test_list_policies_returns_config_policies(config_policy_engine):
response = await list_policies()
assert response.total_count == 1
policy = response.policies[0]
assert policy.policy_name == "config-policy"
assert policy.policy_id == "config-policy"
assert policy.guardrails_add == ["tooling"]
assert policy.policy_definition_location == "config"
@pytest.mark.asyncio
async def test_list_policies_excludes_config_policies_for_version_filters(config_policy_engine):
assert (await list_policies(version_status="draft")).policies == []
assert len((await list_policies(version_status="production")).policies) == 1
@pytest.mark.asyncio
async def test_list_policy_attachments_returns_config_attachments(config_policy_engine):
response = await list_policy_attachments()
assert response.total_count == 1
attachment = response.attachments[0]
assert attachment.policy_name == "config-policy"
assert attachment.scope == "*"
assert attachment.policy_definition_location == "config"
@pytest.mark.asyncio
async def test_get_policy_returns_config_policy(config_policy_engine):
policy = await get_policy("config-policy")
assert policy.policy_name == "config-policy"
assert policy.policy_definition_location == "config"
@pytest.mark.asyncio
async def test_mutating_config_policy_is_rejected(config_policy_engine):
with pytest.raises(HTTPException) as exc_info:
await delete_policy("config-policy")
assert exc_info.value.status_code == 400
assert "config.yaml" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_deleting_config_attachment_is_rejected(config_policy_engine):
attachment_id = (await list_policy_attachments()).attachments[0].attachment_id
with pytest.raises(HTTPException) as exc_info:
await delete_policy_attachment(attachment_id)
assert exc_info.value.status_code == 400
assert "config.yaml" in str(exc_info.value.detail)

View file

@ -0,0 +1,113 @@
"""
Unit tests for PolicyRegistry - config vs DB provenance of in-memory policies.
"""
from datetime import datetime, timezone
from unittest.mock import AsyncMock, MagicMock
import pytest
from litellm.proxy.policy_engine.policy_registry import PolicyRegistry
def _make_row(policy_name, guardrails_add, version_status="production", policy_id="pid-1"):
row = MagicMock()
row.policy_id = policy_id
row.policy_name = policy_name
row.version_number = 1
row.version_status = version_status
row.parent_version_id = None
row.is_latest = True
row.published_at = None
row.production_at = None
row.inherit = None
row.description = None
row.guardrails_add = guardrails_add
row.guardrails_remove = []
row.condition = None
row.pipeline = None
row.created_at = datetime.now(timezone.utc)
row.updated_at = datetime.now(timezone.utc)
row.created_by = None
row.updated_by = None
return row
def _prisma_with_rows(*rows):
prisma = MagicMock()
prisma.db.litellm_policytable.find_many = AsyncMock(
side_effect=lambda where=None, order=None: [
row
for row in rows
if where is None
or (
row.version_status == where.get("version_status")
if isinstance(where.get("version_status"), str)
else row.version_status in where.get("version_status", {}).get("in", [])
)
]
)
return prisma
class TestConfigPoliciesSurviveDbSync:
"""A DB-connected proxy must keep enforcing policies defined in config.yaml."""
@pytest.mark.asyncio
async def test_config_policy_survives_sync(self):
registry = PolicyRegistry()
registry.load_policies({"config-policy": {"guardrails": {"add": ["tooling"]}}})
await registry.sync_policies_from_db(_prisma_with_rows(_make_row("db-policy", ["pii"])))
config_policy = registry.get_policy("config-policy")
assert config_policy is not None
assert config_policy.guardrails.get_add() == ["tooling"]
assert registry.get_policy_source("config-policy") == "config"
db_policy = registry.get_policy("db-policy")
assert db_policy is not None
assert db_policy.guardrails.get_add() == ["pii"]
assert registry.get_policy_source("db-policy") == "db"
@pytest.mark.asyncio
async def test_config_policy_survives_repeated_syncs(self):
registry = PolicyRegistry()
registry.load_policies({"config-policy": {"guardrails": {"add": ["tooling"]}}})
prisma = _prisma_with_rows(_make_row("db-policy", ["pii"]))
await registry.sync_policies_from_db(prisma)
await registry.sync_policies_from_db(prisma)
assert registry.has_policy("config-policy")
assert registry.get_policy_names().count("db-policy") == 1
@pytest.mark.asyncio
async def test_db_policy_shadows_config_policy_of_same_name(self):
registry = PolicyRegistry()
registry.load_policies({"shared": {"guardrails": {"add": ["from-config"]}}})
await registry.sync_policies_from_db(_prisma_with_rows(_make_row("shared", ["from-db"])))
policy = registry.get_policy("shared")
assert policy is not None
assert policy.guardrails.get_add() == ["from-db"]
assert registry.get_policy_source("shared") == "db"
assert "shared" not in registry.get_config_policies()
class TestGetConfigPolicies:
def test_reports_config_policies(self):
registry = PolicyRegistry()
registry.load_policies({"config-policy": {"description": "d", "guardrails": {"add": ["tooling"]}}})
config_policies = registry.get_config_policies()
assert list(config_policies) == ["config-policy"]
assert config_policies["config-policy"].description == "d"
def test_empty_before_any_config_is_loaded(self):
assert PolicyRegistry().get_config_policies() == {}
def test_source_is_none_for_unknown_policy(self):
assert PolicyRegistry().get_policy_source("nope") is None