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

1#!/usr/bin/env python3 

2""" 

3Azure Text Moderation Native Guardrail Integrationfor LiteLLM 

4""" 

5 

6from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, cast 

7 

8from fastapi import HTTPException 

9 

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 

18 

19from .base import AzureGuardrailBase 

20 

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 

29 

30 

31class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardrail): 

32 """ 

33 LiteLLM Built-in Guardrail for Azure Content Safety (Text Moderation). 

34 

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. 

37 

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 """ 

44 

45 use_native_lifecycle_hooks: ClassVar[bool] = True 

46 

47 default_severity_threshold: int = 2 

48 

49 @classmethod 

50 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: 

51 return [ 

52 GuardrailEventHooks.pre_call, 

53 GuardrailEventHooks.post_call, 

54 ] 

55 

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 ) 

69 

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 ) 

79 

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 } 

92 

93 self.severity_threshold = int(severity_threshold) if severity_threshold else None 

94 self.severity_threshold_by_category = severity_threshold_by_category 

95 

96 verbose_proxy_logger.info("Initialized Azure Text Moderation Guardrail: %s", guardrail_name) 

97 

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 ) 

103 

104 return AzureContentSafetyTextModerationConfigModel 

105 

106 async def async_make_request(self, text: str) -> "AzureTextModerationGuardrailResponse": 

107 """ 

108 Make a request to the Azure Text Moderation API. 

109 

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 ) 

119 

120 from .base import AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH 

121 

122 chunks: Final = self.split_text_by_words(text, AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH) 

123 

124 last_response: AzureTextModerationGuardrailResponse | None = None 

125 

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)) 

132 

133 chunk_response = cast(AzureTextModerationGuardrailResponse, response_json) 

134 

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 

146 

147 last_response = chunk_response 

148 

149 # chunks is always non-empty (split_text_by_words guarantees ≥1 element) 

150 assert last_response is not None 

151 return last_response 

152 

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 

165 

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 """ 

172 

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 

217 

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. 

228 

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) 

240 

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 

249 

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 

257 

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 

269 

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 

279 

280 error_returned: Final = json.dumps({"error": e.detail}) 

281 return f"data: {error_returned}\n\n" 

282 

283 

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 ""