Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/zscaler_ai_guard/zscaler_ai_guard.py: 14%
176 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# +-------------------------------------------------------------+
2#
3# Use Zscaler AI Guard for your LLM calls
4#
5# +-------------------------------------------------------------+
6import os
7from typing import TYPE_CHECKING, Final, Literal, Optional
9from fastapi import HTTPException
11from litellm._logging import verbose_proxy_logger
12from litellm.integrations.custom_guardrail import (
13 CustomGuardrail,
14 log_guardrail_information,
15)
16from litellm.llms.custom_httpx.http_handler import (
17 get_async_httpx_client,
18 httpxSpecialProvider,
19)
20from litellm.types.guardrails import GuardrailEventHooks
21from litellm.types.utils import GenericGuardrailAPIInputs
23if TYPE_CHECKING: 23 ↛ 24line 23 didn't jump to line 24 because the condition on line 23 was never true
24 from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
25 from litellm.types.guardrails import LitellmParams
26 from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
28DEFAULT_GUARDRAIL_TIMEOUT: Final = 5.0
31class ZscalerAIGuard(CustomGuardrail):
32 @classmethod
33 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
34 return [
35 GuardrailEventHooks.pre_call,
36 GuardrailEventHooks.post_call,
37 ]
39 def __init__(
40 self,
41 api_key: str | None = None,
42 api_base: str | None = None,
43 policy_id: int | None = None,
44 send_user_api_key_alias: bool | None = None,
45 send_user_api_key_user_id: bool | None = None,
46 send_user_api_key_team_id: bool | None = None,
47 timeout: float | None = None,
48 **kwargs,
49 ):
50 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
51 self.optional_params = kwargs
52 self.zscaler_ai_guard_url = api_base or os.getenv(
53 "ZSCALER_AI_GUARD_URL",
54 "https://api.us1.zseclipse.net/v1/detection/execute-policy",
55 )
56 self.policy_id = policy_id if policy_id is not None else int(os.getenv("ZSCALER_AI_GUARD_POLICY_ID", -1))
57 self.api_key = api_key or os.getenv("ZSCALER_AI_GUARD_API_KEY")
58 self.send_user_api_key_alias = (
59 send_user_api_key_alias
60 if send_user_api_key_alias is not None
61 else os.getenv("SEND_USER_API_KEY_ALIAS", "False").lower() in ("true", "1")
62 )
63 self.send_user_api_key_user_id = (
64 send_user_api_key_user_id
65 if send_user_api_key_user_id is not None
66 else os.getenv("SEND_USER_API_KEY_USER_ID", "False").lower() in ("true", "1")
67 )
68 self.send_user_api_key_team_id = (
69 send_user_api_key_team_id
70 if send_user_api_key_team_id is not None
71 else os.getenv("SEND_USER_API_KEY_TEAM_ID", "False").lower() in ("true", "1")
72 )
73 self.timeout = self._resolve_timeout(timeout)
75 verbose_proxy_logger.debug(
76 "send_user_api_key_alias: %s, \n send_user_api_key_user_id:%s, \n send_user_api_key_team_id:%s",
77 self.send_user_api_key_alias,
78 self.send_user_api_key_user_id,
79 self.send_user_api_key_team_id,
80 )
82 super().__init__(**kwargs)
84 verbose_proxy_logger.debug("ZscalerAIGuard Initializing ...")
86 @staticmethod
87 def _resolve_timeout(timeout: float | None) -> float:
88 """
89 Resolve the effective per-request timeout, falling back to the default
90 when it is unset or non-positive.
91 """
92 if timeout is None:
93 return DEFAULT_GUARDRAIL_TIMEOUT
95 if timeout <= 0:
96 verbose_proxy_logger.warning(
97 "Ignoring non-positive Zscaler AI Guard timeout %s, using %s seconds",
98 timeout,
99 DEFAULT_GUARDRAIL_TIMEOUT,
100 )
101 return DEFAULT_GUARDRAIL_TIMEOUT
103 return timeout
105 def update_in_memory_litellm_params(self, litellm_params: "LitellmParams") -> None:
106 super().update_in_memory_litellm_params(litellm_params)
107 self.timeout = self._resolve_timeout(litellm_params.timeout)
109 @staticmethod
110 def _resolve_metadata_value(request_data: dict | None, key: str) -> str | None:
111 """
112 Resolve metadata value from request_data, checking both metadata locations.
114 During pre-call: metadata is at request_data["metadata"][key]
115 During post-call: metadata is at request_data["litellm_metadata"][key]
116 (set by transform_user_api_key_dict_to_metadata which prefixes keys with 'user_api_key_')
118 Also handles key name mapping for UserAPIKeyAuth fields:
119 - key_alias -> user_api_key_key_alias (in litellm_metadata)
120 - user_id -> user_api_key_user_id
121 - team_id -> user_api_key_team_id
122 """
123 if request_data is None:
124 return None
126 # Check litellm_metadata first (set during post-call by guardrail framework)
127 litellm_metadata: Final = request_data.get("litellm_metadata", {})
128 if litellm_metadata:
129 value = litellm_metadata.get(key)
130 if value is not None:
131 return str(value).strip()
132 # Handle key_alias -> user_api_key_key_alias mapping
133 # transform_user_api_key_dict_to_metadata prefixes "key_alias" -> "user_api_key_key_alias"
134 if key == "user_api_key_alias":
135 value = litellm_metadata.get("user_api_key_key_alias")
136 if value is not None:
137 return str(value).strip()
139 # Then check regular metadata (set during pre-call by proxy_server)
140 metadata: Final = request_data.get("metadata", {})
141 if metadata:
142 value = metadata.get(key)
143 if value is not None:
144 return str(value).strip()
146 return None
148 @log_guardrail_information
149 async def apply_guardrail(
150 self,
151 inputs: "GenericGuardrailAPIInputs",
152 request_data: dict,
153 input_type: Literal["request", "response"],
154 logging_obj: Optional["LiteLLMLoggingObj"] = None,
155 ) -> "GenericGuardrailAPIInputs":
156 """
157 Apply Zscaler AI Guard guardrail to batch of texts.
159 Args:
160 inputs: Dictionary containing texts and optional images
161 request_data: Request data dictionary containing metadata
162 input_type: Whether this is a "request" or "response"
163 logging_obj: Optional logging object
165 Returns:
166 GenericGuardrailAPIInputs - texts unchanged if passed, images unchanged
168 Raises:
169 Exception: If content is blocked by Zscaler AI Guard
170 """
172 texts: Final = inputs.get("texts", [])
173 try:
174 verbose_proxy_logger.debug("ZscalerAIGuard: Checking %s text(s)", len(texts))
175 metadata: Final = request_data.get("metadata", {})
177 user_api_key_metadata: Final = metadata.get("user_api_key_metadata", {}) or {}
178 team_metadata: Final = metadata.get("team_metadata", {}) or {}
180 # Precedence for policy_id:
181 # 1. metadata.zguard_policy_id # request level
182 # 2. user_api_key_metadata.zguard_policy_id # Key level
183 # 3. team_metadata.zguard_policy_id # Team level
184 # 4. self.policy_id (from environment) # Global
185 policy_id: Final = (
186 metadata.get("zguard_policy_id")
187 if "zguard_policy_id" in metadata
188 else (
189 user_api_key_metadata.get("zguard_policy_id")
190 if "zguard_policy_id" in user_api_key_metadata
191 else (
192 team_metadata.get("zguard_policy_id") if "zguard_policy_id" in team_metadata else self.policy_id
193 )
194 )
195 )
196 verbose_proxy_logger.info("policy_id applied: %s", policy_id)
198 kwargs: Final = {}
199 if self.send_user_api_key_alias:
200 kwargs["user_api_key_alias"] = self._resolve_metadata_value(request_data, "user_api_key_alias") or "N/A"
201 if self.send_user_api_key_team_id:
202 kwargs["user_api_key_team_id"] = (
203 self._resolve_metadata_value(request_data, "user_api_key_team_id") or "N/A"
204 )
205 if self.send_user_api_key_user_id:
206 kwargs["user_api_key_user_id"] = (
207 self._resolve_metadata_value(request_data, "user_api_key_user_id") or "N/A"
208 )
209 verbose_proxy_logger.debug("inside apply_guardrail kwargs: %s", kwargs)
211 zscaler_ai_guard_result = None
212 direction: Final = "OUT" if input_type == "response" else "IN"
213 verbose_proxy_logger.debug("direction: %s", direction)
214 # Concatenate all texts and send to Zscaler AI Guard
215 if texts:
216 concatenated_text: Final = " ".join(texts)
217 zscaler_ai_guard_result = await self.make_zscaler_ai_guard_api_call(
218 zscaler_ai_guard_url=self.zscaler_ai_guard_url,
219 api_key=self.api_key,
220 policy_id=policy_id,
221 direction=direction,
222 content=concatenated_text,
223 **kwargs,
224 )
225 verbose_proxy_logger.debug("response from zscaler ai guards: %s", zscaler_ai_guard_result)
226 if zscaler_ai_guard_result and zscaler_ai_guard_result.get("action") == "BLOCK":
227 blocking_info: Final = zscaler_ai_guard_result.get("zscaler_ai_guard_response")
228 error_message = f"Content blocked by Zscaler AI Guard: {self.extract_blocking_info(blocking_info)}"
229 raise HTTPException(status_code=400, detail={"error": error_message})
230 except HTTPException:
231 raise
232 except Exception as e:
233 verbose_proxy_logger.error("ZscalerAIGuard: Failed to apply guardrail: %s", str(e))
234 raise e
236 verbose_proxy_logger.debug("ZscalerAIGuard: Successfully applied guardrail.")
237 return inputs
239 def extract_blocking_info(self, response):
240 """
241 Extracts transaction ID and blocking detector details from a response.
242 """
243 transaction_id: Final = response.get("transactionId", None)
245 # Find which detectors are invoked and blocking
246 blocking_detectors: Final = []
247 detector_responses: Final = response.get("detectorResponses", {})
248 for detector, details in detector_responses.items():
249 if details.get("action") == "BLOCK":
250 blocking_detectors.append(detector)
252 return {
253 "transactionId": transaction_id,
254 "blockingDetectors": blocking_detectors,
255 }
257 def _create_user_facing_error(self, reason: str):
258 """
259 create an error dictionary that return to use
260 """
261 return {
262 "error_type": "Zscaler AI Guard Error",
263 "reason": reason,
264 }
266 def _prepare_headers(self, api_key, **kwargs):
267 headers: Final = {
268 "Content-Type": "application/json",
269 "Authorization": f"Bearer {api_key}",
270 }
271 extra_headers: Final = headers.copy()
272 if self.send_user_api_key_alias:
273 verbose_proxy_logger.debug("kwargs: %s", kwargs)
274 user_api_key_alias: Final = kwargs.get("user_api_key_alias", "N/A")
275 verbose_proxy_logger.debug("kwargs user_api_key_alias: %s", user_api_key_alias)
276 extra_headers.update({"user-api-key-alias": user_api_key_alias})
278 if self.send_user_api_key_team_id:
279 user_api_key_team_id: Final = kwargs.get("user_api_key_team_id", "N/A")
280 extra_headers.update({"user-api-key-team-id": user_api_key_team_id})
282 if self.send_user_api_key_user_id:
283 user_api_key_user_id: Final = kwargs.get("user_api_key_user_id", "N/A")
284 extra_headers.update({"user-api-key-user-id": user_api_key_user_id})
286 verbose_proxy_logger.debug("extra_headers: %s", extra_headers)
287 return extra_headers
289 async def _send_request(self, url, headers, data):
290 async_client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback)
292 response: Final = await async_client.post(
293 f"{url}",
294 headers=headers,
295 json=data,
296 timeout=self.timeout,
297 )
298 response.raise_for_status()
299 return response
301 def _handle_response(self, response, direction):
302 # Raise exceptions on critical errors to stop the request
303 if response.status_code == 429: # Rate limit
304 verbose_proxy_logger.error("Zscaler AI Guard rate limit reached. Blocking request.")
305 user_facing_error = self._create_user_facing_error("Rate limit reached. status_code: 429")
306 # This exception will be caught by the proxy and returned to the user
307 raise HTTPException(status_code=500, detail=user_facing_error)
309 if response.status_code >= 500: # Server error
310 verbose_proxy_logger.error(
311 "Zscaler AI Guard service is unavailable (Status: %s). Blocking request.", response.status_code
312 )
313 user_facing_error = self._create_user_facing_error(f"Service is unavailable (HTTP {response.status_code})")
314 raise HTTPException(status_code=500, detail=user_facing_error)
316 if response.status_code == 200:
317 json_response: Final = response.json()
318 statusCode_in_response: Final = json_response.get("statusCode", None)
319 if statusCode_in_response == 200:
320 guardrail_result: Final = json_response.get("action", None)
321 verbose_proxy_logger.info("Zscaler AI Guard response: %s", json_response)
323 if guardrail_result == "BLOCK":
324 verbose_proxy_logger.info(
325 "Violated Zscaler AI Guard guardrail policy. zscaler_ai_guard_response: %s", json_response
326 )
327 return {
328 "action": "BLOCK",
329 "zscaler_ai_guard_response": json_response,
330 }
331 elif guardrail_result == "ALLOW" or guardrail_result == "DETECT":
332 verbose_proxy_logger.debug(
333 "%s is allowed by Zscaler AI Guard. guardrail_result: %s", direction, guardrail_result
334 )
335 return {
336 "action": "ALLOW",
337 "zscaler_ai_guard_response": json_response,
338 "direction": direction,
339 }
340 else:
341 verbose_proxy_logger.error(
342 "Action field in response is %s, expecting 'ALLOW', 'BLOCK' or 'DETECT'", guardrail_result
343 )
344 user_facing_error = self._create_user_facing_error(
345 f"Action field in response is {guardrail_result}, expecting 'ALLOW', 'BLOCK' or 'DETECT'"
346 )
347 raise HTTPException(status_code=500, detail=user_facing_error)
348 else:
349 errorMsg: Final = json_response.get("errorMsg", None)
350 verbose_proxy_logger.error("statusCode in response: %s, errorMsg: %s", statusCode_in_response, errorMsg)
351 user_facing_error = self._create_user_facing_error(
352 f"statusCode in response: {statusCode_in_response}, errorMsg: {errorMsg}"
353 )
354 raise HTTPException(status_code=500, detail=user_facing_error)
355 else:
356 verbose_proxy_logger.error("Zscaler AI Guard status_code - %s", response.status_code)
357 user_facing_error = self._create_user_facing_error(f"Response status code: {response.status_code}")
358 raise HTTPException(status_code=response.status_code, detail=user_facing_error)
360 async def make_zscaler_ai_guard_api_call(
361 self, zscaler_ai_guard_url, api_key, policy_id, direction, content, **kwargs
362 ):
363 """
364 Makes an API call to the Zscaler AI Guard service and handles retries, errors, and response parsing.
365 """
367 extra_headers: Final = self._prepare_headers(api_key, **kwargs)
369 data: Final = {
370 "direction": direction,
371 "content": content,
372 }
373 # Only include policyId when explicitly configured (policy_id >= 1)
374 # When policy_id is None, 0, or -1 (default), use resolve-and-execute-policy which infers
375 # the policy from headers (e.g., user-api-key-alias)
376 if policy_id is not None and policy_id >= 1:
377 data["policyId"] = policy_id
378 try:
379 response: Final = await self._send_request(zscaler_ai_guard_url, extra_headers, data)
380 return self._handle_response(response, direction)
381 except HTTPException:
382 raise
383 except Exception as e:
384 verbose_proxy_logger.error("%s. Blocking request.", e)
385 user_facing_error: Final = self._create_user_facing_error(f"{e}")
386 raise HTTPException(status_code=500, detail=user_facing_error)
388 @staticmethod
389 def get_config_model() -> type["GuardrailConfigModel"] | None:
390 from litellm.types.proxy.guardrails.guardrail_hooks.zscaler_ai_guard import (
391 ZscalerAIGuardConfigModel,
392 )
394 return ZscalerAIGuardConfigModel