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
« 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# +-------------------------------------------------------------+
8import os
9from collections.abc import AsyncGenerator, AsyncIterable
10from datetime import datetime
11from typing import TYPE_CHECKING, Any, Final, Literal, Optional
13import httpx
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)
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
42GUARDRAIL_NAME: Final = "enkryptai"
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)
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 )
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"
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"
71 # store kwargs as optional_params
72 self.optional_params = kwargs
74 # Set supported event hooks
75 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
77 super().__init__(guardrail_name=guardrail_name, **kwargs)
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 )
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.
93 Args:
94 prompt (str): The text to analyze for potential violations
95 request_data (dict): Optional request data for logging purposes
97 Returns:
98 EnkryptAIResponse: Response from the Enkrypt AI Guardrails API
99 """
100 start_time: Final = datetime.now()
102 payload: Final = {"text": prompt}
104 headers: Final = {"Content-Type": "application/json", "apikey": self.api_key}
106 # Add policy header if policy_name is set
107 if self.policy_name:
108 headers["x-enkrypt-policy"] = self.policy_name
110 verbose_proxy_logger.debug(
111 "EnkryptAI request to %s with payload: %s",
112 self.api_url,
113 payload,
114 )
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()
130 end_time = datetime.now()
131 duration = (end_time - start_time).total_seconds()
133 verbose_proxy_logger.debug(
134 "EnkryptAI response from %s with payload: %s",
135 self.api_url,
136 response_json,
137 )
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 )
152 return response_json
154 except httpx.HTTPError as e:
155 end_time = datetime.now()
156 duration = (end_time - start_time).total_seconds()
158 verbose_proxy_logger.error("EnkryptAI API request failed: %s", str(e))
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 )
172 raise
174 def _process_enkryptai_guardrails_response(self, response: EnkryptAIResponse) -> EnkryptAIProcessedResult:
175 """
176 Process the response from the Enkrypt AI Guardrails API
178 Args:
179 response: The response from the API with 'summary' and 'details' keys
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", {})
187 detected_attacks: Final[list[str]] = []
188 attack_details: Final[dict[str, Any]] = {}
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, {})
203 return {"attacks_detected": detected_attacks, "attack_details": attack_details}
205 def _determine_guardrail_status(self, response_json: EnkryptAIResponse) -> GuardrailStatus:
206 """
207 Determine the guardrail status based on EnkryptAI API response.
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"
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"]
222 if attacks_detected:
223 return "guardrail_intervened"
225 return "success"
227 except Exception as e:
228 verbose_proxy_logger.error("Error determining EnkryptAI guardrail status: %s", str(e))
229 return "guardrail_failed_to_respond"
231 def _create_error_message(self, processed_result: EnkryptAIProcessedResult) -> str:
232 """
233 Create a detailed error message from processed guardrail results.
235 Args:
236 processed_result: Processed response with detected attacks and their details
238 Returns:
239 Formatted error message string
240 """
241 attacks_detected: Final = processed_result["attacks_detected"]
242 attack_details: Final = processed_result["attack_details"]
244 error_message = f"Guardrail failed: {len(attacks_detected)} violation(s) detected\n\n"
246 for attack_type in attacks_detected:
247 error_message += f"- {attack_type.upper()}:\n"
248 details = attack_details.get(attack_type, {})
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"
267 return error_message.strip()
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")
283 from litellm.proxy.common_utils.callback_utils import (
284 add_guardrail_to_applied_guardrails_header,
285 )
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
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 )
301 verbose_proxy_logger.debug("Guardrails async_pre_call_hook result: %s", result)
303 # Process the guardrails response
304 processed_result = self._process_enkryptai_guardrails_response(result)
305 attacks_detected = processed_result["attacks_detected"]
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)
312 # Add guardrail to applied guardrails header
313 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
315 return data
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
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 )
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
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 )
347 verbose_proxy_logger.debug("Guardrails async_moderation_hook result: %s", result)
349 # Process the guardrails response
350 processed_result = self._process_enkryptai_guardrails_response(result)
351 attacks_detected = processed_result["attacks_detected"]
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)
358 # Add guardrail to applied guardrails header
359 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
361 return data
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
372 It can be used to reject a response
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
381 if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call) is not True:
382 return
384 verbose_proxy_logger.debug("async_post_call_success_hook response: %s", response)
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
396 if not has_text_content:
397 verbose_proxy_logger.warning("EnkryptAI: not running guardrail. No output text in response")
398 return
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 )
409 verbose_proxy_logger.debug("Guardrails async_post_call_success_hook result: %s", result)
411 # Process the guardrails response
412 processed_result = self._process_enkryptai_guardrails_response(result)
413 attacks_detected = processed_result["attacks_detected"]
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)
420 # Add guardrail to applied guardrails header
421 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
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.
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
440 Returns:
441 GenericGuardrailAPIInputs - texts unchanged if passed, images unchanged
443 Raises:
444 ValueError: If any attacks are detected
445 """
446 texts: Final = inputs.get("texts", [])
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"]
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)
463 return inputs
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
474 This is useful for guardrails that need to see the entire response, such as PII masking.
476 See Aim guardrail implementation for an example - https://github.com/BerriAI/litellm/blob/d0e022cfacb8e9ebc5409bb652059b6fd97b45c0/litellm/proxy/guardrails/guardrail_hooks/aim.py#L168
478 Triggered by mode: 'post_call'
479 """
480 async for item in response:
481 yield item
483 @staticmethod
484 def get_config_model():
485 from litellm.types.proxy.guardrails.guardrail_hooks.enkryptai import (
486 EnkryptAIGuardrailConfigModel,
487 )
489 return EnkryptAIGuardrailConfigModel
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 ]