From 09314f239c34f35a86ce4155b8c4d445944a6038 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sun, 13 Sep 2026 00:54:29 -0700 Subject: [PATCH] fix(guardrails): hand per-message rewrites back as structured_messages A guardrail that rewrites text per chat message now returns the rewritten rows as structured_messages instead of only texts, so the Responses and chat handlers write the rewrite back through the structured path. The generic guardrail API response accepts an optional structured_messages list, Prompt Security modify builds one from modified_messages, and rows a server echoes back exactly as shown are restored to the original row objects because the request model drops undeclared keys. Texts-only per-message answers keep the named rejection on both endpoints. --- .../base_llm/guardrail_translation/utils.py | 42 ++--- .../guardrail_translation/handler.py | 27 +--- .../generic_guardrail_api.py | 30 +++- .../prompt_security/prompt_security.py | 8 +- .../guardrail_hooks/generic_guardrail_api.py | 15 +- ...test_openai_responses_guardrail_handler.py | 148 +++++++++--------- .../test_generic_guardrail_api.py | 105 +++++++++++++ .../test_prompt_security_guardrails.py | 14 +- 8 files changed, 244 insertions(+), 145 deletions(-) diff --git a/litellm/llms/base_llm/guardrail_translation/utils.py b/litellm/llms/base_llm/guardrail_translation/utils.py index 383e668e45c..1172c93959b 100644 --- a/litellm/llms/base_llm/guardrail_translation/utils.py +++ b/litellm/llms/base_llm/guardrail_translation/utils.py @@ -2,7 +2,6 @@ from __future__ import annotations import json from collections.abc import Callable, Iterator, Mapping, Sequence -from itertools import accumulate from types import MappingProxyType from typing import Final, TypeVar, cast # noqa: TID251 # a rebuilt chat row has no typed constructor across roles @@ -390,7 +389,7 @@ def _part_with_text(part: object, text: str) -> object: return {**part, "text": text} # mutable-ok: content parts stay JSON-plain dicts -def _content_with_slot_texts(content: Sequence[object], texts: Sequence[str]) -> list[object]: +def _content_with_slot_texts(content: Sequence[object], texts: Sequence[str]) -> Sequence[object]: text_part_indices: Final = tuple( index for index, part in enumerate(content) if _content_part_text(part) is not None ) @@ -401,39 +400,18 @@ def _content_with_slot_texts(content: Sequence[object], texts: Sequence[str]) -> ] -def _message_with_slot_texts(message: AllMessageValues, texts: Sequence[str]) -> AllMessageValues: +def message_with_slot_texts(message: AllMessageValues, texts: Sequence[str]) -> AllMessageValues | None: + """Swap one rewritten text into each text slot of a chat row, in order. + + A slot is a string ``content`` or one list part carrying a string ``text``; + images and other parts ride along untouched. Returns None unless the counts + line up exactly, so a rewrite never lands on the wrong slot. + """ + if message_text_slot_count(message) != len(texts): + return None content: Final = message.get("content") if not isinstance(content, (str, list)): return message rewritten_content: Final = texts[0] if isinstance(content, str) else _content_with_slot_texts(content, texts) rewritten: Final = {**message, "content": rewritten_content} # mutable-ok: chat rows stay JSON-plain dicts return cast("AllMessageValues", rewritten) # cast-ok: the same row with only its text slots swapped - - -def message_with_slot_texts(message: AllMessageValues, texts: Sequence[str]) -> AllMessageValues | None: - if message_text_slot_count(message) != len(texts): - return None - return _message_with_slot_texts(message, texts) - - -def messages_with_slot_texts( - messages: Sequence[AllMessageValues], - texts: Sequence[str], -) -> list[AllMessageValues] | None: - """Spread one flat list of rewritten texts over the messages' text slots, in order. - - A slot is a string ``content`` or one list part carrying a string ``text``; - images and other parts ride along untouched. A guardrail that answers one - text per message it saw produces exactly this shape, which stops matching - the endpoint's own per-text extraction as soon as the request carries - instructions or tool items. Returns None unless the counts line up exactly, - so a rewrite never lands on the wrong slot. - """ - slot_counts: Final = tuple(message_text_slot_count(message) for message in messages) - if sum(slot_counts) != len(texts): - return None - offsets: Final = tuple(accumulate(slot_counts, initial=0)) - return [ # mutable-ok: guardrail rows travel as a list - _message_with_slot_texts(message, texts[start:end]) - for message, start, end in zip(messages, offsets, offsets[1:]) - ] diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py index 61903a54cd3..2fe11d9f7bd 100644 --- a/litellm/llms/openai/responses/guardrail_translation/handler.py +++ b/litellm/llms/openai/responses/guardrail_translation/handler.py @@ -53,7 +53,6 @@ from litellm.llms.base_llm.guardrail_translation.base_translation import ( ) from litellm.llms.base_llm.guardrail_translation.utils import ( blocked_responses_stream_usage, - messages_with_slot_texts, stream_item_field, stream_item_fingerprint, stream_item_items, @@ -396,20 +395,6 @@ def _patched_request_fields( ) -def _guardrailed_structured_messages( - structured_messages: Sequence[AllMessageValues] | None, - sent_text_count: int, - guardrailed_inputs: GenericGuardrailAPIInputs, -) -> Sequence[AllMessageValues] | None: - returned: Final = guardrailed_inputs.get("structured_messages") - if returned is not None and returned is not structured_messages: - return returned - rewritten_texts: Final = guardrailed_inputs.get("texts") - if not structured_messages or rewritten_texts is None or len(rewritten_texts) == sent_text_count: - return None - return messages_with_slot_texts(structured_messages, rewritten_texts) - - def _patch_or_convert_request_fields( raw_input: object, instructions: object, @@ -488,8 +473,7 @@ class OpenAIResponsesHandler(BaseTranslation): form.chat_tools for form in LiteLLMCompletionResponsesConfig.responses_tools_to_chat_forms(original_tools) ) extracted: Final = self._extract_guardrail_inputs(data, input_data, flattened_tool_groups) - sent_texts: Final = extracted.inputs.get("texts") - if not sent_texts: + if not extracted.inputs.get("texts"): return data if structured_messages: extracted.inputs["structured_messages"] = structured_messages @@ -502,9 +486,7 @@ class OpenAIResponsesHandler(BaseTranslation): self._apply_guardrailed_tools_to_data( data, original_tools, flattened_tool_groups, guardrailed_inputs.get("tools") ) - written_back: Final = self._written_back_request_fields( - data, structured_messages, len(sent_texts), guardrailed_inputs - ) + written_back: Final = self._written_back_request_fields(data, structured_messages, guardrailed_inputs) if written_back is not None: data["input"] = list(written_back.input) # mutable-ok: JSON body if written_back.instructions is None: @@ -571,11 +553,10 @@ class OpenAIResponsesHandler(BaseTranslation): def _written_back_request_fields( data: Mapping[str, object], structured_messages: Sequence[AllMessageValues] | None, - sent_text_count: int, guardrailed_inputs: GenericGuardrailAPIInputs, ) -> _RequestFields | None: - guardrailed: Final = _guardrailed_structured_messages(structured_messages, sent_text_count, guardrailed_inputs) - if guardrailed is None: + guardrailed: Final = guardrailed_inputs.get("structured_messages") + if guardrailed is None or guardrailed is structured_messages: return None return _patch_or_convert_request_fields( data.get("input"), diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py index d8296003ae9..16159d32a7f 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py @@ -7,7 +7,7 @@ import fnmatch import os -from collections.abc import Mapping +from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, Final, Literal, Optional import httpx @@ -24,7 +24,7 @@ from litellm.llms.custom_httpx.http_handler import ( httpxSpecialProvider, ) from litellm.types.guardrails import GuardrailEventHooks -from litellm.types.llms.openai import ChatCompletionToolParam +from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ( GenericGuardrailAPIMetadata, GenericGuardrailAPIRequest, @@ -150,6 +150,22 @@ def _extract_inbound_headers( return None +def _rows_with_unchanged_originals( + original_rows: Sequence[AllMessageValues] | None, + shown_rows: Sequence[AllMessageValues] | None, + returned_rows: Sequence[AllMessageValues], +) -> tuple[AllMessageValues, ...]: + """The request model drops row keys its message types do not declare, so a + row the server echoes back verbatim is restored to the original row object; + only rows the server actually changed reach the endpoint write-back.""" + if original_rows is None or shown_rows is None or len(returned_rows) != len(original_rows): + return tuple(returned_rows) + return tuple( + original if returned == shown else returned + for original, shown, returned in zip(original_rows, shown_rows, returned_rows) + ) + + class GenericGuardrailAPI(CustomGuardrail): """ Generic Guardrail API integration for LiteLLM. @@ -322,6 +338,8 @@ class GenericGuardrailAPI(CustomGuardrail): texts: list, images: list[str] | None, tools: list[ChatCompletionToolParam] | None, + structured_messages: Sequence[AllMessageValues] | None, + shown_messages: Sequence[AllMessageValues] | None, guardrail_response: GenericGuardrailAPIResponse, ) -> GenericGuardrailAPIInputs: # Action is NONE or no modifications needed @@ -336,6 +354,12 @@ class GenericGuardrailAPI(CustomGuardrail): return_inputs["tools"] = guardrail_response.tools elif tools: return_inputs["tools"] = tools + if guardrail_response.structured_messages: + return_inputs["structured_messages"] = list( # mutable-ok: guardrail inputs take a list + _rows_with_unchanged_originals( + structured_messages, shown_messages, guardrail_response.structured_messages + ) + ) if guardrail_response.stream_holdback_chars is not None: return_inputs["stream_holdback_chars"] = guardrail_response.stream_holdback_chars return return_inputs @@ -473,6 +497,8 @@ class GenericGuardrailAPI(CustomGuardrail): texts=texts, images=images, tools=tools, + structured_messages=structured_messages, + shown_messages=guardrail_request.structured_messages, guardrail_response=guardrail_response, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py b/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py index 72c87a793a5..41d3b202344 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py +++ b/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py @@ -287,7 +287,7 @@ class PromptSecurityGuardrail(CustomGuardrail): structured_messages, modified_messages ) if rewritten_messages is not None: - inputs["structured_messages"] = rewritten_messages + inputs["structured_messages"] = list(rewritten_messages) # mutable-ok: guardrail inputs take a list return inputs @@ -298,7 +298,7 @@ class PromptSecurityGuardrail(CustomGuardrail): self, structured_messages: Sequence[AllMessageValues], modified_messages: Sequence[Mapping[str, object]], - ) -> list[AllMessageValues] | None: + ) -> tuple[AllMessageValues, ...] | None: sent_indices: Final = tuple( index for index, message in enumerate(structured_messages) if self._is_sent_to_protect(message) ) @@ -313,9 +313,7 @@ class PromptSecurityGuardrail(CustomGuardrail): ) if len(replacements) != len(sent_indices): return None - return [ # mutable-ok: guardrail inputs take a list - replacements.get(index, message) for index, message in enumerate(structured_messages) - ] + return tuple(replacements.get(index, message) for index, message in enumerate(structured_messages)) async def _apply_guardrail_on_response( self, diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py index 4a868c48352..44e2cc2404f 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py @@ -1,4 +1,5 @@ -from typing import Any, Final, Literal +from collections.abc import Mapping, Sequence +from typing import Any, Final, Literal, cast # noqa: TID251 # JSON chat rows have no typed constructor across roles from pydantic import BaseModel, ConfigDict, Field from typing_extensions import TypedDict @@ -158,12 +159,21 @@ def coerce_stream_holdback_value(value: Any) -> int: return 0 +def structured_messages_from_response(value: object) -> Sequence[AllMessageValues] | None: + if not isinstance(value, list): + return None + if not all(isinstance(message, Mapping) and isinstance(message.get("role"), str) for message in value): + return None + return cast("Sequence[AllMessageValues]", value) # cast-ok: JSON rows checked for a role, the same trust texts get + + class GenericGuardrailAPIResponse: """Response model for the Generic Guardrail API""" texts: list[str] | None images: list[str] | None tools: list[GuardrailToolParam] | None + structured_messages: Sequence[AllMessageValues] | None action: str blocked_reason: str | None stream_holdback_chars: list[int] | None @@ -176,12 +186,14 @@ class GenericGuardrailAPIResponse: images: list[str] | None = None, tools: list[GuardrailToolParam] | None = None, stream_holdback_chars: list[int] | None = None, + structured_messages: Sequence[AllMessageValues] | None = None, ) -> None: self.action = action self.blocked_reason = blocked_reason self.texts = texts self.images = images self.tools = tools + self.structured_messages = structured_messages # Number of trailing chars, indexed the same as ``texts``, that the # framework must withhold from streaming emission until the next # processing round (word-boundary safety for text transformations). @@ -200,4 +212,5 @@ class GenericGuardrailAPIResponse: images=data.get("images"), tools=data.get("tools"), stream_holdback_chars=stream_holdback_chars, + structured_messages=structured_messages_from_response(data.get("structured_messages")), ) diff --git a/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py b/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py index b4b467773da..394f134e99c 100644 --- a/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py +++ b/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py @@ -8,7 +8,7 @@ with guardrail transformations. import copy from collections.abc import Callable from typing import Any, List, Literal, Optional, Tuple -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch import logging @@ -31,6 +31,7 @@ from litellm.llms.openai.responses.guardrail_translation.handler import ( OpenAIResponsesHandler, ) from litellm.llms.openai.responses.guardrail_translation.tool_merge import merge_guardrailed_tools +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI from litellm.types.llms.openai import ChatCompletionToolCallChunk from litellm.responses.litellm_completion_transformation.transformation import ( LiteLLMCompletionResponsesConfig, @@ -2342,103 +2343,94 @@ SSN = "123-45-6789" REDACTED_SSN = "" -def _slot_texts(message: dict) -> list[str]: - content = message.get("content") - if isinstance(content, str): - return [content] - if isinstance(content, list): - return [part["text"] for part in content if isinstance(part, dict) and isinstance(part.get("text"), str)] - return [] +def _redacted(value: object) -> object: + if isinstance(value, str): + return value.replace(SSN, REDACTED_SSN) + if isinstance(value, list): + return [{**part, "text": _redacted(part["text"])} if "text" in part else part for part in value] + return value -class PerMessageRedactionGuardrail(CustomGuardrail): - """Guardrail that answers one redacted text per message it was shown and hands - back only texts, the way Prompt Security in modify mode and a generic guardrail - API server that scans per message do.""" +def _per_message_guardrail_server(structured_messages_in_answer: bool) -> Callable[..., MagicMock]: + """Answers one redacted text per chat row it was shown, the way a guardrail + that scans per message does, and optionally the rewritten rows themselves.""" - def __init__(self, extra_texts: int = 0): - super().__init__(guardrail_name="per-message-redactor") - self.extra_texts = extra_texts + def post(url: str, json: dict, headers: dict) -> MagicMock: + rows = json["structured_messages"] + answer: dict = { + "action": "GUARDRAIL_INTERVENED", + "texts": [_redacted(row["content"]) if isinstance(row.get("content"), str) else "" for row in rows], + } + if structured_messages_in_answer: + answer["structured_messages"] = [{**row, "content": _redacted(row.get("content"))} for row in rows] + response = MagicMock() + response.json.return_value = answer + response.raise_for_status = MagicMock() + return response - async def apply_guardrail( - self, - inputs: GenericGuardrailAPIInputs, - request_data: dict, - input_type: Literal["request", "response"], - logging_obj: Optional[Any] = None, - ) -> GenericGuardrailAPIInputs: - messages = inputs.get("structured_messages") or [] - texts = [text.replace(SSN, REDACTED_SSN) for message in messages for text in _slot_texts(message)] - return {**inputs, "texts": texts + ["junk"] * self.extra_texts} + return post -class TestPerMessageTextWriteBack: - """A guardrail that rewrites one text per message it saw must land on the - instructions and the input items those messages came from, not be rejected.""" +def _per_message_redactor() -> GenericGuardrailAPI: + return GenericGuardrailAPI( + api_base="https://guardrail.test", + guardrail_name="per-message-redactor", + event_hook="pre_call", + default_on=True, + ) + + +def _tool_replay_request() -> dict: + return { + "model": "gpt-5.6", + "instructions": "Never repeat the SSN " + SSN + " back.", + "input": [ + {"role": "user", "content": "Look up " + SSN + " for me."}, + {"type": "function_call", "call_id": "call_1", "name": "lookup_customer", "arguments": '{"id": "42"}'}, + {"type": "function_call_output", "call_id": "call_1", "output": '{"ssn": "' + SSN + '"}'}, + ], + } + + +class TestPerMessageRewriteWriteBack: + """A guardrail that rewrites per chat row hands the rows back as + structured_messages, and the handler lands them on the instructions and the + input items they came from; the same rewrite handed back as texts alone has + no item to land on and is rejected by name instead of sent unrewritten.""" @pytest.mark.asyncio - async def test_instructions_plus_tool_replay_gets_each_rewrite_in_place(self): - handler = OpenAIResponsesHandler() - function_call_item = { - "type": "function_call", - "call_id": "call_1", - "name": "lookup_customer", - "arguments": '{"query": "' + SSN + '"}', - } - data = { - "model": "gpt-5.6", - "instructions": "Never repeat the SSN " + SSN + " back.", - "input": [ - {"role": "user", "content": "Look up " + SSN + " for me."}, - function_call_item, - {"type": "function_call_output", "call_id": "call_1", "output": '{"ssn": "' + SSN + '"}'}, - ], - } + async def test_structured_rows_land_on_instructions_and_tool_output(self): + guardrail = _per_message_redactor() + data = _tool_replay_request() + function_call_item = data["input"][1] - result = await handler.process_input_messages(data, PerMessageRedactionGuardrail()) + with patch.object(guardrail.async_handler, "post", side_effect=_per_message_guardrail_server(True)): + result = await OpenAIResponsesHandler().process_input_messages(data, guardrail) assert result["instructions"] == "Never repeat the SSN " + REDACTED_SSN + " back." - assert [item.get("type", item.get("role")) for item in result["input"]] == [ - "user", - "function_call", - "function_call_output", - ] - assert _slot_texts(result["input"][0]) == ["Look up " + REDACTED_SSN + " for me."] + assert _texts(result["input"][0]) == ["Look up " + REDACTED_SSN + " for me."] assert result["input"][1] == function_call_item - assert result["input"][2]["output"] == '{"ssn": "' + REDACTED_SSN + '"}' - assert result["input"][2]["call_id"] == "call_1" - - @pytest.mark.asyncio - async def test_string_input_with_instructions_keeps_the_two_apart(self): - handler = OpenAIResponsesHandler() - data = { - "model": "gpt-5.6", - "instructions": "Redact " + SSN + " everywhere.", - "input": "My SSN is " + SSN + ".", + assert result["input"][2] == { + "type": "function_call_output", + "call_id": "call_1", + "output": '{"ssn": "' + REDACTED_SSN + '"}', } - result = await handler.process_input_messages(data, PerMessageRedactionGuardrail()) - - assert result["instructions"] == "Redact " + REDACTED_SSN + " everywhere." - assert [_slot_texts(item) for item in result["input"]] == [["My SSN is " + REDACTED_SSN + "."]] - @pytest.mark.asyncio - async def test_count_matching_neither_texts_nor_messages_is_still_rejected(self): + async def test_texts_only_per_message_answer_is_rejected_by_name(self): from litellm.proxy.policy_engine.pipeline_executor import UnappliableRequestRewrite - handler = OpenAIResponsesHandler() - original_input = [ - {"role": "user", "content": "Look up " + SSN + " for me."}, - {"type": "function_call_output", "call_id": "call_1", "output": '{"ssn": "' + SSN + '"}'}, - ] - data = {"model": "gpt-5.6", "instructions": "Be terse.", "input": copy.deepcopy(original_input)} + guardrail = _per_message_redactor() + data = _tool_replay_request() + original = copy.deepcopy(data) - with pytest.raises(UnappliableRequestRewrite) as excinfo: - await handler.process_input_messages(data, PerMessageRedactionGuardrail(extra_texts=1)) + with patch.object(guardrail.async_handler, "post", side_effect=_per_message_guardrail_server(False)): + with pytest.raises(UnappliableRequestRewrite) as excinfo: + await OpenAIResponsesHandler().process_input_messages(data, guardrail) assert excinfo.value.guardrail_name == "per-message-redactor" - assert data["input"] == original_input - assert data["instructions"] == "Be terse." + assert data["input"] == original["input"] + assert data["instructions"] == original["instructions"] class TestProvenancePatching: diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py index 83cc9ae8bb9..cc9942e0e40 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py @@ -582,6 +582,111 @@ class TestGuardrailActions: assert result_images is None +class TestStructuredMessagesInResponse: + """A guardrail server that rewrites per chat row answers with the rewritten + rows as structured_messages, which the endpoint handlers write back by row.""" + + @pytest.mark.asyncio + async def test_returned_rows_are_handed_back_as_structured_messages( + self, generic_guardrail, mock_request_data_input + ): + rewritten_rows = [ + {"role": "system", "content": "Never repeat an SSN."}, + {"role": "user", "content": "Look up [REDACTED] for me."}, + {"role": "tool", "tool_call_id": "call_1", "content": '{"ssn": "[REDACTED]"}'}, + ] + mock_response = MagicMock() + mock_response.json.return_value = { + "action": "GUARDRAIL_INTERVENED", + "texts": ["Never repeat an SSN.", "Look up [REDACTED] for me.", '{"ssn": "[REDACTED]"}'], + "structured_messages": rewritten_rows, + } + mock_response.raise_for_status = MagicMock() + + with patch.object(generic_guardrail.async_handler, "post", return_value=mock_response): + guardrailed_inputs = await generic_guardrail.apply_guardrail( + inputs={"texts": ["Look up 123-45-6789 for me."]}, + request_data=mock_request_data_input, + input_type="request", + ) + + assert guardrailed_inputs["structured_messages"] == rewritten_rows + assert guardrailed_inputs["texts"] == mock_response.json.return_value["texts"] + + @pytest.mark.asyncio + async def test_rows_echoed_back_as_shown_keep_their_original_keys( + self, generic_guardrail, mock_request_data_input + ): + tool_call_row = { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}, "index": 0} + ], + } + original_rows = [ + {"role": "user", "content": "Look up 123-45-6789 for me.", "name": "pat"}, + tool_call_row, + {"role": "tool", "tool_call_id": "call_1", "content": '{"ssn": "123-45-6789"}'}, + ] + + def echo_with_tool_output_redacted(url, json, headers): + shown_rows = json["structured_messages"] + assert "index" not in shown_rows[1]["tool_calls"][0] + assert "name" not in shown_rows[0] + answer = MagicMock() + answer.json.return_value = { + "action": "GUARDRAIL_INTERVENED", + "texts": ["Look up 123-45-6789 for me."], + "structured_messages": [ + shown_rows[0], + shown_rows[1], + {**shown_rows[2], "content": '{"ssn": "[REDACTED]"}'}, + ], + } + answer.raise_for_status = MagicMock() + return answer + + with patch.object(generic_guardrail.async_handler, "post", side_effect=echo_with_tool_output_redacted): + guardrailed_inputs = await generic_guardrail.apply_guardrail( + inputs={"texts": ["Look up 123-45-6789 for me."], "structured_messages": original_rows}, + request_data=mock_request_data_input, + input_type="request", + ) + + returned_rows = guardrailed_inputs["structured_messages"] + assert returned_rows[0] is original_rows[0] + assert returned_rows[1] is tool_call_row + assert returned_rows[2] == {"role": "tool", "tool_call_id": "call_1", "content": '{"ssn": "[REDACTED]"}'} + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "structured_messages", + [[], [{"content": "a row with no role"}], "not a list"], + ids=["empty", "no_role", "not_a_list"], + ) + async def test_rows_that_are_not_chat_messages_are_ignored( + self, generic_guardrail, mock_request_data_input, structured_messages + ): + mock_response = MagicMock() + mock_response.json.return_value = { + "action": "GUARDRAIL_INTERVENED", + "texts": ["[REDACTED]"], + "structured_messages": structured_messages, + } + mock_response.raise_for_status = MagicMock() + + with patch.object(generic_guardrail.async_handler, "post", return_value=mock_response): + guardrailed_inputs = await generic_guardrail.apply_guardrail( + inputs={"texts": ["Look up 123-45-6789 for me."]}, + request_data=mock_request_data_input, + input_type="request", + ) + + assert "structured_messages" not in guardrailed_inputs + assert guardrailed_inputs["texts"] == ["[REDACTED]"] + + class TestImageSupport: """Test image handling in guardrail requests""" diff --git a/tests/test_litellm/proxy/guardrails/test_prompt_security_guardrails.py b/tests/test_litellm/proxy/guardrails/test_prompt_security_guardrails.py index 9e83098eb04..70083e50f01 100644 --- a/tests/test_litellm/proxy/guardrails/test_prompt_security_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/test_prompt_security_guardrails.py @@ -1,5 +1,6 @@ import asyncio import base64 +from collections.abc import Mapping, Sequence from unittest.mock import AsyncMock, patch import pytest @@ -12,6 +13,7 @@ from litellm.proxy.guardrails.guardrail_hooks.prompt_security.prompt_security im PromptSecurityGuardrailMissingSecrets, ) from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2 +from litellm.types.llms.openai import AllMessageValues def test_prompt_security_guard_config(monkeypatch: pytest.MonkeyPatch): @@ -174,7 +176,7 @@ async def test_apply_guardrail_modify_request(monkeypatch: pytest.MonkeyPatch): assert result["texts"] == ["User prompt with PII: SSN [REDACTED]"] -def _modify_response(modified_messages: list) -> Response: +def _modify_response(modified_messages: Sequence[Mapping[str, object]]) -> Response: mock_response = Response( json={"result": {"prompt": {"action": "modify", "modified_messages": modified_messages}}}, status_code=200, @@ -184,7 +186,7 @@ def _modify_response(modified_messages: list) -> Response: return mock_response -def _tool_replay_messages() -> list: +def _tool_replay_messages() -> list[AllMessageValues]: return [ {"role": "system", "content": "Never echo an SSN like 123-45-6789."}, { @@ -224,7 +226,9 @@ async def test_modify_returns_structured_messages_with_tool_rows_kept(monkeypatc ] with patch.object(guardrail.async_handler, "post", return_value=_modify_response(modified_messages)): - result = await guardrail.apply_guardrail(inputs=inputs, request_data={"messages": messages}, input_type="request") + result = await guardrail.apply_guardrail( + inputs=inputs, request_data={"messages": messages}, input_type="request" + ) assert result["structured_messages"] == [ {"role": "system", "content": "Never echo an SSN like [REDACTED]."}, @@ -257,7 +261,9 @@ async def test_modify_with_unexpected_message_count_keeps_texts_only(monkeypatch modified_messages = [{"role": "user", "content": "Look up [REDACTED]"}] with patch.object(guardrail.async_handler, "post", return_value=_modify_response(modified_messages)): - result = await guardrail.apply_guardrail(inputs=inputs, request_data={"messages": messages}, input_type="request") + result = await guardrail.apply_guardrail( + inputs=inputs, request_data={"messages": messages}, input_type="request" + ) assert result["structured_messages"] is messages assert result["texts"] == ["Look up [REDACTED]"]