From 33fadd70a3b6d5711baaa86a31081710540f38e3 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Mon, 27 Jul 2026 11:22:06 -0700 Subject: [PATCH] fix(guardrails): compress content-parts messages in headroom guardrail Anthropic-format requests translate to messages whose content is a list of part dicts, which the headroom compression service's transforms silently skip (they only rewrite string content), so compression never applied to Anthropic client traffic while the guardrail still reported itself as applied. Flatten all-text part lists to plain strings for /v1/compress and restore the original shapes from the response: untouched rows keep their exact original parts, a rewritten row collapses to one part carrying the last declared cache_control breakpoint (a breakpoint caches the prefix ending at its part, so the last one and its TTL still describe the merged row). Rows with any non-text part are never flattened, since merging text across a non-text part would move a later breakpoint to the other side of it; they pass through the service untouched, matching its own behavior for non-string content. Flattening and write-back use the shared content_text helpers that compresr's breakpoint fix also uses. Resolves LIT-4795 Co-Authored-By: Claude Fable 5 --- .../guardrail_hooks/headroom/headroom.py | 62 ++++- .../guardrail_hooks/test_headroom.py | 231 ++++++++++++++++++ 2 files changed, 292 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py index 7b166185865..2735acd7787 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py +++ b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py @@ -28,6 +28,11 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] httpxSpecialProvider, ) +from litellm.proxy.guardrails.guardrail_hooks.content_text import ( + content_to_text, + is_all_text_parts, + merge_rewritten_text_parts, +) from litellm.secret_managers.main import get_secret_str from litellm.types.guardrails import GuardrailEventHooks, Mode from litellm.types.integrations.custom_logger import AgenticLoopPlan, AgenticLoopRequestPatch @@ -51,6 +56,60 @@ def _is_object_list(value: object) -> TypeGuard[list[object]]: # guard-ok: isin return isinstance(value, list) +def _flatten_messages_for_compression(messages: list[dict[str, object]]) -> list[dict[str, object]]: + """Collapse all-text list-of-parts content to plain strings for /v1/compress. + + The compression service's transforms only rewrite string content and skip + the OpenAI list-of-parts shape, which is what every Anthropic-format + request translates to. Only rows whose parts are ALL text are flattened: + cache_control breakpoints are positional (each caches the prefix ending + at its part), so merging text across a non-text part would move a later + breakpoint to the other side of it. Rows with non-text parts are sent + unchanged and pass through the service untouched. + """ + flattened: list[dict[str, object]] = [] + for msg in messages: + content = msg.get("content") + if is_all_text_parts(content): + text = content_to_text(content) + if text: + flattened.append({**msg, "content": text}) + continue + flattened.append(msg) + return flattened + + +def _restore_content_shapes( + originals: list[dict[str, object]], returned: list[dict[str, object]] +) -> list[dict[str, object]]: + """Write compressed text back into each original row's content shape. + + Rows are matched positionally; the pairing is only trusted when the + service kept the row count and every role lines up. If it restructured + the conversation (e.g. dropped rows), its output is adopted as-is, which + is the pre-flattening behavior. + """ + if len(returned) != len(originals): + return returned + for orig, ret in zip(originals, returned): + if orig.get("role") != ret.get("role"): + return returned + restored: list[dict[str, object]] = [] + for orig, ret in zip(originals, returned): + orig_content = orig.get("content") + ret_content = ret.get("content") + if isinstance(orig_content, list) and isinstance(ret_content, str): + if ret_content == content_to_text(orig_content): + # Untouched row: keep the exact original parts, including + # per-part fields like cache_control on later text parts. + restored.append({**ret, "content": orig_content}) + else: + restored.append({**ret, "content": merge_rewritten_text_parts(orig_content, ret_content)}) + else: + restored.append(ret) + return restored + + def extract_hashes_from_messages(messages: list[dict[str, object]]) -> list[str]: hashes: list[str] = [] for msg in messages: @@ -491,10 +550,11 @@ class HeadroomGuardrail(CustomGuardrail): model = self.headroom_model or request_data.get("model") start_time = time.time() compressed, compression_succeeded, stats = await self._call_compress( - messages=messages, + messages=_flatten_messages_for_compression(messages), model=model if isinstance(model, str) else None, ) end_time = time.time() + compressed = _restore_content_shapes(originals=messages, returned=compressed) from litellm.proxy.common_utils.callback_utils import ( add_guardrail_to_applied_guardrails_header, diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py index 4dc527ca45d..248893ed153 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py @@ -1551,3 +1551,234 @@ async def test_apply_guardrail_litellm_timeout_fail_open_forwards_uncompressed() ) assert result["structured_messages"] == ORIGINAL_MESSAGES + + + + +# --------------------------------------------------------------------------- +# Content-parts flattening (LIT-4795) +# +# Anthropic-format requests translate to messages whose content is a list of +# part dicts. The compression service only rewrites string content, so the +# guardrail flattens ALL-TEXT part lists on the wire and restores the +# original shapes afterwards. Rows with non-text parts are never flattened: +# cache_control breakpoints are positional, and merging text across a +# non-text part would move a later breakpoint to the other side of it. +# --------------------------------------------------------------------------- + +PARTS_MESSAGES = [ + { + "role": "system", + "content": [ + {"type": "text", "text": "You are Claude Code.", "cache_control": {"type": "ephemeral"}}, + { + "type": "text", + "text": "Second system block. " + "B" * 5000, + "cache_control": {"type": "ephemeral", "ttl": "1h"}, + }, + ], + }, + { + "role": "user", + "content": [ + {"type": "text", "text": "Mixed row text."}, + {"type": "image_url", "image_url": {"url": "https://example.com/x.png"}}, + ], + }, + {"role": "tool", "content": "tool output " + "C" * 500}, +] + +FLATTENED_SYSTEM_TEXT = "You are Claude Code.\n\nSecond system block. " + "B" * 5000 + + +def _parts_copy() -> list: + return json.loads(json.dumps(PARTS_MESSAGES)) + + +def _echo_wire_view() -> list: + """What the service receives (and echoes back when it changes nothing).""" + return [ + {"role": "system", "content": FLATTENED_SYSTEM_TEXT}, + json.loads(json.dumps(PARTS_MESSAGES[1])), + {"role": "tool", "content": "tool output " + "C" * 500}, + ] + + +@pytest.mark.asyncio +async def test_apply_guardrail_flattens_all_text_rows_only( + guardrail: HeadroomGuardrail, +): + inputs = GenericGuardrailAPIInputs( + texts=["B" * 5000], + structured_messages=_parts_copy(), + ) + mock_response = _make_compress_response(_echo_wire_view()) + + with patch.object( + guardrail.async_handler, + "post", + new_callable=AsyncMock, + return_value=mock_response, + ) as mock_post: + await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "claude-fable-5"}, + input_type="request", + ) + + wire_messages = mock_post.call_args.kwargs["json"]["messages"] + assert wire_messages[0]["content"] == FLATTENED_SYSTEM_TEXT + # Mixed text+image row is never flattened: merging its text would move a + # later cache_control breakpoint across the image part. + assert isinstance(wire_messages[1]["content"], list) + assert wire_messages[2]["content"] == "tool output " + "C" * 500 + + +@pytest.mark.asyncio +async def test_apply_guardrail_restores_rewritten_all_text_row( + guardrail: HeadroomGuardrail, +): + inputs = GenericGuardrailAPIInputs( + texts=["B" * 5000], + structured_messages=_parts_copy(), + ) + compressed = _echo_wire_view() + compressed[0]["content"] = "compressed system. Retrieve more: hash=b573993006976af767214fac" + mock_response = _make_compress_response(compressed) + + with patch.object( + guardrail.async_handler, + "post", + new_callable=AsyncMock, + return_value=mock_response, + ): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "claude-fable-5"}, + input_type="request", + ) + + messages = result["structured_messages"] + system_content = messages[0]["content"] + # Rewritten all-text row collapses to one part carrying the LAST declared + # breakpoint: an Anthropic breakpoint caches the prefix ending at its + # part, so after the merge the last one (and its TTL) still describes the + # row. + assert isinstance(system_content, list) + assert len(system_content) == 1 + assert system_content[0]["text"] == "compressed system. Retrieve more: hash=b573993006976af767214fac" + assert system_content[0]["cache_control"] == {"type": "ephemeral", "ttl": "1h"} + # Mixed row passes through byte-identical. + assert messages[1]["content"] == PARTS_MESSAGES[1]["content"] + # Hashes inside restored parts still drive retrieve-tool injection. + assert has_headroom_retrieve_tool(result.get("tools") or []) + + +@pytest.mark.asyncio +async def test_apply_guardrail_keeps_originals_when_service_echoes_unchanged( + guardrail: HeadroomGuardrail, +): + inputs = GenericGuardrailAPIInputs( + texts=["B" * 5000], + structured_messages=_parts_copy(), + ) + mock_response = _make_compress_response(_echo_wire_view()) + + with patch.object( + guardrail.async_handler, + "post", + new_callable=AsyncMock, + return_value=mock_response, + ): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "claude-fable-5"}, + input_type="request", + ) + + messages = result["structured_messages"] + assert [m["content"] for m in messages] == [m["content"] for m in PARTS_MESSAGES] + + +@pytest.mark.asyncio +async def test_apply_guardrail_adopts_service_output_when_rows_dropped( + guardrail: HeadroomGuardrail, +): + inputs = GenericGuardrailAPIInputs( + texts=["B" * 5000], + structured_messages=_parts_copy(), + ) + dropped = [ + {"role": "system", "content": FLATTENED_SYSTEM_TEXT}, + {"role": "user", "content": "B" * 50}, + ] + mock_response = _make_compress_response(dropped) + + with patch.object( + guardrail.async_handler, + "post", + new_callable=AsyncMock, + return_value=mock_response, + ): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "claude-fable-5"}, + input_type="request", + ) + + assert result["structured_messages"] == dropped + + +@pytest.mark.asyncio +async def test_apply_guardrail_sends_textless_parts_rows_unflattened( + guardrail: HeadroomGuardrail, +): + image_only = [ + {"role": "user", "content": [{"type": "image_url", "image_url": {"url": "https://example.com/x.png"}}]}, + {"role": "user", "content": "D" * 5000}, + ] + inputs = GenericGuardrailAPIInputs( + texts=["D" * 5000], + structured_messages=json.loads(json.dumps(image_only)), + ) + mock_response = _make_compress_response(json.loads(json.dumps(image_only))) + + with patch.object( + guardrail.async_handler, + "post", + new_callable=AsyncMock, + return_value=mock_response, + ) as mock_post: + await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "claude-fable-5"}, + input_type="request", + ) + + wire_messages = mock_post.call_args.kwargs["json"]["messages"] + assert isinstance(wire_messages[0]["content"], list) + assert wire_messages[1]["content"] == "D" * 5000 + + +@pytest.mark.asyncio +async def test_fail_open_returns_original_parts_shapes(): + guardrail = _make_guardrail(unreachable_fallback="fail_open") + inputs = GenericGuardrailAPIInputs( + texts=["B" * 5000], + structured_messages=_parts_copy(), + ) + + with patch.object( + guardrail.async_handler, + "post", + new_callable=AsyncMock, + side_effect=httpx.ConnectError("boom"), + ): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={}, + input_type="request", + ) + + messages = result["structured_messages"] + assert [m["content"] for m in messages] == [m["content"] for m in PARTS_MESSAGES]