From 004b134a0a42647e39b007460ee64a1213fb0a21 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 30 Jul 2026 19:21:36 +0000 Subject: [PATCH] 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> --- .../policy_engine/attachment_registry.py | 13 +- .../proxy/policy_engine/policy_endpoints.py | 123 ++++++++++++++++-- .../proxy/policy_engine/policy_registry.py | 25 +++- .../proxy/policy_engine/resolver_types.py | 10 +- .../policy_engine/test_attachment_registry.py | 34 +++++ .../test_policy_engine_endpoints.py | 86 ++++++++++++ .../policy_engine/test_policy_registry.py | 113 ++++++++++++++++ 7 files changed, 390 insertions(+), 14 deletions(-) create mode 100644 tests/test_litellm/proxy/policy_engine/test_policy_engine_endpoints.py create mode 100644 tests/test_litellm/proxy/policy_engine/test_policy_registry.py diff --git a/litellm/proxy/policy_engine/attachment_registry.py b/litellm/proxy/policy_engine/attachment_registry.py index 8ef509810ba..0ba4e129117 100644 --- a/litellm/proxy/policy_engine/attachment_registry.py +++ b/litellm/proxy/policy_engine/attachment_registry.py @@ -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( diff --git a/litellm/proxy/policy_engine/policy_endpoints.py b/litellm/proxy/policy_engine/policy_endpoints.py index a879f6b6f7e..86ecebb1605 100644 --- a/litellm/proxy/policy_engine/policy_endpoints.py +++ b/litellm/proxy/policy_engine/policy_endpoints.py @@ -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") diff --git a/litellm/proxy/policy_engine/policy_registry.py b/litellm/proxy/policy_engine/policy_registry.py index e1afbf2f5f2..af47bdd9dfa 100644 --- a/litellm/proxy/policy_engine/policy_registry.py +++ b/litellm/proxy/policy_engine/policy_registry.py @@ -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_ 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( diff --git a/litellm/types/proxy/policy_engine/resolver_types.py b/litellm/types/proxy/policy_engine/resolver_types.py index 2c7e8d5afc9..6edd520e783 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 @@ -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): diff --git a/tests/test_litellm/proxy/policy_engine/test_attachment_registry.py b/tests/test_litellm/proxy/policy_engine/test_attachment_registry.py index 1ae0b4d3d48..d3e5d6769bc 100644 --- a/tests/test_litellm/proxy/policy_engine/test_attachment_registry.py +++ b/tests/test_litellm/proxy/policy_engine/test_attachment_registry.py @@ -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"] diff --git a/tests/test_litellm/proxy/policy_engine/test_policy_engine_endpoints.py b/tests/test_litellm/proxy/policy_engine/test_policy_engine_endpoints.py new file mode 100644 index 00000000000..94be2d1c7b6 --- /dev/null +++ b/tests/test_litellm/proxy/policy_engine/test_policy_engine_endpoints.py @@ -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) diff --git a/tests/test_litellm/proxy/policy_engine/test_policy_registry.py b/tests/test_litellm/proxy/policy_engine/test_policy_registry.py new file mode 100644 index 00000000000..f3e7449c6ae --- /dev/null +++ b/tests/test_litellm/proxy/policy_engine/test_policy_registry.py @@ -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