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

1# +-------------------------------------------------------------+ 

2# 

3# Use Zscaler AI Guard for your LLM calls 

4# 

5# +-------------------------------------------------------------+ 

6import os 

7from typing import TYPE_CHECKING, Final, Literal, Optional 

8 

9from fastapi import HTTPException 

10 

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 

22 

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 

27 

28DEFAULT_GUARDRAIL_TIMEOUT: Final = 5.0 

29 

30 

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 ] 

38 

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) 

74 

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 ) 

81 

82 super().__init__(**kwargs) 

83 

84 verbose_proxy_logger.debug("ZscalerAIGuard Initializing ...") 

85 

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 

94 

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 

102 

103 return timeout 

104 

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) 

108 

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. 

113 

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

117 

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 

125 

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

138 

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

145 

146 return None 

147 

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. 

158 

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 

164 

165 Returns: 

166 GenericGuardrailAPIInputs - texts unchanged if passed, images unchanged 

167 

168 Raises: 

169 Exception: If content is blocked by Zscaler AI Guard 

170 """ 

171 

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", {}) 

176 

177 user_api_key_metadata: Final = metadata.get("user_api_key_metadata", {}) or {} 

178 team_metadata: Final = metadata.get("team_metadata", {}) or {} 

179 

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) 

197 

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) 

210 

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 

235 

236 verbose_proxy_logger.debug("ZscalerAIGuard: Successfully applied guardrail.") 

237 return inputs 

238 

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) 

244 

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) 

251 

252 return { 

253 "transactionId": transaction_id, 

254 "blockingDetectors": blocking_detectors, 

255 } 

256 

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 } 

265 

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

277 

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

281 

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

285 

286 verbose_proxy_logger.debug("extra_headers: %s", extra_headers) 

287 return extra_headers 

288 

289 async def _send_request(self, url, headers, data): 

290 async_client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback) 

291 

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 

300 

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) 

308 

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) 

315 

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) 

322 

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) 

359 

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

366 

367 extra_headers: Final = self._prepare_headers(api_key, **kwargs) 

368 

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) 

387 

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 ) 

393 

394 return ZscalerAIGuardConfigModel