merge litellm_internal_staging

This commit is contained in:
Devin AI 2026-07-31 04:29:13 +00:00
commit cd54183e62
19 changed files with 1335 additions and 82 deletions

View file

@ -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,

View file

@ -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()

View file

@ -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 <your_api_key>\"\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-<index>\").\n\nExample Request:\n```bash\ncurl -X GET \"http://localhost:4000/policies/attachments/list\" \\\n -H \"Authorization: Bearer <your_api_key>\"\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 <your_api_key>\"\ncurl -X GET \"http://localhost:4000/policies/list?version_status=production\" \\\n -H \"Authorization: Bearer <your_api_key>\"\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 <your_api_key>\"\ncurl -X GET \"http://localhost:4000/policies/list?version_status=production\" \\\n -H \"Authorization: Bearer <your_api_key>\"\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": [
{

View file

@ -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)}")

View file

@ -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,13 @@ 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 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".
@ -84,11 +121,27 @@ 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 if db_policy.version_status == "production"
}
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
]
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 +659,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-<index>").
Example Request:
```bash
@ -635,11 +691,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}")

View file

@ -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,23 +306,42 @@ 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}")
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
@ -323,11 +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]
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
@ -501,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:
@ -591,14 +629,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_<uuid> 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 +649,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 +684,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}")
@ -983,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)}")

View file

@ -6,7 +6,7 @@ the final guardrails list.
"""
from datetime import datetime
from typing import Any, Dict, List, Optional
from typing import Any, Dict, List, Literal, Optional
from pydantic import BaseModel, ConfigDict, Field
@ -220,6 +220,10 @@ class PolicyDBResponse(BaseModel):
updated_at: Optional[datetime] = Field(default=None, description="When the policy was last updated.")
created_by: Optional[str] = Field(default=None, description="Who created the policy.")
updated_by: Optional[str] = Field(default=None, description="Who last updated the policy.")
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):

View file

@ -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

View file

@ -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

View file

@ -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,

View file

@ -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,75 @@ 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
@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() == ()

View file

@ -0,0 +1,248 @@
"""
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_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_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")
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"

View file

@ -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,
)
@ -450,3 +452,182 @@ 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"]
@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
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()
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

View file

@ -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 (
<DropdownMenu>
<DropdownMenuTrigger
@ -65,6 +70,8 @@ function AttachmentRowActions({ attachment, isAdmin, onDeleteClick }: Attachment
<DropdownMenuItem
variant="destructive"
data-testid="attachment-action-delete"
disabled={isConfigAttachment}
title={isConfigAttachment ? CONFIG_ATTACHMENT_HINT : undefined}
onClick={() => onDeleteClick(attachment.attachment_id)}
>
<Trash2 />

View file

@ -145,4 +145,34 @@ describe("PolicyTable", () => {
await user.click(screen.getByRole("button", { name: /grouped/ }));
expect(defaultProps.onViewClick).toHaveBeenCalledWith("prod-id");
});
const sameNamedDbDraft: Partial<Policy> = {
policy_name: "config-policy",
policy_id: "db-draft-id",
version_status: "draft",
version_number: 2,
};
const configTwin: Partial<Policy> = {
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(<PolicyTable {...defaultProps} policies={policies} />);
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(<PolicyTable {...defaultProps} policies={policies} />);
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");
});
});

View file

@ -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<PolicyTableProps> = ({
<DataTable
data={rows}
columns={columns}
getRowId={(row) => row.policy_name}
getRowId={(row) => `${row.primaryPolicy.definition_location ?? "db"}:${row.policy_name}`}
sortingMode="client"
sorting={sorting}
onSortingChange={setSorting}

View file

@ -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 <span className="text-muted-foreground">-</span>;
@ -45,6 +48,8 @@ interface PolicyRowActionsProps {
}
function PolicyRowActions({ policy, onEditClick, onDeleteClick }: PolicyRowActionsProps) {
const isConfigPolicy = policy.definition_location === "config";
return (
<DropdownMenu>
<DropdownMenuTrigger
@ -55,7 +60,12 @@ function PolicyRowActions({ policy, onEditClick, onDeleteClick }: PolicyRowActio
<MoreHorizontal className="size-4" />
</DropdownMenuTrigger>
<DropdownMenuContent align="end" className="w-52">
<DropdownMenuItem data-testid="policy-action-edit" onClick={() => onEditClick(policy)}>
<DropdownMenuItem
data-testid="policy-action-edit"
disabled={isConfigPolicy}
title={isConfigPolicy ? CONFIG_POLICY_HINT : undefined}
onClick={() => onEditClick(policy)}
>
<Pencil />
Edit policy
</DropdownMenuItem>
@ -63,6 +73,8 @@ function PolicyRowActions({ policy, onEditClick, onDeleteClick }: PolicyRowActio
<DropdownMenuItem
variant="destructive"
data-testid="policy-action-delete"
disabled={isConfigPolicy}
title={isConfigPolicy ? CONFIG_POLICY_HINT : undefined}
onClick={() => onDeleteClick(policy.policy_id, policy.policy_name || "Unnamed Policy")}
>
<Trash2 />
@ -93,18 +105,23 @@ export const getPolicyTableColumns = ({
header: ({ column }) => <DataTableSortHeader column={column} title="Name" />,
size: 220,
enableSorting: true,
cell: ({ row }) => (
<IdentityCell
title={row.original.policy_name}
titleClassName="max-w-60"
badge={
row.original.versionCount > 1 ? (
<StatusBadge tone="neutral" label={`${row.original.versionCount} versions`} />
) : undefined
}
onClick={() => onViewClick(row.original.primaryPolicy.policy_id)}
/>
),
cell: ({ row }) => {
const isConfigPolicy = row.original.primaryPolicy.definition_location === "config";
const versionBadge =
row.original.versionCount > 1 ? (
<StatusBadge tone="neutral" label={`${row.original.versionCount} versions`} />
) : undefined;
return (
<IdentityCell
title={row.original.policy_name}
titleClassName="max-w-60"
badge={
isConfigPolicy ? <StatusBadge tone="neutral" label="Config" tooltip={CONFIG_POLICY_HINT} /> : versionBadge
}
onClick={isConfigPolicy ? undefined : () => onViewClick(row.original.primaryPolicy.policy_id)}
/>
);
},
},
{
id: "description",

View file

@ -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 {

View file

@ -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-<index>").
*
* 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".
@ -29374,6 +29380,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.
@ -29505,6 +29518,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.
@ -33112,6 +33132,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 */