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
« 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
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)
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
16# Azure Content Safety APIs have a 10,000 character limit per request.
17AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH: Final = 10000
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
23AZURE_CONTENT_SAFETY_DEFAULT_API_VERSION: Final = "2024-09-01"
24JAVELIN_API_VERSION_STORED_BY_OLDER_RELEASES: Final = "v1"
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
33class AzureGuardrailBase:
34 """
35 Base class for Azure guardrails.
37 Provides shared initialisation (API credentials, HTTP client) and
38 utilities (text splitting, authenticated POST) used by all Azure
39 Content Safety guardrails.
40 """
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)
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")
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.
60 Args:
61 endpoint_path: The API action, e.g. ``"text:shieldPrompt"`` or
62 ``"text:analyze"``.
63 request_body: JSON-serialisable request payload.
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 }
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
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.
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.
94 Args:
95 text: The text to split
96 max_length: Maximum character length of each chunk
98 Returns:
99 List of text chunks, each not exceeding max_length
100 """
101 if len(text) <= max_length:
102 return [text]
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)]
109 chunks: Final[list[str]] = []
110 current_chunk = ""
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 = ""
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:]
127 current_chunk = token
129 if current_chunk:
130 chunks.append(current_chunk)
132 return chunks
134 def get_user_prompt(self, messages: list["AllMessageValues"]) -> str | None:
135 """
136 Get the last consecutive block of messages from the user.
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)