From 5b61f8bdef84ff5fe844729ec84f7a6d9d4cd6e9 Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Mon, 5 Oct 2026 12:02:09 +0300 Subject: [PATCH] fix(guardrails): treat a row echoed without its null fields as an echo A guardrail that re-serializes the rows it was sent often drops null fields, such as the thinking_blocks: null an Anthropic assistant row is posted with or the content: null of a chat tool-call row. The echo check compared rows exactly, so such an echo looked like a rewrite of every row: the masked texts were ignored and the unmasked rows written back, and the raw value reached the model. Rows now count as an echo when they match apart from null fields, in both the every-row check and the per-row restore. A field the guardrail sets to null where the caller had a value still counts as a change --- litellm/proxy/guardrails/_content_utils.py | 19 ++++++ .../generic_guardrail_api.py | 7 +- .../test_generic_guardrail_api.py | 67 +++++++++++++++++++ .../proxy/guardrails/test_content_utils.py | 33 +++++++++ 4 files changed, 123 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/guardrails/_content_utils.py b/litellm/proxy/guardrails/_content_utils.py index 2a2ef5217b8..d3c44335baf 100644 --- a/litellm/proxy/guardrails/_content_utils.py +++ b/litellm/proxy/guardrails/_content_utils.py @@ -8,9 +8,13 @@ skip the other shapes — these helpers normalise that so every hook sees every text fragment. """ +import json from collections.abc import Callable, Iterator, Mapping, Sequence +from types import MappingProxyType from typing import Any, Final +from pydantic_core import to_jsonable_python + # Call types whose body carries free-form chat / prompt text that # text-content guardrails (banned keywords, content moderation, secret # detection, …) should inspect. The proxy ingress passes ``route_type`` @@ -307,3 +311,18 @@ def build_inspection_messages(data: dict[str, Any]) -> list[dict[str, str]]: role = message.get("role", "user") or "user" flattened.append({"role": role, "content": text}) return flattened + + +def _null_free_object(pairs: Sequence[tuple[str, object]]) -> Mapping[str, object]: + return MappingProxyType({key: item for key, item in pairs if item is not None}) + + +def _null_free(value: object) -> object: + normalized: Final[object] = json.loads( # pyright: ignore[reportAny] # stdlib parse of our own json.dumps output + json.dumps(value, default=to_jsonable_python), object_pairs_hook=_null_free_object + ) + return normalized + + +def same_json_ignoring_nulls(left: object, right: object) -> bool: + return _null_free(left) == _null_free(right) 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 8e064a38044..de7e55bc4c8 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 @@ -28,6 +28,7 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) +from litellm.proxy.guardrails._content_utils import same_json_ignoring_nulls from litellm.types.guardrails import GuardrailEventHooks from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ( @@ -192,16 +193,16 @@ def _structured_rows_to_write_back( returned_rows: Sequence[AllMessageValues], ) -> tuple[AllMessageValues, ...] | None: """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. + row the server echoes back, null fields aside, is restored to the original row object. A server that echoes every row back unchanged has not rewritten anything per row, so its answer is read from texts, as it was before rows could be returned at all.""" if original_rows is None or shown_rows is None or len(returned_rows) != len(original_rows): return tuple(returned_rows) - if all(returned == shown for shown, returned in zip(shown_rows, returned_rows)): + if all(same_json_ignoring_nulls(returned, shown) for shown, returned in zip(shown_rows, returned_rows)): return None return tuple( - original if returned == shown else returned + original if same_json_ignoring_nulls(returned, shown) else returned for original, shown, returned in zip(original_rows, shown_rows, returned_rows) ) diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py index 792c4759221..173aa01e7e6 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py @@ -766,6 +766,17 @@ def _echo_every_row_and_mask_texts(request_json: Mapping[str, JsonValue]) -> Map } +def _without_null_values(row: Mapping[str, JsonValue]) -> Mapping[str, JsonValue]: + return {key: value for key, value in row.items() if value is not None} + + +def _echo_every_row_without_its_nulls_and_mask_texts(request_json: Mapping[str, JsonValue]) -> Mapping[str, JsonValue]: + rows_without_nulls: Final[JsonValue] = json.loads( # pyright: ignore[reportAny] # stdlib parse of a JSON copy + json.dumps(request_json["structured_messages"]), object_hook=_without_null_values + ) + return {**_echo_every_row_and_mask_texts(request_json), "structured_messages": rows_without_nulls} + + def _echo_first_row_and_mask_the_rest(request_json: Mapping[str, JsonValue]) -> Mapping[str, JsonValue]: first_row, *other_rows = request_json["structured_messages"] return { @@ -774,6 +785,19 @@ def _echo_first_row_and_mask_the_rest(request_json: Mapping[str, JsonValue]) -> } +def _echo_first_row_without_its_nulls_and_mask_the_rest( + request_json: Mapping[str, JsonValue], +) -> Mapping[str, JsonValue]: + first_row, *other_rows = request_json["structured_messages"] + return { + "action": "GUARDRAIL_INTERVENED", + "structured_messages": [ + _without_null_values(first_row), + *({**row, "content": _masked(row["content"])} for row in other_rows), + ], + } + + def _masked_part(part: Mapping[str, JsonValue]) -> Mapping[str, JsonValue]: return {**part, "text": _masked(part["text"])} if part["type"] == "text" else part @@ -896,6 +920,31 @@ class TestEchoedRowsReachingTheLLM: {"role": "user", "content": [{"type": "text", "text": "my ssn is [SSN]"}, {"type": "text", "text": "ok"}]} ], "an unchanged echo of every content block row must leave the rewrite to texts" + @pytest.mark.asyncio + async def test_an_echo_that_drops_null_fields_still_applies_the_masked_texts(self) -> None: + guardrail: Final = _guardrail_answering(_echo_every_row_without_its_nulls_and_mask_texts) + tool_use_turn: Final[AllAnthropicMessageValues] = { + "role": "assistant", + "content": [ + {"type": "text", "text": "calling"}, + {"type": "tool_use", "id": "t1", "name": "f", "input": {}}, + ], + } + tool_result_turn: Final[AllAnthropicMessageValues] = { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": "t1", "content": "r"}], + } + + llm_bound: Final = await _llm_bound_anthropic_messages( + guardrail, [{"role": "user", "content": f"my ssn is {_SSN}"}, tool_use_turn, tool_result_turn] + ) + + assert llm_bound == [ + {"role": "user", "content": "my ssn is [SSN]"}, + tool_use_turn, + tool_result_turn, + ], "an echo without the null fields the rows were posted with must leave the rewrite to texts" + @pytest.mark.asyncio async def test_every_responses_input_text_row_echoed_applies_the_masked_texts(self) -> None: guardrail: Final = _guardrail_answering(_echo_every_row_and_mask_texts) @@ -925,6 +974,24 @@ class TestEchoedRowsReachingTheLLM: {"role": "user", "content": "my ssn is [SSN]"}, ], "the echoed row must keep the caller's keys the request model drops" + @pytest.mark.asyncio + async def test_a_row_echoed_without_its_null_fields_is_restored_to_the_callers_row(self) -> None: + guardrail: Final = _guardrail_answering(_echo_first_row_without_its_nulls_and_mask_the_rest) + tool_call_turn: Final[AllMessageValues] = { + "role": "assistant", + "content": None, + "tool_calls": [{"id": "c1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}], + } + + llm_bound: Final = await _llm_bound_messages( + guardrail, [tool_call_turn, {"role": "tool", "tool_call_id": "c1", "content": f"ssn {_SSN}"}] + ) + + assert llm_bound == [ + tool_call_turn, + {"role": "tool", "tool_call_id": "c1", "content": "ssn [SSN]"}, + ], "a row echoed without the null fields it was posted with must stay the caller's row" + @pytest.mark.asyncio @pytest.mark.parametrize( "unvalidated_part", [_guarded_text_part, _pdf_document_part], ids=["guarded_text", "pdf_document"] diff --git a/tests/unit/proxy/guardrails/test_content_utils.py b/tests/unit/proxy/guardrails/test_content_utils.py index 920ffc77095..7084348d361 100644 --- a/tests/unit/proxy/guardrails/test_content_utils.py +++ b/tests/unit/proxy/guardrails/test_content_utils.py @@ -1,5 +1,7 @@ """Tests for the shared guardrail content extraction helpers.""" +import pytest + from litellm.proxy.guardrails._content_utils import ( apply_redacted_messages_back, build_inspection_messages, @@ -7,6 +9,7 @@ from litellm.proxy.guardrails._content_utils import ( is_non_conversational_call_type, is_string_batch_input, iter_message_text, + same_json_ignoring_nulls, walk_user_text, ) @@ -741,3 +744,33 @@ def test_is_non_conversational_call_type_defaults_to_inspecting_unknown_call_typ """A call type this module has never heard of must still be inspected — failing closed is the point of the deny-list.""" assert is_non_conversational_call_type("some_future_call_type") is False + + +@pytest.mark.parametrize( + ("left", "right", "same"), + [ + ({"role": "assistant", "thinking_blocks": None}, {"role": "assistant"}, True), + ( + {"content": [{"type": "text", "text": "x", "cache_control": None}]}, + {"content": [{"type": "text", "text": "x"}]}, + True, + ), + ({"content": ("a", "b")}, {"content": ["a", "b"]}, True), + ({"content": "x"}, {"content": "y"}, False), + ({"role": "user", "name": "a"}, {"role": "user"}, False), + ([None], [], False), + ], + ids=[ + "dropped_null_key", + "dropped_nested_null_key", + "tuple_as_list", + "changed_value", + "dropped_set_key", + "null_list_item", + ], +) +def test_same_json_ignoring_nulls_treats_only_a_dropped_null_field_as_no_change( + left: object, right: object, same: bool +) -> None: + assert same_json_ignoring_nulls(left, right) is same + assert same_json_ignoring_nulls(right, left) is same