mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
UI - Fix regression where Guardrail Entity Could not be selected and entity was not displayed (#16165)
* fix PiiEntityCategoryMap * fix OpenAIChatCompletionsHandler * fix lint
This commit is contained in:
parent
3c0d530197
commit
ca79ebe7d2
9 changed files with 138 additions and 18 deletions
|
|
@ -285,7 +285,7 @@ class CustomGuardrail(CustomLogger):
|
|||
data, self.event_hook
|
||||
)
|
||||
if result is not None:
|
||||
return result
|
||||
return result
|
||||
return True
|
||||
|
||||
def _event_hook_is_event_type(self, event_type: GuardrailEventHooks) -> bool:
|
||||
|
|
@ -403,6 +403,7 @@ class CustomGuardrail(CustomLogger):
|
|||
text: str,
|
||||
language: Optional[str] = None,
|
||||
entities: Optional[List[PiiEntityType]] = None,
|
||||
request_data: Optional[dict] = None,
|
||||
) -> str:
|
||||
"""
|
||||
Apply your guardrail logic to the given text
|
||||
|
|
@ -411,6 +412,7 @@ class CustomGuardrail(CustomLogger):
|
|||
text: The text to apply the guardrail to
|
||||
language: The language of the text
|
||||
entities: The entities to mask, optional
|
||||
request_data: The request data dictionary to store guardrail metadata
|
||||
|
||||
Any of the custom guardrails can override this method to provide custom guardrail logic
|
||||
|
||||
|
|
|
|||
|
|
@ -63,6 +63,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
tasks=tasks,
|
||||
task_mappings=task_mappings,
|
||||
guardrail_to_apply=guardrail_to_apply,
|
||||
request_data=data,
|
||||
)
|
||||
|
||||
# Step 2: Run all guardrail tasks in parallel
|
||||
|
|
@ -88,6 +89,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
tasks: List,
|
||||
task_mappings: List[Tuple[int, Optional[int]]],
|
||||
guardrail_to_apply: "CustomGuardrail",
|
||||
request_data: Optional[Dict[str, Any]] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Extract text content from a message and create guardrail tasks.
|
||||
|
|
@ -100,7 +102,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
|
||||
if isinstance(content, str):
|
||||
# Simple string content
|
||||
tasks.append(guardrail_to_apply.apply_guardrail(text=content))
|
||||
tasks.append(guardrail_to_apply.apply_guardrail(text=content, request_data=request_data))
|
||||
task_mappings.append((msg_idx, None))
|
||||
|
||||
elif isinstance(content, list):
|
||||
|
|
@ -109,7 +111,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
text_str = content_item.get("text", None)
|
||||
if text_str is None:
|
||||
continue
|
||||
tasks.append(guardrail_to_apply.apply_guardrail(text=text_str))
|
||||
tasks.append(guardrail_to_apply.apply_guardrail(text=text_str, request_data=request_data))
|
||||
task_mappings.append((msg_idx, int(content_idx)))
|
||||
|
||||
async def _apply_guardrail_responses_to_input(
|
||||
|
|
@ -217,6 +219,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
tasks: List,
|
||||
task_mappings: List[Tuple[int, Optional[int]]],
|
||||
guardrail_to_apply: "CustomGuardrail",
|
||||
request_data: Optional[Dict[str, Any]] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Extract text content from a response choice and create guardrail tasks.
|
||||
|
|
@ -233,7 +236,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
if choice.message.content and isinstance(choice.message.content, str):
|
||||
# Simple string content
|
||||
tasks.append(
|
||||
guardrail_to_apply.apply_guardrail(text=choice.message.content)
|
||||
guardrail_to_apply.apply_guardrail(text=choice.message.content, request_data=request_data)
|
||||
)
|
||||
task_mappings.append((choice_idx, None))
|
||||
|
||||
|
|
@ -242,7 +245,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
for content_idx, content_item in enumerate(choice.message.content):
|
||||
content_text = content_item.get("text")
|
||||
if content_text:
|
||||
tasks.append(guardrail_to_apply.apply_guardrail(text=content_text))
|
||||
tasks.append(guardrail_to_apply.apply_guardrail(text=content_text, request_data=request_data))
|
||||
task_mappings.append((choice_idx, int(content_idx)))
|
||||
|
||||
async def _apply_guardrail_responses_to_output(
|
||||
|
|
|
|||
|
|
@ -697,12 +697,15 @@ async def get_guardrail_ui_settings():
|
|||
# Convert the PII_ENTITY_CATEGORIES_MAP to the format expected by the UI
|
||||
category_maps = []
|
||||
for category, entities in PII_ENTITY_CATEGORIES_MAP.items():
|
||||
category_maps.append({"category": category, "entities": entities})
|
||||
category_maps.append({
|
||||
"category": category.value,
|
||||
"entities": [entity.value for entity in entities]
|
||||
})
|
||||
|
||||
return GuardrailUIAddGuardrailSettings(
|
||||
supported_entities=list(PiiEntityType),
|
||||
supported_actions=list(PiiAction),
|
||||
supported_modes=list(GuardrailEventHooks),
|
||||
supported_entities=[entity.value for entity in PiiEntityType],
|
||||
supported_actions=[action.value for action in PiiAction],
|
||||
supported_modes=[mode.value for mode in GuardrailEventHooks],
|
||||
pii_entity_categories=category_maps,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -1156,22 +1156,34 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
text: str,
|
||||
language: Optional[str] = None,
|
||||
entities: Optional[List[PiiEntityType]] = None,
|
||||
request_data: Optional[dict] = None,
|
||||
) -> str:
|
||||
"""
|
||||
Apply Bedrock guardrail to the given text for testing purposes.
|
||||
|
||||
This method allows users to test Bedrock guardrails without making actual LLM calls.
|
||||
It creates a mock request and response to test the guardrail functionality.
|
||||
|
||||
Args:
|
||||
text: The text to analyze
|
||||
language: Optional language parameter (not used by Bedrock)
|
||||
entities: Optional entities parameter (not used by Bedrock)
|
||||
request_data: Optional request data dictionary for logging metadata
|
||||
"""
|
||||
try:
|
||||
verbose_proxy_logger.debug("Bedrock Guardrail: Applying guardrail")
|
||||
mock_messages: List[AllMessageValues] = [
|
||||
ChatCompletionUserMessage(role="user", content=text)
|
||||
]
|
||||
|
||||
# Use provided request_data or create a mock one for testing
|
||||
if request_data is None:
|
||||
request_data = {"messages": mock_messages}
|
||||
|
||||
bedrock_response = await self.make_bedrock_api_request(
|
||||
source="INPUT",
|
||||
messages=mock_messages,
|
||||
request_data={"messages": mock_messages},
|
||||
request_data=request_data,
|
||||
)
|
||||
|
||||
if bedrock_response.get("action") == "BLOCKED":
|
||||
|
|
|
|||
|
|
@ -491,10 +491,11 @@ class EnkryptAIGuardrails(CustomGuardrail):
|
|||
text: str,
|
||||
language: Optional[str] = None,
|
||||
entities: Optional[List[PiiEntityType]] = None,
|
||||
request_data: Optional[dict] = None,
|
||||
) -> str:
|
||||
result = await self._call_enkryptai_guardrails(
|
||||
prompt=text,
|
||||
request_data={},
|
||||
request_data=request_data or {},
|
||||
)
|
||||
# Process the guardrails response
|
||||
processed_result = self._process_enkryptai_guardrails_response(result)
|
||||
|
|
|
|||
|
|
@ -856,6 +856,7 @@ class NomaGuardrail(CustomGuardrail):
|
|||
text: str,
|
||||
language: Optional[str] = None,
|
||||
entities: Optional[List[PiiEntityType]] = None,
|
||||
request_data: Optional[dict] = None,
|
||||
) -> str:
|
||||
"""
|
||||
Apply Noma guardrail to the given text for testing purposes.
|
||||
|
|
@ -867,6 +868,7 @@ class NomaGuardrail(CustomGuardrail):
|
|||
text: The text to analyze
|
||||
language: Optional language parameter (not used by Noma)
|
||||
entities: Optional entities parameter (not used by Noma)
|
||||
request_data: Optional request data dictionary for logging metadata
|
||||
|
||||
Returns:
|
||||
The original text if allowed, or anonymized text if available
|
||||
|
|
@ -884,11 +886,15 @@ class NomaGuardrail(CustomGuardrail):
|
|||
# Create payload for Noma API
|
||||
payload = {"request": {"text": text}}
|
||||
|
||||
# Use provided request_data or create a mock one for testing
|
||||
if request_data is None:
|
||||
request_data = {"messages": [{"role": "user", "content": text}]}
|
||||
|
||||
# Call Noma API
|
||||
response_json = await self._call_noma_api(
|
||||
payload=payload,
|
||||
llm_request_id=None,
|
||||
request_data={"messages": [{"role": "user", "content": text}]},
|
||||
request_data=request_data,
|
||||
user_auth=mock_user_auth,
|
||||
extra_data={},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -702,6 +702,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
text: str,
|
||||
language: Optional[str] = None,
|
||||
entities: Optional[List[PiiEntityType]] = None,
|
||||
request_data: Optional[dict] = None,
|
||||
) -> str:
|
||||
"""
|
||||
UI will call this function to check:
|
||||
|
|
@ -712,7 +713,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
text=text,
|
||||
output_parse_pii=self.output_parse_pii,
|
||||
presidio_config=None,
|
||||
request_data={},
|
||||
request_data=request_data or {},
|
||||
)
|
||||
return text
|
||||
|
||||
|
|
|
|||
|
|
@ -231,8 +231,8 @@ PII_ENTITY_CATEGORIES_MAP = {
|
|||
|
||||
|
||||
class PiiEntityCategoryMap(TypedDict):
|
||||
category: PiiEntityCategory
|
||||
entities: List[PiiEntityType]
|
||||
category: str
|
||||
entities: List[str]
|
||||
|
||||
|
||||
class GuardrailParamUITypes(str, Enum):
|
||||
|
|
@ -611,9 +611,9 @@ class ListGuardrailsResponse(BaseModel):
|
|||
|
||||
|
||||
class GuardrailUIAddGuardrailSettings(BaseModel):
|
||||
supported_entities: List[PiiEntityType]
|
||||
supported_actions: List[PiiAction]
|
||||
supported_modes: List[GuardrailEventHooks]
|
||||
supported_entities: List[str]
|
||||
supported_actions: List[str]
|
||||
supported_modes: List[str]
|
||||
pii_entity_categories: List[PiiEntityCategoryMap]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -533,6 +533,98 @@ async def test_logging_hook_multiple_content_items(presidio_guardrail):
|
|||
print("✓ Logging hook multiple content items test passed")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_presidio_sets_guardrail_information_in_request_data():
|
||||
"""
|
||||
Test that Presidio populates guardrail information into request_data metadata.
|
||||
|
||||
This validates that add_standard_logging_guardrail_information_to_request_data
|
||||
correctly sets the guardrail information that will be used for logging.
|
||||
"""
|
||||
presidio = _OPTIONAL_PresidioPIIMasking(
|
||||
guardrail_name="test_presidio",
|
||||
output_parse_pii=True,
|
||||
)
|
||||
|
||||
request_data = {
|
||||
"messages": [{"role": "user", "content": "Test"}],
|
||||
"model": "gpt-4o",
|
||||
"metadata": {},
|
||||
}
|
||||
|
||||
async def mock_check_pii(text, output_parse_pii, presidio_config, request_data):
|
||||
assert request_data is not None
|
||||
|
||||
presidio.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_provider="presidio",
|
||||
guardrail_json_response=[],
|
||||
request_data=request_data,
|
||||
guardrail_status="success",
|
||||
start_time=1234567890.0,
|
||||
end_time=1234567891.0,
|
||||
duration=1.0,
|
||||
masked_entity_count={"EMAIL_ADDRESS": 1, "PERSON": 1},
|
||||
)
|
||||
|
||||
return text
|
||||
|
||||
with patch.object(presidio, 'check_pii', mock_check_pii):
|
||||
await presidio.apply_guardrail(
|
||||
text="Test message",
|
||||
request_data=request_data,
|
||||
)
|
||||
|
||||
assert "metadata" in request_data
|
||||
assert "standard_logging_guardrail_information" in request_data["metadata"]
|
||||
|
||||
guardrail_info = request_data["metadata"]["standard_logging_guardrail_information"]
|
||||
assert "masked_entity_count" in guardrail_info
|
||||
assert guardrail_info["masked_entity_count"]["EMAIL_ADDRESS"] == 1
|
||||
assert guardrail_info["masked_entity_count"]["PERSON"] == 1
|
||||
|
||||
print("✓ Presidio sets guardrail_information in request_data")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_request_data_flows_to_apply_guardrail():
|
||||
"""
|
||||
Test that request_data is correctly passed to apply_guardrail method.
|
||||
|
||||
This validates the fix where guardrail translation handler passes data
|
||||
as request_data to apply_guardrail so guardrails can store metadata for logging.
|
||||
"""
|
||||
presidio = _OPTIONAL_PresidioPIIMasking(
|
||||
guardrail_name="test_presidio",
|
||||
output_parse_pii=True,
|
||||
)
|
||||
|
||||
request_data = {
|
||||
"messages": [{"role": "user", "content": "Test message"}],
|
||||
"model": "gpt-4o",
|
||||
"metadata": {},
|
||||
}
|
||||
|
||||
async def mock_check_pii(text, output_parse_pii, presidio_config, request_data):
|
||||
assert request_data is not None, "request_data should be passed to check_pii"
|
||||
assert "metadata" in request_data, "request_data should have metadata"
|
||||
|
||||
request_data.setdefault("metadata", {})
|
||||
request_data["metadata"]["test_flag"] = "passed_correctly"
|
||||
|
||||
return text
|
||||
|
||||
with patch.object(presidio, 'check_pii', mock_check_pii):
|
||||
result = await presidio.apply_guardrail(
|
||||
text="Test message",
|
||||
request_data=request_data,
|
||||
)
|
||||
|
||||
assert "metadata" in request_data
|
||||
assert request_data["metadata"].get("test_flag") == "passed_correctly"
|
||||
|
||||
print("✓ request_data correctly passed to apply_guardrail")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Run tests
|
||||
asyncio.run(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue