From 4e9df553926347412a85e6b3bb22d4e5baae2afd Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 30 Jul 2026 12:21:03 -0700 Subject: [PATCH 1/7] fix(policy_engine): preserve config-defined policies across DB sync and expose them via list APIs --- litellm/proxy/_lazy_openapi_snapshot.json | 68 +++++- .../policy_engine/attachment_registry.py | 29 ++- .../proxy/policy_engine/policy_endpoints.py | 74 ++++++- .../proxy/policy_engine/policy_registry.py | 48 ++++- .../proxy/policy_engine/resolver_types.py | 10 +- .../policy_engine/test_attachment_registry.py | 64 ++++++ .../test_policy_engine_endpoints.py | 195 ++++++++++++++++++ .../policy_engine/test_policy_versioning.py | 79 +++++++ .../_components/AttachmentTableColumns.tsx | 7 + .../_components/PolicyTableColumns.tsx | 43 ++-- .../src/components/policies/types.ts | 2 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 35 +++- 12 files changed, 611 insertions(+), 43 deletions(-) create mode 100644 tests/test_litellm/proxy/policy_engine/test_policy_engine_endpoints.py diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 96f6ee89d56..12da0a26708 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -20402,6 +20402,16 @@ "description": "Who created the attachment.", "title": "Created By" }, + "definition_location": { + "default": "db", + "description": "Where this attachment is defined: 'db' (database) or 'config' (config.yaml).", + "enum": [ + "db", + "config" + ], + "title": "Definition Location", + "type": "string" + }, "keys": { "description": "Key patterns.", "items": { @@ -20658,6 +20668,16 @@ "description": "Who created the policy.", "title": "Created By" }, + "definition_location": { + "default": "db", + "description": "Where this policy is defined: 'db' (database) or 'config' (config.yaml).", + "enum": [ + "db", + "config" + ], + "title": "Definition Location", + "type": "string" + }, "description": { "anyOf": [ { @@ -21129,12 +21149,45 @@ "title": "PolicyVersionStatusUpdateRequest", "type": "object" }, + "UsageChartPoint": { + "properties": { + "blocked": { + "title": "Blocked", + "type": "integer" + }, + "date": { + "title": "Date", + "type": "string" + }, + "passed": { + "title": "Passed", + "type": "integer" + }, + "score": { + "anyOf": [ + { + "type": "number" + }, + { + "type": "null" + } + ], + "title": "Score" + } + }, + "required": [ + "date", + "passed", + "blocked" + ], + "title": "UsageChartPoint", + "type": "object" + }, "UsageOverviewResponse": { "properties": { "chart": { "items": { - "additionalProperties": true, - "type": "object" + "$ref": "#/components/schemas/UsageChartPoint" }, "title": "Chart", "type": "array" @@ -21243,6 +21296,13 @@ }, "ValidationError": { "properties": { + "ctx": { + "title": "Context", + "type": "object" + }, + "input": { + "title": "Input" + }, "loc": { "items": { "anyOf": [ @@ -21420,7 +21480,7 @@ }, "/policies/attachments/list": { "get": { - "description": "List all policy attachments from the database.\n\nExample Request:\n```bash\ncurl -X GET \"http://localhost:4000/policies/attachments/list\" \\\n -H \"Authorization: Bearer \"\n```\n\nExample Response:\n```json\n{\n \"attachments\": [\n {\n \"attachment_id\": \"123e4567-e89b-12d3-a456-426614174000\",\n \"policy_name\": \"global-baseline\",\n \"scope\": \"*\",\n \"teams\": [],\n \"keys\": [],\n \"models\": [],\n \"created_at\": \"2024-01-01T00:00:00Z\",\n \"updated_at\": \"2024-01-01T00:00:00Z\"\n }\n ],\n \"total_count\": 1\n}\n```", + "description": "List all policy attachments from the database and config.yaml.\n\nConfig-defined attachments are returned with definition_location \"config\" and a\nsynthetic attachment_id (\"config-\").\n\nExample Request:\n```bash\ncurl -X GET \"http://localhost:4000/policies/attachments/list\" \\\n -H \"Authorization: Bearer \"\n```\n\nExample Response:\n```json\n{\n \"attachments\": [\n {\n \"attachment_id\": \"123e4567-e89b-12d3-a456-426614174000\",\n \"policy_name\": \"global-baseline\",\n \"scope\": \"*\",\n \"teams\": [],\n \"keys\": [],\n \"models\": [],\n \"created_at\": \"2024-01-01T00:00:00Z\",\n \"updated_at\": \"2024-01-01T00:00:00Z\"\n }\n ],\n \"total_count\": 1\n}\n```", "operationId": "list_policy_attachments_policies_attachments_list_get", "responses": { "200": { @@ -21596,7 +21656,7 @@ }, "/policies/list": { "get": { - "description": "List all policies from the database. Optionally filter by version_status.\n\nQuery params:\n- version_status: Optional. One of \"draft\", \"published\", \"production\".\n If omitted, all versions are returned.\n\nExample Request:\n```bash\ncurl -X GET \"http://localhost:4000/policies/list\" \\\n -H \"Authorization: Bearer \"\ncurl -X GET \"http://localhost:4000/policies/list?version_status=production\" \\\n -H \"Authorization: Bearer \"\n```\n\nExample Response:\n```json\n{\n \"policies\": [\n {\n \"policy_id\": \"123e4567-e89b-12d3-a456-426614174000\",\n \"policy_name\": \"global-baseline\",\n \"version_number\": 1,\n \"version_status\": \"production\",\n \"inherit\": null,\n \"description\": \"Base guardrails for all requests\",\n \"guardrails_add\": [\"pii_masking\"],\n \"guardrails_remove\": [],\n \"condition\": null,\n \"created_at\": \"2024-01-01T00:00:00Z\",\n \"updated_at\": \"2024-01-01T00:00:00Z\"\n }\n ],\n \"total_count\": 1\n}\n```", + "description": "List all policies from the database and config.yaml. Optionally filter by version_status.\n\nConfig-defined policies are returned with definition_location \"config\" and are treated\nas production versions. On a name conflict with a DB policy, only the DB policy is returned.\n\nQuery params:\n- version_status: Optional. One of \"draft\", \"published\", \"production\".\n If omitted, all versions are returned.\n\nExample Request:\n```bash\ncurl -X GET \"http://localhost:4000/policies/list\" \\\n -H \"Authorization: Bearer \"\ncurl -X GET \"http://localhost:4000/policies/list?version_status=production\" \\\n -H \"Authorization: Bearer \"\n```\n\nExample Response:\n```json\n{\n \"policies\": [\n {\n \"policy_id\": \"123e4567-e89b-12d3-a456-426614174000\",\n \"policy_name\": \"global-baseline\",\n \"version_number\": 1,\n \"version_status\": \"production\",\n \"inherit\": null,\n \"description\": \"Base guardrails for all requests\",\n \"guardrails_add\": [\"pii_masking\"],\n \"guardrails_remove\": [],\n \"condition\": null,\n \"created_at\": \"2024-01-01T00:00:00Z\",\n \"updated_at\": \"2024-01-01T00:00:00Z\"\n }\n ],\n \"total_count\": 1\n}\n```", "operationId": "list_policies_policies_list_get", "parameters": [ { diff --git a/litellm/proxy/policy_engine/attachment_registry.py b/litellm/proxy/policy_engine/attachment_registry.py index 8ef509810ba..fe9ad3bef6d 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: tuple[PolicyAttachment, ...] = () self._initialized: bool = False def load_attachments(self, attachments_config: List[Dict[str, Any]]) -> None: @@ -62,6 +63,7 @@ class AttachmentRegistry: verbose_proxy_logger.error(f"Error loading attachment: {str(e)}") raise ValueError(f"Invalid attachment: {str(e)}") from e + self._config_attachments = tuple(self._attachments) self._initialized = True verbose_proxy_logger.info(f"Loaded {len(self._attachments)} policy attachments") @@ -173,6 +175,15 @@ class AttachmentRegistry: """ return self._attachments.copy() + def get_config_attachments(self) -> tuple[PolicyAttachment, ...]: + """ + Get the attachments loaded from config.yaml. + + Returns: + Tuple of config-defined PolicyAttachment objects + """ + return self._config_attachments + def get_attachments_for_policy(self, policy_name: str) -> List[PolicyAttachment]: """ Get all attachments for a specific policy. @@ -199,6 +210,7 @@ class AttachmentRegistry: Clear all attachments from the registry. """ self._attachments = [] + self._config_attachments = () self._initialized = False def add_attachment(self, attachment: PolicyAttachment) -> None: @@ -428,6 +440,7 @@ class AttachmentRegistry: ) -> None: """ Sync policy attachments from the database to in-memory registry. + Config-loaded attachments are preserved. Args: prisma_client: The Prisma client instance @@ -435,11 +448,8 @@ class AttachmentRegistry: try: attachments = await self.get_all_attachments_from_db(prisma_client) - # Clear existing attachments and reload from DB - self._attachments = [] - - for attachment_response in attachments: - attachment = PolicyAttachment( + db_attachments = [ + PolicyAttachment( policy=attachment_response.policy_name, scope=attachment_response.scope, teams=(attachment_response.teams if attachment_response.teams else None), @@ -447,10 +457,15 @@ class AttachmentRegistry: models=(attachment_response.models if attachment_response.models else None), tags=attachment_response.tags if attachment_response.tags else None, ) - self._attachments.append(attachment) + for attachment_response in attachments + ] + self._attachments = [*self._config_attachments, *db_attachments] self._initialized = True - verbose_proxy_logger.info(f"Synced {len(attachments)} attachments from DB to in-memory registry") + verbose_proxy_logger.info( + f"Synced {len(attachments)} attachments from DB to in-memory registry " + f"({len(self._config_attachments)} config-defined attachments preserved)" + ) except Exception as e: verbose_proxy_logger.exception(f"Error syncing attachments from DB: {e}") raise Exception(f"Error syncing attachments from DB: {str(e)}") diff --git a/litellm/proxy/policy_engine/policy_endpoints.py b/litellm/proxy/policy_engine/policy_endpoints.py index a879f6b6f7e..787f7069996 100644 --- a/litellm/proxy/policy_engine/policy_endpoints.py +++ b/litellm/proxy/policy_engine/policy_endpoints.py @@ -17,6 +17,8 @@ from litellm.proxy.policy_engine.policy_registry import get_policy_registry from litellm.types.proxy.policy_engine import ( GuardrailPipeline, PipelineTestRequest, + Policy, + PolicyAttachment, PolicyAttachmentCreateRequest, PolicyAttachmentDBResponse, PolicyAttachmentListResponse, @@ -33,6 +35,35 @@ from litellm.types.proxy.policy_engine import ( router = APIRouter() +def _config_policy_to_db_response(policy_name: str, policy: Policy) -> PolicyDBResponse: + return PolicyDBResponse( + policy_id=policy_name, + policy_name=policy_name, + version_number=1, + version_status="production", + 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, + definition_location="config", + ) + + +def _config_attachment_to_db_response(index: int, attachment: PolicyAttachment) -> PolicyAttachmentDBResponse: + return PolicyAttachmentDBResponse( + attachment_id=f"config-{index}", + policy_name=attachment.policy, + scope=attachment.scope, + teams=attachment.teams or [], + keys=attachment.keys or [], + models=attachment.models or [], + tags=attachment.tags or [], + definition_location="config", + ) + + # ───────────────────────────────────────────────────────────────────────────── # Policy CRUD Endpoints # ───────────────────────────────────────────────────────────────────────────── @@ -46,7 +77,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 config.yaml. Optionally filter by version_status. + + Config-defined policies are returned with definition_location "config" and are treated + as production versions. On a name conflict with a DB policy, only the DB policy is returned. Query params: - version_status: Optional. One of "draft", "published", "production". @@ -84,11 +118,25 @@ async def list_policies(version_status: Optional[str] = None): """ from litellm.proxy.proxy_server import prisma_client - if prisma_client is None: - raise HTTPException(status_code=500, detail="Database not connected") - try: - policies = await get_policy_registry().get_all_policies_from_db(prisma_client, version_status=version_status) + registry = get_policy_registry() + db_policies = ( + await registry.get_all_policies_from_db(prisma_client, version_status=version_status) + if prisma_client is not None + else [] + ) + db_policy_names = {db_policy.policy_name for db_policy in db_policies} + include_config = version_status in (None, "production") + config_policies = ( + [ + _config_policy_to_db_response(policy_name, policy) + for policy_name, policy in registry.list_config_policies().items() + if policy_name not in db_policy_names and registry.get_source(policy_name) != "db" + ] + if include_config + else [] + ) + 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}") @@ -606,7 +654,10 @@ async def test_pipeline( ) async def list_policy_attachments(): """ - List all policy attachments from the database. + List all policy attachments from the database and config.yaml. + + Config-defined attachments are returned with definition_location "config" and a + synthetic attachment_id ("config-"). Example Request: ```bash @@ -635,11 +686,14 @@ async def list_policy_attachments(): """ from litellm.proxy.proxy_server import prisma_client - if prisma_client is None: - raise HTTPException(status_code=500, detail="Database not connected") - try: - attachments = await get_attachment_registry().get_all_attachments_from_db(prisma_client) + registry = get_attachment_registry() + db_attachments = await registry.get_all_attachments_from_db(prisma_client) if prisma_client is not None else [] + config_attachments = [ + _config_attachment_to_db_response(index, attachment) + for index, attachment in enumerate(registry.get_config_attachments()) + ] + 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}") diff --git a/litellm/proxy/policy_engine/policy_registry.py b/litellm/proxy/policy_engine/policy_registry.py index e1afbf2f5f2..9456afaa1b9 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, @@ -162,6 +163,8 @@ class PolicyRegistry: def __init__(self): self._policies: dict[str, Policy] = {} + self._config_policies: Mapping[str, Policy] = {} + self._sources: Mapping[str, Literal["db", "config"]] = {} self._policies_by_id: dict[str, tuple[str, Policy]] = {} self._initialized: bool = False @@ -174,6 +177,8 @@ class PolicyRegistry: This is the raw config from the YAML file. """ self._policies = {} + self._config_policies = {} + self._sources = {} self._policies_by_id = {} for policy_name, policy_data in policies_config.items(): @@ -185,6 +190,8 @@ class PolicyRegistry: verbose_proxy_logger.error(f"Error loading policy '{policy_name}': {str(e)}") raise ValueError(f"Invalid policy '{policy_name}': {str(e)}") from e + self._config_policies = dict(self._policies) + self._sources = {policy_name: "config" for policy_name in self._policies} self._initialized = True verbose_proxy_logger.info(f"Loaded {len(self._policies)} policies") @@ -299,17 +306,35 @@ class PolicyRegistry: Clear all policies from the registry. """ self._policies = {} + self._config_policies = {} + self._sources = {} self._initialized = False - def add_policy(self, policy_name: str, policy: Policy) -> None: + def get_source(self, policy_name: str) -> Optional[Literal["db", "config"]]: + """ + Return the provenance of an in-memory policy, or None if unknown. + """ + return self._sources.get(policy_name) + + def list_config_policies(self) -> Mapping[str, Policy]: + """ + Return the policies loaded from config.yaml, keyed by policy name. + """ + return dict(self._config_policies) + + def add_policy(self, policy_name: str, policy: Policy, source: Literal["db", "config"] = "db") -> None: """ Add or update a single policy. Args: policy_name: Name of the policy policy: Policy object to add + source: Provenance of the policy ("db" or "config") """ self._policies[policy_name] = policy + self._sources = {**self._sources, policy_name: source} + if source == "config": + self._config_policies = {**self._config_policies, policy_name: policy} self._initialized = True verbose_proxy_logger.debug(f"Added/updated policy: {policy_name}") @@ -325,6 +350,7 @@ class PolicyRegistry: """ if policy_name in self._policies: del self._policies[policy_name] + self._sources = {name: source for name, source in self._sources.items() if name != policy_name} verbose_proxy_logger.debug(f"Removed policy: {policy_name}") return True return False @@ -591,14 +617,14 @@ class PolicyRegistry: """ Sync policies from the database to in-memory registry. - Production versions are loaded into _policies (by policy name) for resolution. + - Config-loaded policies are preserved; on a name conflict the DB version wins. - 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: - self._policies = {} production = await self.get_all_policies_from_db(prisma_client, version_status="production") - for policy_response in production: - policy = self._parse_policy( + db_policies = { + policy_response.policy_name: self._parse_policy( policy_response.policy_name, { "inherit": policy_response.inherit, @@ -611,7 +637,16 @@ class PolicyRegistry: "pipeline": policy_response.pipeline, }, ) - self.add_policy(policy_response.policy_name, policy) + for policy_response in production + } + for policy_name in set(db_policies) & set(self._config_policies): + verbose_proxy_logger.warning( + f"Policy '{policy_name}' is defined in both config.yaml and the DB; the DB version takes precedence" + ) + config_sources: Mapping[str, Literal["db", "config"]] = {name: "config" for name in self._config_policies} + db_sources: Mapping[str, Literal["db", "config"]] = {name: "db" for name in db_policies} + self._policies = {**self._config_policies, **db_policies} + self._sources = {**config_sources, **db_sources} self._policies_by_id = {} non_production = await _policy_table(prisma_client).find_many( @@ -637,7 +672,8 @@ class PolicyRegistry: self._initialized = True verbose_proxy_logger.info( f"Synced {len(production)} production policies and {len(non_production)} " - "draft/published (by ID) from DB to in-memory registry" + "draft/published (by ID) from DB to in-memory registry " + f"({len(self._config_policies)} config-defined policies preserved)" ) except Exception as e: verbose_proxy_logger.exception(f"Error syncing policies from DB: {e}") diff --git a/litellm/types/proxy/policy_engine/resolver_types.py b/litellm/types/proxy/policy_engine/resolver_types.py index 2c7e8d5afc9..b4096cd2044 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.") + definition_location: Literal["db", "config"] = Field( + default="db", + description="Where this policy is defined: 'db' (database) or 'config' (config.yaml).", + ) 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.") + definition_location: Literal["db", "config"] = Field( + default="db", + description="Where this attachment is defined: 'db' (database) or 'config' (config.yaml).", + ) 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..cf470c66000 100644 --- a/tests/test_litellm/proxy/policy_engine/test_attachment_registry.py +++ b/tests/test_litellm/proxy/policy_engine/test_attachment_registry.py @@ -4,6 +4,9 @@ Unit tests for AttachmentRegistry - tests policy attachment matching. Tests the main entry point: get_attached_policies() """ +from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock + import pytest from litellm.proxy.policy_engine.attachment_registry import ( @@ -389,3 +392,64 @@ class TestAttachmentRegistrySingleton: registry1 = get_attachment_registry() registry2 = get_attachment_registry() assert registry1 is registry2 + + +def _make_db_attachment_row(attachment_id="att-1", policy_name="db-policy", scope=None, teams=None): + row = MagicMock() + row.attachment_id = attachment_id + row.policy_name = policy_name + row.scope = scope + row.teams = teams or [] + row.keys = [] + row.models = [] + row.tags = [] + 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_attachment_rows(rows): + prisma = MagicMock() + prisma.db.litellm_policyattachmenttable.find_many = AsyncMock(return_value=rows) + return prisma + + +class TestConfigAttachmentsPreservedAcrossDbSync: + """Config-defined attachments must survive sync_attachments_from_db (regression for issue #35255).""" + + @pytest.mark.asyncio + async def test_sync_with_empty_db_preserves_config_attachments(self): + registry = AttachmentRegistry() + registry.load_attachments([{"policy": "config-policy", "scope": "*"}]) + + await registry.sync_attachments_from_db(_prisma_with_attachment_rows([])) + + context = PolicyMatchContext(team_alias="any-team", key_alias="any-key", model="gpt-5.2") + assert registry.get_attached_policies(context) == ["config-policy"] + + @pytest.mark.asyncio + async def test_sync_merges_db_attachments_with_config_attachments(self): + registry = AttachmentRegistry() + registry.load_attachments([{"policy": "config-policy", "scope": "*"}]) + db_row = _make_db_attachment_row(policy_name="db-policy", teams=["db-team"]) + + await registry.sync_attachments_from_db(_prisma_with_attachment_rows([db_row])) + + assert len(registry.get_all_attachments()) == 2 + assert len(registry.get_config_attachments()) == 1 + context = PolicyMatchContext(team_alias="db-team", key_alias="k", model="gpt-5.2") + attached = registry.get_attached_policies(context) + assert "config-policy" in attached + assert "db-policy" in attached + + @pytest.mark.asyncio + async def test_repeated_syncs_do_not_duplicate_config_attachments(self): + registry = AttachmentRegistry() + registry.load_attachments([{"policy": "config-policy", "scope": "*"}]) + + await registry.sync_attachments_from_db(_prisma_with_attachment_rows([])) + await registry.sync_attachments_from_db(_prisma_with_attachment_rows([])) + + assert len(registry.get_all_attachments()) == 1 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..78126508c2d --- /dev/null +++ b/tests/test_litellm/proxy/policy_engine/test_policy_engine_endpoints.py @@ -0,0 +1,195 @@ +""" +Unit tests for policy_engine/policy_endpoints.py list endpoints. + +Regression tests for issue #35255: config-defined policies and attachments must be +returned by the list endpoints (marked definition_location="config"), DB rows must keep +their exact shape, and the endpoints must not 500 when no database is connected. +""" + +from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock + +import pytest + +import litellm.proxy.policy_engine.policy_endpoints as policy_endpoints +from litellm.proxy.policy_engine.attachment_registry import AttachmentRegistry +from litellm.proxy.policy_engine.policy_registry import PolicyRegistry + + +def _make_policy_row( + policy_id="uuid-1", + policy_name="db-policy", + version_status="production", + guardrails_add=None, +): + 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 = "db description" + row.guardrails_add = guardrails_add or [] + 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 = "admin" + row.updated_by = "admin" + return row + + +def _make_attachment_row(attachment_id="att-1", policy_name="db-policy", scope="*"): + row = MagicMock() + row.attachment_id = attachment_id + row.policy_name = policy_name + row.scope = scope + row.teams = [] + row.keys = [] + row.models = [] + row.tags = [] + row.created_at = datetime.now(timezone.utc) + row.updated_at = datetime.now(timezone.utc) + row.created_by = "admin" + row.updated_by = "admin" + return row + + +@pytest.fixture +def policy_registry(monkeypatch): + registry = PolicyRegistry() + monkeypatch.setattr(policy_endpoints, "get_policy_registry", lambda: registry) + return registry + + +@pytest.fixture +def attachment_registry(monkeypatch): + registry = AttachmentRegistry() + monkeypatch.setattr(policy_endpoints, "get_attachment_registry", lambda: registry) + return registry + + +def _set_prisma(monkeypatch, prisma): + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma) + + +class TestListPoliciesIncludesConfig: + @pytest.mark.asyncio + async def test_returns_config_policies_without_prisma(self, policy_registry, monkeypatch): + _set_prisma(monkeypatch, None) + policy_registry.load_policies( + {"config-policy": {"description": "from config", "guardrails": {"add": ["tooling"]}}} + ) + + response = await policy_endpoints.list_policies() + + assert response.total_count == 1 + entry = response.policies[0] + assert entry.policy_name == "config-policy" + assert entry.policy_id == "config-policy" + assert entry.definition_location == "config" + assert entry.version_status == "production" + assert entry.guardrails_add == ["tooling"] + assert entry.description == "from config" + assert entry.created_at is None + + @pytest.mark.asyncio + async def test_merges_db_rows_with_config_and_keeps_db_row_shape(self, policy_registry, monkeypatch): + row = _make_policy_row(policy_id="uuid-1", policy_name="db-policy", guardrails_add=["db-guard"]) + prisma = MagicMock() + prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[row]) + _set_prisma(monkeypatch, prisma) + policy_registry.load_policies({"config-policy": {"guardrails": {"add": ["tooling"]}}}) + + response = await policy_endpoints.list_policies() + + assert response.total_count == 2 + db_entry = next(p for p in response.policies if p.policy_name == "db-policy") + assert db_entry.definition_location == "db" + assert db_entry.policy_id == "uuid-1" + assert db_entry.guardrails_add == ["db-guard"] + assert db_entry.description == "db description" + assert db_entry.created_at == row.created_at + assert db_entry.created_by == "admin" + config_entry = next(p for p in response.policies if p.policy_name == "config-policy") + assert config_entry.definition_location == "config" + + @pytest.mark.asyncio + async def test_db_policy_shadows_config_policy_with_same_name(self, policy_registry, monkeypatch): + row = _make_policy_row(policy_id="uuid-1", policy_name="shared-name", guardrails_add=["db-guard"]) + prisma = MagicMock() + prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[row]) + _set_prisma(monkeypatch, prisma) + policy_registry.load_policies({"shared-name": {"guardrails": {"add": ["config-guard"]}}}) + + response = await policy_endpoints.list_policies() + + assert response.total_count == 1 + assert response.policies[0].definition_location == "db" + assert response.policies[0].guardrails_add == ["db-guard"] + + @pytest.mark.asyncio + async def test_version_status_filter_excludes_config_policies(self, policy_registry, monkeypatch): + row = _make_policy_row(policy_id="uuid-1", policy_name="db-policy", version_status="draft") + prisma = MagicMock() + prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[row]) + _set_prisma(monkeypatch, prisma) + policy_registry.load_policies({"config-policy": {"guardrails": {"add": ["tooling"]}}}) + + response = await policy_endpoints.list_policies(version_status="draft") + + assert response.total_count == 1 + assert response.policies[0].policy_name == "db-policy" + assert response.policies[0].definition_location == "db" + + @pytest.mark.asyncio + async def test_production_filter_includes_config_policies(self, policy_registry, monkeypatch): + _set_prisma(monkeypatch, None) + policy_registry.load_policies({"config-policy": {"guardrails": {"add": ["tooling"]}}}) + + response = await policy_endpoints.list_policies(version_status="production") + + assert response.total_count == 1 + assert response.policies[0].definition_location == "config" + + +class TestListAttachmentsIncludesConfig: + @pytest.mark.asyncio + async def test_returns_config_attachments_without_prisma(self, attachment_registry, monkeypatch): + _set_prisma(monkeypatch, None) + attachment_registry.load_attachments([{"policy": "config-policy", "scope": "*"}]) + + response = await policy_endpoints.list_policy_attachments() + + assert response.total_count == 1 + entry = response.attachments[0] + assert entry.attachment_id == "config-0" + assert entry.policy_name == "config-policy" + assert entry.scope == "*" + assert entry.definition_location == "config" + assert entry.created_at is None + + @pytest.mark.asyncio + async def test_merges_db_attachments_with_config_and_keeps_db_row_shape(self, attachment_registry, monkeypatch): + row = _make_attachment_row(attachment_id="att-1", policy_name="db-policy") + prisma = MagicMock() + prisma.db.litellm_policyattachmenttable.find_many = AsyncMock(return_value=[row]) + _set_prisma(monkeypatch, prisma) + attachment_registry.load_attachments([{"policy": "config-policy", "scope": "*"}]) + + response = await policy_endpoints.list_policy_attachments() + + assert response.total_count == 2 + db_entry = next(a for a in response.attachments if a.policy_name == "db-policy") + assert db_entry.attachment_id == "att-1" + assert db_entry.definition_location == "db" + assert db_entry.created_at == row.created_at + assert db_entry.created_by == "admin" + config_entry = next(a for a in response.attachments if a.policy_name == "config-policy") + assert config_entry.attachment_id == "config-0" + assert config_entry.definition_location == "config" 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 dd20021d0e1..5a840979c1b 100644 --- a/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py +++ b/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py @@ -450,3 +450,82 @@ class TestGetPolicyRegistrySingleton: a = get_policy_registry() b = get_policy_registry() assert a is b + + +def _prisma_with_policy_rows(production_rows, non_production_rows=None): + prisma = MagicMock() + prisma.db.litellm_policytable.find_many = AsyncMock(side_effect=[production_rows, non_production_rows or []]) + return prisma + + +class TestConfigPoliciesPreservedAcrossDbSync: + """Config-defined policies must survive sync_policies_from_db (regression for issue #35255).""" + + @pytest.mark.asyncio + async def test_sync_with_empty_db_preserves_config_policies(self): + registry = PolicyRegistry() + registry.load_policies({"config-policy": {"description": "from config", "guardrails": {"add": ["tooling"]}}}) + + await registry.sync_policies_from_db(_prisma_with_policy_rows([])) + + assert registry.has_policy("config-policy") + policy = registry.get_policy("config-policy") + assert policy is not None + assert policy.guardrails.add == ["tooling"] + assert registry.get_source("config-policy") == "config" + + @pytest.mark.asyncio + async def test_sync_merges_db_policies_with_config_policies(self): + registry = PolicyRegistry() + registry.load_policies({"config-policy": {"guardrails": {"add": ["tooling"]}}}) + db_row = _make_row(policy_id="db-1", policy_name="db-policy", guardrails_add=["db-guard"]) + + await registry.sync_policies_from_db(_prisma_with_policy_rows([db_row])) + + assert registry.has_policy("config-policy") + assert registry.has_policy("db-policy") + assert registry.get_source("config-policy") == "config" + assert registry.get_source("db-policy") == "db" + + @pytest.mark.asyncio + async def test_db_wins_on_policy_name_conflict(self): + registry = PolicyRegistry() + registry.load_policies({"shared-name": {"guardrails": {"add": ["config-guard"]}}}) + db_row = _make_row(policy_id="db-1", policy_name="shared-name", guardrails_add=["db-guard"]) + + await registry.sync_policies_from_db(_prisma_with_policy_rows([db_row])) + + policy = registry.get_policy("shared-name") + assert policy is not None + assert policy.guardrails.add == ["db-guard"] + assert registry.get_source("shared-name") == "db" + + @pytest.mark.asyncio + async def test_config_policy_restored_after_conflicting_db_row_deleted(self): + registry = PolicyRegistry() + registry.load_policies({"shared-name": {"guardrails": {"add": ["config-guard"]}}}) + db_row = _make_row(policy_id="db-1", policy_name="shared-name", guardrails_add=["db-guard"]) + + await registry.sync_policies_from_db(_prisma_with_policy_rows([db_row])) + await registry.sync_policies_from_db(_prisma_with_policy_rows([])) + + policy = registry.get_policy("shared-name") + assert policy is not None + assert policy.guardrails.add == ["config-guard"] + assert registry.get_source("shared-name") == "config" + + @pytest.mark.asyncio + async def test_config_policy_resolves_guardrails_after_sync(self): + from litellm.proxy.policy_engine.policy_resolver import PolicyResolver + + registry = PolicyRegistry() + registry.load_policies({"config-policy": {"guardrails": {"add": ["tooling"]}}}) + + await registry.sync_policies_from_db(_prisma_with_policy_rows([])) + + resolved = PolicyResolver.resolve_policy_guardrails( + policy_name="config-policy", + policies=registry.get_all_policies(), + context=None, + ) + assert resolved.guardrails == ["tooling"] diff --git a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/AttachmentTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/AttachmentTableColumns.tsx index 9e0f8d6715d..ded9e3a1e6d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/AttachmentTableColumns.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/AttachmentTableColumns.tsx @@ -41,7 +41,12 @@ interface AttachmentRowActionsProps { onDeleteClick: (attachmentId: string) => void; } +const CONFIG_ATTACHMENT_HINT = + "Config attachments are defined in the config file and cannot be deleted from the dashboard."; + function AttachmentRowActions({ attachment, isAdmin, onDeleteClick }: AttachmentRowActionsProps) { + const isConfigAttachment = attachment.definition_location === "config"; + return ( onDeleteClick(attachment.attachment_id)} > diff --git a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTableColumns.tsx index dd036d83283..488de728bad 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTableColumns.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTableColumns.tsx @@ -22,6 +22,9 @@ export interface PolicyRow { versionCount: number; } +const CONFIG_POLICY_HINT = + "Config policies are defined in the config file and cannot be edited or deleted from the dashboard."; + function GuardrailChips({ guardrails, tone }: { guardrails: string[]; tone: "success" | "error" }) { if (guardrails.length === 0) { return -; @@ -45,6 +48,8 @@ interface PolicyRowActionsProps { } function PolicyRowActions({ policy, onEditClick, onDeleteClick }: PolicyRowActionsProps) { + const isConfigPolicy = policy.definition_location === "config"; + return ( - onEditClick(policy)}> + onEditClick(policy)} + > Edit policy @@ -63,6 +73,8 @@ function PolicyRowActions({ policy, onEditClick, onDeleteClick }: PolicyRowActio onDeleteClick(policy.policy_id, policy.policy_name || "Unnamed Policy")} > @@ -93,18 +105,23 @@ export const getPolicyTableColumns = ({ header: ({ column }) => , size: 220, enableSorting: true, - cell: ({ row }) => ( - 1 ? ( - - ) : undefined - } - onClick={() => onViewClick(row.original.primaryPolicy.policy_id)} - /> - ), + cell: ({ row }) => { + const isConfigPolicy = row.original.primaryPolicy.definition_location === "config"; + const versionBadge = + row.original.versionCount > 1 ? ( + + ) : undefined; + return ( + : versionBadge + } + onClick={isConfigPolicy ? undefined : () => onViewClick(row.original.primaryPolicy.policy_id)} + /> + ); + }, }, { id: "description", diff --git a/ui/litellm-dashboard/src/components/policies/types.ts b/ui/litellm-dashboard/src/components/policies/types.ts index 887781ff943..6ac110e3c0a 100644 --- a/ui/litellm-dashboard/src/components/policies/types.ts +++ b/ui/litellm-dashboard/src/components/policies/types.ts @@ -14,6 +14,7 @@ export interface Policy { updated_at?: string; created_by?: string; updated_by?: string; + definition_location?: "db" | "config"; } export interface PolicyCondition { @@ -47,6 +48,7 @@ export interface PolicyAttachment { updated_at?: string; created_by?: string; updated_by?: string; + definition_location?: "db" | "config"; } export interface PolicyCreateRequest { diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index ed975c6be0a..6543e437c07 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -9379,7 +9379,10 @@ export interface paths { }; /** * List Policy Attachments - * @description List all policy attachments from the database. + * @description List all policy attachments from the database and config.yaml. + * + * Config-defined attachments are returned with definition_location "config" and a + * synthetic attachment_id ("config-"). * * Example Request: * ```bash @@ -9487,7 +9490,10 @@ export interface paths { }; /** * List Policies - * @description List all policies from the database. Optionally filter by version_status. + * @description List all policies from the database and config.yaml. Optionally filter by version_status. + * + * Config-defined policies are returned with definition_location "config" and are treated + * as production versions. On a name conflict with a DB policy, only the DB policy is returned. * * Query params: * - version_status: Optional. One of "draft", "published", "production". @@ -29379,6 +29385,13 @@ export interface components { * @description Who created the attachment. */ created_by?: string | null; + /** + * Definition Location + * @description Where this attachment is defined: 'db' (database) or 'config' (config.yaml). + * @default db + * @enum {string} + */ + definition_location: "db" | "config"; /** * Keys * @description Key patterns. @@ -29510,6 +29523,13 @@ export interface components { * @description Who created the policy. */ created_by?: string | null; + /** + * Definition Location + * @description Where this policy is defined: 'db' (database) or 'config' (config.yaml). + * @default db + * @enum {string} + */ + definition_location: "db" | "config"; /** * Description * @description Policy description. @@ -33117,6 +33137,17 @@ export interface components { */ model?: string | null; }; + /** UsageChartPoint */ + UsageChartPoint: { + /** Blocked */ + blocked: number; + /** Date */ + date: string; + /** Passed */ + passed: number; + /** Score */ + score?: number | null; + }; /** UsageDetailResponse */ UsageDetailResponse: { /** Avglatency */ From 91290c60206e1482bfc8e945ea456161eb935b31 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 30 Jul 2026 15:33:52 -0700 Subject: [PATCH 2/7] fix(policy_engine): only hide config policies behind production DB versions in policies list --- .../proxy/policy_engine/policy_endpoints.py | 9 +++++-- .../policy_engine/test_attachment_registry.py | 11 ++++++++ .../test_policy_engine_endpoints.py | 26 ++++++++++++++++++ .../policy_engine/test_policy_versioning.py | 27 +++++++++++++++++++ 4 files changed, 71 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/policy_engine/policy_endpoints.py b/litellm/proxy/policy_engine/policy_endpoints.py index 787f7069996..718223da5d8 100644 --- a/litellm/proxy/policy_engine/policy_endpoints.py +++ b/litellm/proxy/policy_engine/policy_endpoints.py @@ -80,7 +80,10 @@ async def list_policies(version_status: Optional[str] = None): List all policies from the database and config.yaml. Optionally filter by version_status. Config-defined policies are returned with definition_location "config" and are treated - as production versions. On a name conflict with a DB policy, only the DB policy is returned. + as production versions. On a name conflict with a production DB policy, only the DB policy + is returned, mirroring runtime resolution where only production DB versions override config. + A draft or published DB version does not hide the config policy, since the config version + is still the one being enforced. Query params: - version_status: Optional. One of "draft", "published", "production". @@ -125,7 +128,9 @@ async def list_policies(version_status: Optional[str] = None): if prisma_client is not None else [] ) - db_policy_names = {db_policy.policy_name for db_policy in db_policies} + db_policy_names = { + db_policy.policy_name for db_policy in db_policies if db_policy.version_status == "production" + } include_config = version_status in (None, "production") config_policies = ( [ 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 cf470c66000..cc231e383a3 100644 --- a/tests/test_litellm/proxy/policy_engine/test_attachment_registry.py +++ b/tests/test_litellm/proxy/policy_engine/test_attachment_registry.py @@ -453,3 +453,14 @@ class TestConfigAttachmentsPreservedAcrossDbSync: await registry.sync_attachments_from_db(_prisma_with_attachment_rows([])) assert len(registry.get_all_attachments()) == 1 + + @pytest.mark.asyncio + async def test_clear_removes_config_snapshot_so_sync_does_not_resurrect(self): + registry = AttachmentRegistry() + registry.load_attachments([{"policy": "config-policy", "scope": "*"}]) + + registry.clear() + await registry.sync_attachments_from_db(_prisma_with_attachment_rows([])) + + assert registry.get_all_attachments() == [] + assert registry.get_config_attachments() == () 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 index 78126508c2d..9c486540b3d 100644 --- a/tests/test_litellm/proxy/policy_engine/test_policy_engine_endpoints.py +++ b/tests/test_litellm/proxy/policy_engine/test_policy_engine_endpoints.py @@ -133,6 +133,32 @@ class TestListPoliciesIncludesConfig: assert response.policies[0].definition_location == "db" assert response.policies[0].guardrails_add == ["db-guard"] + @pytest.mark.asyncio + async def test_draft_db_policy_does_not_hide_enforced_config_policy(self, policy_registry, monkeypatch): + """ + Runtime sync only lets production DB versions override a config policy, + so a draft or published DB version sharing the name must not suppress + the config entry: the config version is still the one being enforced, + and hiding it makes the list API disagree with actual enforcement. + """ + row = _make_policy_row( + policy_id="uuid-1", policy_name="shared-name", version_status="draft", guardrails_add=["db-guard"] + ) + prisma = MagicMock() + prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[row]) + _set_prisma(monkeypatch, prisma) + policy_registry.load_policies({"shared-name": {"guardrails": {"add": ["config-guard"]}}}) + + response = await policy_endpoints.list_policies() + + assert response.total_count == 2 + config_entry = next(p for p in response.policies if p.definition_location == "config") + assert config_entry.policy_name == "shared-name" + assert config_entry.version_status == "production" + assert config_entry.guardrails_add == ["config-guard"] + db_entry = next(p for p in response.policies if p.definition_location == "db") + assert db_entry.version_status == "draft" + @pytest.mark.asyncio async def test_version_status_filter_excludes_config_policies(self, policy_registry, monkeypatch): row = _make_policy_row(policy_id="uuid-1", policy_name="db-policy", version_status="draft") 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 5a840979c1b..ae80373ee23 100644 --- a/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py +++ b/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py @@ -13,8 +13,10 @@ from litellm.proxy.policy_engine.policy_registry import ( get_policy_registry, ) from litellm.types.proxy.policy_engine import ( + Policy, PolicyCreateRequest, PolicyDBResponse, + PolicyGuardrails, PolicyUpdateRequest, ) @@ -529,3 +531,28 @@ class TestConfigPoliciesPreservedAcrossDbSync: context=None, ) assert resolved.guardrails == ["tooling"] + + @pytest.mark.asyncio + async def test_add_policy_with_config_source_survives_sync(self): + registry = PolicyRegistry() + registry.add_policy( + "late-config-policy", + Policy(guardrails=PolicyGuardrails(add=["tooling"])), + source="config", + ) + + await registry.sync_policies_from_db(_prisma_with_policy_rows([])) + + assert registry.has_policy("late-config-policy") + assert registry.get_source("late-config-policy") == "config" + + @pytest.mark.asyncio + async def test_clear_removes_config_snapshot_so_sync_does_not_resurrect(self): + registry = PolicyRegistry() + registry.load_policies({"config-policy": {"guardrails": {"add": ["tooling"]}}}) + + registry.clear() + await registry.sync_policies_from_db(_prisma_with_policy_rows([])) + + assert not registry.has_policy("config-policy") + assert registry.get_source("config-policy") is None From ec016d1bd86664112561bd1c7f5fd5a0009ada2e Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 30 Jul 2026 16:01:57 -0700 Subject: [PATCH 3/7] fix(policy_engine): decide config policy suppression from fresh db query only --- .../proxy/policy_engine/policy_endpoints.py | 2 +- .../test_policy_engine_endpoints.py | 27 +++++++++++++++++++ 2 files changed, 28 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/policy_engine/policy_endpoints.py b/litellm/proxy/policy_engine/policy_endpoints.py index 718223da5d8..cff1378c676 100644 --- a/litellm/proxy/policy_engine/policy_endpoints.py +++ b/litellm/proxy/policy_engine/policy_endpoints.py @@ -136,7 +136,7 @@ async def list_policies(version_status: Optional[str] = None): [ _config_policy_to_db_response(policy_name, policy) for policy_name, policy in registry.list_config_policies().items() - if policy_name not in db_policy_names and registry.get_source(policy_name) != "db" + if policy_name not in db_policy_names ] if include_config else [] 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 index 9c486540b3d..1ca830dc1e6 100644 --- a/tests/test_litellm/proxy/policy_engine/test_policy_engine_endpoints.py +++ b/tests/test_litellm/proxy/policy_engine/test_policy_engine_endpoints.py @@ -159,6 +159,33 @@ class TestListPoliciesIncludesConfig: db_entry = next(p for p in response.policies if p.definition_location == "db") assert db_entry.version_status == "draft" + @pytest.mark.asyncio + async def test_stale_registry_provenance_does_not_hide_config_policy(self, policy_registry, monkeypatch): + """ + Another proxy instance can delete or demote the production DB override + between registry syncs. The endpoint's fresh DB query is the source of + truth for conflicts; stale in-memory provenance from the last sync must + not suppress the config entry once no production override exists. + """ + policy_registry.load_policies({"shared-name": {"guardrails": {"add": ["config-guard"]}}}) + production_row = _make_policy_row(policy_id="uuid-1", policy_name="shared-name", guardrails_add=["db-guard"]) + sync_prisma = MagicMock() + sync_prisma.db.litellm_policytable.find_many = AsyncMock(side_effect=[[production_row], []]) + await policy_registry.sync_policies_from_db(sync_prisma) + assert policy_registry.get_source("shared-name") == "db" + + fresh_prisma = MagicMock() + fresh_prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[]) + _set_prisma(monkeypatch, fresh_prisma) + + response = await policy_endpoints.list_policies() + + assert response.total_count == 1 + entry = response.policies[0] + assert entry.policy_name == "shared-name" + assert entry.definition_location == "config" + assert entry.guardrails_add == ["config-guard"] + @pytest.mark.asyncio async def test_version_status_filter_excludes_config_policies(self, policy_registry, monkeypatch): row = _make_policy_row(policy_id="uuid-1", policy_name="db-policy", version_status="draft") From 35f770f43e10b87483ab28d9f3ad0c67cb0021fd Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 30 Jul 2026 19:23:50 -0700 Subject: [PATCH 4/7] fix(policy_engine): restore config policy immediately when its DB override is removed and keep same-named DB drafts reachable in the UI --- .../proxy/policy_engine/policy_registry.py | 32 +++++++--- .../policy_engine/test_policy_versioning.py | 63 +++++++++++++++++++ .../policies/_components/PolicyTable.test.tsx | 30 +++++++++ .../policies/_components/PolicyTable.tsx | 15 +++-- 4 files changed, 125 insertions(+), 15 deletions(-) diff --git a/litellm/proxy/policy_engine/policy_registry.py b/litellm/proxy/policy_engine/policy_registry.py index 9456afaa1b9..7d2b439270a 100644 --- a/litellm/proxy/policy_engine/policy_registry.py +++ b/litellm/proxy/policy_engine/policy_registry.py @@ -340,7 +340,8 @@ class PolicyRegistry: def remove_policy(self, policy_name: str) -> bool: """ - Remove a policy by name. + Remove a policy by name. If a config-defined policy shares the name, + it is restored immediately instead of waiting for the next DB sync. Args: policy_name: Name of the policy to remove @@ -348,12 +349,18 @@ class PolicyRegistry: Returns: True if policy was removed, False if it didn't exist """ - if policy_name in self._policies: - del self._policies[policy_name] - self._sources = {name: source for name, source in self._sources.items() if name != policy_name} - verbose_proxy_logger.debug(f"Removed policy: {policy_name}") + if policy_name not in self._policies: + return False + config_fallback = self._config_policies.get(policy_name) + if config_fallback is not None: + self._policies[policy_name] = config_fallback + self._sources = {**self._sources, policy_name: "config"} + verbose_proxy_logger.debug(f"Removed policy: {policy_name}; restored config-defined version") return True - return False + del self._policies[policy_name] + self._sources = {name: source for name, source in self._sources.items() if name != policy_name} + verbose_proxy_logger.debug(f"Removed policy: {policy_name}") + return True # ───────────────────────────────────────────────────────────────────────── # Database CRUD Methods @@ -527,10 +534,15 @@ class PolicyRegistry: # 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." - ) + if self.get_source(policy_name) == "config": + result["warning"] = ( + "Production version was deleted. The config-defined policy with the same name is active again." + ) + else: + 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: 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 ae80373ee23..41d856e7baf 100644 --- a/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py +++ b/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py @@ -556,3 +556,66 @@ class TestConfigPoliciesPreservedAcrossDbSync: assert not registry.has_policy("config-policy") assert registry.get_source("config-policy") is None + + +class TestRemovePolicyRestoresConfigFallback: + """Deleting a same-named DB override must re-activate the config policy immediately, not at the next sync.""" + + def test_remove_policy_restores_config_version_immediately(self): + registry = PolicyRegistry() + registry.load_policies({"shared-name": {"guardrails": {"add": ["config-guard"]}}}) + registry.add_policy("shared-name", Policy(guardrails=PolicyGuardrails(add=["db-guard"])), source="db") + + assert registry.remove_policy("shared-name") is True + + policy = registry.get_policy("shared-name") + assert policy is not None + assert policy.guardrails.add == ["config-guard"] + assert registry.get_source("shared-name") == "config" + + def test_remove_policy_without_config_fallback_removes_entirely(self): + registry = PolicyRegistry() + registry.add_policy("db-only", Policy(guardrails=PolicyGuardrails(add=["db-guard"]))) + + assert registry.remove_policy("db-only") is True + + assert not registry.has_policy("db-only") + assert registry.get_source("db-only") is None + + def test_remove_missing_policy_returns_false(self): + registry = PolicyRegistry() + + assert registry.remove_policy("missing") is False + + @pytest.mark.asyncio + async def test_delete_production_override_reactivates_config_policy_and_says_so(self): + registry = PolicyRegistry() + registry.load_policies({"shared-name": {"guardrails": {"add": ["config-guard"]}}}) + registry.add_policy("shared-name", Policy(guardrails=PolicyGuardrails(add=["db-guard"])), source="db") + prisma = MagicMock() + prod_row = _make_row(policy_id="prod-1", policy_name="shared-name", 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 "config" in result["warning"] + policy = registry.get_policy("shared-name") + assert policy is not None + assert policy.guardrails.add == ["config-guard"] + assert registry.get_source("shared-name") == "config" + + @pytest.mark.asyncio + async def test_delete_all_versions_reactivates_config_policy(self): + registry = PolicyRegistry() + registry.load_policies({"shared-name": {"guardrails": {"add": ["config-guard"]}}}) + registry.add_policy("shared-name", Policy(guardrails=PolicyGuardrails(add=["db-guard"])), source="db") + prisma = MagicMock() + prisma.db.litellm_policytable.delete_many = AsyncMock() + + await registry.delete_all_versions(policy_name="shared-name", prisma_client=prisma) + + assert registry.get_source("shared-name") == "config" + policy = registry.get_policy("shared-name") + assert policy is not None + assert policy.guardrails.add == ["config-guard"] diff --git a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTable.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTable.test.tsx index 06c939aa151..9be1bd60ec8 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTable.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTable.test.tsx @@ -145,4 +145,34 @@ describe("PolicyTable", () => { await user.click(screen.getByRole("button", { name: /grouped/ })); expect(defaultProps.onViewClick).toHaveBeenCalledWith("prod-id"); }); + + const sameNamedDbDraft: Partial = { + policy_name: "config-policy", + policy_id: "db-draft-id", + version_status: "draft", + version_number: 2, + }; + const configTwin: Partial = { + policy_name: "config-policy", + policy_id: "config-policy", + version_status: "production", + definition_location: "config", + }; + + it("should render a config policy and a same-named DB draft as separate rows", () => { + const policies = [makePolicy(sameNamedDbDraft), makePolicy(configTwin)]; + renderWithProviders(); + expect(screen.getAllByText("config-policy")).toHaveLength(2); + expect(screen.getByText("Config")).toBeInTheDocument(); + }); + + it("should keep a same-named DB draft reachable next to a config policy", async () => { + const user = userEvent.setup(); + const policies = [makePolicy(sameNamedDbDraft), makePolicy(configTwin)]; + renderWithProviders(); + await user.click(screen.getByRole("button", { name: "config-policy" })); + expect(defaultProps.onViewClick).toHaveBeenCalledWith("db-draft-id"); + await user.click(screen.getByTestId("policy-actions-db-draft-id")); + expect(await screen.findByTestId("policy-action-edit")).not.toHaveAttribute("data-disabled"); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTable.tsx index d6e841c2119..3405ac6b6bb 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTable.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTable.tsx @@ -9,16 +9,21 @@ import { Policy } from "@/components/policies/types"; import { getPolicyTableColumns, PolicyRow } from "./PolicyTableColumns"; -/** One row per policy name; primaryPolicy is used for display and for Edit (FlowBuilder loads all versions) */ +/** One row per DB policy name plus one row per config policy, so a config policy never hides same-named DB versions; primaryPolicy is used for display and for Edit (FlowBuilder loads all versions) */ function groupPoliciesByName(policies: Policy[]): PolicyRow[] { - const names = Array.from(new Set(policies.map((policy) => policy.policy_name || "(unnamed)"))); - return names.map((policyName) => { - const versions = policies.filter((policy) => (policy.policy_name || "(unnamed)") === policyName); + const dbPolicies = policies.filter((policy) => policy.definition_location !== "config"); + const names = Array.from(new Set(dbPolicies.map((policy) => policy.policy_name || "(unnamed)"))); + const dbRows = names.map((policyName) => { + const versions = dbPolicies.filter((policy) => (policy.policy_name || "(unnamed)") === policyName); const primary = versions.find((version) => version.version_status === "production") ?? [...versions].sort((a, b) => (b.version_number ?? 0) - (a.version_number ?? 0))[0]; return { policy_name: policyName, primaryPolicy: primary, versionCount: versions.length }; }); + const configRows = policies + .filter((policy) => policy.definition_location === "config") + .map((policy) => ({ policy_name: policy.policy_name || "(unnamed)", primaryPolicy: policy, versionCount: 1 })); + return [...dbRows, ...configRows]; } interface PolicyTableProps { @@ -67,7 +72,7 @@ const PolicyTable: React.FC = ({ row.policy_name} + getRowId={(row) => `${row.primaryPolicy.definition_location ?? "db"}:${row.policy_name}`} sortingMode="client" sorting={sorting} onSortingChange={setSorting} From b42ef469cfb3a87096c455e446bed4d4ff5dd43e Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 30 Jul 2026 19:48:35 -0700 Subject: [PATCH 5/7] fix(policy_engine): warn that the config-defined policy reactivates when all DB versions are deleted --- litellm/proxy/policy_engine/policy_registry.py | 12 ++++++++++-- .../proxy/policy_engine/test_policy_versioning.py | 14 +++++++++++++- 2 files changed, 23 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/policy_engine/policy_registry.py b/litellm/proxy/policy_engine/policy_registry.py index 7d2b439270a..01b88836387 100644 --- a/litellm/proxy/policy_engine/policy_registry.py +++ b/litellm/proxy/policy_engine/policy_registry.py @@ -1031,12 +1031,20 @@ class PolicyRegistry: prisma_client: The Prisma client instance Returns: - Dict with success message + Dict with "message" and optional "warning" if a config-defined policy took over. """ try: await _policy_table(prisma_client).delete_many(where={"policy_name": policy_name}) self.remove_policy(policy_name) - return {"message": f"All versions of policy '{policy_name}' deleted successfully"} + message = f"All versions of policy '{policy_name}' deleted successfully" + if self.get_source(policy_name) == "config": + return { + "message": message, + "warning": ( + "All DB versions were deleted. The config-defined policy with the same name is active again." + ), + } + return {"message": message} except Exception as e: verbose_proxy_logger.exception(f"Error deleting all versions: {e}") raise Exception(f"Error deleting all versions: {str(e)}") 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 41d856e7baf..ebebfde5cd3 100644 --- a/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py +++ b/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py @@ -613,9 +613,21 @@ class TestRemovePolicyRestoresConfigFallback: prisma = MagicMock() prisma.db.litellm_policytable.delete_many = AsyncMock() - await registry.delete_all_versions(policy_name="shared-name", prisma_client=prisma) + result = await registry.delete_all_versions(policy_name="shared-name", prisma_client=prisma) assert registry.get_source("shared-name") == "config" policy = registry.get_policy("shared-name") assert policy is not None assert policy.guardrails.add == ["config-guard"] + assert "config" in result["warning"] + + async def test_delete_all_versions_without_config_twin_has_no_warning(self): + registry = PolicyRegistry() + registry.add_policy("db-only", Policy(guardrails=PolicyGuardrails(add=["db-guard"])), source="db") + prisma = MagicMock() + prisma.db.litellm_policytable.delete_many = AsyncMock() + + result = await registry.delete_all_versions(policy_name="db-only", prisma_client=prisma) + + assert registry.get_policy("db-only") is None + assert "warning" not in result From 47ebc964eb8e3ed266a3d7a5473299050d262a36 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 30 Jul 2026 20:55:32 -0700 Subject: [PATCH 6/7] test: patch unified guardrail mapping global instead of loader to fix order-dependent flake --- .../test_passthrough_post_call_guardrails.py | 25 +++++++++++-------- 1 file changed, 14 insertions(+), 11 deletions(-) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py index a2f7476abd1..470179a0429 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py @@ -294,19 +294,22 @@ class TestUnifiedGuardrailCallTypeResolution: response_body = {"candidates": [{"content": {"parts": [{"text": "hello"}]}}]} - with patch( - "litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail.load_guardrail_translation_mappings" - ) as mock_load: - mock_handler_instance = AsyncMock() - mock_handler_instance.process_output_response = AsyncMock( - return_value=response_body - ) - mock_handler_class = MagicMock(return_value=mock_handler_instance) + mock_handler_instance = AsyncMock() + mock_handler_instance.process_output_response = AsyncMock( + return_value=response_body + ) + mock_handler_class = MagicMock(return_value=mock_handler_instance) - from litellm.types.utils import CallTypes - - mock_load.return_value = {CallTypes.pass_through: mock_handler_class} + from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail import ( + unified_guardrail as unified_guardrail_module, + ) + from litellm.types.utils import CallTypes + with patch.object( + unified_guardrail_module, + "endpoint_guardrail_translation_mappings", + {CallTypes.pass_through: mock_handler_class}, + ): result = await unified.async_post_call_success_hook( data=data, user_api_key_dict=user_api_key_dict, From ed21c2e3023df92556004d1a0852d9b087d9e2a2 Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Thu, 30 Jul 2026 21:18:33 -0700 Subject: [PATCH 7/7] feat(s3): support SSE-KMS encryption params on both S3 logging paths (#35291) * feat(s3): support SSE-KMS encryption params on both S3 logging paths * fix(s3): ignore non-string SSE config values instead of crashing logger init * Update litellm/integrations/s3.py Co-authored-by: devin-ai-integration[bot] <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(s3): invalidate only the mistyped SSE field instead of dropping both --------- Co-authored-by: devin-ai-integration[bot] <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/integrations/s3.py | 44 +++ litellm/integrations/s3_v2.py | 30 ++- tests/test_litellm/integrations/test_s3.py | 156 +++++++++++ tests/test_litellm/integrations/test_s3_v2.py | 250 ++++++++++++++++++ 4 files changed, 468 insertions(+), 12 deletions(-) create mode 100644 tests/test_litellm/integrations/test_s3.py diff --git a/litellm/integrations/s3.py b/litellm/integrations/s3.py index e8252d87572..07bd957b5a3 100644 --- a/litellm/integrations/s3.py +++ b/litellm/integrations/s3.py @@ -24,6 +24,8 @@ class S3Logger: s3_aws_secret_access_key=None, s3_aws_session_token=None, s3_config=None, + s3_server_side_encryption: str | None = None, + s3_sse_kms_key_id: str | None = None, **kwargs, ): import boto3 @@ -50,11 +52,16 @@ class S3Logger: s3_aws_session_token = litellm.s3_callback_params.get("s3_aws_session_token") s3_config = litellm.s3_callback_params.get("s3_config") s3_path = litellm.s3_callback_params.get("s3_path") + s3_server_side_encryption = litellm.s3_callback_params.get("s3_server_side_encryption") + s3_sse_kms_key_id = litellm.s3_callback_params.get("s3_sse_kms_key_id") # done reading litellm.s3_callback_params s3_use_team_prefix = bool(litellm.s3_callback_params.get("s3_use_team_prefix", False)) self.s3_use_team_prefix = s3_use_team_prefix self.bucket_name = s3_bucket_name self.s3_path = s3_path + self.s3_server_side_encryption, self.s3_sse_kms_key_id = resolve_sse_params( + s3_server_side_encryption, s3_sse_kms_key_id + ) verbose_logger.debug(f"s3 logger using endpoint url {s3_endpoint_url}") # Create an S3 client with custom endpoint URL self.s3_client = boto3.client( @@ -136,6 +143,15 @@ class S3Logger: print_verbose(f"\ns3 Logger - Logging payload = {payload_str}") + sse_params = { + key: value + for key, value in { + "ServerSideEncryption": self.s3_server_side_encryption, + "SSEKMSKeyId": self.s3_sse_kms_key_id, + }.items() + if value + } + response = self.s3_client.put_object( Bucket=self.bucket_name, Key=s3_object_key, @@ -144,6 +160,7 @@ class S3Logger: ContentLanguage="en", ContentDisposition=f'inline; filename="{s3_object_download_filename}"', CacheControl="private, immutable, max-age=31536000, s-maxage=0", + **sse_params, ) print_verbose(f"Response from s3:{str(response)}") @@ -155,6 +172,33 @@ class S3Logger: pass +def _validated_sse_value(name: str, value: str | None) -> str | None: + if value is None or isinstance(value, str): + return value + verbose_logger.warning( + f"s3 logging: ignoring {name} because it has invalid type {type(value).__name__}; expected a string" + ) + return None + + +def resolve_sse_params( + server_side_encryption: str | None, + sse_kms_key_id: str | None, +) -> tuple[str | None, str | None]: + valid_sse = _validated_sse_value("s3_server_side_encryption", server_side_encryption) + valid_key_id = _validated_sse_value("s3_sse_kms_key_id", sse_kms_key_id) + algorithm = valid_sse or ("aws:kms" if valid_key_id else None) + if algorithm is None: + return None, None + if valid_key_id and not algorithm.startswith("aws:kms"): + verbose_logger.warning( + f"s3 logging: ignoring s3_sse_kms_key_id because s3_server_side_encryption is {algorithm}; " + "set it to aws:kms to encrypt with the KMS key" + ) + return algorithm, None + return algorithm, valid_key_id + + def get_s3_object_key( s3_path: str, prefix: str, diff --git a/litellm/integrations/s3_v2.py b/litellm/integrations/s3_v2.py index 5b953035cfd..7fa78f39460 100644 --- a/litellm/integrations/s3_v2.py +++ b/litellm/integrations/s3_v2.py @@ -8,13 +8,14 @@ NOTE 1: S3 does not provide a BATCH PUT API endpoint, so we create tasks to uplo import asyncio import time +from collections.abc import Mapping from datetime import datetime from typing import List, Optional, cast import litellm from litellm._logging import print_verbose, verbose_logger from litellm.constants import DEFAULT_S3_BATCH_SIZE, DEFAULT_S3_FLUSH_INTERVAL_SECONDS -from litellm.integrations.s3 import get_s3_object_key +from litellm.integrations.s3 import get_s3_object_key, resolve_sse_params from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM @@ -55,6 +56,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): s3_use_key_prefix: bool = False, s3_use_virtual_hosted_style: bool = False, s3_server_side_encryption: Optional[str] = None, + s3_sse_kms_key_id: str | None = None, s3_callback_params_override: Optional[dict] = None, **kwargs, ): @@ -94,6 +96,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): s3_use_key_prefix=s3_use_key_prefix, s3_use_virtual_hosted_style=s3_use_virtual_hosted_style, s3_server_side_encryption=s3_server_side_encryption, + s3_sse_kms_key_id=s3_sse_kms_key_id, ) verbose_logger.debug(f"s3 logger using endpoint url {s3_endpoint_url}") @@ -148,6 +151,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): s3_use_key_prefix: bool = False, s3_use_virtual_hosted_style: bool = False, s3_server_side_encryption: Optional[str] = None, + s3_sse_kms_key_id: str | None = None, params_source: Optional[dict] = None, ): """ @@ -197,10 +201,20 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): bool(params.get("s3_use_virtual_hosted_style", False)) or s3_use_virtual_hosted_style ) - self.s3_server_side_encryption = params.get("s3_server_side_encryption") or s3_server_side_encryption + self.s3_server_side_encryption, self.s3_sse_kms_key_id = resolve_sse_params( + params.get("s3_server_side_encryption") or s3_server_side_encryption, + params.get("s3_sse_kms_key_id") or s3_sse_kms_key_id, + ) return + def _sse_headers(self) -> Mapping[str, str]: + candidates = { + "x-amz-server-side-encryption": self.s3_server_side_encryption, + "x-amz-server-side-encryption-aws-kms-key-id": self.s3_sse_kms_key_id, + } + return {key: value for key, value in candidates.items() if value} + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): await self._async_log_event_base( kwargs=kwargs, @@ -335,11 +349,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): "Content-Language": "en", "Content-Disposition": f'inline; filename="{batch_logging_element.s3_object_download_filename}"', "Cache-Control": "private, immutable, max-age=31536000, s-maxage=0", - **( - {"x-amz-server-side-encryption": self.s3_server_side_encryption} - if self.s3_server_side_encryption - else {} - ), + **self._sse_headers(), } req = requests.Request("PUT", url, data=json_string, headers=headers) prepped = req.prepare() @@ -510,11 +520,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): "Content-Language": "en", "Content-Disposition": f'inline; filename="{batch_logging_element.s3_object_download_filename}"', "Cache-Control": "private, immutable, max-age=31536000, s-maxage=0", - **( - {"x-amz-server-side-encryption": self.s3_server_side_encryption} - if self.s3_server_side_encryption - else {} - ), + **self._sse_headers(), } req = requests.Request("PUT", url, data=json_string, headers=headers) prepped = req.prepare() diff --git a/tests/test_litellm/integrations/test_s3.py b/tests/test_litellm/integrations/test_s3.py new file mode 100644 index 00000000000..7e997870852 --- /dev/null +++ b/tests/test_litellm/integrations/test_s3.py @@ -0,0 +1,156 @@ +from datetime import datetime +from unittest.mock import MagicMock, patch + +import litellm +from litellm.integrations.s3 import S3Logger + +TEST_KMS_KEY_ARN = "arn:aws:kms:us-east-1:111122223333:key/test-key-id" + + +def _standard_logging_payload() -> dict: + return { + "id": "chatcmpl-test-id", + "metadata": {"user_api_key_team_alias": None}, + } + + +def _log_event_kwargs() -> dict: + return { + "litellm_params": {"metadata": {}}, + "standard_logging_object": _standard_logging_payload(), + } + + +def _run_log_event(callback_params: dict) -> MagicMock: + original = litellm.s3_callback_params + litellm.s3_callback_params = callback_params + try: + with patch("boto3.client") as mock_boto3_client: + mock_s3_client = MagicMock() + mock_boto3_client.return_value = mock_s3_client + logger = S3Logger() + logger.log_event( + kwargs=_log_event_kwargs(), + response_obj={}, + start_time=datetime(2026, 7, 30, 12, 0, 0), + end_time=datetime(2026, 7, 30, 12, 0, 1), + print_verbose=lambda *args, **kwargs: None, + ) + return mock_s3_client + finally: + litellm.s3_callback_params = original + + +def test_put_object_includes_sse_kms_params_when_configured(): + """ + When s3_server_side_encryption and s3_sse_kms_key_id are set in + s3_callback_params, put_object must receive ServerSideEncryption and + SSEKMSKeyId so objects land encrypted with the customer-managed key. + """ + mock_s3_client = _run_log_event( + { + "s3_bucket_name": "test-bucket", + "s3_region_name": "us-east-1", + "s3_server_side_encryption": "aws:kms", + "s3_sse_kms_key_id": TEST_KMS_KEY_ARN, + } + ) + + put_object_kwargs = mock_s3_client.put_object.call_args.kwargs + assert put_object_kwargs["ServerSideEncryption"] == "aws:kms" + assert put_object_kwargs["SSEKMSKeyId"] == TEST_KMS_KEY_ARN + + +def test_put_object_supports_sse_s3_without_key_id(): + """SSE-S3 (AES256) needs only ServerSideEncryption, no key id.""" + mock_s3_client = _run_log_event( + { + "s3_bucket_name": "test-bucket", + "s3_region_name": "us-east-1", + "s3_server_side_encryption": "AES256", + } + ) + + put_object_kwargs = mock_s3_client.put_object.call_args.kwargs + assert put_object_kwargs["ServerSideEncryption"] == "AES256" + assert "SSEKMSKeyId" not in put_object_kwargs + + +def test_put_object_omits_sse_params_by_default(): + """Without SSE config, put_object kwargs must stay unchanged.""" + mock_s3_client = _run_log_event( + { + "s3_bucket_name": "test-bucket", + "s3_region_name": "us-east-1", + } + ) + + put_object_kwargs = mock_s3_client.put_object.call_args.kwargs + assert "ServerSideEncryption" not in put_object_kwargs + assert "SSEKMSKeyId" not in put_object_kwargs + + +def test_put_object_infers_aws_kms_when_only_key_id_set(): + """A key id without an algorithm must infer aws:kms instead of sending an invalid request.""" + mock_s3_client = _run_log_event( + { + "s3_bucket_name": "test-bucket", + "s3_region_name": "us-east-1", + "s3_sse_kms_key_id": TEST_KMS_KEY_ARN, + } + ) + + put_object_kwargs = mock_s3_client.put_object.call_args.kwargs + assert put_object_kwargs["ServerSideEncryption"] == "aws:kms" + assert put_object_kwargs["SSEKMSKeyId"] == TEST_KMS_KEY_ARN + + +def test_put_object_drops_key_id_when_algorithm_is_not_kms(): + """AES256 plus a key id is invalid for S3; the key id must be dropped, not sent.""" + mock_s3_client = _run_log_event( + { + "s3_bucket_name": "test-bucket", + "s3_region_name": "us-east-1", + "s3_server_side_encryption": "AES256", + "s3_sse_kms_key_id": TEST_KMS_KEY_ARN, + } + ) + + put_object_kwargs = mock_s3_client.put_object.call_args.kwargs + assert put_object_kwargs["ServerSideEncryption"] == "AES256" + assert "SSEKMSKeyId" not in put_object_kwargs + + +def test_non_string_algorithm_is_dropped_and_valid_key_id_is_rescued(): + """ + A YAML boolean in s3_server_side_encryption must not crash logger init and + must not discard the valid key id; aws:kms is inferred from the key id. + """ + mock_s3_client = _run_log_event( + { + "s3_bucket_name": "test-bucket", + "s3_region_name": "us-east-1", + "s3_server_side_encryption": True, + "s3_sse_kms_key_id": TEST_KMS_KEY_ARN, + } + ) + + put_object_kwargs = mock_s3_client.put_object.call_args.kwargs + assert put_object_kwargs["ServerSideEncryption"] == "aws:kms" + assert put_object_kwargs["SSEKMSKeyId"] == TEST_KMS_KEY_ARN + + +def test_non_string_key_id_is_dropped_and_valid_algorithm_is_kept(): + """A mistyped key id (unquoted YAML number) must not disable the valid algorithm.""" + mock_s3_client = _run_log_event( + { + "s3_bucket_name": "test-bucket", + "s3_region_name": "us-east-1", + "s3_server_side_encryption": "aws:kms", + "s3_sse_kms_key_id": 12345, + } + ) + + put_object_kwargs = mock_s3_client.put_object.call_args.kwargs + assert put_object_kwargs["ServerSideEncryption"] == "aws:kms" + assert "SSEKMSKeyId" not in put_object_kwargs diff --git a/tests/test_litellm/integrations/test_s3_v2.py b/tests/test_litellm/integrations/test_s3_v2.py index f0a33f2ebfc..3977daae92f 100644 --- a/tests/test_litellm/integrations/test_s3_v2.py +++ b/tests/test_litellm/integrations/test_s3_v2.py @@ -1388,3 +1388,253 @@ def test_s3_server_side_encryption_read_from_callback_params(): assert logger.s3_server_side_encryption == "aws:kms" finally: litellm.s3_callback_params = original + + +@pytest.mark.asyncio +async def test_async_upload_sets_sse_kms_key_id_header_when_configured(): + """ + When s3_sse_kms_key_id is set alongside aws:kms, the PUT must carry + x-amz-server-side-encryption-aws-kms-key-id so objects are encrypted + with the customer-managed KMS key instead of the bucket default. + """ + from unittest.mock import AsyncMock, MagicMock + + from litellm.types.integrations.s3_v2 import s3BatchLoggingElement + + logger = S3Logger( + s3_bucket_name="test-bucket", + s3_aws_access_key_id="test-key", + s3_aws_secret_access_key="test-secret", + s3_region_name="us-east-1", + s3_server_side_encryption="aws:kms", + s3_sse_kms_key_id="arn:aws:kms:us-east-1:111122223333:key/test-key-id", + ) + + test_element = s3BatchLoggingElement( + s3_object_key="2025-09-14/test-sse-kms.json", + payload={"test": "sse-kms"}, + s3_object_download_filename="test-sse-kms.json", + ) + + response = MagicMock() + response.status_code = 200 + response.raise_for_status = MagicMock() + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put.return_value = response + + await logger.async_upload_data_to_s3(test_element) + + headers = logger.async_httpx_client.put.call_args.kwargs["headers"] + assert headers["x-amz-server-side-encryption"] == "aws:kms" + assert headers["x-amz-server-side-encryption-aws-kms-key-id"] == ( + "arn:aws:kms:us-east-1:111122223333:key/test-key-id" + ) + + +def test_sync_upload_sets_sse_kms_key_id_header_when_configured(): + """The sync upload path must carry the same SSE-KMS headers.""" + from unittest.mock import MagicMock + + from litellm.types.integrations.s3_v2 import s3BatchLoggingElement + + logger = S3Logger( + s3_bucket_name="test-bucket", + s3_aws_access_key_id="test-key", + s3_aws_secret_access_key="test-secret", + s3_region_name="us-east-1", + s3_server_side_encryption="aws:kms", + s3_sse_kms_key_id="arn:aws:kms:us-east-1:111122223333:key/test-key-id", + ) + + test_element = s3BatchLoggingElement( + s3_object_key="2025-09-14/test-sync-sse-kms.json", + payload={"test": "sync-sse-kms"}, + s3_object_download_filename="test-sync-sse-kms.json", + ) + + response = MagicMock() + response.status_code = 200 + response.raise_for_status = MagicMock() + mock_sync_client = MagicMock() + mock_sync_client.put.return_value = response + + with patch( + "litellm.integrations.s3_v2._get_httpx_client", + return_value=mock_sync_client, + ): + logger.upload_data_to_s3(test_element) + + headers = mock_sync_client.put.call_args.kwargs["headers"] + assert headers["x-amz-server-side-encryption"] == "aws:kms" + assert headers["x-amz-server-side-encryption-aws-kms-key-id"] == ( + "arn:aws:kms:us-east-1:111122223333:key/test-key-id" + ) + + +@pytest.mark.asyncio +async def test_async_upload_omits_kms_key_id_header_when_not_configured(): + """SSE without a key id must not emit the KMS key id header.""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.types.integrations.s3_v2 import s3BatchLoggingElement + + logger = S3Logger( + s3_bucket_name="test-bucket", + s3_aws_access_key_id="test-key", + s3_aws_secret_access_key="test-secret", + s3_region_name="us-east-1", + s3_server_side_encryption="AES256", + ) + + test_element = s3BatchLoggingElement( + s3_object_key="2025-09-14/test-aes256.json", + payload={"test": "aes256"}, + s3_object_download_filename="test-aes256.json", + ) + + response = MagicMock() + response.status_code = 200 + response.raise_for_status = MagicMock() + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put.return_value = response + + await logger.async_upload_data_to_s3(test_element) + + headers = logger.async_httpx_client.put.call_args.kwargs["headers"] + assert headers["x-amz-server-side-encryption"] == "AES256" + assert "x-amz-server-side-encryption-aws-kms-key-id" not in headers + + +def test_s3_sse_kms_key_id_read_from_callback_params(): + """s3_sse_kms_key_id can be configured via s3_callback_params.""" + import litellm + + original = litellm.s3_callback_params + litellm.s3_callback_params = { + "s3_bucket_name": "from-global", + "s3_server_side_encryption": "aws:kms", + "s3_sse_kms_key_id": "arn:aws:kms:us-east-1:111122223333:key/test-key-id", + } + try: + logger = S3Logger() + assert logger.s3_sse_kms_key_id == ("arn:aws:kms:us-east-1:111122223333:key/test-key-id") + finally: + litellm.s3_callback_params = original + + +@pytest.mark.asyncio +async def test_async_upload_infers_aws_kms_when_only_key_id_set(): + """ + Setting only s3_sse_kms_key_id must not produce an invalid request + (S3 rejects a key id without an algorithm); aws:kms is inferred. + """ + from unittest.mock import AsyncMock, MagicMock + + from litellm.types.integrations.s3_v2 import s3BatchLoggingElement + + logger = S3Logger( + s3_bucket_name="test-bucket", + s3_aws_access_key_id="test-key", + s3_aws_secret_access_key="test-secret", + s3_region_name="us-east-1", + s3_sse_kms_key_id="arn:aws:kms:us-east-1:111122223333:key/test-key-id", + ) + + test_element = s3BatchLoggingElement( + s3_object_key="2025-09-14/test-kms-only.json", + payload={"test": "kms-only"}, + s3_object_download_filename="test-kms-only.json", + ) + + response = MagicMock() + response.status_code = 200 + response.raise_for_status = MagicMock() + logger.async_httpx_client = AsyncMock() + logger.async_httpx_client.put.return_value = response + + await logger.async_upload_data_to_s3(test_element) + + headers = logger.async_httpx_client.put.call_args.kwargs["headers"] + assert headers["x-amz-server-side-encryption"] == "aws:kms" + assert headers["x-amz-server-side-encryption-aws-kms-key-id"] == ( + "arn:aws:kms:us-east-1:111122223333:key/test-key-id" + ) + + +def test_s3_sse_kms_key_id_read_from_audit_override_params(): + """The audit-log override path must honor s3_sse_kms_key_id too.""" + import litellm + + original = litellm.s3_callback_params + litellm.s3_callback_params = {"s3_bucket_name": "normal-logs-bucket"} + try: + logger = S3Logger( + s3_callback_params_override={ + "s3_bucket_name": "audit-logs-bucket", + "s3_sse_kms_key_id": "arn:aws:kms:us-east-1:111122223333:key/audit-key-id", + } + ) + assert logger.s3_bucket_name == "audit-logs-bucket" + assert logger.s3_sse_kms_key_id == ("arn:aws:kms:us-east-1:111122223333:key/audit-key-id") + finally: + litellm.s3_callback_params = original + + +def test_kms_key_id_dropped_when_algorithm_is_not_kms(): + """ + AES256 plus a KMS key id is an invalid S3 combination; the key id must be + dropped at init so uploads keep working instead of silently 400ing. + """ + import litellm + + original = litellm.s3_callback_params + litellm.s3_callback_params = { + "s3_bucket_name": "from-global", + "s3_server_side_encryption": "AES256", + "s3_sse_kms_key_id": "arn:aws:kms:us-east-1:111122223333:key/test-key-id", + } + try: + logger = S3Logger() + assert logger.s3_server_side_encryption == "AES256" + assert logger.s3_sse_kms_key_id is None + finally: + litellm.s3_callback_params = original + + +def test_non_string_algorithm_is_dropped_and_valid_key_id_is_rescued(): + """ + A YAML boolean in s3_server_side_encryption must not crash logger init and + must not discard the valid key id; aws:kms is inferred from the key id. + """ + import litellm + + original = litellm.s3_callback_params + litellm.s3_callback_params = { + "s3_bucket_name": "from-global", + "s3_server_side_encryption": True, + "s3_sse_kms_key_id": "arn:aws:kms:us-east-1:111122223333:key/test-key-id", + } + try: + logger = S3Logger() + assert logger.s3_server_side_encryption == "aws:kms" + assert logger.s3_sse_kms_key_id == ("arn:aws:kms:us-east-1:111122223333:key/test-key-id") + finally: + litellm.s3_callback_params = original + + +def test_non_string_key_id_is_dropped_and_valid_algorithm_is_kept(): + """A mistyped key id (unquoted YAML number) must not disable the valid algorithm.""" + import litellm + + original = litellm.s3_callback_params + litellm.s3_callback_params = { + "s3_bucket_name": "from-global", + "s3_server_side_encryption": "aws:kms", + "s3_sse_kms_key_id": 12345, + } + try: + logger = S3Logger() + assert logger.s3_server_side_encryption == "aws:kms" + assert logger.s3_sse_kms_key_id is None + finally: + litellm.s3_callback_params = original