Merge pull request #34660 from BerriAI/litellm_lit4804_compresr_cache_control

fix(guardrails): preserve cache_control breakpoints in compresr write-back
This commit is contained in:
tin-berri 2026-07-27 11:17:51 -07:00 • committed by GitHub
commit 10cd4288b6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 140 additions and 43 deletions

View file

@ -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)

View file

@ -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]

View file

@ -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 ─────────────────────────────────────────────