From 528460a89566115d00f60fd96b2a212ff1a13232 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Fri, 23 Jan 2026 12:24:06 -0800 Subject: [PATCH] fix: endpoints to resolve guardrails --- .../proxy/policy_engine/policy_endpoints.py | 63 +++++++++++++ .../proxy/policy_engine/policy_registry.py | 94 +++++++++++++++---- 2 files changed, 139 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/policy_engine/policy_endpoints.py b/litellm/proxy/policy_engine/policy_endpoints.py index a4aac7caf2c..98c1887cfcc 100644 --- a/litellm/proxy/policy_engine/policy_endpoints.py +++ b/litellm/proxy/policy_engine/policy_endpoints.py @@ -292,6 +292,69 @@ async def delete_policy(policy_id: str): raise HTTPException(status_code=500, detail=str(e)) +@router.get( + "/policies/{policy_id}/resolved-guardrails", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], +) +async def get_resolved_guardrails(policy_id: str): + """ + Get the resolved guardrails for a policy (including inherited guardrails). + + This endpoint resolves the full inheritance chain and returns the final + set of guardrails that would be applied for this policy. + + Example Request: + ```bash + curl -X GET "http://localhost:4000/policies/123e4567-e89b-12d3-a456-426614174000/resolved-guardrails" \\ + -H "Authorization: Bearer " + ``` + + Example Response: + ```json + { + "policy_id": "123e4567-e89b-12d3-a456-426614174000", + "policy_name": "healthcare-compliance", + "resolved_guardrails": ["pii_masking", "prompt_injection", "toxicity_filter"] + } + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + # Get the policy + policy = await POLICY_REGISTRY.get_policy_by_id_from_db( + policy_id=policy_id, + prisma_client=prisma_client, + ) + if policy is None: + raise HTTPException( + status_code=404, detail=f"Policy with ID {policy_id} not found" + ) + + # Resolve guardrails + resolved = await POLICY_REGISTRY.resolve_guardrails_from_db( + policy_name=policy.policy_name, + prisma_client=prisma_client, + ) + + return { + "policy_id": policy.policy_id, + "policy_name": policy.policy_name, + "resolved_guardrails": resolved, + } + except HTTPException: + raise + except ValueError as e: + raise HTTPException(status_code=400, detail=str(e)) + except Exception as e: + verbose_proxy_logger.exception(f"Error resolving guardrails: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + # ───────────────────────────────────────────────────────────────────────────── # Policy Attachment CRUD Endpoints # ───────────────────────────────────────────────────────────────────────────── diff --git a/litellm/proxy/policy_engine/policy_registry.py b/litellm/proxy/policy_engine/policy_registry.py index b83c4d3e246..cb4d688fce6 100644 --- a/litellm/proxy/policy_engine/policy_registry.py +++ b/litellm/proxy/policy_engine/policy_registry.py @@ -208,23 +208,30 @@ class PolicyRegistry: PolicyDBResponse with the created policy """ try: - condition_json = None - if policy_request.condition: - condition_json = safe_dumps(policy_request.condition.model_dump()) + from prisma import Json + + # Build data dict, only include condition if it's set + data: Dict[str, Any] = { + "policy_name": policy_request.policy_name, + "guardrails_add": policy_request.guardrails_add or [], + "guardrails_remove": policy_request.guardrails_remove or [], + "created_at": datetime.now(timezone.utc), + "updated_at": datetime.now(timezone.utc), + } + + # Only add optional fields if they have values + if policy_request.inherit is not None: + data["inherit"] = policy_request.inherit + if policy_request.description is not None: + data["description"] = policy_request.description + if created_by is not None: + data["created_by"] = created_by + data["updated_by"] = created_by + if policy_request.condition is not None: + data["condition"] = Json(policy_request.condition.model_dump()) created_policy = await prisma_client.db.litellm_policytable.create( - data={ - "policy_name": policy_request.policy_name, - "inherit": policy_request.inherit, - "description": policy_request.description, - "guardrails_add": policy_request.guardrails_add or [], - "guardrails_remove": policy_request.guardrails_remove or [], - "condition": condition_json, - "created_at": datetime.now(timezone.utc), - "updated_at": datetime.now(timezone.utc), - "created_by": created_by, - "updated_by": created_by, - } + data=data ) # Also add to in-memory registry @@ -298,9 +305,8 @@ class PolicyRegistry: if policy_request.guardrails_remove is not None: update_data["guardrails_remove"] = policy_request.guardrails_remove if policy_request.condition is not None: - update_data["condition"] = safe_dumps( - policy_request.condition.model_dump() - ) + from prisma import Json + update_data["condition"] = Json(policy_request.condition.model_dump()) updated_policy = await prisma_client.db.litellm_policytable.update( where={"policy_id": policy_id}, @@ -489,6 +495,58 @@ class PolicyRegistry: verbose_proxy_logger.exception(f"Error syncing policies from DB: {e}") raise Exception(f"Error syncing policies from DB: {str(e)}") + async def resolve_guardrails_from_db( + self, + policy_name: str, + prisma_client: "PrismaClient", + ) -> List[str]: + """ + Resolve all guardrails for a policy from the database. + + Uses the existing PolicyResolver to handle inheritance chain resolution. + + Args: + policy_name: Name of the policy to resolve + prisma_client: The Prisma client instance + + Returns: + List of resolved guardrail names + """ + from litellm.proxy.policy_engine.policy_resolver import PolicyResolver + + try: + # Load all policies from DB to ensure we have the full inheritance chain + policies = await self.get_all_policies_from_db(prisma_client) + + # Build a temporary in-memory map for resolution + temp_policies = {} + for policy_response in policies: + policy = self._parse_policy( + policy_response.policy_name, + { + "inherit": policy_response.inherit, + "description": policy_response.description, + "guardrails": { + "add": policy_response.guardrails_add, + "remove": policy_response.guardrails_remove, + }, + "condition": policy_response.condition, + }, + ) + temp_policies[policy_response.policy_name] = policy + + # Use the existing PolicyResolver to resolve guardrails + resolved_policy = PolicyResolver.resolve_policy_guardrails( + policy_name=policy_name, + policies=temp_policies, + context=None, # No context needed for simple resolution + ) + + return sorted(resolved_policy.guardrails) + except Exception as e: + verbose_proxy_logger.exception(f"Error resolving guardrails from DB: {e}") + raise Exception(f"Error resolving guardrails from DB: {str(e)}") + # Global singleton instance _policy_registry: Optional[PolicyRegistry] = None