[UI] Allow adding Bedrock, Presidio, Lakera, AIM guardrails on UI (#10874)

* ui fix bedrock guard

* polish: logo should appear after selecting provider

* fix ui config bedrock

* fix: refactor - use specific configs per provider

* fix: refactor - use specific configs per provider

* feat: ui, show provider specific params for guardrails

* fix: updated type of LiteLLM params for guardrails

* fix: updated type of LiteLLM params for guardrails

* ui, use endpoint for adding presidio, bedrock guardrails

* fix: linting error

* add llama guard and secret detector on UI

* add aim on ui

* allow adding lakera AI on litellm ui

* fix: fixes for params to init guardrails

* test: test_guardrail_info_response

* test: test_initialize_presidio_guardrail

* fix: init guardrails

* fix: init guardrails

* add showSearch

* working bedrock guard
This commit is contained in:
Ishaan Jaff 2025-05-15 21:22:56 -07:00 • committed by GitHub
parent c6d36e8912
commit dc16e47df6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
19 changed files with 525 additions and 210 deletions

View file

@ -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

View file

@ -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,
)
#########################################################

View file

@ -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)

View file

@ -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

View file

@ -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,

View file

@ -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"

View file

@ -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):

View file

@ -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",
),

View file

@ -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"

Binary file not shown.

After

Width:  |  Height:  |  Size: 48 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 3.7 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 2.6 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 24 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 48 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 15 KiB

View file

@ -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<AddGuardrailFormProps> = ({
visible,
onClose,
@ -48,22 +62,29 @@ const AddGuardrailForm: React.FC<AddGuardrailFormProps> = ({
const [selectedEntities, setSelectedEntities] = useState<string[]>([]);
const [selectedActions, setSelectedActions] = useState<{[key: string]: string}>({});
const [currentStep, setCurrentStep] = useState(0);
const [providerParams, setProviderParams] = useState<ProviderParamsResponse | null>(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<AddGuardrailFormProps> = ({
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<AddGuardrailFormProps> = ({
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<AddGuardrailFormProps> = ({
}
}
/******************************
* 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<string>(
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<AddGuardrailFormProps> = ({
<Select
placeholder="Select a guardrail provider"
onChange={handleProviderChange}
labelInValue={false}
optionLabelProp="label"
dropdownRender={menu => menu}
showSearch={true}
>
{Object.entries(GuardrailProviders).map(([key, value]) => (
<Option
key={key}
value={key}
label={value}
label={
<div style={{ display: 'flex', alignItems: 'center' }}>
{guardrailLogoMap[value] && (
<img
src={guardrailLogoMap[value]}
alt=""
style={{
height: '20px',
width: '20px',
marginRight: '8px',
objectFit: 'contain'
}}
onError={(e) => {
// Hide broken image icon if image fails to load
e.currentTarget.style.display = 'none';
}}
/>
)}
<span>{value}</span>
</div>
}
>
<div style={{ display: 'flex', alignItems: 'center' }}>
{guardrailLogoMap[value] && (
@ -285,7 +341,11 @@ const AddGuardrailForm: React.FC<AddGuardrailFormProps> = ({
</Form.Item>
{/* Use the GuardrailProviderFields component to render provider-specific fields */}
<GuardrailProviderFields selectedProvider={selectedProvider} />
<GuardrailProviderFields
selectedProvider={selectedProvider}
accessToken={accessToken}
providerParams={providerParams}
/>
<Form.Item
name="mode"

View file

@ -1,9 +1,19 @@
export enum GuardrailProviders {
PresidioPII = "Presidio PII",
Bedrock = "Bedrock Guardrail",
LLMGuard = "LLM Guard Endpoint",
SecretDetector = "Secret Detector",
AIM = "AIM Guardrail",
Lakera = "Lakera"
}
export const guardrail_provider_map: Record<string, string> = {
PresidioPII: "presidio",
Bedrock: "bedrock",
LLMGuard: "llmguard_moderations",
SecretDetector: "hide_secrets",
AIM: "aim",
Lakera: "lakera"
};
@ -21,7 +31,12 @@ export const shouldRenderPIIConfigSettings = (provider: string | null) => {
const asset_logos_folder = '../ui/assets/logos/';
export const guardrailLogoMap: Record<string, string> = {
[GuardrailProviders.PresidioPII]: `${asset_logos_folder}presidio.png`
[GuardrailProviders.PresidioPII]: `${asset_logos_folder}presidio.png`,
[GuardrailProviders.Bedrock]: `${asset_logos_folder}bedrock.svg`,
[GuardrailProviders.LLMGuard]: `${asset_logos_folder}llm_guard.png`,
[GuardrailProviders.SecretDetector]: `${asset_logos_folder}secret_detect.png`,
[GuardrailProviders.AIM]: `${asset_logos_folder}aim_logo.jpeg`,
[GuardrailProviders.Lakera]: `${asset_logos_folder}lakeraai.jpeg`
};
export const getGuardrailLogoAndName = (guardrailValue: string): { logo: string, displayName: string } => {

View file

@ -1,87 +1,127 @@
import React from "react";
import { Form, Select } from "antd";
import React, { useState, useEffect } from "react";
import { Form, Select, Spin } from "antd";
import { TextInput } from "@tremor/react";
import { GuardrailProviders } from './guardrail_info_helpers';
import { GuardrailProviders, guardrail_provider_map } from './guardrail_info_helpers';
import { getGuardrailProviderSpecificParams } from "../networking";
interface GuardrailProviderFieldsProps {
selectedProvider: string | null;
accessToken?: string | null;
providerParams?: ProviderParamsResponse | null;
}
interface ProviderField {
key: string;
label: string;
placeholder?: string;
tooltip?: string;
required?: boolean;
type?: "text" | "password" | "select";
interface ProviderParam {
param: string;
description: string;
required: boolean;
default_value?: string;
options?: string[];
defaultValue?: string;
type?: string;
}
// Define fields for each guardrail provider
const GUARDRAIL_PROVIDER_FIELDS: Record<string, ProviderField[]> = {
// Presidio PII fields
PresidioPII: [
{
key: "presidio_analyzer_api_base",
label: "Presidio Analyzer API Base",
placeholder: "https://your-analyzer-api-url",
tooltip: "The base URL for your Presidio Analyzer API",
required: true
},
{
key: "presidio_anonymizer_api_base",
label: "Presidio Anonymizer API Base",
placeholder: "https://your-anonymizer-api-url",
tooltip: "The base URL for your Presidio Anonymizer API",
required: true
}
],
// Add more provider specific fields here as needed
Bedrock: [
{
key: "config",
label: "Configuration",
placeholder: '{"guardrail_id": "...", "guardrail_version": "..."}',
tooltip: "JSON configuration for Bedrock guardrail including guardrail_id and guardrail_version",
type: "text",
required: false
}
]
};
interface ProviderParamsResponse {
[provider: string]: ProviderParam[];
}
const GuardrailProviderFields: React.FC<GuardrailProviderFieldsProps> = ({
selectedProvider
selectedProvider,
accessToken,
providerParams: providerParamsProp = null
}) => {
// If no provider is selected or if the provider doesn't have specific fields, return null
if (!selectedProvider || !GUARDRAIL_PROVIDER_FIELDS[selectedProvider]) {
const [loading, setLoading] = useState(false);
const [providerParams, setProviderParams] = useState<ProviderParamsResponse | null>(providerParamsProp);
const [error, setError] = useState<string | null>(null);
// Fetch provider-specific parameters when component mounts
useEffect(() => {
if (providerParamsProp) {
// Props updated externally
setProviderParams(providerParamsProp);
return;
}
const fetchProviderParams = async () => {
if (!accessToken) return;
setLoading(true);
setError(null);
try {
const data = await getGuardrailProviderSpecificParams(accessToken);
console.log("Provider params API response:", data);
setProviderParams(data);
} catch (error) {
console.error("Error fetching provider params:", error);
setError("Failed to load provider parameters");
} finally {
setLoading(false);
}
};
// Only fetch if not provided via props
if (!providerParamsProp) {
fetchProviderParams();
}
}, [accessToken, providerParamsProp]);
// If no provider is selected, don't render anything
if (!selectedProvider) {
return null;
}
const providerFields = GUARDRAIL_PROVIDER_FIELDS[selectedProvider];
// Show loading state
if (loading) {
return <Spin tip="Loading provider parameters..." />;
}
// Show error state
if (error) {
return <div className="text-red-500">{error}</div>;
}
// Get the provider key matching the selected provider in the guardrail_provider_map
const providerKey = guardrail_provider_map[selectedProvider]?.toLowerCase();
// Get parameters for the selected provider
const providerFields = providerParams && providerParams[providerKey];
console.log("Provider key:", providerKey);
console.log("Provider fields:", providerFields);
if (!providerFields || providerFields.length === 0) {
return <div>No configuration fields available for this provider.</div>;
}
return (
<>
{providerFields.map((field) => (
<Form.Item
key={field.key}
name={field.key}
label={field.label}
tooltip={field.tooltip}
rules={field.required ? [{ required: true, message: `Please enter ${field.label}` }] : undefined}
key={field.param}
name={field.param}
label={field.param}
tooltip={field.description}
rules={field.required ? [{ required: true, message: `${field.param} is required` }] : undefined}
>
{field.type === "select" ? (
<Select placeholder={field.placeholder} defaultValue={field.defaultValue}>
{field.options?.map((option) => (
{field.type === "select" && field.options ? (
<Select
placeholder={field.description}
defaultValue={field.default_value}
>
{field.options.map((option) => (
<Select.Option key={option} value={option}>
{option}
</Select.Option>
))}
</Select>
) : field.param.includes("password") || field.param.includes("secret") ? (
<TextInput
placeholder={field.description}
type="password"
/>
) : (
<TextInput
placeholder={field.placeholder}
type={field.type === "password" ? "password" : "text"}
placeholder={field.description}
type="text"
/>
)}
</Form.Item>

View file

@ -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;
}
};