mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
71b825a7f0
commit
004b134a0a
7 changed files with 390 additions and 14 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
113
tests/test_litellm/proxy/policy_engine/test_policy_registry.py
Normal file
113
tests/test_litellm/proxy/policy_engine/test_policy_registry.py
Normal 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
|
||||
Loading…
Add table
Reference in a new issue