diff --git a/litellm/proxy/guardrails/guardrail_hooks/pangea/pangea.py b/litellm/proxy/guardrails/guardrail_hooks/pangea/pangea.py index 07e9fa760fd..be4052e4eca 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/pangea/pangea.py +++ b/litellm/proxy/guardrails/guardrail_hooks/pangea/pangea.py @@ -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