Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/policy_engine/policy_registry.py: 64%
374 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2Policy Registry - In-memory storage for policies.
4Handles storing, retrieving, and managing policies.
6Policies define WHAT guardrails to apply. WHERE they apply is defined
7by policy_attachments (see AttachmentRegistry).
8"""
10import json
11from collections.abc import Mapping, Sequence
12from datetime import datetime, timezone
13from typing import (
14 TYPE_CHECKING,
15 Any,
16 Final,
17 Literal,
18 Optional,
19 Protocol,
20 TypedDict,
21 Union,
22 cast, # noqa: TID251 # prisma types the condition/pipeline Json columns as str, but reads return decoded values
23)
25from litellm._logging import verbose_proxy_logger
26from litellm.repositories.prisma_protocols import TableActions
27from litellm.repositories.table_repositories import PolicyRepository
28from litellm.types.proxy.policy_engine import (
29 GuardrailPipeline,
30 PipelineStep,
31 Policy,
32 PolicyCondition,
33 PolicyCreateRequest,
34 PolicyDBResponse,
35 PolicyGuardrails,
36 PolicyUpdateRequest,
37 PolicyVersionCompareResponse,
38 PolicyVersionListResponse,
39)
41if TYPE_CHECKING: 41 ↛ 42line 41 didn't jump to line 42 because the condition on line 41 was never true
42 from litellm.proxy.utils import PrismaClient
44# Prefix for policy version IDs in request body. Use policy_<uuid> to execute a specific version.
45POLICY_VERSION_ID_PREFIX: Final = "policy_"
48class _RawPipelineStep(TypedDict):
49 guardrail: str
52class _RawPipelineConfig(TypedDict, total=False):
53 mode: str
54 steps: Sequence[Union[PipelineStep, "_RawPipelineStep"]]
57class _PolicyRow(Protocol):
58 policy_id: str
59 policy_name: str
60 version_number: int
61 version_status: str
62 parent_version_id: str | None
63 is_latest: bool
64 published_at: datetime | None
65 production_at: datetime | None
66 inherit: str | None
67 description: str | None
68 guardrails_add: list[str] | None
69 guardrails_remove: list[str] | None
70 condition: dict[str, object] | None
71 pipeline: dict[str, object] | None
72 created_at: datetime
73 updated_at: datetime
74 created_by: str | None
75 updated_by: str | None
78class _PolicyVersionSourceRow(Protocol):
79 @property
80 def policy_id(self) -> str: ... 80 ↛ exitline 80 didn't return from function 'policy_id' because
82 @property
83 def policy_name(self) -> str: ... 83 ↛ exitline 83 didn't return from function 'policy_name' because
85 @property
86 def version_number(self) -> int: ... 86 ↛ exitline 86 didn't return from function 'version_number' because
88 @property
89 def inherit(self) -> str | None: ... 89 ↛ exitline 89 didn't return from function 'inherit' because
91 @property
92 def description(self) -> str | None: ... 92 ↛ exitline 92 didn't return from function 'description' because
94 @property
95 def guardrails_add(self) -> Sequence[str] | None: ... 95 ↛ exitline 95 didn't return from function 'guardrails_add' because
97 @property
98 def guardrails_remove(self) -> Sequence[str] | None: ... 98 ↛ exitline 98 didn't return from function 'guardrails_remove' because
100 @property
101 def condition(self) -> Mapping[str, object] | str | None: ... 101 ↛ exitline 101 didn't return from function 'condition' because
103 @property
104 def pipeline(self) -> Mapping[str, object] | str | None: ... 104 ↛ exitline 104 didn't return from function 'pipeline' because
107class _PolicyTableClient(Protocol):
108 async def create(self, data: Mapping[str, object]) -> _PolicyRow: ... 108 ↛ exitline 108 didn't return from function 'create' because
110 async def find_unique(self, where: Mapping[str, object]) -> _PolicyRow | None: ... 110 ↛ exitline 110 didn't return from function 'find_unique' because
112 async def find_many( 112 ↛ exitline 112 didn't return from function 'find_many' because
113 self,
114 where: Mapping[str, object] | None = None,
115 order: Mapping[str, str] | None = None,
116 ) -> Sequence[_PolicyRow]: ...
118 async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> _PolicyRow: ... 118 ↛ exitline 118 didn't return from function 'update' because
120 async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: ... 120 ↛ exitline 120 didn't return from function 'update_many' because
122 async def delete(self, where: Mapping[str, object]) -> _PolicyRow | None: ... 122 ↛ exitline 122 didn't return from function 'delete' because
124 async def delete_many(self, where: Mapping[str, object]) -> int: ... 124 ↛ exitline 124 didn't return from function 'delete_many' because
127def _policy_table(prisma_client: "PrismaClient") -> _PolicyTableClient:
128 table: Final = PolicyRepository(prisma_client).table
129 return cast( # cast-ok: prisma types Json columns as str; the client hands back the decoded condition/pipeline
130 "_PolicyTableClient", table
131 )
134def _policy_version_source_table(prisma_client: "PrismaClient") -> "TableActions[_PolicyVersionSourceRow]":
135 table: Final[TableActions[_PolicyVersionSourceRow]] = PolicyRepository(prisma_client).table
136 return table
139def _row_to_policy_db_response(row: _PolicyRow) -> PolicyDBResponse:
140 """Build PolicyDBResponse from a Prisma LiteLLM_PolicyTable row."""
141 return PolicyDBResponse(
142 policy_id=row.policy_id,
143 policy_name=row.policy_name,
144 version_number=getattr(row, "version_number", 1),
145 version_status=getattr(row, "version_status", "production"),
146 parent_version_id=getattr(row, "parent_version_id", None),
147 is_latest=getattr(row, "is_latest", True),
148 published_at=getattr(row, "published_at", None),
149 production_at=getattr(row, "production_at", None),
150 inherit=row.inherit,
151 description=row.description,
152 guardrails_add=row.guardrails_add or [],
153 guardrails_remove=row.guardrails_remove or [],
154 condition=row.condition,
155 pipeline=row.pipeline,
156 created_at=row.created_at,
157 updated_at=row.updated_at,
158 created_by=row.created_by,
159 updated_by=row.updated_by,
160 )
163class PolicyRegistry:
164 """
165 In-memory registry for storing and managing policies.
167 This is a singleton that holds all loaded policies and provides
168 methods to access them.
170 Policies define WHAT guardrails to apply:
171 - Base guardrails via guardrails.add/remove
172 - Inheritance via inherit field
173 - Conditional guardrails via condition.model
174 """
176 def __init__(self):
177 self._policies: dict[str, Policy] = {}
178 self._config_policies: Mapping[str, Policy] = {}
179 self._sources: Mapping[str, Literal["db", "config"]] = {}
180 self._policies_by_id: dict[str, tuple[str, Policy]] = {}
181 self._initialized: bool = False
183 def load_policies(self, policies_config: Mapping[str, dict[str, object]]) -> None:
184 """
185 Load policies from a configuration dictionary.
187 Args:
188 policies_config: Dictionary mapping policy names to policy definitions.
189 This is the raw config from the YAML file.
190 """
191 self._policies = {}
192 self._config_policies = {}
193 self._sources = {}
194 self._policies_by_id = {}
196 for policy_name, policy_data in policies_config.items():
197 try:
198 policy = self._parse_policy(policy_name, policy_data)
199 self._policies[policy_name] = policy
200 verbose_proxy_logger.debug("Loaded policy: %s", policy_name)
201 except Exception as e:
202 verbose_proxy_logger.error("Error loading policy '%s': %s", policy_name, e)
203 raise ValueError(f"Invalid policy '{policy_name}': {e}") from e
205 self._config_policies = dict(self._policies)
206 self._sources = {policy_name: "config" for policy_name in self._policies}
207 self._initialized = True
208 verbose_proxy_logger.info("Loaded %s policies", len(self._policies))
210 def _parse_policy(self, policy_name: str, policy_data: dict[str, Any]) -> Policy:
211 """
212 Parse a policy from raw configuration data.
214 Args:
215 policy_name: Name of the policy
216 policy_data: Raw policy configuration
218 Returns:
219 Parsed Policy object
220 """
221 # Parse guardrails
222 guardrails_data: Final = policy_data.get("guardrails", {})
223 if isinstance(guardrails_data, dict): 223 ↛ 230line 223 didn't jump to line 230 because the condition on line 223 was always true
224 guardrails = PolicyGuardrails(
225 add=guardrails_data.get("add"),
226 remove=guardrails_data.get("remove"),
227 )
228 else:
229 # Handle legacy format where guardrails might be a list
230 guardrails = PolicyGuardrails(add=guardrails_data if guardrails_data else None)
232 # Parse condition (simple model-based condition)
233 condition = None
234 condition_data: Final = policy_data.get("condition")
235 if condition_data:
236 condition = PolicyCondition(model=condition_data.get("model"))
238 # Parse pipeline (optional ordered guardrail execution)
239 pipeline: Final = PolicyRegistry._parse_pipeline(policy_data.get("pipeline"))
241 return Policy(
242 inherit=policy_data.get("inherit"),
243 description=policy_data.get("description"),
244 guardrails=guardrails,
245 condition=condition,
246 pipeline=pipeline,
247 )
249 @staticmethod
250 def _parse_pipeline(
251 pipeline_data: Optional["_RawPipelineConfig"],
252 ) -> GuardrailPipeline | None:
253 """Parse a pipeline configuration from raw data."""
254 if pipeline_data is None: 254 ↛ 257line 254 didn't jump to line 257 because the condition on line 254 was always true
255 return None
257 steps_data: Final[Sequence[PipelineStep | _RawPipelineStep]] = pipeline_data.get("steps", [])
258 steps = [PipelineStep(**step_data) if isinstance(step_data, dict) else step_data for step_data in steps_data]
260 return GuardrailPipeline(
261 mode=pipeline_data.get("mode", "pre_call"),
262 steps=steps,
263 )
265 def get_policy(self, policy_name: str) -> Policy | None:
266 """
267 Get a policy by name.
269 Args:
270 policy_name: Name of the policy to retrieve
272 Returns:
273 Policy object if found, None otherwise
274 """
275 return self._policies.get(policy_name)
277 def get_all_policies(self) -> dict[str, Policy]:
278 """
279 Get all loaded policies.
281 Returns:
282 Dictionary mapping policy names to Policy objects
283 """
284 return self._policies.copy()
286 def get_policy_names(self) -> list[str]:
287 """
288 Get list of all policy names.
290 Returns:
291 List of policy names
292 """
293 return list(self._policies.keys())
295 def has_policy(self, policy_name: str) -> bool:
296 """
297 Check if a policy exists.
299 Args:
300 policy_name: Name of the policy to check
302 Returns:
303 True if policy exists, False otherwise
304 """
305 return policy_name in self._policies
307 def is_initialized(self) -> bool:
308 """
309 Check if the registry has been initialized with policies.
311 Returns:
312 True if policies have been loaded, False otherwise
313 """
314 return self._initialized
316 def clear(self) -> None:
317 """
318 Clear all policies from the registry.
319 """
320 self._policies = {}
321 self._config_policies = {}
322 self._sources = {}
323 self._initialized = False
325 def get_source(self, policy_name: str) -> Literal["db", "config"] | None:
326 """
327 Return the provenance of an in-memory policy, or None if unknown.
328 """
329 return self._sources.get(policy_name)
331 def list_config_policies(self) -> Mapping[str, Policy]:
332 """
333 Return the policies loaded from config.yaml, keyed by policy name.
334 """
335 return dict(self._config_policies)
337 def add_policy(self, policy_name: str, policy: Policy, source: Literal["db", "config"] = "db") -> None:
338 """
339 Add or update a single policy.
341 Args:
342 policy_name: Name of the policy
343 policy: Policy object to add
344 source: Provenance of the policy ("db" or "config")
345 """
346 self._policies[policy_name] = policy
347 self._sources = {**self._sources, policy_name: source}
348 if source == "config": 348 ↛ 349line 348 didn't jump to line 349 because the condition on line 348 was never true
349 self._config_policies = {**self._config_policies, policy_name: policy}
350 self._initialized = True
351 verbose_proxy_logger.debug("Added/updated policy: %s", policy_name)
353 def remove_policy(self, policy_name: str) -> bool:
354 """
355 Remove a policy by name. If a config-defined policy shares the name,
356 it is restored immediately instead of waiting for the next DB sync.
358 Args:
359 policy_name: Name of the policy to remove
361 Returns:
362 True if policy was removed, False if it didn't exist
363 """
364 if policy_name not in self._policies:
365 return False
366 config_fallback: Final = self._config_policies.get(policy_name)
367 if config_fallback is not None: 367 ↛ 368line 367 didn't jump to line 368 because the condition on line 367 was never true
368 self._policies[policy_name] = config_fallback
369 self._sources = {**self._sources, policy_name: "config"}
370 verbose_proxy_logger.debug("Removed policy: %s; restored config-defined version", policy_name)
371 return True
372 del self._policies[policy_name]
373 self._sources = {name: source for name, source in self._sources.items() if name != policy_name}
374 verbose_proxy_logger.debug("Removed policy: %s", policy_name)
375 return True
377 # ─────────────────────────────────────────────────────────────────────────
378 # Database CRUD Methods
379 # ─────────────────────────────────────────────────────────────────────────
381 async def add_policy_to_db(
382 self,
383 policy_request: PolicyCreateRequest,
384 prisma_client: "PrismaClient",
385 created_by: str | None = None,
386 ) -> PolicyDBResponse:
387 """
388 Add a policy to the database.
390 Args:
391 policy_request: The policy creation request
392 prisma_client: The Prisma client instance
393 created_by: User who created the policy
395 Returns:
396 PolicyDBResponse with the created policy
397 """
398 try:
399 now: Final = datetime.now(timezone.utc)
400 # Build data dict; new policy is v1 production
401 data: Final[dict[str, object]] = {
402 "policy_name": policy_request.policy_name,
403 "version_number": 1,
404 "version_status": "production",
405 "is_latest": True,
406 "production_at": now,
407 "guardrails_add": policy_request.guardrails_add or [],
408 "guardrails_remove": policy_request.guardrails_remove or [],
409 "created_at": now,
410 "updated_at": now,
411 }
413 # Only add optional fields if they have values
414 if policy_request.inherit is not None:
415 data["inherit"] = policy_request.inherit
416 if policy_request.description is not None:
417 data["description"] = policy_request.description
418 if created_by is not None: 418 ↛ 421line 418 didn't jump to line 421 because the condition on line 418 was always true
419 data["created_by"] = created_by
420 data["updated_by"] = created_by
421 if policy_request.condition is not None:
422 data["condition"] = json.dumps(policy_request.condition.model_dump())
423 if policy_request.pipeline is not None:
424 validated_pipeline: Final = GuardrailPipeline(**policy_request.pipeline)
425 data["pipeline"] = json.dumps(validated_pipeline.model_dump())
427 created_policy: Final = await _policy_table(prisma_client).create(data=data)
429 # Also add to in-memory registry
430 policy: Final = self._parse_policy(
431 policy_request.policy_name,
432 {
433 "inherit": policy_request.inherit,
434 "description": policy_request.description,
435 "guardrails": {
436 "add": policy_request.guardrails_add,
437 "remove": policy_request.guardrails_remove,
438 },
439 "condition": (policy_request.condition.model_dump() if policy_request.condition else None),
440 "pipeline": policy_request.pipeline,
441 },
442 )
443 self.add_policy(policy_request.policy_name, policy)
445 return _row_to_policy_db_response(created_policy)
446 except Exception as e:
447 verbose_proxy_logger.exception("Error adding policy to DB: %s", e)
448 raise Exception(f"Error adding policy to DB: {e}")
450 async def update_policy_in_db(
451 self,
452 policy_id: str,
453 policy_request: PolicyUpdateRequest,
454 prisma_client: "PrismaClient",
455 updated_by: str | None = None,
456 ) -> PolicyDBResponse:
457 """
458 Update a policy in the database. Only draft versions can be updated.
460 Args:
461 policy_id: The ID of the policy to update
462 policy_request: The policy update request
463 prisma_client: The Prisma client instance
464 updated_by: User who updated the policy
466 Returns:
467 PolicyDBResponse with the updated policy
469 Raises:
470 Exception: If policy is not in draft status (only drafts are editable).
471 """
472 try:
473 existing: Final = await _policy_table(prisma_client).find_unique(where={"policy_id": policy_id})
474 if existing is None:
475 raise Exception(f"Policy with ID {policy_id} not found")
476 version_status: Final = getattr(existing, "version_status", "production")
477 if version_status != "draft":
478 raise Exception(f"Only draft versions can be updated. This policy has status '{version_status}'.")
480 # Build update data - only include fields that are set
481 update_data: Final[dict[str, object]] = {
482 "updated_at": datetime.now(timezone.utc),
483 "updated_by": updated_by,
484 }
486 if policy_request.policy_name is not None:
487 update_data["policy_name"] = policy_request.policy_name
488 if policy_request.inherit is not None:
489 update_data["inherit"] = policy_request.inherit
490 if policy_request.description is not None:
491 update_data["description"] = policy_request.description
492 if policy_request.guardrails_add is not None:
493 update_data["guardrails_add"] = policy_request.guardrails_add
494 if policy_request.guardrails_remove is not None:
495 update_data["guardrails_remove"] = policy_request.guardrails_remove
496 if policy_request.condition is not None:
497 update_data["condition"] = json.dumps(policy_request.condition.model_dump())
498 if policy_request.pipeline is not None:
499 validated_pipeline: Final = GuardrailPipeline(**policy_request.pipeline)
500 update_data["pipeline"] = json.dumps(validated_pipeline.model_dump())
502 updated_policy: Final = await _policy_table(prisma_client).update(
503 where={"policy_id": policy_id},
504 data=update_data,
505 )
507 # Do NOT update in-memory registry: drafts are not loaded into memory.
509 return _row_to_policy_db_response(updated_policy)
510 except Exception as e:
511 verbose_proxy_logger.exception("Error updating policy in DB: %s", e)
512 raise Exception(f"Error updating policy in DB: {e}")
514 async def delete_policy_from_db(
515 self,
516 policy_id: str,
517 prisma_client: "PrismaClient",
518 ) -> Mapping[str, str]:
519 """
520 Delete a policy version from the database.
522 If the deleted version was production, it is removed from the in-memory
523 registry. No other version is auto-promoted; admin must explicitly promote.
525 Args:
526 policy_id: The ID of the policy version to delete
527 prisma_client: The Prisma client instance
529 Returns:
530 Dict with "message" and optional "warning" if production was deleted.
531 """
532 try:
533 policy: Final = await _policy_table(prisma_client).find_unique(where={"policy_id": policy_id})
535 if policy is None: 535 ↛ 536line 535 didn't jump to line 536 because the condition on line 535 was never true
536 raise Exception(f"Policy with ID {policy_id} not found")
538 version_status: Final = getattr(policy, "version_status", "production")
539 policy_name: Final = policy.policy_name
541 # Delete from DB
542 await _policy_table(prisma_client).delete(where={"policy_id": policy_id})
544 result: Final[dict[str, str]] = {"message": f"Policy {policy_id} deleted successfully"}
546 # Remove from in-memory registry only if this was the production version
547 if version_status == "production": 547 ↛ 559line 547 didn't jump to line 559 because the condition on line 547 was always true
548 self.remove_policy(policy_name)
549 if self.get_source(policy_name) == "config": 549 ↛ 550line 549 didn't jump to line 550 because the condition on line 549 was never true
550 result["warning"] = (
551 "Production version was deleted. The config-defined policy with the same name is active again."
552 )
553 else:
554 result["warning"] = (
555 "Production version was deleted. No other version was promoted. "
556 "Promote another version to production if this policy should remain active."
557 )
559 return result
560 except Exception as e:
561 verbose_proxy_logger.exception("Error deleting policy from DB: %s", e)
562 raise Exception(f"Error deleting policy from DB: {e}")
564 async def get_policy_by_id_from_db(
565 self,
566 policy_id: str,
567 prisma_client: "PrismaClient",
568 ) -> PolicyDBResponse | None:
569 """
570 Get a policy by ID from the database.
572 Args:
573 policy_id: The ID of the policy to retrieve
574 prisma_client: The Prisma client instance
576 Returns:
577 PolicyDBResponse if found, None otherwise
578 """
579 try:
580 policy: Final = await _policy_table(prisma_client).find_unique(where={"policy_id": policy_id})
582 if policy is None:
583 return None
585 return _row_to_policy_db_response(policy)
586 except Exception as e:
587 verbose_proxy_logger.exception("Error getting policy from DB: %s", e)
588 raise Exception(f"Error getting policy from DB: {e}")
590 def get_policy_by_id_for_request(self, policy_id: str) -> tuple[str, Policy] | None:
591 """
592 Return a policy version by ID from in-memory cache (no DB access).
594 Used when the request body specifies policy_<uuid> to execute a specific version
595 (e.g. published or draft). The cache is populated by sync_policies_from_db,
596 which loads draft and published versions keyed by policy_id.
598 Args:
599 policy_id: The policy version ID (raw UUID, no prefix)
601 Returns:
602 (policy_name, Policy) if found, None otherwise
603 """
604 return self._policies_by_id.get(policy_id)
606 async def get_all_policies_from_db(
607 self,
608 prisma_client: "PrismaClient",
609 version_status: str | None = None,
610 ) -> list[PolicyDBResponse]:
611 """
612 Get all policies from the database, optionally filtered by version_status.
614 Args:
615 prisma_client: The Prisma client instance
616 version_status: If set, only return policies with this status
617 ("draft", "published", "production").
619 Returns:
620 List of PolicyDBResponse objects
621 """
622 try:
623 where: Final[dict[str, str]] = {}
624 if version_status is not None:
625 where["version_status"] = version_status
627 policies: Final = await _policy_table(prisma_client).find_many(
628 where=where if where else None,
629 order={"created_at": "desc"},
630 )
632 return [_row_to_policy_db_response(p) for p in policies]
633 except Exception as e:
634 verbose_proxy_logger.exception("Error getting policies from DB: %s", e)
635 raise Exception(f"Error getting policies from DB: {e}")
637 async def sync_policies_from_db(
638 self,
639 prisma_client: "PrismaClient",
640 ) -> None:
641 """
642 Sync policies from the database to in-memory registry.
643 - Production versions are loaded into _policies (by policy name) for resolution.
644 - Config-loaded policies are preserved; on a name conflict the DB version wins.
645 - Draft and published versions are loaded into _policies_by_id so request-body
646 policy_<uuid> overrides can be resolved without DB access in the hot path.
647 """
648 try:
649 production: Final = await self.get_all_policies_from_db(prisma_client, version_status="production")
650 db_policies: Final = {
651 policy_response.policy_name: self._parse_policy(
652 policy_response.policy_name,
653 {
654 "inherit": policy_response.inherit,
655 "description": policy_response.description,
656 "guardrails": {
657 "add": policy_response.guardrails_add,
658 "remove": policy_response.guardrails_remove,
659 },
660 "condition": policy_response.condition,
661 "pipeline": policy_response.pipeline,
662 },
663 )
664 for policy_response in production
665 }
666 for policy_name in set(db_policies) & set(self._config_policies): 666 ↛ 667line 666 didn't jump to line 667 because the loop on line 666 never started
667 verbose_proxy_logger.warning(
668 "Policy '%s' is defined in both config.yaml and the DB; the DB version takes precedence",
669 policy_name,
670 )
671 config_sources: Mapping[str, Literal["db", "config"]] = {name: "config" for name in self._config_policies}
672 db_sources: Final[Mapping[str, Literal["db", "config"]]] = {name: "db" for name in db_policies}
673 self._policies = {**self._config_policies, **db_policies}
674 self._sources = {**config_sources, **db_sources}
676 self._policies_by_id = {}
677 non_production: Final = await _policy_table(prisma_client).find_many(
678 where={"version_status": {"in": ["draft", "published"]}},
679 order={"created_at": "desc"},
680 )
681 for row in non_production:
682 policy = self._parse_policy(
683 row.policy_name,
684 {
685 "inherit": row.inherit,
686 "description": row.description,
687 "guardrails": {
688 "add": row.guardrails_add or [],
689 "remove": row.guardrails_remove or [],
690 },
691 "condition": row.condition,
692 "pipeline": row.pipeline,
693 },
694 )
695 self._policies_by_id[row.policy_id] = (row.policy_name, policy)
697 self._initialized = True
698 verbose_proxy_logger.info(
699 "Synced %s production policies and %s draft/published (by ID) from DB to in-memory registry (%s config-defined policies preserved)",
700 len(production),
701 len(non_production),
702 len(self._config_policies),
703 )
704 except Exception as e:
705 verbose_proxy_logger.exception("Error syncing policies from DB: %s", e)
706 raise Exception(f"Error syncing policies from DB: {e}")
708 async def resolve_guardrails_from_db(
709 self,
710 policy_name: str,
711 prisma_client: "PrismaClient",
712 ) -> list[str]:
713 """
714 Resolve all guardrails for a policy from the database.
716 Uses the existing PolicyResolver to handle inheritance chain resolution.
718 Args:
719 policy_name: Name of the policy to resolve
720 prisma_client: The Prisma client instance
722 Returns:
723 List of resolved guardrail names
724 """
725 from litellm.proxy.policy_engine.policy_resolver import PolicyResolver
727 try:
728 # Load only production versions so inheritance resolves against production
729 policies: Final = await self.get_all_policies_from_db(prisma_client, version_status="production")
731 # Build a temporary in-memory map for resolution
732 temp_policies: Final = {}
733 for policy_response in policies:
734 policy = self._parse_policy(
735 policy_response.policy_name,
736 {
737 "inherit": policy_response.inherit,
738 "description": policy_response.description,
739 "guardrails": {
740 "add": policy_response.guardrails_add,
741 "remove": policy_response.guardrails_remove,
742 },
743 "condition": policy_response.condition,
744 "pipeline": policy_response.pipeline,
745 },
746 )
747 temp_policies[policy_response.policy_name] = policy
749 # Use the existing PolicyResolver to resolve guardrails
750 resolved_policy: Final = PolicyResolver.resolve_policy_guardrails(
751 policy_name=policy_name,
752 policies=temp_policies,
753 context=None, # No context needed for simple resolution
754 )
756 return sorted(resolved_policy.guardrails)
757 except Exception as e:
758 verbose_proxy_logger.exception("Error resolving guardrails from DB: %s", e)
759 raise Exception(f"Error resolving guardrails from DB: {e}")
761 async def get_versions_by_policy_name(
762 self,
763 policy_name: str,
764 prisma_client: "PrismaClient",
765 ) -> PolicyVersionListResponse:
766 """
767 Get all versions of a policy by name, ordered by version_number descending.
769 Args:
770 policy_name: Name of the policy
771 prisma_client: The Prisma client instance
773 Returns:
774 PolicyVersionListResponse with policy_name and list of versions
775 """
776 try:
777 rows: Final = await _policy_table(prisma_client).find_many(
778 where={"policy_name": policy_name},
779 order={"version_number": "desc"},
780 )
781 versions: Final = [_row_to_policy_db_response(r) for r in rows]
782 return PolicyVersionListResponse(
783 policy_name=policy_name,
784 versions=versions,
785 total_count=len(versions),
786 )
787 except Exception as e:
788 verbose_proxy_logger.exception("Error getting versions: %s", e)
789 raise Exception(f"Error getting versions: {e}")
791 async def create_new_version(
792 self,
793 policy_name: str,
794 prisma_client: "PrismaClient",
795 source_policy_id: str | None = None,
796 created_by: str | None = None,
797 ) -> PolicyDBResponse:
798 """
799 Create a new draft version of a policy. Copies all fields from the source.
800 Source is current production if source_policy_id is None.
802 Args:
803 policy_name: Name of the policy
804 prisma_client: The Prisma client instance
805 source_policy_id: Policy ID to clone from; if None, use current production
806 created_by: User who created the version
808 Returns:
809 PolicyDBResponse for the new draft version
810 """
811 try:
812 if source_policy_id is not None:
813 source = await _policy_version_source_table(prisma_client).find_unique(
814 where={"policy_id": source_policy_id}
815 )
816 if source is None: 816 ↛ 818line 816 didn't jump to line 818 because the condition on line 816 was always true
817 raise Exception(f"Source policy {source_policy_id} not found")
818 if source.policy_name != policy_name:
819 raise Exception(f"Source policy name '{source.policy_name}' does not match '{policy_name}'")
820 else:
821 # Find current production version for this policy_name
822 prod: Final = await _policy_version_source_table(prisma_client).find_first(
823 where={
824 "policy_name": policy_name,
825 "version_status": "production",
826 }
827 )
828 if prod is None:
829 raise Exception(f"No production version found for policy '{policy_name}'")
830 source = prod
832 # Next version number
833 latest: Final = await _policy_version_source_table(prisma_client).find_first(
834 where={"policy_name": policy_name},
835 order={"version_number": "desc"},
836 )
837 next_num: Final = (latest.version_number + 1) if latest else 1
839 now: Final = datetime.now(timezone.utc)
840 # Set is_latest=False on all existing versions for this policy_name
841 await _policy_table(prisma_client).update_many(
842 where={"policy_name": policy_name},
843 data={"is_latest": False},
844 )
846 data: Final[dict[str, object]] = {
847 "policy_name": policy_name,
848 "version_number": next_num,
849 "version_status": "draft",
850 "parent_version_id": source.policy_id,
851 "is_latest": True,
852 "published_at": None,
853 "production_at": None,
854 "inherit": source.inherit,
855 "description": source.description,
856 "guardrails_add": source.guardrails_add or [],
857 "guardrails_remove": source.guardrails_remove or [],
858 "created_at": now,
859 "updated_at": now,
860 "created_by": created_by,
861 "updated_by": created_by,
862 }
863 # Prisma expects Json fields as JSON strings on create (same as add_policy_to_db)
864 if source.condition is not None: 864 ↛ 868line 864 didn't jump to line 868 because the condition on line 864 was always true
865 data["condition"] = (
866 json.dumps(source.condition) if isinstance(source.condition, dict) else source.condition
867 )
868 if source.pipeline is not None: 868 ↛ 869line 868 didn't jump to line 869 because the condition on line 868 was never true
869 data["pipeline"] = json.dumps(source.pipeline) if isinstance(source.pipeline, dict) else source.pipeline
871 created: Final = await _policy_table(prisma_client).create(data=data)
872 return _row_to_policy_db_response(created)
873 except Exception as e:
874 verbose_proxy_logger.exception("Error creating new version: %s", e)
875 raise Exception(f"Error creating new version: {e}")
877 async def update_version_status(
878 self,
879 policy_id: str,
880 new_status: str,
881 prisma_client: "PrismaClient",
882 updated_by: str | None = None,
883 ) -> PolicyDBResponse:
884 """
885 Update a policy version's status. Valid transitions:
886 - draft -> published (sets published_at)
887 - published -> production (sets production_at, demotes current production to published, updates in-memory)
888 - production -> published (demotes, removes from in-memory)
889 - draft -> production: NOT allowed (must publish first)
890 - published -> draft: NOT allowed
892 Args:
893 policy_id: The policy version ID
894 new_status: "published" or "production"
895 prisma_client: The Prisma client instance
896 updated_by: User who updated
898 Returns:
899 PolicyDBResponse for the updated version
900 """
901 try:
902 if new_status not in ("published", "production"):
903 raise Exception(f"Invalid status '{new_status}'. Use 'published' or 'production'.")
905 row: Final = await _policy_table(prisma_client).find_unique(where={"policy_id": policy_id})
906 if row is None:
907 raise Exception(f"Policy with ID {policy_id} not found")
909 current: Final = getattr(row, "version_status", "production")
910 policy_name: Final = row.policy_name
911 now: Final = datetime.now(timezone.utc)
913 if new_status == "published": 913 ↛ 914line 913 didn't jump to line 914 because the condition on line 913 was never true
914 if current != "draft":
915 raise Exception(f"Only draft versions can be published. Current status: '{current}'.")
916 updated = await _policy_table(prisma_client).update(
917 where={"policy_id": policy_id},
918 data={
919 "version_status": "published",
920 "published_at": now,
921 "updated_at": now,
922 "updated_by": updated_by,
923 },
924 )
925 return _row_to_policy_db_response(updated)
927 # new_status == "production"
928 if current not in ("draft", "published"): 928 ↛ 933line 928 didn't jump to line 933 because the condition on line 928 was always true
929 raise Exception(
930 f"Only draft or published versions can be promoted to production. Current: '{current}'."
931 )
932 # Plan: "draft -> production" NOT allowed
933 if current == "draft":
934 raise Exception("Cannot promote draft directly to production. Publish the version first.")
936 # Demote current production to published
937 await _policy_table(prisma_client).update_many(
938 where={
939 "policy_name": policy_name,
940 "version_status": "production",
941 },
942 data={
943 "version_status": "published",
944 "updated_at": now,
945 "updated_by": updated_by,
946 },
947 )
949 # Promote this version to production
950 updated = await _policy_table(prisma_client).update(
951 where={"policy_id": policy_id},
952 data={
953 "version_status": "production",
954 "production_at": now,
955 "updated_at": now,
956 "updated_by": updated_by,
957 },
958 )
960 # Update in-memory registry: remove old production (by name), add this one
961 self.remove_policy(policy_name)
962 policy: Final = self._parse_policy(
963 policy_name,
964 {
965 "inherit": updated.inherit,
966 "description": updated.description,
967 "guardrails": {
968 "add": updated.guardrails_add or [],
969 "remove": updated.guardrails_remove or [],
970 },
971 "condition": updated.condition,
972 "pipeline": updated.pipeline,
973 },
974 )
975 self.add_policy(policy_name, policy)
977 return _row_to_policy_db_response(updated)
978 except Exception as e:
979 verbose_proxy_logger.exception("Error updating version status: %s", e)
980 raise Exception(f"Error updating version status: {e}")
982 async def compare_versions(
983 self,
984 policy_id_a: str,
985 policy_id_b: str,
986 prisma_client: "PrismaClient",
987 ) -> PolicyVersionCompareResponse:
988 """
989 Compare two policy versions and return field-by-field diffs.
991 Args:
992 policy_id_a: First policy version ID
993 policy_id_b: Second policy version ID
994 prisma_client: The Prisma client instance
996 Returns:
997 PolicyVersionCompareResponse with both versions and field_diffs
998 """
999 try:
1000 a: Final = await _policy_table(prisma_client).find_unique(where={"policy_id": policy_id_a})
1001 b: Final = await _policy_table(prisma_client).find_unique(where={"policy_id": policy_id_b})
1002 if a is None: 1002 ↛ 1004line 1002 didn't jump to line 1004 because the condition on line 1002 was always true
1003 raise Exception(f"Policy {policy_id_a} not found")
1004 if b is None:
1005 raise Exception(f"Policy {policy_id_b} not found")
1007 resp_a: Final = _row_to_policy_db_response(a)
1008 resp_b: Final = _row_to_policy_db_response(b)
1010 # Compare fields that are part of policy content (not metadata)
1011 compare_fields: Final = (
1012 "inherit",
1013 "description",
1014 "guardrails_add",
1015 "guardrails_remove",
1016 "condition",
1017 "pipeline",
1018 )
1019 field_diffs: Final[dict[str, dict[str, object]]] = {}
1020 for field in compare_fields:
1021 val_a = getattr(resp_a, field)
1022 val_b = getattr(resp_b, field)
1023 if val_a != val_b:
1024 field_diffs[field] = {"version_a": val_a, "version_b": val_b}
1026 return PolicyVersionCompareResponse(
1027 version_a=resp_a,
1028 version_b=resp_b,
1029 field_diffs=field_diffs,
1030 )
1031 except Exception as e:
1032 verbose_proxy_logger.exception("Error comparing versions: %s", e)
1033 raise Exception(f"Error comparing versions: {e}")
1035 async def delete_all_versions(
1036 self,
1037 policy_name: str,
1038 prisma_client: "PrismaClient",
1039 ) -> Mapping[str, str]:
1040 """
1041 Delete all versions of a policy. Also removes from in-memory registry.
1043 Args:
1044 policy_name: Name of the policy
1045 prisma_client: The Prisma client instance
1047 Returns:
1048 Dict with "message" and optional "warning" if a config-defined policy took over.
1049 """
1050 try:
1051 await _policy_table(prisma_client).delete_many(where={"policy_name": policy_name})
1052 self.remove_policy(policy_name)
1053 message: Final = f"All versions of policy '{policy_name}' deleted successfully"
1054 if self.get_source(policy_name) == "config": 1054 ↛ 1055line 1054 didn't jump to line 1055 because the condition on line 1054 was never true
1055 return {
1056 "message": message,
1057 "warning": (
1058 "All DB versions were deleted. The config-defined policy with the same name is active again."
1059 ),
1060 }
1061 return {"message": message}
1062 except Exception as e:
1063 verbose_proxy_logger.exception("Error deleting all versions: %s", e)
1064 raise Exception(f"Error deleting all versions: {e}")
1067# Global singleton instance
1068_policy_registry: PolicyRegistry | None = None
1071def get_policy_registry() -> PolicyRegistry:
1072 """
1073 Get the global PolicyRegistry singleton.
1075 Returns:
1076 The global PolicyRegistry instance
1077 """
1078 global _policy_registry
1079 if _policy_registry is None:
1080 _policy_registry = PolicyRegistry()
1081 return _policy_registry