diff --git a/litellm/proxy/guardrails/guardrail_hooks/compresr/compresr.py b/litellm/proxy/guardrails/guardrail_hooks/compresr/compresr.py index a95bdb670c3..e512be23fc9 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/compresr/compresr.py +++ b/litellm/proxy/guardrails/guardrail_hooks/compresr/compresr.py @@ -47,6 +47,11 @@ from litellm.llms.custom_httpx.http_handler import ( httpxSpecialProvider, ) from litellm.proxy._types import UserAPIKeyAuth +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 ( @@ -144,48 +149,20 @@ def _is_object_list(value: object) -> TypeGuard[list[object]]: # guard-ok: isin return isinstance(value, list) -def _content_to_text(content: object) -> str: - """Collapse a message ``content`` (str or list-of-parts) to plain text. - - For the multimodal list shape, joins ``{type: "text", text: ...}`` parts - with blank-line separators; non-text parts are ignored. - """ - if isinstance(content, str): - return content - if isinstance(content, list): - parts: list[str] = [] - for part in content: - if isinstance(part, dict) and part.get("type") == "text": - text = part.get("text") - if isinstance(text, str): - parts.append(text) - return "\n\n".join(parts) - return "" - - def _replace_text_in_content(content: object, new_text: str) -> object: """Write ``new_text`` back into a ``content`` value, preserving shape. - ``str`` content is replaced directly. For list-of-parts content the first - text part carries ``new_text``, later text parts are dropped, and - non-text parts (images, audio, files) pass through untouched. + ``str`` content is replaced directly. An all-text part list collapses to a + single part carrying the last declared cache_control breakpoint. Anything + else is returned unchanged: breakpoints are positional, so one compressed + string cannot be written back across a non-text part without moving text + to the other side of it. """ if isinstance(content, str): return new_text - if isinstance(content, list): - out: list[object] = [] - replaced = False - for part in content: - if isinstance(part, dict) and part.get("type") == "text": - if not replaced: - out.append({**part, "text": new_text}) - replaced = True - continue - out.append(part) - if not replaced: - out.insert(0, {"type": "text", "text": new_text}) - return out - return new_text + if _is_object_list(content) and is_all_text_parts(content): + return merge_rewritten_text_parts(content, new_text) + return content def _render_tool_intent(fn: dict[str, object]) -> str: @@ -422,7 +399,7 @@ def _assistant_text_from_response(response: object) -> str | None: if isinstance(choices, list) and choices: message = get_attribute_or_key(choices[0], "message", None) if message is not None: - text = _content_to_text(get_attribute_or_key(message, "content", None)) + text = content_to_text(get_attribute_or_key(message, "content", None)) if text: return text content = get_attribute_or_key(response, "content", None) @@ -905,7 +882,10 @@ class CompresrGuardrail(CustomGuardrail): continue else: continue - if len(_content_to_text(msg.get("content"))) < self.min_chars_to_compress: + content = msg.get("content") + if _is_object_list(content) and not is_all_text_parts(content): + continue + if len(content_to_text(content)) < self.min_chars_to_compress: continue targets.append(idx) return targets @@ -916,7 +896,7 @@ class CompresrGuardrail(CustomGuardrail): ) -> tuple[str, int | None]: for idx in range(len(messages) - 1, -1, -1): if messages[idx].get("role") == "user": - return _content_to_text(messages[idx].get("content")), idx + return content_to_text(messages[idx].get("content")), idx return "", None def _apply_compression_results( @@ -1034,7 +1014,7 @@ class CompresrGuardrail(CustomGuardrail): verbose_proxy_logger.debug("Compresr: no messages eligible for compression") return inputs - contexts = [_content_to_text(messages[idx].get("content")) for idx in targets] + contexts = [content_to_text(messages[idx].get("content")) for idx in targets] start_time = time.monotonic() results = await self._call_compress(contexts=contexts, queries=queries) diff --git a/litellm/proxy/guardrails/guardrail_hooks/content_text.py b/litellm/proxy/guardrails/guardrail_hooks/content_text.py new file mode 100644 index 00000000000..f4211e67512 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/content_text.py @@ -0,0 +1,55 @@ +"""Shared content-part helpers for compression guardrails (headroom, compresr). + +Compression services only transform plain-string message content: every +transform in the service pipeline gates on ``isinstance(content, str)`` and +silently skips the OpenAI list-of-parts shape. Guardrails that send messages +to such a service collapse text-bearing part lists to strings here, and write +the rewritten text back through ``merge_rewritten_text_parts``. + +Anthropic ``cache_control`` breakpoints are positional: each one caches the +prefix ending at the part that carries it. A single compressed string can +therefore only be written back over a run of text parts, never across a +non-text part, which is what ``is_all_text_parts`` gates. +""" + +from collections.abc import Sequence + + +def content_to_text(content: object) -> str: + """Collapse a message ``content`` (str or list-of-parts) to plain text. + + For the multimodal list shape, joins ``{type: "text", text: ...}`` parts + with blank-line separators; non-text parts are ignored. + """ + if isinstance(content, str): + return content + if isinstance(content, list): + parts: list[str] = [] + for part in content: + if isinstance(part, dict) and part.get("type") == "text": + text = part.get("text") + if isinstance(text, str): + parts.append(text) + return "\n\n".join(parts) + return "" + + +def is_all_text_parts(content: object) -> bool: + """True when ``content`` is a non-empty part list holding only text parts.""" + if not isinstance(content, list) or not content: + return False + return all(isinstance(part, dict) and part.get("type") == "text" for part in content) + + +def merge_rewritten_text_parts(parts: Sequence[object], new_text: str) -> list[object]: + """Collapse a rewritten all-text part list into one part carrying ``new_text``. + + Only all-text rows are ever flattened, so the merged part IS the whole row: + it keeps the first part's fields and the LAST declared cache_control + breakpoint. A breakpoint caches the prefix ending at its part, so after the + merge the last one (and its TTL) is the one that still describes the row. + """ + dict_parts = tuple(part for part in parts if isinstance(part, dict)) + breakpoints = tuple(part["cache_control"] for part in dict_parts if part.get("cache_control") is not None) + base = {**dict_parts[0], "text": new_text} if dict_parts else {"type": "text", "text": new_text} + return [{**base, "cache_control": breakpoints[-1]} if breakpoints else base] diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_compresr.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_compresr.py index f6f29eee5bc..feb7090c2e3 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_compresr.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_compresr.py @@ -6,7 +6,8 @@ Tests cover: resolved via tool_call_id, falling back to the last user message) - target selection: tool outputs by default, system/history opt-in, min-chars threshold, targets without a derivable query are left uncompressed -- multimodal content: text parts replaced, non-text parts preserved +- multimodal content: all-text rows merge into one part carrying the last + cache_control breakpoint, rows holding a non-text part are left uncompressed - recovery: hash marker appended, compresr_retrieve tool injected, originals stored per litellm_call_id, agentic loop returns the original content and rejects hashes not issued for the current request @@ -560,7 +561,7 @@ async def test_short_messages_skipped(guardrail: CompresrGuardrail): @pytest.mark.asyncio -async def test_multimodal_text_replaced_non_text_preserved( +async def test_multimodal_row_with_non_text_part_is_not_compressed( guardrail: CompresrGuardrail, ): image_part = {"type": "image_url", "image_url": {"url": "https://example.com/x.png"}} @@ -572,6 +573,35 @@ async def test_multimodal_text_replaced_non_text_preserved( "content": [{"type": "text", "text": TOOL_OUTPUT}, image_part], }, ] + expected = json.loads(json.dumps(messages[1]["content"])) + mock_post = AsyncMock(return_value=_make_single_compress_response()) + + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=_apply_inputs(messages), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + + mock_post.assert_not_called() + assert result["structured_messages"][1]["content"] == expected + + +@pytest.mark.asyncio +async def test_all_text_row_merges_and_keeps_last_cache_control( + guardrail: CompresrGuardrail, +): + messages = [ + {"role": "user", "content": USER_QUESTION}, + { + "role": "tool", + "tool_call_id": "c1", + "content": [ + {"type": "text", "text": TOOL_OUTPUT, "cache_control": {"type": "ephemeral"}}, + {"type": "text", "text": TOOL_OUTPUT, "cache_control": {"type": "ephemeral", "ttl": "1h"}}, + ], + }, + ] mock_post = AsyncMock(return_value=_make_single_compress_response()) with patch.object(guardrail.async_handler, "post", mock_post): @@ -583,9 +613,41 @@ async def test_multimodal_text_replaced_non_text_preserved( content = result["structured_messages"][1]["content"] assert isinstance(content, list) + assert len(content) == 1 assert content[0]["type"] == "text" assert content[0]["text"].startswith("compressed summary") - assert content[1] == image_part + assert content[0]["cache_control"] == {"type": "ephemeral", "ttl": "1h"} + + +@pytest.mark.asyncio +async def test_text_around_non_text_part_is_never_relocated( + guardrail: CompresrGuardrail, +): + image_part = {"type": "image_url", "image_url": {"url": "https://example.com/x.png"}} + messages = [ + {"role": "user", "content": USER_QUESTION}, + { + "role": "tool", + "tool_call_id": "c1", + "content": [ + {"type": "text", "text": TOOL_OUTPUT}, + image_part, + {"type": "text", "text": TOOL_OUTPUT, "cache_control": {"type": "ephemeral"}}, + ], + }, + ] + expected = json.loads(json.dumps(messages[1]["content"])) + mock_post = AsyncMock(return_value=_make_single_compress_response()) + + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=_apply_inputs(messages), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + + mock_post.assert_not_called() + assert result["structured_messages"][1]["content"] == expected # ── passthrough / bypass ─────────────────────────────────────────────