Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/azure/base.py: 26%

53 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1import re 

2from typing import TYPE_CHECKING, Any, Final 

3 

4from litellm._logging import verbose_proxy_logger 

5from litellm.litellm_core_utils.prompt_templates.common_utils import ( 

6 get_last_user_message, 

7) 

8from litellm.llms.custom_httpx.http_handler import ( 

9 get_async_httpx_client, 

10 httpxSpecialProvider, 

11) 

12 

13if TYPE_CHECKING: 13 ↛ 14line 13 didn't jump to line 14 because the condition on line 13 was never true

14 from litellm.types.llms.openai import AllMessageValues 

15 

16# Azure Content Safety APIs have a 10,000 character limit per request. 

17AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH: Final = 10000 

18 

19# Azure Content Safety bills text in 1,000-character "text records"; a submitted 

20# chunk of N characters consumes ceil(N / 1000) text records. 

21AZURE_CONTENT_SAFETY_TEXT_RECORD_LENGTH: Final = 1000 

22 

23AZURE_CONTENT_SAFETY_DEFAULT_API_VERSION: Final = "2024-09-01" 

24JAVELIN_API_VERSION_STORED_BY_OLDER_RELEASES: Final = "v1" 

25 

26 

27def resolve_content_safety_api_version(configured: str | None) -> str: 

28 if not configured or configured == JAVELIN_API_VERSION_STORED_BY_OLDER_RELEASES: 

29 return AZURE_CONTENT_SAFETY_DEFAULT_API_VERSION 

30 return configured 

31 

32 

33class AzureGuardrailBase: 

34 """ 

35 Base class for Azure guardrails. 

36 

37 Provides shared initialisation (API credentials, HTTP client) and 

38 utilities (text splitting, authenticated POST) used by all Azure 

39 Content Safety guardrails. 

40 """ 

41 

42 def __init__( 

43 self, 

44 api_key: str, 

45 api_base: str, 

46 **kwargs: Any, 

47 ): 

48 # Forward remaining kwargs to the next class in the MRO 

49 # (typically CustomGuardrail). 

50 super().__init__(**kwargs) 

51 

52 self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) 

53 self.api_key = api_key 

54 self.api_base = api_base 

55 self.api_version: str | None = kwargs.get("api_version") 

56 

57 async def _post_to_content_safety(self, endpoint_path: str, request_body: dict[str, object]) -> dict[str, Any]: 

58 """POST to an Azure Content Safety endpoint with standard auth headers. 

59 

60 Args: 

61 endpoint_path: The API action, e.g. ``"text:shieldPrompt"`` or 

62 ``"text:analyze"``. 

63 request_body: JSON-serialisable request payload. 

64 

65 Returns: 

66 Parsed JSON response dict. 

67 """ 

68 api_version: Final = resolve_content_safety_api_version(self.api_version) 

69 url: Final = f"{self.api_base}/contentsafety/{endpoint_path}?api-version={api_version}" 

70 headers: Final = { 

71 "Ocp-Apim-Subscription-Key": self.api_key, 

72 "Content-Type": "application/json", 

73 } 

74 

75 verbose_proxy_logger.debug("Azure Content Safety request [%s]: %s", endpoint_path, request_body) 

76 response: Final = await self.async_handler.post( 

77 url=url, 

78 headers=headers, 

79 json=request_body, 

80 ) 

81 response_json: Final[dict[str, Any]] = response.json() 

82 verbose_proxy_logger.debug("Azure Content Safety response [%s]: %s", endpoint_path, response_json) 

83 return response_json 

84 

85 @staticmethod 

86 def split_text_by_words(text: str, max_length: int) -> list[str]: 

87 """ 

88 Split text into chunks at word boundaries without breaking words. 

89 

90 Always returns at least one chunk. Short text (≤ max_length) is 

91 returned as a single-element list so callers can use a uniform 

92 loop without branching on length. 

93 

94 Args: 

95 text: The text to split 

96 max_length: Maximum character length of each chunk 

97 

98 Returns: 

99 List of text chunks, each not exceeding max_length 

100 """ 

101 if len(text) <= max_length: 

102 return [text] 

103 

104 # Tokenize into alternating non-whitespace and whitespace runs so 

105 # that original newlines, tabs, and multiple spaces are preserved 

106 # within each chunk. 

107 tokens: Final = [match.group(0) for match in re.finditer(r"\S+|\s+", text)] 

108 

109 chunks: Final[list[str]] = [] 

110 current_chunk = "" 

111 

112 for token in tokens: 

113 # Would appending this token exceed the limit? 

114 if len(current_chunk) + len(token) <= max_length: 

115 current_chunk += token 

116 else: 

117 # Flush whatever we have accumulated so far 

118 if current_chunk: 

119 chunks.append(current_chunk) 

120 current_chunk = "" 

121 

122 # Force-split any single token longer than max_length 

123 while len(token) > max_length: 

124 chunks.append(token[:max_length]) 

125 token = token[max_length:] 

126 

127 current_chunk = token 

128 

129 if current_chunk: 

130 chunks.append(current_chunk) 

131 

132 return chunks 

133 

134 def get_user_prompt(self, messages: list["AllMessageValues"]) -> str | None: 

135 """ 

136 Get the last consecutive block of messages from the user. 

137 

138 Example: 

139 messages = [ 

140 {"role": "user", "content": "Hello, how are you?"}, 

141 {"role": "assistant", "content": "I'm good, thank you!"}, 

142 {"role": "user", "content": "What is the weather in Tokyo?"}, 

143 ] 

144 get_user_prompt(messages) -> "What is the weather in Tokyo?" 

145 """ 

146 return get_last_user_message(messages)