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

1import traceback 

2from typing import Final 

3 

4from fastapi import HTTPException 

5 

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) 

15 

16 

17class _PROXY_AzureContentSafety( 

18 CustomLogger 

19): # https://docs.litellm.ai/docs/observability/custom_callback#callback-class 

20 # Class variables or attributes 

21 

22 enforces_request_content: bool = True 

23 

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 

44 

45 self.thresholds = self._configure_thresholds(thresholds) 

46 

47 self.client = ContentSafetyClient(self.endpoint, AzureKeyCredential(self.api_key)) 

48 

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 } 

56 

57 if thresholds is None: 

58 return default_thresholds 

59 

60 for key, default in default_thresholds.items(): 

61 if key not in thresholds: 

62 thresholds[key] = default 

63 

64 return thresholds 

65 

66 def _compute_result(self, response): 

67 result: Final = {} 

68 

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 } 

77 

78 return result 

79 

80 async def test_violation(self, content: str, source: str | None = None): 

81 verbose_proxy_logger.debug("Testing Azure Content-Safety for: %s", content) 

82 

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 ) 

88 

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 

96 

97 result: Final = self._compute_result(response) 

98 verbose_proxy_logger.debug("Azure Content-Safety Result: %s", result) 

99 

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 ) 

111 

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

124 

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

132 

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 

142 

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

150 

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