Update Pangea Guardrail to support new AIDR endpoint

This commit is contained in:
Ryan Means 2025-07-30 17:03:02 -07:00
parent 169a17400f
commit 223587179f

View file

@ -1,6 +1,6 @@
# litellm/proxy/guardrails/guardrail_hooks/pangea.py
import os
from typing import TYPE_CHECKING, Any, Optional, Protocol, Type
from typing import TYPE_CHECKING, Any, Optional, Type
from fastapi import HTTPException
@ -19,7 +19,7 @@ from litellm.proxy.common_utils.callback_utils import (
add_guardrail_to_applied_guardrails_header,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import LLMResponseTypes, ModelResponse, TextCompletionResponse
from litellm.types.utils import Choices, LLMResponseTypes, ModelResponse, TextCompletionResponse
if TYPE_CHECKING:
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
@ -31,14 +31,6 @@ class PangeaGuardrailMissingSecrets(Exception):
pass
class _Transformer(Protocol):
def get_messages(self) -> list[dict]: # noqa: E704
...
def update_original_body(self, prompt_messages: list[dict]) -> Any: # noqa: E704
...
class _TextCompletionRequest:
def __init__(self, body):
self.body = body
@ -53,109 +45,6 @@ class _TextCompletionRequest:
return self.body
class _TextCompletionResponse:
def __init__(self, body):
self.body = body
def get_messages(self) -> list[dict]:
messages = []
for choice in self.body["choices"]:
messages.append({"role": "assistant", "content": choice["text"]})
return messages
def update_original_body(self, prompt_messages: list[dict]) -> Any:
assert len(prompt_messages) == len(self.body["choices"])
for choice, prompt_message in zip(self.body["choices"], prompt_messages):
choice["text"] = prompt_message["content"]
return self.body
class _ChatCompletionRequest:
def __init__(self, body):
self.body = body
def get_messages(self) -> list[dict]:
messages = []
for message in self.body["messages"]:
role = message["role"]
content = message["content"]
if isinstance(content, str):
messages.append({"role": role, "content": content})
if isinstance(content, list):
for content_part in content:
if content_part["type"] == "text":
messages.append({"role": role, "content": content_part["text"]})
return messages
def update_original_body(self, prompt_messages: list[dict]) -> Any:
count = 0
for message in self.body["messages"]:
content = message["content"]
if isinstance(content, str):
message["content"] = prompt_messages[count]["content"]
count += 1
if isinstance(content, list):
for content_part in content:
if content_part["type"] == "text":
content_part["text"] = prompt_messages[count]["content"]
count += 1
assert len(prompt_messages) == count
return self.body
class _ChatCompletionResponse:
def __init__(self, body):
self.body = body
def get_messages(self) -> list[dict]:
messages = []
for choice in self.body["choices"]:
messages.append(
{
"role": choice["message"]["role"],
"content": choice["message"]["content"],
}
)
return messages
def update_original_body(self, prompt_messages: list[dict]) -> Any:
assert len(prompt_messages) == len(self.body["choices"])
for choice, prompt_message in zip(self.body["choices"], prompt_messages):
choice["message"]["content"] = prompt_message["content"]
return self.body
def _get_transformer_for_request(body, call_type) -> Optional[_Transformer]:
match call_type:
case "text_completion" | "atext_completion":
return _TextCompletionRequest(body)
case "completion" | "acompletion":
return _ChatCompletionRequest(body)
return None
def _get_transformer_for_response(body) -> Optional[_Transformer]:
match body:
case TextCompletionResponse():
return _TextCompletionResponse(body)
case ModelResponse():
return _ChatCompletionResponse(body)
return None
class PangeaHandler(CustomGuardrail):
"""
Pangea AI Guardrail handler to interact with the Pangea AI Guard service.
@ -200,7 +89,6 @@ class PangeaHandler(CustomGuardrail):
)
self.pangea_input_recipe = pangea_input_recipe
self.pangea_output_recipe = pangea_output_recipe
self.guardrail_endpoint = f"{self.api_base}/v1/text/guard"
# Pass relevant kwargs to the parent class
super().__init__(guardrail_name=guardrail_name, **kwargs)
@ -208,7 +96,9 @@ class PangeaHandler(CustomGuardrail):
f"Initialized Pangea Guardrail: name={guardrail_name}, recipe={pangea_input_recipe}, api_base={self.api_base}"
)
async def _call_pangea_guard(self, payload: dict, hook_name: str) -> dict:
async def _call_pangea_ai_guard(
self, api: str, payload: dict, hook_name: str
) -> dict:
"""
Makes the API call to the Pangea AI Guard endpoint.
The function itself will raise an error in the case that a response
@ -216,6 +106,7 @@ class PangeaHandler(CustomGuardrail):
should act on.
Args:
api (str): Which API to use (text/guard or v1beta/guard)
payload (dict): The request payload.
request_data (dict): Original request data (used for logging/headers).
hook_name (str): Name of the hook calling this function (for logging).
@ -227,62 +118,84 @@ class PangeaHandler(CustomGuardrail):
Returns:
list[dict]: The original response body
"""
endpoint = f"{self.api_base}/{api}"
headers = {
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
}
try:
verbose_proxy_logger.debug(
f"Pangea Guardrail ({hook_name}): Calling endpoint {self.guardrail_endpoint} with payload: {payload}"
)
response = await self.async_handler.post(
url=self.guardrail_endpoint, json=payload, headers=headers
)
response.raise_for_status() # Raise HTTPError for bad responses (4xx or 5xx)
result = response.json()
verbose_proxy_logger.debug(
f"Pangea Guardrail ({hook_name}): Received response: {result}"
verbose_proxy_logger.debug(
f"Pangea Guardrail ({hook_name}): Calling endpoint {endpoint} with payload: {payload}"
)
response = await self.async_handler.post(
url=endpoint, json=payload, headers=headers
)
response.raise_for_status()
result = response.json()
if result.get("result", {}).get("blocked"):
verbose_proxy_logger.warning(
f"Pangea Guardrail ({hook_name}): Request blocked. Response: {result}"
)
# Check if the request was blocked
if result.get("result", {}).get("blocked") is True:
verbose_proxy_logger.warning(
f"Pangea Guardrail ({hook_name}): Request blocked. Response: {result}"
)
raise HTTPException(
status_code=400, # Bad Request, indicating violation
detail={
"error": "Violated Pangea guardrail policy",
"guardrail_name": self.guardrail_name,
"pangea_response": result.get("result"),
},
)
else:
verbose_proxy_logger.info(
f"Pangea Guardrail ({hook_name}): Request passed. Response: {result.get('result', {}).get('detectors')}"
)
return result
except HTTPException as e:
# Re-raise HTTPException if it's the one we raised for blocking
raise e
except Exception as e:
verbose_proxy_logger.error(
f"Pangea Guardrail ({hook_name}): Error calling API: {e}. Response text: {getattr(e, 'response', None) and getattr(e.response, 'text', None)}" # type: ignore
)
# Decide if you want to block by default on error, or allow through
# Raising an exception here will block the request.
# To allow through on error, you might just log and return.
raise HTTPException(
status_code=500,
status_code=400, # Bad Request, indicating violation
detail={
"error": "Error communicating with Pangea Guardrail",
"error": "Violated Pangea guardrail policy",
"guardrail_name": self.guardrail_name,
"exception": str(e),
},
) from e
)
verbose_proxy_logger.info(
f"Pangea Guardrail ({hook_name}): Request passed. Response: {result.get('result', {}).get('detectors')}"
)
return result
async def _async_pre_call_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
cache: DualCache,
data: dict,
call_type: str
):
transformer = None
messages: Any = None
if call_type == "text_completion" or call_type == "atext_completion":
transformer = _TextCompletionRequest(data)
messages = transformer.get_messages()
else:
messages = data.get("messages")
ai_guard_payload = {
"debug": False,
"input": {
"messages": messages, # type: ignore
"tools": data.get("tools")
},
"event_type": "input",
}
if self.pangea_input_recipe:
ai_guard_payload["recipe"] = self.pangea_input_recipe
ai_guard_response = await self._call_pangea_ai_guard(
"v1beta/guard", ai_guard_payload, "async_pre_call_hook"
)
add_guardrail_to_applied_guardrails_header(
request_data=data, guardrail_name=self.guardrail_name
)
if not ai_guard_response.get("result", {}).get("transformed"):
return
output = ai_guard_response.get("result", {}).get("output", {})
if call_type == "text_completion" or call_type == "atext_completion":
data = transformer.update_original_body(output["messages"]) # type: ignore
else:
data["messages"] = output["messages"]
return data
@log_guardrail_information
async def async_pre_call_hook(
@ -299,50 +212,75 @@ class PangeaHandler(CustomGuardrail):
)
return data
transformer = _get_transformer_for_request(data, call_type)
if not transformer:
verbose_proxy_logger.warning(
f"Pangea Guardrail (async_pre_call_hook): Skipping guardrail {self.guardrail_name}"
f" because we cannot determine type of request: call_type '{call_type}'"
)
return
messages = transformer.get_messages()
if not messages:
verbose_proxy_logger.warning(
f"Pangea Guardrail (async_pre_call_hook): Skipping guardrail {self.guardrail_name}"
" because messages is empty."
)
return
ai_guard_payload = {
"debug": False, # Or make this configurable if needed
"messages": messages,
}
if self.pangea_input_recipe:
ai_guard_payload["recipe"] = self.pangea_input_recipe
ai_guard_response = await self._call_pangea_guard(
ai_guard_payload, "async_pre_call_hook"
)
# Add guardrail name to header if passed
add_guardrail_to_applied_guardrails_header(
request_data=data, guardrail_name=self.guardrail_name
)
prompt_messages = ai_guard_response.get("result", {}).get("prompt_messages", [])
try:
return transformer.update_original_body(prompt_messages)
return await self._async_pre_call_hook(user_api_key_dict, cache, data, call_type)
except HTTPException:
raise
except Exception as e:
raise HTTPException(
status_code=500,
detail={
"error": "Failed to update original request body",
"error": "Error in Pangea Guardrail",
"guardrail_name": self.guardrail_name,
"exceptions": str(e),
},
}
) from e
async def _async_post_call_success_hook(
self,
data: dict,
user_api_key_dict: UserAPIKeyAuth,
# This union isn't actually correct -- it can get other response types depending on the API called
response: LLMResponseTypes,
):
if isinstance(response, TextCompletionResponse):
# Assume the earlier call type as well
input_messages = _TextCompletionRequest(data).get_messages()
if not isinstance(response, ModelResponse):
return
else:
input_messages = data.get("messages")
if choices := response.get("choices"):
if isinstance(choices, list):
serialized_choices = []
for c in choices:
if isinstance(c, Choices):
try:
serialized_choices.append(c.model_dump())
except Exception:
serialized_choices.append(c.dict())
else:
serialized_choices.append(c)
choices = serialized_choices
ai_guard_payload = {
"debug": False,
"input": {
"messages": input_messages,
"tools": data.get("tools"),
"choices": choices,
},
"event_type": "output",
}
if self.pangea_output_recipe:
ai_guard_payload["recipe"] = self.pangea_output_recipe
ai_guard_response = await self._call_pangea_ai_guard(
"v1beta/guard", ai_guard_payload, "async_pre_call_hook"
)
add_guardrail_to_applied_guardrails_header(
request_data=data, guardrail_name=self.guardrail_name
)
if not ai_guard_response.get("result", {}).get("transformed"):
return
output = ai_guard_response.get("result", {}).get("output", {})
response.choices = output["choices"]
return data
@log_guardrail_information
async def async_post_call_success_hook(
self,
@ -365,39 +303,18 @@ class PangeaHandler(CustomGuardrail):
f"Pangea Guardrail (async_pre_call_hook): Guardrail is disabled {self.guardrail_name}."
)
return data
transformer = _get_transformer_for_response(response)
if not transformer:
verbose_proxy_logger.warning(
f"Pangea Guardrail (async_post_call_success_hook): Skipping guardrail {self.guardrail_name}"
" because we cannot determine type of request"
)
return
messages = transformer.get_messages()
verbose_proxy_logger.warning(f"GOT MESSAGES: {messages}")
ai_guard_payload = {
"debug": False, # Or make this configurable if needed
"messages": messages,
}
if self.pangea_output_recipe:
ai_guard_payload["recipe"] = self.pangea_output_recipe
ai_guard_response = await self._call_pangea_guard(
ai_guard_payload, "post_call_success_hook"
)
prompt_messages = ai_guard_response.get("result", {}).get("prompt_messages", [])
try:
return transformer.update_original_body(prompt_messages)
return await self._async_post_call_success_hook(data, user_api_key_dict, response)
except HTTPException:
raise
except Exception as e:
raise HTTPException(
status_code=500,
detail={
"error": "Failed to update original response body",
"error": "Error in Pangea Guardrail",
"guardrail_name": self.guardrail_name,
"exceptions": str(e),
},
}
) from e
@staticmethod