mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
fix(cato): guardrail all completion choices on output
When n > 1, only choices[0] was analyzed and redacted. Iterate every Choices entry so block and anonymize actions apply to all completions. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
8747c174c4
commit
a982643c07
2 changed files with 130 additions and 24 deletions
|
|
@ -261,32 +261,32 @@ class CatoNetworksGuardrail(CustomGuardrail):
|
|||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: Union[Any, ModelResponse, EmbeddingResponse, ImageResponse],
|
||||
) -> Any:
|
||||
if (
|
||||
isinstance(response, ModelResponse)
|
||||
and response.choices
|
||||
and isinstance(response.choices[0], Choices)
|
||||
):
|
||||
content = response.choices[0].message.content or ""
|
||||
cato_output_guardrail_result = await self.call_cato_guardrail_on_output(
|
||||
data,
|
||||
content,
|
||||
hook="output",
|
||||
key_alias=user_api_key_dict.key_alias,
|
||||
user_email=self._resolve_cato_user_email(user_api_key_dict),
|
||||
)
|
||||
if cato_output_guardrail_result and cato_output_guardrail_result.get(
|
||||
"detection_message"
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=cato_output_guardrail_result.get("detection_message"),
|
||||
if isinstance(response, ModelResponse) and response.choices:
|
||||
user_email = self._resolve_cato_user_email(user_api_key_dict)
|
||||
for choice in response.choices:
|
||||
if not isinstance(choice, Choices):
|
||||
continue
|
||||
content = choice.message.content or ""
|
||||
cato_output_guardrail_result = await self.call_cato_guardrail_on_output(
|
||||
data,
|
||||
content,
|
||||
hook="output",
|
||||
key_alias=user_api_key_dict.key_alias,
|
||||
user_email=user_email,
|
||||
)
|
||||
if cato_output_guardrail_result and cato_output_guardrail_result.get(
|
||||
"redacted_output"
|
||||
):
|
||||
response.choices[0].message.content = cato_output_guardrail_result.get(
|
||||
if cato_output_guardrail_result and cato_output_guardrail_result.get(
|
||||
"detection_message"
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=cato_output_guardrail_result.get("detection_message"),
|
||||
)
|
||||
if cato_output_guardrail_result and cato_output_guardrail_result.get(
|
||||
"redacted_output"
|
||||
)
|
||||
):
|
||||
choice.message.content = cato_output_guardrail_result.get(
|
||||
"redacted_output"
|
||||
)
|
||||
return response
|
||||
|
||||
async def async_post_call_streaming_iterator_hook(
|
||||
|
|
|
|||
|
|
@ -690,6 +690,112 @@ async def test_post_call_success_hook_no_action_keeps_content():
|
|||
assert result.choices[0].message.content == "all good"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_success_hook_block_action_raises_on_later_choice():
|
||||
guard = _make_guardrail()
|
||||
request_data = {"messages": [{"role": "user", "content": "hi"}]}
|
||||
block_response = _make_response(
|
||||
{
|
||||
"analysis_result": {"policy_drill_down": {"PII": {}}},
|
||||
"required_action": {
|
||||
"action_type": "block_action",
|
||||
"detection_message": "blocked output",
|
||||
"policy_name": "PII",
|
||||
},
|
||||
}
|
||||
)
|
||||
llm_response = ModelResponse(
|
||||
choices=[
|
||||
{
|
||||
"finish_reason": "stop",
|
||||
"index": 0,
|
||||
"message": {"content": "safe", "role": "assistant"},
|
||||
},
|
||||
{
|
||||
"finish_reason": "stop",
|
||||
"index": 1,
|
||||
"message": {"content": "secret", "role": "assistant"},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
async def mock_post_side_effect(url, *args, **kwargs):
|
||||
request_body = kwargs.get("json", {})
|
||||
assistant_content = request_body["messages"][-1]["content"]
|
||||
if assistant_content == "safe":
|
||||
return response_without_detections
|
||||
return block_response
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
side_effect=mock_post_side_effect,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guard.async_post_call_success_hook(
|
||||
data=request_data,
|
||||
response=llm_response,
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
)
|
||||
assert exc_info.value.status_code == 400
|
||||
assert exc_info.value.detail == "blocked output"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_success_hook_anonymize_action_redacts_all_choices():
|
||||
guard = _make_guardrail()
|
||||
request_data = {"messages": [{"role": "user", "content": "hi"}]}
|
||||
|
||||
def anonymize_response_for(content: str) -> Response:
|
||||
return _make_response(
|
||||
{
|
||||
"analysis_result": {"policy_drill_down": {"PII": {}}},
|
||||
"required_action": {
|
||||
"action_type": "anonymize_action",
|
||||
"policy_name": "PII",
|
||||
},
|
||||
"redacted_chat": {
|
||||
"all_redacted_messages": [
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": f"redacted {content}"},
|
||||
]
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
llm_response = ModelResponse(
|
||||
choices=[
|
||||
{
|
||||
"finish_reason": "stop",
|
||||
"index": 0,
|
||||
"message": {"content": "Hello Brian", "role": "assistant"},
|
||||
},
|
||||
{
|
||||
"finish_reason": "stop",
|
||||
"index": 1,
|
||||
"message": {"content": "Hi Alice", "role": "assistant"},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
async def mock_post_side_effect(url, *args, **kwargs):
|
||||
request_body = kwargs.get("json", {})
|
||||
assistant_content = request_body["messages"][-1]["content"]
|
||||
return anonymize_response_for(assistant_content)
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
side_effect=mock_post_side_effect,
|
||||
):
|
||||
result = await guard.async_post_call_success_hook(
|
||||
data=request_data,
|
||||
response=llm_response,
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
)
|
||||
assert isinstance(result, ModelResponse)
|
||||
assert result.choices[0].message.content == "redacted Hello Brian"
|
||||
assert result.choices[1].message.content == "redacted Hi Alice"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_success_hook_skips_non_model_response():
|
||||
guard = _make_guardrail()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue