From a982643c07bfd3d2db710dd30e7c057e47b2f9d0 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Thu, 28 May 2026 22:09:24 +0530 Subject: [PATCH] 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 --- .../cato_networks/cato_networks.py | 48 ++++---- .../guardrail_hooks/test_cato_networks.py | 106 ++++++++++++++++++ 2 files changed, 130 insertions(+), 24 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py b/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py index 683df69238e..3ec1d00964d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py +++ b/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py @@ -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( diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cato_networks.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cato_networks.py index daae0d59a55..517c05c0ade 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cato_networks.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cato_networks.py @@ -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()