Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/hooks/azure_content_safety.py: 0%
79 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
1import traceback
2from typing import Final
4from fastapi import HTTPException
6import litellm
7from litellm._logging import verbose_proxy_logger
8from litellm.caching.caching import DualCache
9from litellm.integrations.custom_logger import CustomLogger
10from litellm.proxy._types import UserAPIKeyAuth
11from litellm.proxy.guardrails._content_utils import (
12 is_text_content_call_type,
13 iter_message_text,
14)
17class _PROXY_AzureContentSafety(
18 CustomLogger
19): # https://docs.litellm.ai/docs/observability/custom_callback#callback-class
20 # Class variables or attributes
22 enforces_request_content: bool = True
24 def __init__(self, endpoint, api_key, thresholds=None):
25 try:
26 from azure.ai.contentsafety.aio import ContentSafetyClient
27 from azure.ai.contentsafety.models import (
28 AnalyzeTextOptions,
29 AnalyzeTextOutputType,
30 TextCategory,
31 )
32 from azure.core.credentials import AzureKeyCredential
33 from azure.core.exceptions import HttpResponseError
34 except Exception as e:
35 raise Exception(
36 f"\033[91mAzure Content-Safety not installed, try running 'pip install azure-ai-contentsafety' to fix this error: {e}\n{traceback.format_exc()}\033[0m"
37 )
38 self.endpoint = endpoint
39 self.api_key = api_key
40 self.text_category = TextCategory
41 self.analyze_text_options = AnalyzeTextOptions
42 self.analyze_text_output_type = AnalyzeTextOutputType
43 self.azure_http_error = HttpResponseError
45 self.thresholds = self._configure_thresholds(thresholds)
47 self.client = ContentSafetyClient(self.endpoint, AzureKeyCredential(self.api_key))
49 def _configure_thresholds(self, thresholds=None):
50 default_thresholds: Final = {
51 self.text_category.HATE: 4,
52 self.text_category.SELF_HARM: 4,
53 self.text_category.SEXUAL: 4,
54 self.text_category.VIOLENCE: 4,
55 }
57 if thresholds is None:
58 return default_thresholds
60 for key, default in default_thresholds.items():
61 if key not in thresholds:
62 thresholds[key] = default
64 return thresholds
66 def _compute_result(self, response):
67 result: Final = {}
69 category_severity: Final = {item.category: item.severity for item in response.categories_analysis}
70 for category in self.text_category:
71 severity = category_severity.get(category)
72 if severity is not None:
73 result[category] = {
74 "filtered": severity >= self.thresholds[category],
75 "severity": severity,
76 }
78 return result
80 async def test_violation(self, content: str, source: str | None = None):
81 verbose_proxy_logger.debug("Testing Azure Content-Safety for: %s", content)
83 # Construct a request
84 request: Final = self.analyze_text_options(
85 text=content,
86 output_type=self.analyze_text_output_type.EIGHT_SEVERITY_LEVELS,
87 )
89 # Analyze text
90 try:
91 response: Final = await self.client.analyze_text(request)
92 except self.azure_http_error:
93 verbose_proxy_logger.debug("Error in Azure Content-Safety: %s", traceback.format_exc())
94 verbose_proxy_logger.debug(traceback.format_exc())
95 raise
97 result: Final = self._compute_result(response)
98 verbose_proxy_logger.debug("Azure Content-Safety Result: %s", result)
100 for key, value in result.items():
101 if value["filtered"]:
102 raise HTTPException(
103 status_code=400,
104 detail={
105 "error": "Violated content safety policy",
106 "source": source,
107 "category": key,
108 "severity": value["severity"],
109 },
110 )
112 async def async_pre_call_hook(
113 self,
114 user_api_key_dict: UserAPIKeyAuth,
115 cache: DualCache,
116 data: dict,
117 call_type: str, # "completion", "embeddings", "image_generation", "moderation"
118 ):
119 verbose_proxy_logger.debug("Inside Azure Content-Safety Pre-Call Hook")
120 try:
121 if is_text_content_call_type(call_type):
122 for text in iter_message_text(data):
123 await self.test_violation(content=text, source="input")
125 except HTTPException as e:
126 raise e
127 except Exception as e:
128 verbose_proxy_logger.error(
129 "litellm.proxy.hooks.azure_content_safety.py::async_pre_call_hook(): Exception occured - %s", e
130 )
131 verbose_proxy_logger.debug(traceback.format_exc())
133 async def async_post_call_success_hook(
134 self,
135 data: dict,
136 user_api_key_dict: UserAPIKeyAuth,
137 response,
138 ):
139 verbose_proxy_logger.debug("Inside Azure Content-Safety Post-Call Hook")
140 if not isinstance(response, litellm.ModelResponse):
141 return
143 for choice in response.choices:
144 if not isinstance(choice, litellm.utils.Choices):
145 continue
146 message = getattr(choice, "message", None)
147 content = getattr(message, "content", None)
148 if isinstance(content, str):
149 await self.test_violation(content=content, source="output")
151 # async def async_post_call_streaming_hook(
152 # self,
153 # user_api_key_dict: UserAPIKeyAuth,
154 # response: str,
155 # ):
156 # verbose_proxy_logger.debug("Inside Azure Content-Safety Call-Stream Hook")
157 # await self.test_violation(content=response, source="output")