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

194 statements  

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

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

2# 

3# Use EnkryptAI Guardrails for your LLM calls 

4# https://enkryptai.com 

5# 

6# +-------------------------------------------------------------+ 

7 

8import os 

9from collections.abc import AsyncGenerator, AsyncIterable 

10from datetime import datetime 

11from typing import TYPE_CHECKING, Any, Final, Literal, Optional 

12 

13import httpx 

14 

15import litellm 

16from litellm._logging import verbose_proxy_logger 

17from litellm.caching.caching import DualCache 

18from litellm.integrations.custom_guardrail import ( 

19 CustomGuardrail, 

20 log_guardrail_information, 

21) 

22from litellm.llms.custom_httpx.http_handler import ( 

23 get_async_httpx_client, 

24 httpxSpecialProvider, 

25) 

26from litellm.proxy._types import UserAPIKeyAuth 

27from litellm.types.guardrails import GuardrailEventHooks 

28from litellm.types.proxy.guardrails.guardrail_hooks.enkryptai import ( 

29 EnkryptAIProcessedResult, 

30 EnkryptAIResponse, 

31) 

32from litellm.types.utils import ( 

33 CallTypesLiteral, 

34 GenericGuardrailAPIInputs, 

35 GuardrailStatus, 

36 ModelResponseStream, 

37) 

38 

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

40 from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj 

41 

42GUARDRAIL_NAME: Final = "enkryptai" 

43 

44 

45class EnkryptAIGuardrails(CustomGuardrail): 

46 def __init__( 

47 self, 

48 guardrail_name: str = "litellm_test", 

49 api_key: str | None = None, 

50 api_base: str | None = None, 

51 policy_name: str | None = None, 

52 **kwargs, 

53 ): 

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

55 

56 # Set API configuration 

57 self.api_key = api_key or os.getenv("ENKRYPTAI_API_KEY") 

58 if not self.api_key: 

59 raise ValueError( 

60 "EnkryptAI API key is required. Set ENKRYPTAI_API_KEY environment variable or pass api_key parameter." 

61 ) 

62 

63 self.api_base = api_base or os.getenv("ENKRYPTAI_API_BASE", "https://api.enkryptai.com") 

64 self.api_url = f"{self.api_base}/guardrails/policy/detect" 

65 

66 # Policy name can be passed as parameter or use guardrail_name 

67 self.policy_name = policy_name 

68 self.guardrail_name = guardrail_name 

69 self.guardrail_provider = "enkryptai" 

70 

71 # store kwargs as optional_params 

72 self.optional_params = kwargs 

73 

74 # Set supported event hooks 

75 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) 

76 

77 super().__init__(guardrail_name=guardrail_name, **kwargs) 

78 

79 verbose_proxy_logger.debug( 

80 "EnkryptAI Guardrail initialized with guardrail_name: %s, policy_name: %s", 

81 self.guardrail_name, 

82 self.policy_name, 

83 ) 

84 

85 async def _call_enkryptai_guardrails( 

86 self, 

87 prompt: str, 

88 request_data: dict | None = None, 

89 ) -> EnkryptAIResponse: 

90 """ 

91 Call Enkrypt AI Guardrails API to detect potential issues in the given prompt. 

92 

93 Args: 

94 prompt (str): The text to analyze for potential violations 

95 request_data (dict): Optional request data for logging purposes 

96 

97 Returns: 

98 EnkryptAIResponse: Response from the Enkrypt AI Guardrails API 

99 """ 

100 start_time: Final = datetime.now() 

101 

102 payload: Final = {"text": prompt} 

103 

104 headers: Final = {"Content-Type": "application/json", "apikey": self.api_key} 

105 

106 # Add policy header if policy_name is set 

107 if self.policy_name: 

108 headers["x-enkrypt-policy"] = self.policy_name 

109 

110 verbose_proxy_logger.debug( 

111 "EnkryptAI request to %s with payload: %s", 

112 self.api_url, 

113 payload, 

114 ) 

115 

116 try: 

117 verbose_proxy_logger.debug( 

118 "EnkryptAI request to %s with payload: %s", 

119 self.api_url, 

120 payload, 

121 ) 

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

123 url=self.api_url, 

124 json=payload, 

125 headers=headers, 

126 ) 

127 response.raise_for_status() 

128 response_json: Final = response.json() 

129 

130 end_time = datetime.now() 

131 duration = (end_time - start_time).total_seconds() 

132 

133 verbose_proxy_logger.debug( 

134 "EnkryptAI response from %s with payload: %s", 

135 self.api_url, 

136 response_json, 

137 ) 

138 

139 # Add guardrail information to request trace 

140 if request_data: 

141 guardrail_status: Final = self._determine_guardrail_status(response_json) 

142 self.add_standard_logging_guardrail_information_to_request_data( 

143 guardrail_provider=self.guardrail_provider, 

144 guardrail_json_response=response_json, 

145 request_data=request_data, 

146 guardrail_status=guardrail_status, 

147 start_time=start_time.timestamp(), 

148 end_time=end_time.timestamp(), 

149 duration=duration, 

150 ) 

151 

152 return response_json 

153 

154 except httpx.HTTPError as e: 

155 end_time = datetime.now() 

156 duration = (end_time - start_time).total_seconds() 

157 

158 verbose_proxy_logger.error("EnkryptAI API request failed: %s", str(e)) 

159 

160 # Add guardrail information with failure status 

161 if request_data: 

162 self.add_standard_logging_guardrail_information_to_request_data( 

163 guardrail_provider=self.guardrail_provider, 

164 guardrail_json_response={"error": str(e)}, 

165 request_data=request_data, 

166 guardrail_status="guardrail_failed_to_respond", 

167 start_time=start_time.timestamp(), 

168 end_time=end_time.timestamp(), 

169 duration=duration, 

170 ) 

171 

172 raise 

173 

174 def _process_enkryptai_guardrails_response(self, response: EnkryptAIResponse) -> EnkryptAIProcessedResult: 

175 """ 

176 Process the response from the Enkrypt AI Guardrails API 

177 

178 Args: 

179 response: The response from the API with 'summary' and 'details' keys 

180 

181 Returns: 

182 EnkryptAIProcessedResult: Processed response with detected attacks and their details 

183 """ 

184 summary: Final = response.get("summary", {}) 

185 details: Final = response.get("details", {}) 

186 

187 detected_attacks: Final[list[str]] = [] 

188 attack_details: Final[dict[str, Any]] = {} 

189 

190 for key, value in summary.items(): 

191 # Check if attack is detected 

192 # For toxicity, it's a list (non-empty list means detected) 

193 # For others, it's 1 for detected, 0 for not detected 

194 if key == "toxicity": 

195 if isinstance(value, list) and len(value) > 0: 

196 detected_attacks.append(key) 

197 attack_details[key] = details.get(key, {}) 

198 else: 

199 if value == 1: 

200 detected_attacks.append(key) 

201 attack_details[key] = details.get(key, {}) 

202 

203 return {"attacks_detected": detected_attacks, "attack_details": attack_details} 

204 

205 def _determine_guardrail_status(self, response_json: EnkryptAIResponse) -> GuardrailStatus: 

206 """ 

207 Determine the guardrail status based on EnkryptAI API response. 

208 

209 Returns: 

210 "success": Content allowed through with no violations 

211 "guardrail_intervened": Content blocked due to policy violations 

212 "guardrail_failed_to_respond": Technical error or API failure 

213 """ 

214 try: 

215 if not isinstance(response_json, dict): 

216 return "guardrail_failed_to_respond" 

217 

218 # Process the response to check for violations 

219 processed_result: Final = self._process_enkryptai_guardrails_response(response_json) 

220 attacks_detected: Final = processed_result["attacks_detected"] 

221 

222 if attacks_detected: 

223 return "guardrail_intervened" 

224 

225 return "success" 

226 

227 except Exception as e: 

228 verbose_proxy_logger.error("Error determining EnkryptAI guardrail status: %s", str(e)) 

229 return "guardrail_failed_to_respond" 

230 

231 def _create_error_message(self, processed_result: EnkryptAIProcessedResult) -> str: 

232 """ 

233 Create a detailed error message from processed guardrail results. 

234 

235 Args: 

236 processed_result: Processed response with detected attacks and their details 

237 

238 Returns: 

239 Formatted error message string 

240 """ 

241 attacks_detected: Final = processed_result["attacks_detected"] 

242 attack_details: Final = processed_result["attack_details"] 

243 

244 error_message = f"Guardrail failed: {len(attacks_detected)} violation(s) detected\n\n" 

245 

246 for attack_type in attacks_detected: 

247 error_message += f"- {attack_type.upper()}:\n" 

248 details = attack_details.get(attack_type, {}) 

249 

250 # Format details based on attack type 

251 if attack_type == "policy_violation": 

252 error_message += f" Policy: {details.get('violating_policy', 'N/A')}\n" 

253 error_message += f" Explanation: {details.get('explanation', 'N/A')}\n" 

254 elif attack_type == "pii": 

255 error_message += f" PII Detected: {details.get('pii', {})}\n" 

256 elif attack_type == "toxicity": 

257 toxic_types = [k for k, v in details.items() if isinstance(v, (int, float)) and v > 0.5] 

258 error_message += f" Types: {', '.join(toxic_types)}\n" 

259 elif attack_type == "keyword_detected": 

260 error_message += f" Keywords: {details.get('detected_keywords', [])}\n" 

261 elif attack_type == "bias": 

262 error_message += f" Bias Detected: {details.get('bias_detected', False)}\n" 

263 else: 

264 error_message += f" Details: {details}\n" 

265 error_message += "\n" 

266 

267 return error_message.strip() 

268 

269 async def async_pre_call_hook( 

270 self, 

271 user_api_key_dict: UserAPIKeyAuth, 

272 cache: DualCache, 

273 data: dict, 

274 call_type: CallTypesLiteral, 

275 ) -> Exception | str | dict | None: 

276 """ 

277 Runs before the LLM API call 

278 Runs on only Input 

279 Use this if you want to MODIFY the input 

280 """ 

281 verbose_proxy_logger.debug("Running EnkryptAI pre-call hook") 

282 

283 from litellm.proxy.common_utils.callback_utils import ( 

284 add_guardrail_to_applied_guardrails_header, 

285 ) 

286 

287 event_type: Final[GuardrailEventHooks] = GuardrailEventHooks.pre_call 

288 if self.should_run_guardrail(data=data, event_type=event_type) is not True: 

289 return data 

290 

291 _messages: Final = data.get("messages") 

292 if _messages: 

293 for message in _messages: 

294 _content = message.get("content") 

295 if isinstance(_content, str): 

296 result = await self._call_enkryptai_guardrails( 

297 prompt=_content, 

298 request_data=data, 

299 ) 

300 

301 verbose_proxy_logger.debug("Guardrails async_pre_call_hook result: %s", result) 

302 

303 # Process the guardrails response 

304 processed_result = self._process_enkryptai_guardrails_response(result) 

305 attacks_detected = processed_result["attacks_detected"] 

306 

307 # If any attacks are detected, raise an error 

308 if attacks_detected: 

309 error_message = self._create_error_message(processed_result) 

310 raise ValueError(error_message) 

311 

312 # Add guardrail to applied guardrails header 

313 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name) 

314 

315 return data 

316 

317 async def async_moderation_hook( 

318 self, 

319 data: dict, 

320 user_api_key_dict: UserAPIKeyAuth, 

321 call_type: CallTypesLiteral, 

322 ): 

323 """ 

324 Runs in parallel to LLM API call 

325 Runs on only Input 

326 

327 This can NOT modify the input, only used to reject or accept a call before going to LLM API 

328 """ 

329 from litellm.proxy.common_utils.callback_utils import ( 

330 add_guardrail_to_applied_guardrails_header, 

331 ) 

332 

333 event_type: Final[GuardrailEventHooks] = GuardrailEventHooks.during_call 

334 if self.should_run_guardrail(data=data, event_type=event_type) is not True: 

335 return 

336 

337 _messages: Final = data.get("messages") 

338 if _messages: 

339 for message in _messages: 

340 _content = message.get("content") 

341 if isinstance(_content, str): 

342 result = await self._call_enkryptai_guardrails( 

343 prompt=_content, 

344 request_data=data, 

345 ) 

346 

347 verbose_proxy_logger.debug("Guardrails async_moderation_hook result: %s", result) 

348 

349 # Process the guardrails response 

350 processed_result = self._process_enkryptai_guardrails_response(result) 

351 attacks_detected = processed_result["attacks_detected"] 

352 

353 # If any attacks are detected, raise an error 

354 if attacks_detected: 

355 error_message = self._create_error_message(processed_result) 

356 raise ValueError(error_message) 

357 

358 # Add guardrail to applied guardrails header 

359 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name) 

360 

361 return data 

362 

363 async def async_post_call_success_hook( 

364 self, 

365 data: dict, 

366 user_api_key_dict: UserAPIKeyAuth, 

367 response, 

368 ): 

369 """ 

370 Runs on response from LLM API call 

371 

372 It can be used to reject a response 

373 

374 Uses Enkrypt AI guardrails to check the response for policy violations, PII, and injection attacks 

375 """ 

376 from litellm.proxy.common_utils.callback_utils import ( 

377 add_guardrail_to_applied_guardrails_header, 

378 ) 

379 from litellm.types.guardrails import GuardrailEventHooks 

380 

381 if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call) is not True: 

382 return 

383 

384 verbose_proxy_logger.debug("async_post_call_success_hook response: %s", response) 

385 

386 # Check if the ModelResponse has text content in its choices 

387 # to avoid sending empty content to EnkryptAI (e.g., during tool calls) 

388 if isinstance(response, litellm.ModelResponse): 

389 has_text_content = False 

390 for choice in response.choices: 

391 if isinstance(choice, litellm.Choices): 

392 if choice.message.content and isinstance(choice.message.content, str): 

393 has_text_content = True 

394 break 

395 

396 if not has_text_content: 

397 verbose_proxy_logger.warning("EnkryptAI: not running guardrail. No output text in response") 

398 return 

399 

400 for choice in response.choices: 

401 if isinstance(choice, litellm.Choices): 

402 verbose_proxy_logger.debug("async_post_call_success_hook choice: %s", choice) 

403 if choice.message.content and isinstance(choice.message.content, str): 

404 result = await self._call_enkryptai_guardrails( 

405 prompt=choice.message.content, 

406 request_data=data, 

407 ) 

408 

409 verbose_proxy_logger.debug("Guardrails async_post_call_success_hook result: %s", result) 

410 

411 # Process the guardrails response 

412 processed_result = self._process_enkryptai_guardrails_response(result) 

413 attacks_detected = processed_result["attacks_detected"] 

414 

415 # If any attacks are detected, raise an error 

416 if attacks_detected: 

417 error_message = self._create_error_message(processed_result) 

418 raise ValueError(error_message) 

419 

420 # Add guardrail to applied guardrails header 

421 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name) 

422 

423 @log_guardrail_information 

424 async def apply_guardrail( 

425 self, 

426 inputs: "GenericGuardrailAPIInputs", 

427 request_data: dict, 

428 input_type: Literal["request", "response"], 

429 logging_obj: Optional["LiteLLMLoggingObj"] = None, 

430 ) -> "GenericGuardrailAPIInputs": 

431 """ 

432 Apply EnkryptAI guardrail to a batch of texts. 

433 

434 Args: 

435 inputs: Dictionary containing texts and optional images 

436 request_data: Request data dictionary containing metadata 

437 input_type: Whether this is a "request" or "response" 

438 logging_obj: Optional logging object 

439 

440 Returns: 

441 GenericGuardrailAPIInputs - texts unchanged if passed, images unchanged 

442 

443 Raises: 

444 ValueError: If any attacks are detected 

445 """ 

446 texts: Final = inputs.get("texts", []) 

447 

448 # Check each text for attacks 

449 for text in texts: 

450 result = await self._call_enkryptai_guardrails( 

451 prompt=text, 

452 request_data=request_data, 

453 ) 

454 # Process the guardrails response 

455 processed_result = self._process_enkryptai_guardrails_response(result) 

456 attacks_detected = processed_result["attacks_detected"] 

457 

458 # If any attacks are detected, raise an error 

459 if attacks_detected: 

460 error_message = self._create_error_message(processed_result) 

461 raise ValueError(error_message) 

462 

463 return inputs 

464 

465 async def async_post_call_streaming_iterator_hook( 

466 self, 

467 user_api_key_dict: UserAPIKeyAuth, 

468 response: AsyncIterable[ModelResponseStream], 

469 request_data: dict, 

470 ) -> AsyncGenerator[ModelResponseStream, None]: 

471 """ 

472 Passes the entire stream to the guardrail 

473 

474 This is useful for guardrails that need to see the entire response, such as PII masking. 

475 

476 See Aim guardrail implementation for an example - https://github.com/BerriAI/litellm/blob/d0e022cfacb8e9ebc5409bb652059b6fd97b45c0/litellm/proxy/guardrails/guardrail_hooks/aim.py#L168 

477 

478 Triggered by mode: 'post_call' 

479 """ 

480 async for item in response: 

481 yield item 

482 

483 @staticmethod 

484 def get_config_model(): 

485 from litellm.types.proxy.guardrails.guardrail_hooks.enkryptai import ( 

486 EnkryptAIGuardrailConfigModel, 

487 ) 

488 

489 return EnkryptAIGuardrailConfigModel 

490 

491 @classmethod 

492 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: 

493 return [ 

494 GuardrailEventHooks.pre_call, 

495 GuardrailEventHooks.post_call, 

496 GuardrailEventHooks.during_call, 

497 ]