diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index cb4f01e7195..80c25683cfc 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -2,7 +2,7 @@ CRUD ENDPOINTS FOR GUARDRAILS """ -from typing import Dict, List, Optional, cast +from typing import Any, Dict, List, Optional, Type, cast from fastapi import APIRouter, Depends, HTTPException from pydantic import BaseModel @@ -12,6 +12,7 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry from litellm.types.guardrails import ( PII_ENTITY_CATEGORIES_MAP, + BedrockGuardrailConfigModel, Guardrail, GuardrailEventHooks, GuardrailInfoResponse, @@ -19,6 +20,8 @@ from litellm.types.guardrails import ( ListGuardrailsResponse, PiiAction, PiiEntityType, + PresidioConfigModel, + SupportedGuardrailIntegrations, ) #### GUARDRAILS ENDPOINTS #### @@ -455,3 +458,80 @@ async def get_guardrail_ui_settings(): supported_modes=list(GuardrailEventHooks), pii_entity_categories=category_maps, ) + + +def _get_fields_from_model(model_class: Type[BaseModel]) -> List[Dict[str, Any]]: + """ + Get the fields from a Pydantic model + """ + fields = [] + for field_name, field in model_class.model_fields.items(): + # Get field metadata + description = field.description or field_name + + # Check if this field is in the required_fields class variable + required = field.is_required() + + fields.append( + { + "param": field_name, + "description": description, + "required": required, + } + ) + return fields + + +@router.get( + "/guardrails/ui/provider_specific_params", + tags=["Guardrails"], + dependencies=[Depends(user_api_key_auth)], +) +async def get_provider_specific_params(): + """ + Get provider-specific parameters for different guardrail types. + + Returns a dictionary mapping guardrail providers to their specific parameters, + including parameter names, descriptions, and whether they are required. + + Example Response: + ```json + { + "bedrock": [ + { + "param": "guardrailIdentifier", + "description": "The ID of your guardrail on Bedrock", + "required": true + }, + { + "param": "guardrailVersion", + "description": "The version of your Bedrock guardrail (e.g., DRAFT or version number)", + "required": true + } + ], + "presidio": [ + { + "param": "presidio_analyzer_api_base", + "description": "Base URL for the Presidio analyzer API", + "required": true + }, + { + "param": "presidio_anonymizer_api_base", + "description": "Base URL for the Presidio anonymizer API", + "required": true + } + ] + } + ``` + """ + # Get fields from the models + bedrock_fields = _get_fields_from_model(BedrockGuardrailConfigModel) + presidio_fields = _get_fields_from_model(PresidioConfigModel) + + # Return the provider-specific parameters + provider_params = { + SupportedGuardrailIntegrations.BEDROCK.value: bedrock_fields, + SupportedGuardrailIntegrations.PRESIDIO.value: presidio_fields, + } + + return provider_params diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 408a451464c..d3d3692dd35 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -64,6 +64,12 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): super().__init__(**kwargs) BaseAWSLLM.__init__(self) + verbose_proxy_logger.debug( + "Bedrock Guardrail initialized with guardrailIdentifier: %s, guardrailVersion: %s", + self.guardrailIdentifier, + self.guardrailVersion, + ) + def convert_to_bedrock_format( self, messages: Optional[List[AllMessageValues]] = None, @@ -74,9 +80,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): if messages: for message in messages: - message_text_content: Optional[List[str]] = ( - self.get_content_for_message(message=message) - ) + message_text_content: Optional[ + List[str] + ] = self.get_content_for_message(message=message) if message_text_content is None: continue for text_content in message_text_content: @@ -313,11 +319,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ######################################################### ########## 2. Update the messages with the guardrail response ########## ######################################################### - data["messages"] = ( - self._update_messages_with_updated_bedrock_guardrail_response( - messages=new_messages, - bedrock_guardrail_response=bedrock_guardrail_response, - ) + data[ + "messages" + ] = self._update_messages_with_updated_bedrock_guardrail_response( + messages=new_messages, + bedrock_guardrail_response=bedrock_guardrail_response, ) ######################################################### @@ -367,11 +373,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ######################################################### ########## 2. Update the messages with the guardrail response ########## ######################################################### - data["messages"] = ( - self._update_messages_with_updated_bedrock_guardrail_response( - messages=new_messages, - bedrock_guardrail_response=bedrock_guardrail_response, - ) + data[ + "messages" + ] = self._update_messages_with_updated_bedrock_guardrail_response( + messages=new_messages, + bedrock_guardrail_response=bedrock_guardrail_response, ) ######################################################### @@ -421,11 +427,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ######################################################### ########## 2. Update the messages with the guardrail response ########## ######################################################### - data["messages"] = ( - self._update_messages_with_updated_bedrock_guardrail_response( - messages=new_messages, - bedrock_guardrail_response=bedrock_guardrail_response, - ) + data[ + "messages" + ] = self._update_messages_with_updated_bedrock_guardrail_response( + messages=new_messages, + bedrock_guardrail_response=bedrock_guardrail_response, ) ######################################################### diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index be2981a68f5..68b0dadc961 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -3,106 +3,117 @@ import litellm from litellm.types.guardrails import * -def initialize_aporia(litellm_params, guardrail): +def initialize_aporia( + litellm_params: LitellmParams, + guardrail: Guardrail, +): from litellm.proxy.guardrails.guardrail_hooks.aporia_ai import AporiaGuardrail _aporia_callback = AporiaGuardrail( - api_base=litellm_params["api_base"], - api_key=litellm_params["api_key"], - guardrail_name=guardrail["guardrail_name"], - event_hook=litellm_params["mode"], - default_on=litellm_params["default_on"], + api_base=litellm_params.api_base, + api_key=litellm_params.api_key, + guardrail_name=guardrail.get("guardrail_name", ""), + event_hook=litellm_params.mode, + default_on=litellm_params.default_on, ) litellm.logging_callback_manager.add_litellm_callback(_aporia_callback) -def initialize_bedrock(litellm_params, guardrail): +def initialize_bedrock(litellm_params: LitellmParams, guardrail: Guardrail): from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import ( BedrockGuardrail, ) _bedrock_callback = BedrockGuardrail( - guardrail_name=guardrail["guardrail_name"], - event_hook=litellm_params["mode"], - guardrailIdentifier=litellm_params["guardrailIdentifier"], - guardrailVersion=litellm_params["guardrailVersion"], - default_on=litellm_params["default_on"], - mask_request_content=litellm_params.get("mask_request_content", None), - mask_response_content=litellm_params.get("mask_response_content", None), + guardrail_name=guardrail.get("guardrail_name", ""), + event_hook=litellm_params.mode, + guardrailIdentifier=litellm_params.guardrailIdentifier, + guardrailVersion=litellm_params.guardrailVersion, + default_on=litellm_params.default_on, + mask_request_content=litellm_params.mask_request_content, + mask_response_content=litellm_params.mask_response_content, + aws_region_name=litellm_params.aws_region_name, + aws_access_key_id=litellm_params.aws_access_key_id, + aws_secret_access_key=litellm_params.aws_secret_access_key, + aws_session_token=litellm_params.aws_session_token, + aws_session_name=litellm_params.aws_session_name, + aws_profile_name=litellm_params.aws_profile_name, + aws_role_name=litellm_params.aws_role_name, + aws_web_identity_token=litellm_params.aws_web_identity_token, + aws_sts_endpoint=litellm_params.aws_sts_endpoint, + aws_bedrock_runtime_endpoint=litellm_params.aws_bedrock_runtime_endpoint, ) litellm.logging_callback_manager.add_litellm_callback(_bedrock_callback) -def initialize_lakera(litellm_params, guardrail): +def initialize_lakera(litellm_params: LitellmParams, guardrail: Guardrail): from litellm.proxy.guardrails.guardrail_hooks.lakera_ai import lakeraAI_Moderation _lakera_callback = lakeraAI_Moderation( - api_base=litellm_params["api_base"], - api_key=litellm_params["api_key"], - guardrail_name=guardrail["guardrail_name"], - event_hook=litellm_params["mode"], - category_thresholds=litellm_params.get("category_thresholds"), - default_on=litellm_params["default_on"], + api_base=litellm_params.api_base, + api_key=litellm_params.api_key, + guardrail_name=guardrail.get("guardrail_name", ""), + event_hook=litellm_params.mode, + category_thresholds=litellm_params.category_thresholds, + default_on=litellm_params.default_on, ) litellm.logging_callback_manager.add_litellm_callback(_lakera_callback) -def initialize_aim(litellm_params, guardrail): +def initialize_aim(litellm_params: LitellmParams, guardrail: Guardrail): from litellm.proxy.guardrails.guardrail_hooks.aim import AimGuardrail _aim_callback = AimGuardrail( - api_base=litellm_params["api_base"], - api_key=litellm_params["api_key"], - guardrail_name=guardrail["guardrail_name"], - event_hook=litellm_params["mode"], - default_on=litellm_params["default_on"], + api_base=litellm_params.api_base, + api_key=litellm_params.api_key, + guardrail_name=guardrail.get("guardrail_name", ""), + event_hook=litellm_params.mode, + default_on=litellm_params.default_on, ) litellm.logging_callback_manager.add_litellm_callback(_aim_callback) -def initialize_presidio(litellm_params, guardrail): +def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail): from litellm.proxy.guardrails.guardrail_hooks.presidio import ( _OPTIONAL_PresidioPIIMasking, ) _presidio_callback = _OPTIONAL_PresidioPIIMasking( - guardrail_name=guardrail["guardrail_name"], - event_hook=litellm_params["mode"], - output_parse_pii=litellm_params["output_parse_pii"], - presidio_ad_hoc_recognizers=litellm_params["presidio_ad_hoc_recognizers"], - mock_redacted_text=litellm_params.get("mock_redacted_text") or None, - default_on=litellm_params["default_on"], - pii_entities_config=litellm_params.get("pii_entities_config"), - presidio_analyzer_api_base=litellm_params.get("presidio_analyzer_api_base"), - presidio_anonymizer_api_base=litellm_params.get("presidio_anonymizer_api_base"), + guardrail_name=guardrail.get("guardrail_name", ""), + event_hook=litellm_params.mode, + output_parse_pii=litellm_params.output_parse_pii, + presidio_ad_hoc_recognizers=litellm_params.presidio_ad_hoc_recognizers, + mock_redacted_text=litellm_params.mock_redacted_text, + default_on=litellm_params.default_on, + pii_entities_config=litellm_params.pii_entities_config, + presidio_analyzer_api_base=litellm_params.presidio_analyzer_api_base, + presidio_anonymizer_api_base=litellm_params.presidio_anonymizer_api_base, ) litellm.logging_callback_manager.add_litellm_callback(_presidio_callback) - if litellm_params["output_parse_pii"]: + if litellm_params.output_parse_pii: _success_callback = _OPTIONAL_PresidioPIIMasking( output_parse_pii=True, - guardrail_name=guardrail["guardrail_name"], + guardrail_name=guardrail.get("guardrail_name", ""), event_hook=GuardrailEventHooks.post_call.value, - presidio_ad_hoc_recognizers=litellm_params["presidio_ad_hoc_recognizers"], - default_on=litellm_params["default_on"], - presidio_analyzer_api_base=litellm_params.get("presidio_analyzer_api_base"), - presidio_anonymizer_api_base=litellm_params.get( - "presidio_anonymizer_api_base" - ), + presidio_ad_hoc_recognizers=litellm_params.presidio_ad_hoc_recognizers, + default_on=litellm_params.default_on, + presidio_analyzer_api_base=litellm_params.presidio_analyzer_api_base, + presidio_anonymizer_api_base=litellm_params.presidio_anonymizer_api_base, ) litellm.logging_callback_manager.add_litellm_callback(_success_callback) -def initialize_hide_secrets(litellm_params, guardrail): +def initialize_hide_secrets(litellm_params: LitellmParams, guardrail: Guardrail): from litellm_enterprise.enterprise_callbacks.secret_detection import ( _ENTERPRISE_SecretDetection, ) _secret_detection_object = _ENTERPRISE_SecretDetection( - detect_secrets_config=litellm_params.get("detect_secrets_config"), - event_hook=litellm_params["mode"], - guardrail_name=guardrail["guardrail_name"], - default_on=litellm_params["default_on"], + detect_secrets_config=litellm_params.detect_secrets_config, + event_hook=litellm_params.mode, + guardrail_name=guardrail.get("guardrail_name", ""), + default_on=litellm_params.default_on, ) litellm.logging_callback_manager.add_litellm_callback(_secret_detection_object) @@ -110,16 +121,16 @@ def initialize_hide_secrets(litellm_params, guardrail): def initialize_guardrails_ai(litellm_params, guardrail): from litellm.proxy.guardrails.guardrail_hooks.guardrails_ai import GuardrailsAI - _guard_name = litellm_params.get("guard_name") + _guard_name = litellm_params.guard_name if not _guard_name: raise Exception( "GuardrailsAIException - Please pass the Guardrails AI guard name via 'litellm_params::guard_name'" ) _guardrails_ai_callback = GuardrailsAI( - api_base=litellm_params.get("api_base"), + api_base=litellm_params.api_base, guard_name=_guard_name, guardrail_name=SupportedGuardrailIntegrations.GURDRAILS_AI.value, - default_on=litellm_params["default_on"], + default_on=litellm_params.default_on, ) litellm.logging_callback_manager.add_litellm_callback(_guardrails_ai_callback) diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 6ad48521f0d..15c5d9d664a 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -71,7 +71,7 @@ class GuardrailRegistry: """ try: guardrail_name = guardrail.get("guardrail_name") - litellm_params: str = safe_dumps(guardrail.get("litellm_params", {})) + litellm_params: str = safe_dumps(dict(guardrail.get("litellm_params", {}))) guardrail_info: str = safe_dumps(guardrail.get("guardrail_info", {})) # Create guardrail in DB @@ -117,7 +117,7 @@ class GuardrailRegistry: """ try: guardrail_name = guardrail.get("guardrail_name") - litellm_params = guardrail.get("litellm_params", {}) + litellm_params: str = safe_dumps(dict(guardrail.get("litellm_params", {}))) guardrail_info = guardrail.get("guardrail_info", {}) # Update in DB diff --git a/litellm/proxy/guardrails/init_guardrails.py b/litellm/proxy/guardrails/init_guardrails.py index dd7a46388d9..2864f0d2684 100644 --- a/litellm/proxy/guardrails/init_guardrails.py +++ b/litellm/proxy/guardrails/init_guardrails.py @@ -120,11 +120,7 @@ class InitializeGuardrails: litellm_params_data = guardrail["litellm_params"] verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data) - _litellm_params_kwargs = { - k: litellm_params_data.get(k) for k in LitellmParams.__annotations__.keys() - } - - litellm_params = LitellmParams(**_litellm_params_kwargs) # type: ignore + litellm_params = LitellmParams(**litellm_params_data) if ( "category_thresholds" in litellm_params_data @@ -133,17 +129,17 @@ class InitializeGuardrails: lakera_category_thresholds = LakeraCategoryThresholds( **litellm_params_data["category_thresholds"] ) - litellm_params["category_thresholds"] = lakera_category_thresholds + litellm_params.category_thresholds = lakera_category_thresholds - api_key: Optional[str] = litellm_params.get("api_key") - if api_key and api_key.startswith("os.environ/"): - litellm_params["api_key"] = str(get_secret(litellm_params["api_key"])) # type: ignore + if litellm_params.api_key and litellm_params.api_key.startswith("os.environ/"): + litellm_params.api_key = str(get_secret(litellm_params.api_key)) - api_base: Optional[str] = litellm_params.get("api_base") - if api_base and api_base.startswith("os.environ/"): - litellm_params["api_base"] = str(get_secret(litellm_params["api_base"])) # type: ignore + if litellm_params.api_base and litellm_params.api_base.startswith( + "os.environ/" + ): + litellm_params.api_base = str(get_secret(litellm_params.api_base)) - guardrail_type: Optional[str] = litellm_params.get("guardrail") + guardrail_type = litellm_params.guardrail if guardrail_type is None: raise ValueError("guardrail_type is required") @@ -206,13 +202,13 @@ class InitializeGuardrails: spec.loader.exec_module(module) # type: ignore _guardrail_class = getattr(module, _class_name) - mode: Optional[str] = litellm_params.get("mode") + mode = litellm_params.mode if mode is None: raise ValueError( f"mode is required for guardrail {guardrail_type} please set mode to one of the following: {', '.join(GuardrailEventHooks)}" ) - default_on: Optional[bool] = litellm_params.get("default_on") + default_on = litellm_params.default_on _guardrail_callback = _guardrail_class( guardrail_name=guardrail["guardrail_name"], event_hook=mode, diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index f58f2cc0ac3..999638d058f 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -2,9 +2,14 @@ model_list: - model_name: openai/gpt-4o litellm_params: model: openai/gpt-4o - api_key: os.environ/OPENAI_API_KEY + api_key: any_key + api_base: https://exampleopenaiendpoint-production.up.railway.app/ -litellm_settings: - callbacks: ["smtp_email"] - +guardrails: + - guardrail_name: "bedrock-pre-guard" + litellm_params: + guardrail: bedrock # supported values: "aporia", "bedrock", "lakera" + mode: "during_call" + guardrailIdentifier: ff6ujrregl1q + guardrailVersion: "DRAFT" \ No newline at end of file diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 0d1c1329528..5e04f2e24bc 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -210,43 +210,118 @@ class PiiEntityCategoryMap(TypedDict): entities: List[PiiEntityType] -class LitellmParams(TypedDict, total=False): - guardrail: str - mode: str - api_key: Optional[str] - api_base: Optional[str] +class PresidioConfigModel(BaseModel): + """Configuration parameters for the Presidio PII masking guardrail""" + + presidio_analyzer_api_base: Optional[str] = Field( + default=None, + description="Base URL for the Presidio analyzer API", + ) + presidio_anonymizer_api_base: Optional[str] = Field( + default=None, + description="Base URL for the Presidio anonymizer API", + ) + pii_entities_config: Optional[Dict[PiiEntityType, PiiAction]] = Field( + default=None, description="Configuration for PII entity types and actions" + ) + output_parse_pii: Optional[bool] = Field( + default=None, description="Whether to parse PII in model outputs" + ) + presidio_ad_hoc_recognizers: Optional[str] = Field( + default=None, + description="Path to a JSON file containing ad-hoc recognizers for Presidio", + ) + mock_redacted_text: Optional[dict] = Field( + default=None, description="Mock redacted text for testing" + ) + + +class BedrockGuardrailConfigModel(BaseModel): + """Configuration parameters for the AWS Bedrock guardrail""" + + guardrailIdentifier: Optional[str] = Field( + default=None, description="The ID of your guardrail on Bedrock" + ) + guardrailVersion: Optional[str] = Field( + default=None, + description="The version of your Bedrock guardrail (e.g., DRAFT or version number)", + ) + aws_region_name: Optional[str] = Field( + default=None, description="AWS region where your guardrail is deployed" + ) + aws_access_key_id: Optional[str] = Field( + default=None, description="AWS access key ID for authentication" + ) + aws_secret_access_key: Optional[str] = Field( + default=None, description="AWS secret access key for authentication" + ) + aws_session_token: Optional[str] = Field( + default=None, description="AWS session token for temporary credentials" + ) + aws_session_name: Optional[str] = Field( + default=None, description="Name of the AWS session" + ) + aws_profile_name: Optional[str] = Field( + default=None, description="AWS profile name for credential retrieval" + ) + aws_role_name: Optional[str] = Field( + default=None, description="AWS role name for assuming roles" + ) + aws_web_identity_token: Optional[str] = Field( + default=None, description="Web identity token for AWS role assumption" + ) + aws_sts_endpoint: Optional[str] = Field( + default=None, description="AWS STS endpoint URL" + ) + aws_bedrock_runtime_endpoint: Optional[str] = Field( + default=None, description="AWS Bedrock runtime endpoint URL" + ) + + +class LitellmParams( + PresidioConfigModel, + BedrockGuardrailConfigModel, +): + guardrail: str = Field(description="The type of guardrail integration to use") + mode: str = Field( + description="When to apply the guardrail (pre_call, post_call, during_call, logging_only)" + ) + api_key: Optional[str] = Field( + default=None, description="API key for the guardrail service" + ) + api_base: Optional[str] = Field( + default=None, description="Base URL for the guardrail service API" + ) # Lakera specific params - category_thresholds: Optional[LakeraCategoryThresholds] - - # Bedrock specific params - guardrailIdentifier: Optional[str] - guardrailVersion: Optional[str] - - # Presidio params - output_parse_pii: Optional[bool] - presidio_ad_hoc_recognizers: Optional[str] - mock_redacted_text: Optional[dict] - # PII control params - pii_entities_config: Optional[Dict[PiiEntityType, PiiAction]] - presidio_analyzer_api_base: Optional[str] - presidio_anonymizer_api_base: Optional[str] + category_thresholds: Optional[LakeraCategoryThresholds] = Field( + default=None, + description="Threshold configuration for Lakera guardrail categories", + ) # hide secrets params - detect_secrets_config: Optional[dict] + detect_secrets_config: Optional[dict] = Field( + default=None, description="Configuration for detect-secrets guardrail" + ) # guardrails ai params - guard_name: Optional[str] - default_on: Optional[bool] + guard_name: Optional[str] = Field( + default=None, description="Name of the guardrail in guardrails.ai" + ) + default_on: Optional[bool] = Field( + default=None, description="Whether the guardrail is enabled by default" + ) ################## PII control params ################# ######################################################## - mask_request_content: Optional[ - bool - ] # will mask request content if guardrail makes any changes - mask_response_content: Optional[ - bool - ] # will mask response content if guardrail makes any changes + mask_request_content: Optional[bool] = Field( + default=None, + description="Will mask request content if guardrail makes any changes", + ) + mask_response_content: Optional[bool] = Field( + default=None, + description="Will mask response content if guardrail makes any changes", + ) class Guardrail(TypedDict, total=False): diff --git a/tests/local_testing/test_guardrails_config.py b/tests/guardrails_tests/test_guardrails_config.py similarity index 97% rename from tests/local_testing/test_guardrails_config.py rename to tests/guardrails_tests/test_guardrails_config.py index 66406ac5f4b..c2a90220431 100644 --- a/tests/local_testing/test_guardrails_config.py +++ b/tests/guardrails_tests/test_guardrails_config.py @@ -93,11 +93,11 @@ def test_guardrail_list_of_event_hooks(): def test_guardrail_info_response(): - from litellm.types.guardrails import GuardrailInfoResponse, LitellmParams + from litellm.types.guardrails import GuardrailInfoResponse, LitellmParams, GuardrailLiteLLMParamsResponse guardrail_info = GuardrailInfoResponse( guardrail_name="aporia-pre-guard", - litellm_params=LitellmParams( + litellm_params=GuardrailLiteLLMParamsResponse( guardrail="aporia", mode="pre_call", ), diff --git a/tests/litellm/proxy/guardrails/test_init_guardrails.py b/tests/litellm/proxy/guardrails/test_init_guardrails.py index 2d4e06c110d..782f9d69aee 100644 --- a/tests/litellm/proxy/guardrails/test_init_guardrails.py +++ b/tests/litellm/proxy/guardrails/test_init_guardrails.py @@ -36,7 +36,7 @@ def test_initialize_presidio_guardrail(): assert result["guardrail_name"] == "test_presidio_guardrail" assert ( - result["litellm_params"]["guardrail"] + result["litellm_params"].guardrail == SupportedGuardrailIntegrations.PRESIDIO.value ) - assert result["litellm_params"]["mode"] == "pre_call" + assert result["litellm_params"].mode == "pre_call" diff --git a/ui/litellm-dashboard/out/assets/logos/llm_guard.png b/ui/litellm-dashboard/out/assets/logos/llm_guard.png new file mode 100644 index 00000000000..ed01c0f044e Binary files /dev/null and b/ui/litellm-dashboard/out/assets/logos/llm_guard.png differ diff --git a/ui/litellm-dashboard/public/assets/logos/aim_logo.jpeg b/ui/litellm-dashboard/public/assets/logos/aim_logo.jpeg new file mode 100644 index 00000000000..60fc2a9295c Binary files /dev/null and b/ui/litellm-dashboard/public/assets/logos/aim_logo.jpeg differ diff --git a/ui/litellm-dashboard/public/assets/logos/lakeraai.jpeg b/ui/litellm-dashboard/public/assets/logos/lakeraai.jpeg new file mode 100644 index 00000000000..b30d3ede6be Binary files /dev/null and b/ui/litellm-dashboard/public/assets/logos/lakeraai.jpeg differ diff --git a/ui/litellm-dashboard/public/assets/logos/litellm.jpg b/ui/litellm-dashboard/public/assets/logos/litellm.jpg new file mode 100644 index 00000000000..a10a1d24969 Binary files /dev/null and b/ui/litellm-dashboard/public/assets/logos/litellm.jpg differ diff --git a/ui/litellm-dashboard/public/assets/logos/llm_guard.png b/ui/litellm-dashboard/public/assets/logos/llm_guard.png new file mode 100644 index 00000000000..ed01c0f044e Binary files /dev/null and b/ui/litellm-dashboard/public/assets/logos/llm_guard.png differ diff --git a/ui/litellm-dashboard/public/assets/logos/secret_detect.png b/ui/litellm-dashboard/public/assets/logos/secret_detect.png new file mode 100644 index 00000000000..b7e09d33077 Binary files /dev/null and b/ui/litellm-dashboard/public/assets/logos/secret_detect.png differ diff --git a/ui/litellm-dashboard/src/components/guardrails/add_guardrail_form.tsx b/ui/litellm-dashboard/src/components/guardrails/add_guardrail_form.tsx index b6fe36fea04..27ce2039e05 100644 --- a/ui/litellm-dashboard/src/components/guardrails/add_guardrail_form.tsx +++ b/ui/litellm-dashboard/src/components/guardrails/add_guardrail_form.tsx @@ -3,7 +3,7 @@ import { Card, Form, Typography, Select, Input, Switch, Tooltip, Modal, message, import { Button, TextInput } from '@tremor/react'; import type { FormInstance } from 'antd'; import { GuardrailProviders, guardrail_provider_map, shouldRenderPIIConfigSettings, guardrailLogoMap } from './guardrail_info_helpers'; -import { createGuardrailCall, getGuardrailUISettings } from '../networking'; +import { createGuardrailCall, getGuardrailUISettings, getGuardrailProviderSpecificParams } from '../networking'; import PiiConfiguration from './pii_configuration'; import GuardrailProviderFields from './guardrail_provider_fields'; @@ -35,6 +35,20 @@ interface LiteLLMParams { [key: string]: any; // Allow additional properties for specific guardrails } +// Mapping of provider -> list of param descriptors +interface ProviderParam { + param: string; + description: string; + required: boolean; + default_value?: string; + options?: string[]; + type?: string; +} + +interface ProviderParamsResponse { + [provider: string]: ProviderParam[]; +} + const AddGuardrailForm: React.FC = ({ visible, onClose, @@ -48,22 +62,29 @@ const AddGuardrailForm: React.FC = ({ const [selectedEntities, setSelectedEntities] = useState([]); const [selectedActions, setSelectedActions] = useState<{[key: string]: string}>({}); const [currentStep, setCurrentStep] = useState(0); + const [providerParams, setProviderParams] = useState(null); - // Fetch guardrail settings when the component mounts + // Fetch guardrail UI settings + provider params on mount / accessToken change useEffect(() => { - const fetchGuardrailSettings = async () => { + if (!accessToken) return; + + const fetchData = async () => { try { - if (!accessToken) return; - - const data = await getGuardrailUISettings(accessToken); - setGuardrailSettings(data); + // Parallel requests for speed + const [uiSettings, providerParamsResp] = await Promise.all([ + getGuardrailUISettings(accessToken), + getGuardrailProviderSpecificParams(accessToken), + ]); + + setGuardrailSettings(uiSettings); + setProviderParams(providerParamsResp); } catch (error) { - console.error('Error fetching guardrail settings:', error); - message.error('Failed to load guardrail settings'); + console.error('Error fetching guardrail data:', error); + message.error('Failed to load guardrail configuration'); } }; - - fetchGuardrailSettings(); + + fetchData(); }, [accessToken]); const handleProviderChange = (value: string) => { @@ -109,10 +130,7 @@ const AddGuardrailForm: React.FC = ({ if (selectedProvider === 'PresidioPII') { fieldsToValidate.push('presidio_analyzer_api_base', 'presidio_anonymizer_api_base'); - } else if (selectedProvider === 'Bedrock') { - fieldsToValidate.push('config'); } - await form.validateFields(fieldsToValidate); } } @@ -193,18 +211,7 @@ const AddGuardrailForm: React.FC = ({ try { const configObj = JSON.parse(values.config); // For some guardrails, the config values need to be in litellm_params - // Especially for providers like Bedrock that need guardrailIdentifier and guardrailVersion - if (values.provider === 'Bedrock' && configObj) { - if (configObj.guardrail_id) { - guardrailData.litellm_params.guardrailIdentifier = configObj.guardrail_id; - } - if (configObj.guardrail_version) { - guardrailData.litellm_params.guardrailVersion = configObj.guardrail_version; - } - } else { - // For other providers, add the config to guardrail_info - guardrailData.guardrail_info = configObj; - } + guardrailData.guardrail_info = configObj; } catch (error) { message.error('Invalid JSON in configuration'); setLoading(false); @@ -212,6 +219,32 @@ const AddGuardrailForm: React.FC = ({ } } + /****************************** + * Add provider-specific params + * ---------------------------------- + * The backend exposes exactly which extra parameters a provider + * accepts via `/guardrails/ui/provider_specific_params`. + * Instead of copying every unknown form field, we fetch the list for + * the selected provider and ONLY pass those recognised params. + ******************************/ + + // Use pre-fetched provider params to copy recognised params + if (providerParams && selectedProvider) { + const providerKey = guardrail_provider_map[selectedProvider]?.toLowerCase(); + const providerSpecificParams = providerParams[providerKey] || []; + + const allowedParams = new Set( + providerSpecificParams.map((p) => p.param) + ); + + allowedParams.forEach((paramName) => { + const paramValue = values[paramName]; + if (paramValue !== undefined && paramValue !== null && paramValue !== '') { + guardrailData.litellm_params[paramName] = paramValue; + } + }); + } + if (!accessToken) { throw new Error("No access token available"); } @@ -252,13 +285,36 @@ const AddGuardrailForm: React.FC = ({ - {field.options?.map((option) => ( + {field.type === "select" && field.options ? ( + + ) : field.param.includes("password") || field.param.includes("secret") ? ( + ) : ( )} diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 5d12e0c026a..b4320d59f26 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -5096,3 +5096,30 @@ export const getGuardrailUISettings = async (accessToken: string) => { throw error; } }; + +export const getGuardrailProviderSpecificParams = async (accessToken: string) => { + try { + const url = proxyBaseUrl ? `${proxyBaseUrl}/guardrails/ui/provider_specific_params` : `/guardrails/ui/provider_specific_params`; + + const response = await fetch(url, { + method: "GET", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + }); + + if (!response.ok) { + const errorData = await response.text(); + handleError(errorData); + throw new Error("Failed to get guardrail provider specific parameters"); + } + + const data = await response.json(); + console.log("Guardrail provider specific params response:", data); + return data; + } catch (error) { + console.error("Failed to get guardrail provider specific parameters:", error); + throw error; + } +};