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:
Sameer Kankute 2026-05-28 22:09:24 +05:30
parent 8747c174c4
commit a982643c07
No known key found for this signature in database
2 changed files with 130 additions and 24 deletions

View file

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

View file

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