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 8989fabc321..e495f2e99d4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py @@ -128,6 +128,17 @@ def _collect_tool_arguments(message: MutableRequest, slots: _SlotSink) -> None: _collect(legacy, "arguments", slots) +def _collect_system(data: MutableRequest, slots: _SlotSink) -> None: + """Anthropic's /v1/messages carries its system prompt at the top level.""" + system: Final = data.get("system") + if isinstance(system, str): + _collect(data, "system", slots) + return + for part in system if isinstance(system, list) else (): + if isinstance(part, dict): + _collect(part, "text", slots) + + def _collect_responses_fields(data: MutableRequest, slots: _SlotSink) -> None: """The Responses API sends text outside `messages`, in `instructions` and `input`.""" _collect(data, "instructions", slots) @@ -135,7 +146,11 @@ def _collect_responses_fields(data: MutableRequest, slots: _SlotSink) -> None: if isinstance(request_input, str): _collect(data, "input", slots) return - for item in request_input if isinstance(request_input, list) else (): + for index, item in enumerate(request_input if isinstance(request_input, list) else ()): + if isinstance(item, str): + # The embeddings and moderations shape: `input` as an array of strings. + _collect_entry(request_input, index, slots) + continue if not isinstance(item, dict): continue _collect_content(item, slots) @@ -285,6 +300,7 @@ class LLMShieldGuardrail(CustomGuardrail): _collect_tool_arguments(message, slots) _collect_responses_fields(data, slots) _collect_prompt(data, slots) + _collect_system(data, slots) return tuple(slots) # --- hooks -------------------------------------------------------------------- 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 ffd6ede28b6..f16584fbd1d 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 @@ -292,6 +292,44 @@ class TestRequestCoverage: assert data["input"][0]["arguments"] == '{"email": "[EMAIL_1]"}' assert data["input"][1]["output"] == "sent to [EMAIL_1]" + @pytest.mark.asyncio + async def test_anthropic_system_prompt_is_redacted(self): + """/v1/messages carries its system prompt at the top level, not in messages.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["the user is [EMAIL_1]"]}) + + data = {"system": "the user is jane.doe@example.com", "messages": []} + await guardrail.async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="anthropic_messages" + ) + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["the user is jane.doe@example.com"] + assert data["system"] == "the user is [EMAIL_1]" + + @pytest.mark.asyncio + async def test_anthropic_system_blocks_are_redacted(self): + """`system` also accepts a list of text blocks.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = {"system": [{"type": "text", "text": "jane.doe@example.com"}], "messages": []} + await guardrail.async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="anthropic_messages" + ) + + assert data["system"][0]["text"] == "[EMAIL_1]" + + @pytest.mark.asyncio + async def test_string_array_input_is_redacted(self): + """Embeddings and moderations send `input` as an array of bare strings.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]", "[PHONE_1]"]}) + + data = {"input": ["jane.doe@example.com", "555-0100"]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aembedding") + + assert data["input"] == ["[EMAIL_1]", "[PHONE_1]"] + @pytest.mark.asyncio async def test_every_shape_in_one_request_is_redacted(self): guardrail = _guardrail()