diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index 7b19e8c8b13..08cebe1a0c1 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -42,7 +42,7 @@ else: class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callback#callback-class # Class variables or attributes - def __init__(self, message_logging: bool = True) -> None: + def __init__(self, message_logging: bool = True, **kwargs) -> None: self.message_logging = message_logging pass diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index a71718e9b89..d9736220c16 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -17,6 +17,7 @@ from litellm.types.guardrails import ( GuardrailEventHooks, GuardrailInfoResponse, GuardrailUIAddGuardrailSettings, + LakeraV2GuardrailConfigModel, ListGuardrailsResponse, PiiAction, PiiEntityType, @@ -592,11 +593,13 @@ async def get_provider_specific_params(): # Get fields from the models bedrock_fields = _get_fields_from_model(BedrockGuardrailConfigModel) presidio_fields = _get_fields_from_model(PresidioConfigModel) + lakera_v2_fields = _get_fields_from_model(LakeraV2GuardrailConfigModel) # Return the provider-specific parameters provider_params = { SupportedGuardrailIntegrations.BEDROCK.value: bedrock_fields, SupportedGuardrailIntegrations.PRESIDIO.value: presidio_fields, + SupportedGuardrailIntegrations.LAKERA_V2.value: lakera_v2_fields, } return provider_params diff --git a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py index 01b6b9db4d9..e7d7d3b5aaa 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py @@ -62,7 +62,9 @@ class LakeraAIGuardrail(CustomGuardrail): super().__init__(**kwargs) async def call_v2_guard( - self, messages: List[AllMessageValues] + self, + messages: List[AllMessageValues], + request_data: Dict, ) -> Tuple[LakeraAIResponse, Dict]: """ Call the Lakera AI v2 guard API. @@ -116,7 +118,7 @@ class LakeraAIGuardrail(CustomGuardrail): self.add_standard_logging_guardrail_information_to_request_data( guardrail_json_response=guardrail_json_response, guardrail_status=status, - request_data=dict(request) or {}, + request_data=request_data, start_time=start_time.timestamp(), end_time=datetime.now().timestamp(), duration=(datetime.now() - start_time).total_seconds(), @@ -137,11 +139,8 @@ class LakeraAIGuardrail(CustomGuardrail): if not payload: return messages - # Copy so we don’t edit the originals - masked = [msg.copy() for msg in messages] - # For each message, find its detections on the fly - for idx, msg in enumerate(masked): + for idx, msg in enumerate(messages): content = msg.get("content", "") if not content: continue @@ -175,7 +174,7 @@ class LakeraAIGuardrail(CustomGuardrail): masked_entity_count[typ] = masked_entity_count.get(typ, 0) + 1 msg["content"] = content - return masked + return messages async def async_pre_call_hook( self, @@ -197,8 +196,13 @@ class LakeraAIGuardrail(CustomGuardrail): add_guardrail_to_applied_guardrails_header, ) - event_type: GuardrailEventHooks = GuardrailEventHooks.during_call + verbose_proxy_logger.debug("Lakera AI: pre_call_hook") + + event_type: GuardrailEventHooks = GuardrailEventHooks.pre_call if self.should_run_guardrail(data=data, event_type=event_type) is not True: + verbose_proxy_logger.debug( + "Lakera AI: not running guardrail. Guardrail is disabled." + ) return data new_messages: Optional[List[AllMessageValues]] = data.get("messages") @@ -212,7 +216,8 @@ class LakeraAIGuardrail(CustomGuardrail): ########## 1. Make the Lakera AI v2 guard API request ########## ######################################################### lakera_guardrail_response, masked_entity_count = await self.call_v2_guard( - messages=new_messages + messages=new_messages, + request_data=data, ) ######################################################### @@ -277,7 +282,8 @@ class LakeraAIGuardrail(CustomGuardrail): ########## 1. Make the Lakera AI v2 guard API request ########## ######################################################### lakera_guardrail_response, masked_entity_count = await self.call_v2_guard( - messages=new_messages + messages=new_messages, + request_data=data, ) ######################################################### diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 15c5d9d664a..78332ebecde 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -16,6 +16,7 @@ from .guardrail_initializers import ( initialize_guardrails_ai, initialize_hide_secrets, initialize_lakera, + initialize_lakera_v2, initialize_presidio, ) @@ -23,6 +24,7 @@ guardrail_initializer_registry = { SupportedGuardrailIntegrations.APORIA.value: initialize_aporia, SupportedGuardrailIntegrations.BEDROCK.value: initialize_bedrock, SupportedGuardrailIntegrations.LAKERA.value: initialize_lakera, + SupportedGuardrailIntegrations.LAKERA_V2.value: initialize_lakera_v2, SupportedGuardrailIntegrations.AIM.value: initialize_aim, SupportedGuardrailIntegrations.PRESIDIO.value: initialize_presidio, SupportedGuardrailIntegrations.HIDE_SECRETS.value: initialize_hide_secrets, diff --git a/ui/litellm-dashboard/src/components/guardrails/guardrail_info_helpers.tsx b/ui/litellm-dashboard/src/components/guardrails/guardrail_info_helpers.tsx index 5755df533da..9b2818740db 100644 --- a/ui/litellm-dashboard/src/components/guardrails/guardrail_info_helpers.tsx +++ b/ui/litellm-dashboard/src/components/guardrails/guardrail_info_helpers.tsx @@ -1,9 +1,9 @@ export enum GuardrailProviders { PresidioPII = "Presidio PII", Bedrock = "Bedrock Guardrail", - LLMGuard = "LLM Guard Endpoint", - SecretDetector = "Secret Detector", - AIM = "AIM Guardrail", + // LLMGuard = "LLM Guard Endpoint", + // SecretDetector = "Secret Detector", + // AIM = "AIM Guardrail", Lakera = "Lakera" } @@ -13,7 +13,7 @@ export const guardrail_provider_map: Record = { LLMGuard: "llmguard_moderations", SecretDetector: "hide_secrets", AIM: "aim", - Lakera: "lakera" + Lakera: "lakera_v2" };