Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/pangea/pangea.py: 21%
126 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# litellm/proxy/guardrails/guardrail_hooks/pangea.py
2import os
3from typing import TYPE_CHECKING, Final
5from fastapi import HTTPException
7from litellm._logging import verbose_proxy_logger
8from litellm.caching.dual_cache import DualCache
9from litellm.integrations.custom_guardrail import (
10 CustomGuardrail,
11 log_guardrail_information,
12)
13from litellm.llms.custom_httpx.http_handler import (
14 get_async_httpx_client,
15 httpxSpecialProvider,
16)
17from litellm.proxy._types import UserAPIKeyAuth
18from litellm.proxy.common_utils.callback_utils import (
19 add_guardrail_to_applied_guardrails_header,
20)
21from litellm.types.guardrails import GuardrailEventHooks
22from litellm.types.utils import (
23 Choices,
24 LLMResponseTypes,
25 ModelResponse,
26 TextCompletionResponse,
27)
29if TYPE_CHECKING: 29 ↛ 30line 29 didn't jump to line 30 because the condition on line 29 was never true
30 from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
33class PangeaGuardrailMissingSecrets(Exception):
34 """Custom exception for missing Pangea secrets."""
37class _TextCompletionRequest:
38 def __init__(self, body: dict[str, object]) -> None:
39 self.body = body
41 def get_messages(self) -> list[dict]:
42 return [{"role": "user", "content": self.body["prompt"]}]
44 # This mutates the original dict, but we'll still return it anyways
45 def update_original_body(self, prompt_messages: list[dict]) -> dict[str, object]:
46 assert len(prompt_messages) == 1
47 self.body["prompt"] = prompt_messages[0]["content"]
48 return self.body
51class PangeaHandler(CustomGuardrail):
52 """
53 Pangea AI Guardrail handler to interact with the Pangea AI Guard service.
55 This class implements the necessary hooks to call the Pangea AI Guard API
56 for input and output scanning based on the configured recipe.
57 """
59 def __init__(
60 self,
61 guardrail_name: str,
62 pangea_input_recipe: str | None = None,
63 pangea_output_recipe: str | None = None,
64 api_key: str | None = None,
65 api_base: str | None = None,
66 **kwargs,
67 ):
68 """
69 Initializes the PangeaHandler.
71 Args:
72 guardrail_name (str): The name of the guardrail instance.
73 pangea_recipe (str): The Pangea recipe key to use for scanning.
74 api_key (Optional[str]): The Pangea API key. Reads from PANGEA_API_KEY env var if None.
75 api_base (Optional[str]): The Pangea API base URL. Reads from PANGEA_API_BASE env var or uses default if None.
76 **kwargs: Additional arguments passed to the CustomGuardrail base class.
77 """
78 self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
79 self.api_key = api_key or os.environ.get("PANGEA_API_KEY")
80 if not self.api_key:
81 raise PangeaGuardrailMissingSecrets(
82 "Pangea API Key not found. Set PANGEA_API_KEY environment variable or pass it in litellm_params."
83 )
85 # Default Pangea base URL if not provided
86 self.api_base = api_base or os.environ.get("PANGEA_API_BASE") or "https://ai-guard.aws.us.pangea.cloud"
87 self.pangea_input_recipe = pangea_input_recipe
88 self.pangea_output_recipe = pangea_output_recipe
90 # Pass relevant kwargs to the parent class
91 super().__init__(
92 guardrail_name=guardrail_name,
93 supported_event_hooks=list(self.get_supported_event_hooks()),
94 **kwargs,
95 )
96 verbose_proxy_logger.debug(
97 "Initialized Pangea Guardrail: name=%s, recipe=%s, api_base=%s",
98 guardrail_name,
99 pangea_input_recipe,
100 self.api_base,
101 )
103 async def _call_pangea_ai_guard(self, api: str, payload: dict, hook_name: str) -> dict:
104 """
105 Makes the API call to the Pangea AI Guard endpoint.
106 The function itself will raise an error in the case that a response
107 should be blocked, but will return a list of redacted messages that the caller
108 should act on.
110 Args:
111 api (str): Which API to use (text/guard or v1beta/guard)
112 payload (dict): The request payload.
113 request_data (dict): Original request data (used for logging/headers).
114 hook_name (str): Name of the hook calling this function (for logging).
116 Raises:
117 HTTPException: If the Pangea API returns a 'blocked: true' response.
118 Exception: For other API call failures.
120 Returns:
121 list[dict]: The original response body
122 """
123 endpoint: Final = f"{self.api_base}/{api}"
125 headers: Final = {
126 "Authorization": f"Bearer {self.api_key}",
127 "Content-Type": "application/json",
128 }
130 verbose_proxy_logger.debug(
131 "Pangea Guardrail (%s): Calling endpoint %s with payload: %s", hook_name, endpoint, payload
132 )
134 response: Final = await self.async_handler.post(url=endpoint, json=payload, headers=headers)
135 response.raise_for_status()
137 result: Final = response.json()
139 if result.get("result", {}).get("blocked"):
140 verbose_proxy_logger.warning("Pangea Guardrail (%s): Request blocked. Response: %s", hook_name, result)
141 raise HTTPException(
142 status_code=400, # Bad Request, indicating violation
143 detail={
144 "error": "Violated Pangea guardrail policy",
145 "guardrail_name": self.guardrail_name,
146 },
147 )
148 verbose_proxy_logger.debug(
149 "Pangea Guardrail (%s): Request passed. Response: %s", hook_name, result.get("result", {}).get("detectors")
150 )
152 return result
154 async def _async_pre_call_hook(
155 self,
156 user_api_key_dict: UserAPIKeyAuth,
157 cache: DualCache,
158 data: dict,
159 call_type: str,
160 ):
161 transformer = None
162 messages: object = None
163 if call_type == "text_completion" or call_type == "atext_completion":
164 transformer = _TextCompletionRequest(data)
165 messages = transformer.get_messages()
166 else:
167 messages = data.get("messages")
169 ai_guard_payload: Final = {
170 "debug": False,
171 "input": {"messages": messages, "tools": data.get("tools")},
172 "event_type": "input",
173 }
174 if self.pangea_input_recipe:
175 ai_guard_payload["recipe"] = self.pangea_input_recipe
177 ai_guard_response = await self._call_pangea_ai_guard("v1beta/guard", ai_guard_payload, "async_pre_call_hook")
178 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
180 if not ai_guard_response.get("result", {}).get("transformed"):
181 return
183 output: Final = ai_guard_response.get("result", {}).get("output", {})
184 if call_type == "text_completion" or call_type == "atext_completion":
185 data = transformer.update_original_body(output["messages"])
186 else:
187 data["messages"] = output["messages"]
188 return data
190 @log_guardrail_information
191 async def async_pre_call_hook(
192 self,
193 user_api_key_dict: UserAPIKeyAuth,
194 cache: DualCache,
195 data: dict,
196 call_type: str,
197 ):
198 event_type: Final = GuardrailEventHooks.pre_call
199 if self.should_run_guardrail(data=data, event_type=event_type) is not True:
200 verbose_proxy_logger.debug(
201 "Pangea Guardrail (async_pre_call_hook): Guardrail is disabled %s.", self.guardrail_name
202 )
203 return data
205 try:
206 return await self._async_pre_call_hook(user_api_key_dict, cache, data, call_type)
207 except HTTPException:
208 raise
209 except Exception as e:
210 raise HTTPException(
211 status_code=500,
212 detail={
213 "error": "Error in Pangea Guardrail",
214 "guardrail_name": self.guardrail_name,
215 "exceptions": str(e),
216 },
217 ) from e
219 async def _async_post_call_success_hook(
220 self,
221 data: dict,
222 user_api_key_dict: UserAPIKeyAuth,
223 # This union isn't actually correct -- it can get other response types depending on the API called
224 response: LLMResponseTypes,
225 ):
226 if isinstance(response, TextCompletionResponse):
227 # Assume the earlier call type as well
228 input_messages = _TextCompletionRequest(data).get_messages()
229 elif isinstance(response, ModelResponse):
230 messages: Final = data.get("messages")
231 if messages is None:
232 return # No messages to check
233 input_messages = messages
234 else:
235 return
237 if choices := response.get("choices"):
238 if isinstance(choices, list):
239 serialized_choices: Final = []
240 for c in choices:
241 if isinstance(c, Choices):
242 try:
243 serialized_choices.append(c.model_dump())
244 except Exception:
245 serialized_choices.append(c.dict())
246 else:
247 serialized_choices.append(c)
248 choices = serialized_choices
250 ai_guard_payload: Final = {
251 "debug": False,
252 "input": {
253 "messages": input_messages,
254 "tools": data.get("tools"),
255 "choices": choices,
256 },
257 "event_type": "output",
258 }
260 if self.pangea_output_recipe:
261 ai_guard_payload["recipe"] = self.pangea_output_recipe
263 ai_guard_response = await self._call_pangea_ai_guard("v1beta/guard", ai_guard_payload, "async_pre_call_hook")
264 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
266 if not ai_guard_response.get("result", {}).get("transformed"):
267 return
269 output: Final = ai_guard_response.get("result", {}).get("output", {})
270 response.choices = output["choices"]
271 return response
273 @log_guardrail_information
274 async def async_post_call_success_hook(
275 self,
276 data: dict,
277 user_api_key_dict: UserAPIKeyAuth,
278 # This union isn't actually correct -- it can get other response types depending on the API called
279 response: LLMResponseTypes,
280 ):
281 """
282 Guardrail hook run after a successful LLM call (scans output).
284 Args:
285 data (dict): The original request data.
286 user_api_key_dict (UserAPIKeyAuth): User API key details.
287 response (LLMResponseTypes): The response object from the LLM call.
288 """
289 event_type: Final = GuardrailEventHooks.post_call
290 if self.should_run_guardrail(data=data, event_type=event_type) is not True:
291 verbose_proxy_logger.debug(
292 "Pangea Guardrail (async_pre_call_hook): Guardrail is disabled %s.", self.guardrail_name
293 )
294 return data
295 try:
296 return await self._async_post_call_success_hook(data, user_api_key_dict, response)
297 except HTTPException:
298 raise
299 except Exception as e:
300 raise HTTPException(
301 status_code=500,
302 detail={
303 "error": "Error in Pangea Guardrail",
304 "guardrail_name": self.guardrail_name,
305 "exceptions": str(e),
306 },
307 ) from e
309 @staticmethod
310 def get_config_model() -> type["GuardrailConfigModel"] | None:
311 from litellm.types.proxy.guardrails.guardrail_hooks.pangea import (
312 PangeaGuardrailConfigModel,
313 )
315 return PangeaGuardrailConfigModel
317 @classmethod
318 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
319 return [
320 GuardrailEventHooks.pre_call,
321 GuardrailEventHooks.post_call,
322 ]