Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py: 48%
462 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 MANAGEMENT
4All /policy management endpoints
6/policy/validate - Validate a policy configuration
7/policy/list - List all loaded policies
8/policy/info - Get information about a specific policy
9/policy/templates - Get policy templates (GitHub with local fallback)
10"""
12import copy
13import json
14import os
15from collections.abc import AsyncGenerator, AsyncIterator
16from typing import TYPE_CHECKING, Final, Literal, cast
18from fastapi import APIRouter, Depends, HTTPException, Request
19from fastapi.responses import Response, StreamingResponse
20from pydantic import BaseModel, Field
21from typing_extensions import TypedDict
23import litellm
24from litellm._logging import verbose_proxy_logger
25from litellm.constants import (
26 COMPETITOR_LLM_TEMPERATURE,
27 DEFAULT_COMPETITOR_DISCOVERY_MODEL,
28 MAX_COMPETITOR_NAMES,
29)
30from litellm.integrations.custom_guardrail import CustomGuardrail
31from litellm.llms.openai.chat.guardrail_translation.handler import (
32 OpenAIChatCompletionsHandler,
33)
34from litellm.proxy._types import UserAPIKeyAuth
35from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
36from litellm.proxy.common_utils.sse_keepalive import (
37 SSE_COMMENT_PING,
38 wrap_sse_stream_with_keepalive_pings,
39)
40from litellm.proxy.guardrails.guardrail_hooks.custom_code import (
41 RESPONSE_REJECTION_GUARDRAIL_CODE,
42 CustomCodeGuardrail,
43)
44from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry
45from litellm.proxy.management_helpers.utils import management_endpoint_wrapper
46from litellm.proxy.policy_engine.policy_registry import get_policy_registry
47from litellm.proxy.policy_engine.policy_resolver import PolicyResolver
48from litellm.types.proxy.policy_engine import (
49 PolicyGuardrailsResponse,
50 PolicyInfoResponse,
51 PolicyListResponse,
52 PolicyMatchContext,
53 PolicyScopeResponse,
54 PolicySummaryItem,
55 PolicyTestResponse,
56 PolicyValidateRequest,
57 PolicyValidationResponse,
58)
59from litellm.types.utils import GenericGuardrailAPIInputs, ModelResponse
61if TYPE_CHECKING: 61 ↛ 62line 61 didn't jump to line 62 because the condition on line 61 was never true
62 from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
64router: Final = APIRouter()
67class GuardrailApplyError(Exception):
68 """
69 Raised when a guardrail's apply_guardrail fails during apply_policies.
71 Consumers (e.g. Compliance UI) can use guardrail_name and message to show
72 which guardrail triggered and the error reason.
73 """
75 def __init__(self, guardrail_name: str, message: str) -> None:
76 self.guardrail_name = guardrail_name
77 self.message = message
78 super().__init__(f"Guardrail '{guardrail_name}' failed: {message}")
81class GuardrailErrorEntry(TypedDict):
82 """One guardrail failure for ApplyPoliciesResult.guardrail_errors."""
84 guardrail_name: str
85 message: str
88class _ApplyPoliciesResultBase(TypedDict):
89 """Base result of apply_policies: inputs plus any guardrail failures."""
91 inputs: GenericGuardrailAPIInputs
92 guardrail_errors: list[GuardrailErrorEntry]
95class ApplyPoliciesResult(_ApplyPoliciesResultBase, total=False):
96 """Result of apply_policies. agent_response set when agent_id provided."""
98 agent_response: object
101class _ApplyPoliciesPerItemResultBase(TypedDict):
102 """Base result for one input when using inputs_list."""
104 inputs: GenericGuardrailAPIInputs
105 guardrail_errors: list[GuardrailErrorEntry]
108class ApplyPoliciesPerItemResult(_ApplyPoliciesPerItemResultBase, total=False):
109 """Result for one input when using inputs_list. agent_response set when agent_id provided."""
111 agent_response: object
114class ApplyPoliciesListResult(TypedDict):
115 """Result when using inputs_list: one result per input."""
117 results: list[ApplyPoliciesPerItemResult]
120async def apply_policies(
121 policy_names: list[str] | None,
122 inputs: GenericGuardrailAPIInputs,
123 request_data: dict,
124 input_type: Literal["request", "response"],
125 proxy_logging_obj: "LiteLLMLoggingObj",
126 guardrail_names: list[str] | None = None,
127) -> ApplyPoliciesResult:
128 """
129 Apply guardrails to inputs from policy names and/or a direct list of guardrail names.
131 Runs all guardrails in order; if one fails, the error is recorded and execution
132 continues so that all inputs can complete testing and all guardrail failures are
133 collected. No exception is raised; failures are returned in guardrail_errors.
135 Guardrails can be specified in two ways (both can be used together; names are merged):
136 - policy_names: resolve guardrails from the policy registry (with inheritance).
137 - guardrail_names: use this list of guardrail names directly (no policy registry needed).
139 Returns:
140 ApplyPoliciesResult with "inputs" (final GenericGuardrailAPIInputs) and
141 "guardrail_errors" (list of {"guardrail_name", "message"} for each failure).
142 """
143 guardrail_errors: Final[list[GuardrailErrorEntry]] = []
145 guardrail_name_set: Final[set[str]] = set()
147 if guardrail_names: 147 ↛ 148line 147 didn't jump to line 148 because the condition on line 147 was never true
148 guardrail_name_set.update(guardrail_names)
150 if policy_names:
151 registry: Final = get_policy_registry()
152 if not registry.is_initialized(): 152 ↛ 153line 152 didn't jump to line 153 because the condition on line 152 was never true
153 verbose_proxy_logger.debug(
154 "apply_policies: policy engine not initialized, skipping policy-resolved guardrails"
155 )
156 else:
157 policies: Final = registry.get_all_policies()
158 for policy_name in policy_names:
159 resolved = PolicyResolver.resolve_policy_guardrails(
160 policy_name=policy_name,
161 policies=policies,
162 context=None,
163 )
164 guardrail_name_set.update(resolved.guardrails)
166 if not guardrail_name_set: 166 ↛ 169line 166 didn't jump to line 169 because the condition on line 166 was always true
167 return {"inputs": inputs, "guardrail_errors": guardrail_errors}
169 guardrail_registry: Final = GuardrailRegistry()
170 current_inputs = cast(GenericGuardrailAPIInputs, dict(inputs))
172 for guardrail_name in sorted(guardrail_name_set):
173 callback = guardrail_registry.get_initialized_guardrail_callback(guardrail_name=guardrail_name)
174 if callback is None:
175 verbose_proxy_logger.debug(
176 "apply_policies: guardrail '%s' not found, skipping",
177 guardrail_name,
178 )
179 continue
180 if not isinstance(callback, CustomGuardrail):
181 continue
182 if "apply_guardrail" not in type(callback).__dict__:
183 verbose_proxy_logger.debug(
184 "apply_policies: guardrail '%s' has no apply_guardrail, skipping",
185 guardrail_name,
186 )
187 continue
189 try:
190 current_inputs = await callback.apply_guardrail(
191 inputs=current_inputs,
192 request_data=request_data,
193 input_type=input_type,
194 logging_obj=proxy_logging_obj,
195 )
196 except Exception as e:
197 error_reason = str(e)
198 verbose_proxy_logger.debug(
199 "apply_policies: guardrail '%s' failed: %s",
200 guardrail_name,
201 error_reason,
202 )
203 guardrail_errors.append(
204 GuardrailErrorEntry(
205 guardrail_name=guardrail_name,
206 message=error_reason,
207 )
208 )
209 # Continue to next guardrail; current_inputs unchanged for this failure
211 return {"inputs": current_inputs, "guardrail_errors": guardrail_errors}
214def _chat_body_from_inputs(inputs: GenericGuardrailAPIInputs, agent_id: str, request_data: dict) -> dict:
215 """Build a chat completion request body from guardrail inputs and agent_id."""
216 messages: list[dict]
217 structured: Final = inputs.get("structured_messages")
218 texts: Final = inputs.get("texts")
219 if structured:
220 messages = list(structured)
221 elif texts: 221 ↛ 222line 221 didn't jump to line 222 because the condition on line 221 was never true
222 if len(texts) == 1:
223 messages = [{"role": "user", "content": texts[0]}]
224 else:
225 messages = [{"role": "user", "content": "\n".join(texts)}]
226 else:
227 messages = [{"role": "user", "content": "Hello"}]
228 body: Final[dict] = {"model": agent_id, "messages": messages, "stream": False}
229 if request_data: 229 ↛ 230line 229 didn't jump to line 230 because the condition on line 229 was never true
230 body.setdefault("metadata", {}).update(request_data)
231 return body
234def _request_with_json_body(body: dict) -> Request:
235 """Create a Starlette Request that will return the given dict as parsed JSON body."""
236 body_bytes: Final = json.dumps(body).encode()
237 received: Final[list[bool]] = [False]
239 async def receive() -> dict:
240 if received[0]: 240 ↛ 241line 240 didn't jump to line 241 because the condition on line 240 was never true
241 return {"type": "http.disconnect"}
242 received[0] = True
243 return {"type": "http.request", "body": body_bytes, "more_body": False}
245 scope: Final[dict] = {
246 "type": "http",
247 "method": "POST",
248 "path": "/v1/chat/completions",
249 "query_string": b"",
250 "headers": [(b"content-type", b"application/json")],
251 "scheme": "http",
252 "server": ("localhost", 8000),
253 "client": ("127.0.0.1", 0),
254 "root_path": "",
255 "app": None,
256 "asgi": {"version": "3.0", "spec_version": "2.0"},
257 }
258 return Request(scope, receive=receive)
261class TestPoliciesAndGuardrailsRequest(BaseModel):
262 """Request body for POST /utils/test_policies_and_guardrails."""
264 policy_names: list[str] | None = Field(default=None, description="Policy names to resolve guardrails from")
265 guardrail_names: list[str] | None = Field(default=None, description="Guardrail names to apply directly")
266 inputs_list: list[GenericGuardrailAPIInputs] = Field(
267 default=[],
268 description="List of GenericGuardrailAPIInputs; each item processed separately (for batch compliance testing).",
269 )
270 request_data: dict = Field(default_factory=dict, description="Request context (model, user_id, etc.)")
271 input_type: Literal["request", "response"] = Field(
272 default="request", description="Whether inputs are request or response"
273 )
274 agent_id: str | None = Field(
275 default=None,
276 description="When set, call chat completion with this model/agent for each input and include the response in the result.",
277 )
280@router.post(
281 "/utils/test_policies_and_guardrails",
282 tags=["utils"],
283 dependencies=[Depends(user_api_key_auth)],
284)
285@management_endpoint_wrapper
286async def test_policies_and_guardrails(
287 request: Request,
288 data: TestPoliciesAndGuardrailsRequest,
289 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
290):
291 """
292 Apply policies and/or guardrails to inputs (for compliance UI testing).
294 Use inputs_list for batch testing: each input is processed as a separate call so
295 per-input block/allow and errors are returned.
297 Use inputs for a single call (legacy).
298 """
299 from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
300 from litellm.proxy.proxy_server import chat_completion, proxy_logging_obj
301 from litellm.proxy.utils import handle_exception_on_proxy
303 def _serialize_chat_response(response: object) -> object:
304 if isinstance(response, BaseModel):
305 return response.model_dump(exclude_unset=True)
306 if isinstance(response, dict):
307 return response
308 return response
310 async def _get_agent_response(
311 inputs: GenericGuardrailAPIInputs,
312 agent_id: str,
313 user_api_key_dict: UserAPIKeyAuth,
314 ) -> object:
315 body: Final = _chat_body_from_inputs(inputs, agent_id, data.request_data)
316 req: Final = _request_with_json_body(body)
317 resp: Final = Response()
318 result: Final = await chat_completion(
319 request=req,
320 fastapi_response=resp,
321 model=agent_id,
322 user_api_key_dict=user_api_key_dict,
323 )
324 return _serialize_chat_response(result)
326 try:
327 logging_obj: Final = cast(LiteLLMLoggingObj, proxy_logging_obj)
329 results: Final[list[ApplyPoliciesPerItemResult]] = []
330 for inp in data.inputs_list:
331 item_result = await apply_policies(
332 policy_names=data.policy_names,
333 inputs=inp,
334 request_data=data.request_data,
335 input_type=data.input_type,
336 proxy_logging_obj=logging_obj,
337 guardrail_names=data.guardrail_names,
338 )
339 item: ApplyPoliciesPerItemResult = {
340 "inputs": item_result["inputs"],
341 "guardrail_errors": item_result["guardrail_errors"],
342 }
343 if data.agent_id is not None:
344 item["agent_response"] = await _get_agent_response(
345 item_result["inputs"],
346 data.agent_id,
347 user_api_key_dict,
348 )
349 # run response through response_rejection_guardrail (reuses handler extraction + apply)
350 response_rejection_guardrail = CustomCodeGuardrail(
351 custom_code=RESPONSE_REJECTION_GUARDRAIL_CODE,
352 guardrail_name="response_rejection",
353 )
354 try:
355 model_response = ModelResponse.model_validate(item["agent_response"])
356 handler = OpenAIChatCompletionsHandler()
357 await handler.process_output_response(
358 response=model_response,
359 guardrail_to_apply=response_rejection_guardrail,
360 litellm_logging_obj=logging_obj,
361 user_api_key_dict=user_api_key_dict,
362 )
363 except Exception as guardrail_err:
364 item["guardrail_errors"] = list(item["guardrail_errors"])
365 detail = getattr(guardrail_err, "detail", None)
366 if isinstance(detail, dict) and "error" in detail:
367 err_msg = detail["error"]
368 else:
369 err_msg = str(detail if detail is not None else guardrail_err)
370 item["guardrail_errors"].append(
371 GuardrailErrorEntry(
372 guardrail_name="response_rejection",
373 message=err_msg,
374 )
375 )
376 results.append(item)
377 return ApplyPoliciesListResult(results=results)
378 raise ValueError("Either inputs or inputs_list must be provided")
379 except Exception as e:
380 raise handle_exception_on_proxy(e)
383@router.post(
384 "/policy/validate",
385 tags=["policy management"],
386 dependencies=[Depends(user_api_key_auth)],
387 response_model=PolicyValidationResponse,
388)
389@management_endpoint_wrapper
390async def validate_policy(
391 request: Request,
392 data: PolicyValidateRequest,
393 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
394) -> PolicyValidationResponse:
395 """
396 Validate a policy configuration before applying it.
398 Checks:
399 - All referenced guardrails exist in the guardrail registry
400 - All non-wildcard team aliases exist in the database
401 - All non-wildcard key aliases exist in the database
402 - Inheritance chains are valid (no cycles, parents exist)
403 - Scope patterns are syntactically valid
405 Returns:
406 - valid: True if the policy configuration is valid (no blocking errors)
407 - errors: List of blocking validation errors
408 - warnings: List of non-blocking validation warnings
410 Example request:
411 ```json
412 {
413 "policies": {
414 "global-baseline": {
415 "guardrails": {
416 "add": ["pii_blocker", "phi_blocker"]
417 },
418 "scope": {
419 "teams": ["*"],
420 "keys": ["*"],
421 "models": ["*"]
422 }
423 },
424 "healthcare-compliance": {
425 "inherit": "global-baseline",
426 "guardrails": {
427 "add": ["hipaa_audit"]
428 },
429 "scope": {
430 "teams": ["healthcare-team"]
431 }
432 }
433 }
434 }
435 ```
436 """
437 from litellm.proxy.policy_engine.policy_validator import PolicyValidator
438 from litellm.proxy.proxy_server import prisma_client
440 verbose_proxy_logger.debug("Validating policy configuration with %s policies", len(data.policies))
442 validator: Final = PolicyValidator(prisma_client=prisma_client)
444 result: Final = await validator.validate_policy_config(
445 data.policies,
446 validate_db=prisma_client is not None,
447 )
449 return result
452@router.get(
453 "/policy/list",
454 tags=["policy management"],
455 dependencies=[Depends(user_api_key_auth)],
456 response_model=PolicyListResponse,
457)
458@management_endpoint_wrapper
459async def list_policies(
460 request: Request,
461 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
462) -> PolicyListResponse:
463 """
464 List all loaded policies with their resolved guardrails.
466 Returns information about each policy including:
467 - Inheritance configuration
468 - Scope (teams, keys, models)
469 - Guardrails to add/remove
470 - Resolved guardrails (after inheritance)
471 - Inheritance chain
472 """
473 from litellm.proxy.policy_engine.init_policies import get_policies_summary
475 summary: Final = get_policies_summary()
476 return PolicyListResponse(
477 policies={
478 name: PolicySummaryItem(
479 inherit=data.get("inherit"),
480 scope=PolicyScopeResponse(**data.get("scope", {})),
481 guardrails=PolicyGuardrailsResponse(**data.get("guardrails", {})),
482 resolved_guardrails=data.get("resolved_guardrails", []),
483 inheritance_chain=data.get("inheritance_chain", []),
484 )
485 for name, data in summary.get("policies", {}).items()
486 },
487 total_count=summary.get("total_count", 0),
488 )
491@router.get(
492 "/policy/info/{policy_name}",
493 tags=["policy management"],
494 dependencies=[Depends(user_api_key_auth)],
495 response_model=PolicyInfoResponse,
496)
497@management_endpoint_wrapper
498async def get_policy_info(
499 request: Request,
500 policy_name: str,
501 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
502) -> PolicyInfoResponse:
503 """
504 Get detailed information about a specific policy.
506 Returns:
507 - Policy configuration
508 - Resolved guardrails (after inheritance)
509 - Inheritance chain
510 """
511 from litellm.proxy.policy_engine.policy_registry import get_policy_registry
512 from litellm.proxy.policy_engine.policy_resolver import PolicyResolver
514 registry: Final = get_policy_registry()
516 if not registry.is_initialized(): 516 ↛ 517line 516 didn't jump to line 517 because the condition on line 516 was never true
517 raise HTTPException(
518 status_code=404,
519 detail="Policy engine not initialized. No policies loaded.",
520 )
522 policy: Final = registry.get_policy(policy_name)
523 if policy is None:
524 raise HTTPException(
525 status_code=404,
526 detail=f"Policy '{policy_name}' not found",
527 )
529 resolved = PolicyResolver.resolve_policy_guardrails(policy_name=policy_name, policies=registry.get_all_policies())
531 return PolicyInfoResponse(
532 policy_name=policy_name,
533 inherit=policy.inherit,
534 scope=PolicyScopeResponse(
535 teams=[],
536 keys=[],
537 models=[],
538 ),
539 guardrails=PolicyGuardrailsResponse(
540 add=policy.guardrails.get_add(),
541 remove=policy.guardrails.get_remove(),
542 ),
543 resolved_guardrails=resolved.guardrails,
544 inheritance_chain=resolved.inheritance_chain,
545 )
548@router.post(
549 "/policy/test",
550 tags=["policy management"],
551 dependencies=[Depends(user_api_key_auth)],
552 response_model=PolicyTestResponse,
553)
554@management_endpoint_wrapper
555async def test_policy_matching(
556 request: Request,
557 context: PolicyMatchContext,
558 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
559) -> PolicyTestResponse:
560 """
561 Test which policies would match a given request context.
563 This is useful for debugging and understanding policy behavior.
565 Request body:
566 ```json
567 {
568 "team_alias": "healthcare-team",
569 "key_alias": "my-api-key",
570 "model": "gpt-4"
571 }
572 ```
574 Returns:
575 - matching_policies: List of policy names that match
576 - resolved_guardrails: Final list of guardrails that would be applied
577 """
578 from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher
579 from litellm.proxy.policy_engine.policy_registry import get_policy_registry
580 from litellm.proxy.policy_engine.policy_resolver import PolicyResolver
582 registry: Final = get_policy_registry()
584 if not registry.is_initialized(): 584 ↛ 585line 584 didn't jump to line 585 because the condition on line 584 was never true
585 return PolicyTestResponse(
586 context=context,
587 matching_policies=[],
588 resolved_guardrails=[],
589 message="Policy engine not initialized. No policies loaded.",
590 )
592 policies: Final = registry.get_all_policies()
594 # Get matching policies
595 matching_policy_names: Final = PolicyMatcher.get_matching_policies(context=context)
597 # Resolve guardrails
598 resolved_guardrails: Final = PolicyResolver.resolve_guardrails_for_context(context=context, policies=policies)
600 return PolicyTestResponse(
601 context=context,
602 matching_policies=matching_policy_names,
603 resolved_guardrails=resolved_guardrails,
604 )
607POLICY_TEMPLATES_GITHUB_URL: Final = "https://raw.githubusercontent.com/BerriAI/litellm/main/policy_templates.json"
610def _load_policy_templates_from_local_backup() -> list:
611 """Load policy templates from local backup file (litellm/policy_templates_backup.json)."""
612 backup_path: Final = os.path.join(
613 os.path.dirname(__file__),
614 "..",
615 "..",
616 "..",
617 "policy_templates_backup.json",
618 )
619 path: Final = os.path.abspath(backup_path)
620 if not os.path.exists(path): 620 ↛ 621line 620 didn't jump to line 621 because the condition on line 620 was never true
621 return []
622 with open(path, "r") as f:
623 return json.load(f)
626@router.get(
627 "/policy/templates",
628 tags=["policy management"],
629 dependencies=[Depends(user_api_key_auth)],
630)
631@management_endpoint_wrapper
632async def get_policy_templates(
633 request: Request,
634 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
635) -> list:
636 """
637 Get policy templates for the UI (pre-configured guardrail combinations).
639 Fetches from GitHub with automatic fallback to local backup on failure.
640 Set LITELLM_LOCAL_POLICY_TEMPLATES=true to skip GitHub and use local backup only.
641 """
642 use_local: Final = os.getenv("LITELLM_LOCAL_POLICY_TEMPLATES", "").strip().lower() in (
643 "true",
644 "1",
645 "yes",
646 )
647 if use_local: 647 ↛ 648line 647 didn't jump to line 648 because the condition on line 647 was never true
648 return _load_policy_templates_from_local_backup()
650 try:
651 from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
652 from litellm.types.llms.custom_http import httpxSpecialProvider
654 async_client: Final = get_async_httpx_client(
655 llm_provider=httpxSpecialProvider.UI,
656 params={"timeout": 10.0},
657 )
658 response: Final = await async_client.get(POLICY_TEMPLATES_GITHUB_URL)
659 if response.status_code == 200: 659 ↛ 664line 659 didn't jump to line 664 because the condition on line 659 was always true
660 return response.json()
661 except Exception as e:
662 verbose_proxy_logger.debug("Failed to fetch policy templates from GitHub, using local backup: %s", e)
664 return _load_policy_templates_from_local_backup()
667class EnrichTemplateRequest(BaseModel):
668 template_id: str
669 parameters: dict
670 model: str | None = None
671 competitors: list[str] | None = Field(
672 default=None,
673 max_length=MAX_COMPETITOR_NAMES,
674 description="Optional list of competitor names",
675 )
676 instruction: str | None = Field(
677 default=None,
678 description="Refinement instruction for modifying the competitor list (e.g. 'add 10 more from Asia')",
679 )
682def _validate_enrichment_request(data: EnrichTemplateRequest) -> tuple[dict, dict, str]:
683 """
684 Validate enrichment request and return (template, llm_enrichment, brand_name).
686 Raises HTTPException on validation failure.
687 """
688 templates: Final = _load_policy_templates_from_local_backup()
689 template: Final = next((t for t in templates if t.get("id") == data.template_id), None)
690 if template is None:
691 raise HTTPException(status_code=404, detail=f"Template '{data.template_id}' not found")
693 llm_enrichment: Final = template.get("llm_enrichment")
694 if llm_enrichment is None:
695 raise HTTPException(status_code=400, detail="Template does not support LLM enrichment")
697 # Validate competitors list size if provided
698 if data.competitors and len(data.competitors) > MAX_COMPETITOR_NAMES: 698 ↛ 699line 698 didn't jump to line 699 because the condition on line 698 was never true
699 raise HTTPException(
700 status_code=400,
701 detail=f"competitors list exceeds maximum of {MAX_COMPETITOR_NAMES}",
702 )
704 brand_name: Final = data.parameters.get(llm_enrichment["parameter"], "")
705 if not brand_name: 705 ↛ 711line 705 didn't jump to line 711 because the condition on line 705 was always true
706 raise HTTPException(
707 status_code=400,
708 detail=f"Parameter '{llm_enrichment['parameter']}' is required",
709 )
711 return template, llm_enrichment, brand_name
714@router.post(
715 "/policy/templates/enrich",
716 tags=["policy management"],
717 dependencies=[Depends(user_api_key_auth)],
718)
719@management_endpoint_wrapper
720async def enrich_policy_template(
721 data: EnrichTemplateRequest,
722 request: Request,
723 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
724) -> dict:
725 """
726 Enrich a policy template with LLM-discovered data (e.g. competitor names).
728 Calls an onboarded LLM to discover competitors for the given brand name,
729 then returns enriched guardrailDefinitions with the discovered data populated.
730 """
731 template, llm_enrichment, brand_name = _validate_enrichment_request(data)
732 model: Final = data.model or DEFAULT_COMPETITOR_DISCOVERY_MODEL
734 if data.competitors:
735 competitors = data.competitors
736 else:
737 prompt: Final = llm_enrichment["prompt"].replace("{{" + llm_enrichment["parameter"] + "}}", brand_name)
738 competitors = await _discover_competitors_via_llm(prompt, model=model)
740 variations_map: Final = await _generate_competitor_variations(competitors, model=model)
742 enriched_definitions: Final = _build_competitor_guardrail_definitions(
743 template.get("guardrailDefinitions", []),
744 competitors,
745 brand_name,
746 variations_map,
747 )
749 return {
750 "guardrailDefinitions": enriched_definitions,
751 "competitors": competitors,
752 "competitor_variations": variations_map,
753 }
756def _build_refinement_prompt(
757 instruction: str,
758 existing_competitors: list[str],
759 brand_name: str,
760) -> str:
761 """Build a prompt for refining the competitor list based on user instruction."""
762 existing_list: Final = ", ".join(existing_competitors)
763 return (
764 f"I have a brand called '{brand_name}' and the following competitor list:\n"
765 f"{existing_list}\n\n"
766 f"User instruction: {instruction}\n\n"
767 "Return ONLY the NEW names to add (not the existing ones), one per line, "
768 "no numbering, no explanations. If the instruction asks to remove names, "
769 "return nothing."
770 )
773async def _stream_llm_competitor_names(
774 prompt: str,
775 model: str,
776 existing: list[str],
777) -> AsyncIterator[tuple[str | None, bool]]:
778 """
779 Stream competitor names from LLM. Yields (name, is_error) tuples.
781 Deduplicates against existing names (case-insensitive).
782 """
783 from litellm.proxy.proxy_server import llm_router
785 if llm_router is None:
786 raise ValueError("LLM router not initialized")
788 existing_lower: Final = {n.lower() for n in existing}
789 response: Final = await llm_router.acompletion(
790 model=model,
791 messages=[{"role": "user", "content": prompt}],
792 temperature=COMPETITOR_LLM_TEMPERATURE,
793 stream=True,
794 )
795 buffer = ""
796 count = len(existing)
797 async for chunk in response:
798 delta = chunk.choices[0].delta.content or ""
799 buffer += delta
800 while "\n" in buffer:
801 line, buffer = buffer.split("\n", 1)
802 name = _clean_competitor_line(line)
803 if name and name.lower() not in existing_lower and count < MAX_COMPETITOR_NAMES:
804 existing_lower.add(name.lower())
805 count += 1
806 yield name, False
807 # Handle remaining buffer
808 name = _clean_competitor_line(buffer)
809 if name and name.lower() not in existing_lower and count < MAX_COMPETITOR_NAMES:
810 yield name, False
813async def _stream_competitor_events(
814 data: EnrichTemplateRequest,
815 template: dict,
816 llm_enrichment: dict,
817 brand_name: str,
818 model: str,
819) -> AsyncGenerator[str, None]:
820 """Stream competitor names as SSE events, then emit a final 'done' event."""
821 competitors: Final[list[str]] = list(data.competitors or [])
823 if data.instruction and competitors:
824 # Refinement mode: keep existing, stream only new names
825 for comp in competitors:
826 yield f"data: {json.dumps({'type': 'competitor', 'name': comp})}\n\n"
828 refinement_prompt: Final = _build_refinement_prompt(data.instruction, competitors, brand_name)
829 try:
830 async for name, _ in _stream_llm_competitor_names(refinement_prompt, model, competitors):
831 if name:
832 competitors.append(name)
833 yield f"data: {json.dumps({'type': 'competitor', 'name': name})}\n\n"
834 except Exception as e:
835 verbose_proxy_logger.error("LLM competitor refinement failed: %s", e)
836 yield f"data: {json.dumps({'type': 'error', 'message': str(e)})}\n\n"
837 return
838 elif data.competitors and not data.instruction:
839 # Free-form mode (no instruction): just emit existing
840 for comp in competitors:
841 yield f"data: {json.dumps({'type': 'competitor', 'name': comp})}\n\n"
842 else:
843 # Initial discovery mode
844 prompt: Final = llm_enrichment["prompt"].replace("{{" + llm_enrichment["parameter"] + "}}", brand_name)
845 try:
846 async for name, _ in _stream_llm_competitor_names(prompt, model, []):
847 if name:
848 competitors.append(name)
849 yield f"data: {json.dumps({'type': 'competitor', 'name': name})}\n\n"
850 except Exception as e:
851 verbose_proxy_logger.error("LLM competitor streaming failed: %s", e)
852 yield f"data: {json.dumps({'type': 'error', 'message': str(e)})}\n\n"
853 return
855 yield f"data: {json.dumps({'type': 'status', 'message': f'Generating alternate spellings for {len(competitors)} competitors...'})}\n\n"
856 variations_map: Final = await _generate_competitor_variations(competitors, model=model)
858 total_variations: Final = sum(len(v) for v in variations_map.values())
859 yield f"data: {json.dumps({'type': 'status', 'message': f'Building guardrail definitions with {total_variations} variations...'})}\n\n"
860 enriched_definitions: Final = _build_competitor_guardrail_definitions(
861 template.get("guardrailDefinitions", []),
862 competitors,
863 brand_name,
864 variations_map,
865 )
867 yield f"data: {json.dumps({'type': 'done', 'competitors': competitors, 'competitor_variations': variations_map, 'guardrailDefinitions': enriched_definitions})}\n\n"
870@router.post(
871 "/policy/templates/enrich/stream",
872 tags=["policy management"],
873 dependencies=[Depends(user_api_key_auth)],
874)
875async def enrich_policy_template_stream(
876 data: EnrichTemplateRequest,
877 request: Request,
878 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
879):
880 """
881 Stream competitor names as SSE events as the LLM generates them.
883 Events:
884 - data: {"type": "competitor", "name": "..."} — each competitor as discovered
885 - data: {"type": "done", "competitors": [...], "competitor_variations": {...}, "guardrailDefinitions": [...]}
886 """
887 template, llm_enrichment, brand_name = _validate_enrichment_request(data)
888 model: Final = data.model or DEFAULT_COMPETITOR_DISCOVERY_MODEL
890 return StreamingResponse(
891 wrap_sse_stream_with_keepalive_pings(
892 _stream_competitor_events(data, template, llm_enrichment, brand_name, model),
893 ping_interval_seconds=litellm.sse_keepalive_ping_interval_seconds,
894 ping_chunk=SSE_COMMENT_PING,
895 ),
896 media_type="text/event-stream",
897 headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
898 )
901def _clean_competitor_line(line: str) -> str | None:
902 """Strip numbering, bullets, and whitespace from a competitor name line."""
903 name: Final = line.strip().strip(".-) ").strip()
904 return name if name and len(name) > 1 else None
907async def _generate_competitor_variations(competitors: list, model: str = DEFAULT_COMPETITOR_DISCOVERY_MODEL) -> dict:
908 """Generate common misspellings, abbreviations, and alternate names for each competitor."""
909 if not competitors:
910 return {}
912 # Cap the list to prevent oversized prompts
913 capped: Final = competitors[:MAX_COMPETITOR_NAMES]
914 names_list: Final = "\n".join(capped)
915 prompt: Final = (
916 "For each company/brand name below, list 3-5 common misspellings, abbreviations, "
917 "and alternate names that people might type. Include typos, missing spaces, "
918 "wrong suffixes (e.g. 'Airlines' vs 'Airways' vs 'Airline'), and common shortcuts.\n\n"
919 f"Names:\n{names_list}\n\n"
920 "Return the result as one line per variation in the format:\n"
921 "OriginalName: variation1, variation2, variation3\n"
922 "Use the EXACT original name before the colon. No numbering, no extra text."
923 )
925 try:
926 from litellm.proxy.proxy_server import llm_router
928 if llm_router is None:
929 raise ValueError("LLM router not initialized")
930 response: Final = await llm_router.acompletion(
931 model=model,
932 messages=[{"role": "user", "content": prompt}],
933 temperature=COMPETITOR_LLM_TEMPERATURE,
934 )
935 raw: Final = response.choices[0].message.content or ""
936 return _parse_variations_response(raw, capped)
937 except Exception as e:
938 verbose_proxy_logger.error("LLM competitor variation generation failed: %s", e)
939 return {}
942def _parse_variations_response(raw: str, competitors: list) -> dict[str, list[str]]:
943 """Parse the LLM response for competitor variations into a name -> variations map."""
944 # Build a lowercase lookup for case-insensitive matching
945 lower_to_canonical: Final = {comp.lower(): comp for comp in competitors}
946 variations_map: Final[dict[str, list[str]]] = {}
948 for line in raw.strip().split("\n"):
949 if ":" not in line:
950 continue
951 name, _, variations_str = line.partition(":")
952 canonical = lower_to_canonical.get(name.strip().lower())
953 if canonical is None:
954 continue
955 variations = [
956 v.strip() for v in variations_str.split(",") if v.strip() and v.strip().lower() != canonical.lower()
957 ]
958 variations_map[canonical] = variations
960 return variations_map
963async def _discover_competitors_via_llm(prompt: str, model: str = DEFAULT_COMPETITOR_DISCOVERY_MODEL) -> list:
964 """Call an onboarded LLM to discover competitor names."""
965 try:
966 from litellm.proxy.proxy_server import llm_router
968 if llm_router is None:
969 raise ValueError("LLM router not initialized")
970 response: Final = await llm_router.acompletion(
971 model=model,
972 messages=[{"role": "user", "content": prompt}],
973 temperature=COMPETITOR_LLM_TEMPERATURE,
974 )
975 raw: Final = response.choices[0].message.content or ""
976 competitors = [name for line in raw.strip().split("\n") if (name := _clean_competitor_line(line)) is not None]
977 return competitors[:MAX_COMPETITOR_NAMES]
978 except Exception as e:
979 verbose_proxy_logger.error("LLM competitor discovery failed: %s", e)
980 return []
983def _build_all_names_per_competitor(
984 competitors: list[str], variations_map: dict[str, list[str]]
985) -> dict[str, list[str]]:
986 """Build canonical + variation name lists for each competitor."""
987 return {comp: [comp] + variations_map.get(comp, []) for comp in competitors}
990def _build_competitor_guardrail_definitions(
991 definitions: list,
992 competitors: list,
993 brand_name: str,
994 variations_map: dict | None = None,
995) -> list:
996 """Build enriched guardrailDefinitions with competitor names and variations populated."""
997 variations_map = variations_map or {}
998 enriched: Final = copy.deepcopy(definitions)
999 all_names: Final = _build_all_names_per_competitor(competitors, variations_map)
1001 output_blocked: Final = _build_name_blocked_words(competitors, all_names)
1002 recommendation_blocked: Final = _build_recommendation_blocked_words(competitors, all_names)
1003 comparison_blocked: Final = _build_comparison_blocked_words(competitors, all_names, brand_name)
1005 blocked_words_map: Final = {
1006 "competitor-output-blocker": output_blocked,
1007 "competitor-input-blocker": output_blocked,
1008 "competitor-name-blocker": output_blocked,
1009 "competitor-name-input-blocker": output_blocked,
1010 "competitor-name-output-blocker": output_blocked,
1011 "competitor-recommendation-filter": recommendation_blocked,
1012 "competitor-recommendation-input-filter": recommendation_blocked,
1013 "competitor-recommendation-output-filter": recommendation_blocked,
1014 "competitor-comparison-filter": comparison_blocked,
1015 "competitor-comparison-input-filter": comparison_blocked,
1016 "competitor-comparison-output-filter": comparison_blocked,
1017 }
1019 for defn in enriched:
1020 guardrail_name = defn.get("guardrail_name", "")
1021 if guardrail_name in blocked_words_map:
1022 defn["litellm_params"]["blocked_words"] = blocked_words_map[guardrail_name]
1024 return enriched
1027def _build_name_blocked_words(competitors: list[str], all_names: dict[str, list[str]]) -> list[dict]:
1028 """Build blocked word entries for direct competitor name mentions."""
1029 result: Final = []
1030 for comp in competitors:
1031 for name in all_names[comp]:
1032 desc = f"Competitor: {comp}" if name == comp else f"Competitor variation ({comp}): {name}"
1033 result.append({"keyword": name, "action": "BLOCK", "description": desc})
1034 return result
1037def _build_recommendation_blocked_words(competitors: list[str], all_names: dict[str, list[str]]) -> list[dict]:
1038 """Build blocked word entries for competitor recommendations."""
1039 result: Final = []
1040 for comp in competitors:
1041 for name in all_names[comp]:
1042 for prefix in ["try", "use", "switch to", "consider"]:
1043 result.append(
1044 {
1045 "keyword": f"{prefix} {name}",
1046 "action": "BLOCK",
1047 "description": f"Recommendation to competitor ({comp})",
1048 }
1049 )
1050 return result
1053def _build_comparison_blocked_words(
1054 competitors: list[str], all_names: dict[str, list[str]], brand_name: str
1055) -> list[dict]:
1056 """Build blocked word entries for unfavorable competitor comparisons."""
1057 result: Final = []
1058 for comp in competitors:
1059 for name in all_names[comp]:
1060 result.append(
1061 {
1062 "keyword": f"{name} is better",
1063 "action": "BLOCK",
1064 "description": f"Unfavorable comparison ({comp})",
1065 }
1066 )
1068 # Brand-level comparisons (only need one entry each, not per-competitor)
1069 result.append(
1070 {
1071 "keyword": f"better than {brand_name}",
1072 "action": "BLOCK",
1073 "description": "Unfavorable comparison",
1074 }
1075 )
1076 result.append(
1077 {
1078 "keyword": f"{brand_name} is worse",
1079 "action": "BLOCK",
1080 "description": "Unfavorable comparison",
1081 }
1082 )
1084 return result
1087class SuggestTemplatesRequest(BaseModel):
1088 attack_examples: list[str] = Field(default_factory=list)
1089 description: str = Field(default="")
1090 model: str | None = None
1093@router.post(
1094 "/policy/templates/suggest",
1095 tags=["policy management"],
1096 dependencies=[Depends(user_api_key_auth)],
1097)
1098@management_endpoint_wrapper
1099async def suggest_policy_templates(
1100 data: SuggestTemplatesRequest,
1101 request: Request,
1102 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
1103) -> dict:
1104 """
1105 Use AI to suggest policy templates based on attack examples and descriptions.
1107 Calls an LLM with tool calling to match user requirements to available templates.
1108 """
1109 from litellm.proxy.management_endpoints.policy_endpoints.ai_policy_suggester import (
1110 AiPolicySuggester,
1111 )
1113 templates: Final = _load_policy_templates_from_local_backup()
1114 suggester: Final = AiPolicySuggester()
1115 return await suggester.suggest(
1116 templates=templates,
1117 attack_examples=data.attack_examples,
1118 description=data.description,
1119 model=data.model,
1120 )
1123class GuardrailTestResultEntry(TypedDict):
1124 guardrail_name: str
1125 action: str # "passed" | "blocked" | "masked" | "unsupported"
1126 output_text: str
1127 details: str
1130class TestPolicyTemplateRequest(BaseModel):
1131 guardrail_definitions: list[dict] = Field(description="All guardrailDefinitions from the policy template")
1132 text: str = Field(description="Test input text to run guardrails against")
1135class TestPolicyTemplateResponse(TypedDict):
1136 overall_action: str # worst-case across all guardrails
1137 results: list[GuardrailTestResultEntry]
1140@router.post(
1141 "/policy/templates/test",
1142 tags=["policy management"],
1143 dependencies=[Depends(user_api_key_auth)],
1144)
1145@management_endpoint_wrapper
1146async def test_policy_template(
1147 data: TestPolicyTemplateRequest,
1148 request: Request,
1149 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
1150) -> TestPolicyTemplateResponse:
1151 """
1152 Test a policy template's guardrails against a text input without creating them.
1154 Instantiates temporary guardrails from the template definitions, runs them
1155 against the provided text, and returns per-guardrail results so users can
1156 verify the template solves their problem before creating it.
1157 """
1158 from litellm.proxy.utils import handle_exception_on_proxy
1160 try:
1161 results: Final = await _test_guardrail_definitions(
1162 guardrail_definitions=data.guardrail_definitions,
1163 text=data.text,
1164 )
1165 overall: Final = _compute_overall_action(results)
1166 return TestPolicyTemplateResponse(
1167 overall_action=overall,
1168 results=results,
1169 )
1170 except Exception as e:
1171 raise handle_exception_on_proxy(e)
1174async def _test_guardrail_definitions(
1175 guardrail_definitions: list[dict],
1176 text: str,
1177) -> list[GuardrailTestResultEntry]:
1178 """Instantiate and run each guardrail definition against the text."""
1179 from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import (
1180 ContentFilterGuardrail,
1181 )
1183 results: Final[list[GuardrailTestResultEntry]] = []
1185 for guardrail_def in guardrail_definitions:
1186 guardrail_name = guardrail_def.get("guardrail_name", "unknown")
1187 litellm_params = guardrail_def.get("litellm_params", {})
1188 guardrail_type = litellm_params.get("guardrail", "")
1190 if guardrail_type != "litellm_content_filter": 1190 ↛ 1201line 1190 didn't jump to line 1201 because the condition on line 1190 was always true
1191 results.append(
1192 GuardrailTestResultEntry(
1193 guardrail_name=guardrail_name,
1194 action="unsupported",
1195 output_text=text,
1196 details=f"Preview not available for guardrail type: {guardrail_type}",
1197 )
1198 )
1199 continue
1201 try:
1202 guardrail = ContentFilterGuardrail(
1203 guardrail_name=guardrail_name,
1204 patterns=litellm_params.get("patterns"),
1205 blocked_words=litellm_params.get("blocked_words"),
1206 categories=litellm_params.get("categories"),
1207 pattern_redaction_format=litellm_params.get("pattern_redaction_format"),
1208 default_on=litellm_params.get("default_on", False),
1209 )
1211 output = await guardrail.apply_guardrail(
1212 inputs={"texts": [text]},
1213 request_data={},
1214 input_type="request",
1215 )
1216 output_text = output.get("texts", [text])[0] if output.get("texts") else text
1218 if output_text != text:
1219 action = "masked"
1220 details = "Content was modified (masked)"
1221 else:
1222 action = "passed"
1223 details = "No issues detected"
1225 results.append(
1226 GuardrailTestResultEntry(
1227 guardrail_name=guardrail_name,
1228 action=action,
1229 output_text=output_text,
1230 details=details,
1231 )
1232 )
1233 except HTTPException as e:
1234 detail = e.detail if hasattr(e, "detail") else str(e)
1235 if isinstance(detail, dict):
1236 detail = detail.get("error", str(detail))
1237 results.append(
1238 GuardrailTestResultEntry(
1239 guardrail_name=guardrail_name,
1240 action="blocked",
1241 output_text="",
1242 details=str(detail),
1243 )
1244 )
1245 except Exception as e:
1246 results.append(
1247 GuardrailTestResultEntry(
1248 guardrail_name=guardrail_name,
1249 action="error",
1250 output_text=text,
1251 details=str(e),
1252 )
1253 )
1255 return results
1258def _compute_overall_action(results: list[GuardrailTestResultEntry]) -> str:
1259 """Return the worst-case action: blocked > masked > error > unsupported > passed."""
1260 priority: Final = {"blocked": 4, "masked": 3, "error": 2, "unsupported": 1, "passed": 0}
1261 worst = "passed"
1262 for r in results:
1263 if priority.get(r["action"], 0) > priority.get(worst, 0):
1264 worst = r["action"]
1265 return worst