mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
Update Pangea Guardrail to support new AIDR endpoint
This commit is contained in:
parent
169a17400f
commit
223587179f
1 changed files with 139 additions and 222 deletions
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue