mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
commit
10cd4288b6
3 changed files with 140 additions and 43 deletions
|
|
@ -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)
|
||||
|
|
|
|||
55
litellm/proxy/guardrails/guardrail_hooks/content_text.py
Normal file
55
litellm/proxy/guardrails/guardrail_hooks/content_text.py
Normal 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]
|
||||
|
|
@ -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 ─────────────────────────────────────────────
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue