diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py index 4199ad5ca65..125bea5590e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py @@ -227,6 +227,9 @@ class LLMShieldGuardrail(CustomGuardrail): if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call) is not True: return response + if self._is_anthropic_message_response(response): + return await self._restore_anthropic_response(response, data) + choices = getattr(response, "choices", None) if not choices: return response @@ -246,6 +249,34 @@ class LLMShieldGuardrail(CustomGuardrail): message.content = replacement return response + @staticmethod + def _is_anthropic_message_response(response: Any) -> bool: + """Anthropic's native /v1/messages reply arrives as a plain dict.""" + return ( + isinstance(response, dict) + and response.get("type") == "message" + and isinstance(response.get("content"), list) + ) + + async def _restore_anthropic_response(self, response: dict, data: dict) -> dict: + """Restores text blocks in an Anthropic native message reply. + + This shape has no `choices`, so without its own branch the reply would go + back to the caller still carrying placeholders. + """ + blocks = [ + block + for block in response["content"] + if isinstance(block, dict) and block.get("type") == "text" and isinstance(block.get("text"), str) + ] + if not blocks: + return response + + restored = await self._rehydrate([block["text"] for block in blocks], self._session_id(data)) + for block, replacement in zip(blocks, restored): + block["text"] = replacement + return response + async def async_post_call_streaming_iterator_hook( self, user_api_key_dict: UserAPIKeyAuth, diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py index 160c2300690..c2d50bf5fa7 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py @@ -11,7 +11,7 @@ from litellm.proxy.guardrails.guardrail_hooks.llm_shield.llm_shield import ( ) from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2 from litellm.types.guardrails import GuardrailEventHooks -from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices +from litellm.types.utils import Choices, Delta, Message, ModelResponse, ModelResponseStream, StreamingChoices def _guardrail(**overrides: object) -> LLMShieldGuardrail: @@ -156,6 +156,62 @@ class TestRedaction: assert len(sessions) == 1 +class TestRestoration: + @pytest.mark.asyncio + async def test_openai_shape_is_restored(self): + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + response = ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content="[EMAIL_1]"))]) + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result.choices[0].message.content == "a@b.com" + + @pytest.mark.asyncio + async def test_anthropic_message_shape_is_restored(self): + """The /v1/messages reply is a plain dict with no choices. + + Measured against a live provider: without its own branch the reply went + back to the caller still carrying the placeholder, even though the + request had been redacted correctly. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + response = { + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "[EMAIL_1]"}], + } + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result["content"][0]["text"] == "a@b.com" + + @pytest.mark.asyncio + async def test_anthropic_non_text_blocks_are_left_alone(self): + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + response = { + "type": "message", + "role": "assistant", + "content": [ + {"type": "text", "text": "[EMAIL_1]"}, + {"type": "tool_use", "id": "t1", "name": "lookup", "input": {}}, + ], + } + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result["content"][0]["text"] == "a@b.com" + assert result["content"][1] == {"type": "tool_use", "id": "t1", "name": "lookup", "input": {}} + + class TestFailClosed: @pytest.mark.asyncio async def test_unreachable_shield_blocks_the_request(self):