Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py: 19%
110 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#!/usr/bin/env python3
2"""
3Azure Text Moderation Native Guardrail Integrationfor LiteLLM
4"""
6from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, cast
8from fastapi import HTTPException
10from litellm._logging import verbose_proxy_logger
11from litellm.integrations.custom_guardrail import (
12 CustomGuardrail,
13 log_guardrail_information,
14)
15from litellm.proxy._types import UserAPIKeyAuth
16from litellm.types.guardrails import GuardrailEventHooks
17from litellm.types.utils import CallTypesLiteral, GenericGuardrailAPIInputs, LLMResponseTypes
19from .base import AzureGuardrailBase
21if TYPE_CHECKING: 21 ↛ 22line 21 didn't jump to line 22 because the condition on line 21 was never true
22 from litellm.caching.caching import DualCache
23 from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
24 from litellm.types.llms.openai import AllMessageValues
25 from litellm.types.proxy.guardrails.guardrail_hooks.azure.azure_text_moderation import (
26 AzureTextModerationGuardrailResponse,
27 )
28 from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
31class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardrail):
32 """
33 LiteLLM Built-in Guardrail for Azure Content Safety (Text Moderation).
35 This guardrail scans prompts and responses using the Azure Text Moderation API to detect
36 malicious content and policy violations based on severity thresholds.
38 Configuration:
39 guardrail_name: Name of the guardrail instance
40 api_key: Azure Text Moderation API key
41 api_base: Azure Text Moderation API endpoint
42 default_on: Whether to enable by default
43 """
45 use_native_lifecycle_hooks: ClassVar[bool] = True
47 default_severity_threshold: int = 2
49 @classmethod
50 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
51 return [
52 GuardrailEventHooks.pre_call,
53 GuardrailEventHooks.post_call,
54 ]
56 def __init__(
57 self,
58 guardrail_name: str,
59 api_key: str,
60 api_base: str,
61 severity_threshold: int | None = None,
62 severity_threshold_by_category: dict[str, int] | None = None,
63 **kwargs,
64 ):
65 """Initialize Azure Text Moderation guardrail handler."""
66 from litellm.types.proxy.guardrails.guardrail_hooks.azure.azure_text_moderation import (
67 AzureTextModerationRequestBodyOptionalParams,
68 )
70 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
71 # AzureGuardrailBase.__init__ stores api_key, api_base, api_version,
72 # async_handler and forwards the rest to CustomGuardrail.
73 super().__init__(
74 api_key=api_key,
75 api_base=api_base,
76 guardrail_name=guardrail_name,
77 **kwargs,
78 )
80 self.optional_params_request_body: AzureTextModerationRequestBodyOptionalParams = {
81 "categories": kwargs.get("categories")
82 or [
83 "Hate",
84 "Sexual",
85 "SelfHarm",
86 "Violence",
87 ],
88 "blocklistNames": cast(list[str] | None, kwargs.get("blocklistNames") or None),
89 "haltOnBlocklistHit": kwargs.get("haltOnBlocklistHit") or False,
90 "outputType": kwargs.get("outputType") or "FourSeverityLevels",
91 }
93 self.severity_threshold = int(severity_threshold) if severity_threshold else None
94 self.severity_threshold_by_category = severity_threshold_by_category
96 verbose_proxy_logger.info("Initialized Azure Text Moderation Guardrail: %s", guardrail_name)
98 @staticmethod
99 def get_config_model() -> type["GuardrailConfigModel"] | None:
100 from litellm.types.proxy.guardrails.guardrail_hooks.azure.azure_text_moderation import (
101 AzureContentSafetyTextModerationConfigModel,
102 )
104 return AzureContentSafetyTextModerationConfigModel
106 async def async_make_request(self, text: str) -> "AzureTextModerationGuardrailResponse":
107 """
108 Make a request to the Azure Text Moderation API.
110 Long texts are automatically split at word boundaries into chunks
111 that respect the Azure Content Safety 10 000-character limit. Each
112 chunk is analysed independently; a severity-threshold violation in
113 *any* chunk raises an HTTPException immediately.
114 """
115 from litellm.types.proxy.guardrails.guardrail_hooks.azure.azure_text_moderation import (
116 AzureTextModerationGuardrailRequestBody,
117 AzureTextModerationGuardrailResponse,
118 )
120 from .base import AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH
122 chunks: Final = self.split_text_by_words(text, AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH)
124 last_response: AzureTextModerationGuardrailResponse | None = None
126 for chunk in chunks:
127 request_body = AzureTextModerationGuardrailRequestBody(
128 text=chunk,
129 **self.optional_params_request_body,
130 )
131 response_json = await self._post_to_content_safety("text:analyze", cast(dict, request_body))
133 chunk_response = cast(AzureTextModerationGuardrailResponse, response_json)
135 # For multi-chunk texts the callers only see the final response,
136 # so we must check every intermediate chunk here to avoid silently
137 # swallowing a violation that appears in an earlier chunk.
138 try:
139 self.check_severity_threshold(response=chunk_response)
140 except HTTPException:
141 verbose_proxy_logger.warning(
142 "Azure Text Moderation: Violation detected in chunk of length %d",
143 len(chunk),
144 )
145 raise
147 last_response = chunk_response
149 # chunks is always non-empty (split_text_by_words guarantees ≥1 element)
150 assert last_response is not None
151 return last_response
153 @log_guardrail_information
154 async def apply_guardrail(
155 self,
156 inputs: GenericGuardrailAPIInputs,
157 request_data: dict,
158 input_type: Literal["request", "response"],
159 logging_obj: "LiteLLMLoggingObj | None" = None,
160 ) -> GenericGuardrailAPIInputs:
161 for text in inputs.get("texts") or ():
162 if text:
163 await self.async_make_request(text=text)
164 return inputs
166 def check_severity_threshold(self, response: "AzureTextModerationGuardrailResponse") -> Literal[True]:
167 """
168 - Check if threshold set by category
169 - Check if general severity threshold set
170 - If both none, use default_severity_threshold
171 """
173 if self.severity_threshold_by_category:
174 for category in response["categoriesAnalysis"]:
175 severity_category_threshold_item = self.severity_threshold_by_category.get(category["category"])
176 if (
177 severity_category_threshold_item is not None
178 and category["severity"] >= severity_category_threshold_item
179 ):
180 raise HTTPException(
181 status_code=400,
182 detail={
183 "error": "Azure Content Safety Guardrail: {} crossed severity {}, Got severity: {}".format(
184 category["category"],
185 self.severity_threshold_by_category.get(category["category"]),
186 category["severity"],
187 )
188 },
189 )
190 if self.severity_threshold:
191 for category in response["categoriesAnalysis"]:
192 if category["severity"] >= self.severity_threshold:
193 raise HTTPException(
194 status_code=400,
195 detail={
196 "error": "Azure Content Safety Guardrail: {} crossed severity {}, Got severity: {}".format(
197 category["category"],
198 self.severity_threshold,
199 category["severity"],
200 )
201 },
202 )
203 if self.severity_threshold is None and self.severity_threshold_by_category is None:
204 for category in response["categoriesAnalysis"]:
205 if category["severity"] >= self.default_severity_threshold:
206 raise HTTPException(
207 status_code=400,
208 detail={
209 "error": "Azure Content Safety Guardrail: {} crossed severity {}, Got severity: {}".format(
210 category["category"],
211 self.default_severity_threshold,
212 category["severity"],
213 )
214 },
215 )
216 return True
218 @log_guardrail_information
219 async def async_pre_call_hook(
220 self,
221 user_api_key_dict: "UserAPIKeyAuth",
222 cache: "DualCache",
223 data: dict[str, Any],
224 call_type: CallTypesLiteral,
225 ) -> dict[str, object] | None:
226 """
227 Pre-call hook to scan user prompts before sending to LLM.
229 Raises HTTPException if content should be blocked.
230 """
231 verbose_proxy_logger.info(
232 "Azure Text Moderation: Running pre-call prompt scan, on call_type: %s",
233 call_type,
234 )
235 new_messages: Final[list[AllMessageValues] | None] = data.get("messages")
236 if new_messages is None:
237 verbose_proxy_logger.warning("Azure Text Moderation: not running guardrail. No messages in data")
238 return data
239 user_prompt: Final = self.get_user_prompt(new_messages)
241 if user_prompt:
242 verbose_proxy_logger.info("Azure Text Moderation: User prompt: %s", user_prompt)
243 await self.async_make_request(
244 text=user_prompt,
245 )
246 else:
247 verbose_proxy_logger.warning("Azure Text Moderation: No text found")
248 return None
250 async def async_post_call_success_hook(
251 self,
252 data: dict,
253 user_api_key_dict: "UserAPIKeyAuth",
254 response: LLMResponseTypes,
255 ) -> LLMResponseTypes:
256 from litellm.types.utils import Choices, ModelResponse
258 if isinstance(response, ModelResponse) and response.choices:
259 for choice in response.choices:
260 if not isinstance(choice, Choices):
261 continue
262 content = _message_content_to_text(choice.message.content)
263 if not content:
264 continue
265 await self.async_make_request(
266 text=content,
267 )
268 return response
270 async def async_post_call_streaming_hook(self, user_api_key_dict: UserAPIKeyAuth, response: str) -> str:
271 try:
272 if response is not None and len(response) > 0:
273 await self.async_make_request(
274 text=response,
275 )
276 return response
277 except HTTPException as e:
278 import json
280 error_returned: Final = json.dumps({"error": e.detail})
281 return f"data: {error_returned}\n\n"
284def _message_content_to_text(content: object) -> str:
285 if isinstance(content, str):
286 return content
287 if isinstance(content, list):
288 text_parts: Final = [
289 item.get("text") for item in content if isinstance(item, dict) and isinstance(item.get("text"), str)
290 ]
291 return "\n".join(part for part in text_parts if part)
292 return ""