Merge pull request #40315 from rad-p44/fix/headroom-protect-cache-control-rows

fix(headroom): protect cache_control-marked rows anywhere in history
This commit is contained in:
Mateo Wang 2026-09-14 21:03:16 -07:00 committed by GitHub
commit 94f08636c7
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 142 additions and 16 deletions

View file

@ -205,21 +205,35 @@ def _extract_anthropic_tool_exchange_spans(
return spans, None
def _message_has_cache_control(message: Mapping[str, object]) -> bool:
if message.get("cache_control") is not None:
return True
content: Final = message.get("content")
if isinstance(content, list):
return any(isinstance(part, Mapping) and part.get("cache_control") is not None for part in content)
return False
def get_protected_indices(messages: Sequence[Mapping[str, object]]) -> tuple[int, ...]:
"""
Return indices of messages that must never be compressed:
- All system messages
- The last user message
- The last assistant message
- Any message carrying an Anthropic cache_control breakpoint
The last user message is what the model is being asked to act on right now,
so compressing it replaces the live instruction with a marker. Compression
guardrails share this policy; see the Headroom guardrail.
guardrails share this policy; see the Headroom guardrail. A cache_control
breakpoint pins the provider's prompt-cache prefix to that row's exact
bytes, so rewriting a marked row anywhere in history turns the next
request's cache read into a cache write.
"""
system_indices: Final = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "system")
last_user: Final = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "user")[-1:]
last_assistant = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "assistant")[-1:]
return system_indices + last_user + last_assistant
assistant_indices: Final = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "assistant")
cache_control_indices: Final = tuple(index for index, msg in enumerate(messages) if _message_has_cache_control(msg))
return tuple(dict.fromkeys(system_indices + last_user + assistant_indices[-1:] + cache_control_indices))
def _combine_scores(
@ -421,7 +435,7 @@ def compress(
combined_scores = bm25_scores
# Protected messages are never compressed
protected_indices: Final = get_protected_indices(normalized_messages)
protected_indices: Final = get_protected_indices(original_messages)
kept_indices: set[int] = set(protected_indices)
tool_exchange_spans: list[set[int]] = []

View file

@ -6,7 +6,8 @@ never rewrite. It is consumed by compress() and by the Headroom guardrail, so
the two agree on what "never compress this" means.
"""
from litellm.compression.compress import get_protected_indices
from litellm.compression.compress import compress, get_protected_indices
from litellm.types.utils import CallTypes
def test_protects_system_last_user_and_last_assistant():
@ -53,3 +54,94 @@ def test_every_system_row_is_protected():
def test_no_user_or_assistant_rows():
assert sorted(get_protected_indices([{"role": "system", "content": "sys"}])) == [0]
assert get_protected_indices([]) == ()
def test_mid_history_cache_control_part_is_protected():
# A large cached tool result from a few turns back, not the last user or
# last assistant row -- exactly the row a provider prompt-cache pins to
# exact bytes. Rewriting it (even leaving the marker on) changes those
# bytes and turns the next request's cache read into a cache write.
messages = [
{"role": "user", "content": "old question"},
{"role": "assistant", "content": "old answer"},
{
"role": "user",
"content": [
{"type": "text", "text": "a large cached tool result", "cache_control": {"type": "ephemeral"}},
],
},
{"role": "assistant", "content": "ack"},
{"role": "user", "content": "live instruction"},
]
# index 3 = last assistant, index 4 = last user (both protected by role
# regardless), index 2 = the cache_control-marked row itself.
assert sorted(get_protected_indices(messages)) == [2, 3, 4]
def test_cache_control_directly_on_message_is_protected():
messages = [
{"role": "user", "content": "old question", "cache_control": {"type": "ephemeral"}},
{"role": "assistant", "content": "old answer"},
{"role": "user", "content": "live instruction"},
]
assert sorted(get_protected_indices(messages)) == [0, 1, 2]
def test_cache_control_protection_does_not_duplicate_already_protected_rows():
# The last user row is already protected by role; marking it too must not
# produce a duplicate index.
messages = [
{"role": "system", "content": "sys"},
{"role": "user", "content": "live", "cache_control": {"type": "ephemeral"}},
]
protected = get_protected_indices(messages)
assert sorted(protected) == [0, 1]
assert len(protected) == len(set(protected))
def test_content_that_is_not_a_list_of_mappings_is_not_treated_as_cache_control():
# Defensive: a plain string content, or a list of non-dict items, must not
# raise or be misread as carrying a breakpoint.
messages = [
{"role": "assistant", "content": "plain string content"},
{"role": "user", "content": ["not", "a", "dict", "list"]},
{"role": "user", "content": "live instruction"},
]
assert sorted(get_protected_indices(messages)) == [0, 2]
def test_compress_keeps_part_level_cache_control_row_verbatim():
# compress() scores text-only copies of the rows, where a part-level marker
# is gone; protection has to read the original rows or the pinned row is stubbed.
stale_log = {"role": "user", "content": [{"type": "text", "text": "stale log line " * 2000}]}
pinned = {
"role": "user",
"content": [
{"type": "text", "text": "cached tool result " * 2000, "cache_control": {"type": "ephemeral"}},
],
}
messages = [
stale_log,
{"role": "assistant", "content": "old answer"},
pinned,
{"role": "assistant", "content": "ack"},
{"role": "user", "content": "live instruction"},
]
result = compress(
messages,
model="gpt-4o",
call_type=CallTypes.anthropic_messages,
compression_trigger=1000,
compression_target=500,
)
assert len(result["messages"]) == len(messages)
assert result["messages"][2] == pinned
assert result["messages"][0] != stale_log
assert len(result["cache"]) >= 1

View file

@ -1797,12 +1797,8 @@ PARTS_MESSAGES = [
{
"role": "user",
"content": [
{"type": "text", "text": "Earlier turn.", "cache_control": {"type": "ephemeral"}},
{
"type": "text",
"text": "Second block. " + "B" * 5000,
"cache_control": {"type": "ephemeral", "ttl": "1h"},
},
{"type": "text", "text": "Earlier turn."},
{"type": "text", "text": "Second block. " + "B" * 5000},
],
},
{
@ -1891,14 +1887,9 @@ async def test_apply_guardrail_restores_rewritten_all_text_row(
messages = result["structured_messages"]
history_content = messages[1]["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(history_content, list)
assert len(history_content) == 1
assert history_content[0]["text"] == "compressed history. Retrieve more: hash=b573993006976af767214fac"
assert history_content[0]["cache_control"] == {"type": "ephemeral", "ttl": "1h"}
# Mixed row passes through byte-identical.
assert messages[2]["content"] == PARTS_MESSAGES[2]["content"]
# The service-declared hash still drives retrieve-tool injection on a restored row.
@ -2523,6 +2514,35 @@ async def test_history_is_still_compressed(guardrail: HeadroomGuardrail):
assert messages[3] == compressed_history[1]
CACHE_MARKED_HISTORY_MESSAGES = [
{"role": "system", "content": "You are Claude Code. " + "S" * 5000},
{"role": "user", "content": "old question " + "Q" * 5000},
{
"role": "assistant",
"content": "Reading the file now.",
"tool_calls": [{"id": "old_1", "type": "function", "function": {"name": "Read", "arguments": "{}"}}],
},
{
"role": "tool",
"tool_call_id": "old_1",
"content": [{"type": "text", "text": "large cached file body " + "F" * 5000}],
"cache_control": {"type": "ephemeral"},
},
{"role": "assistant", "content": "Summarized the file for you."},
{"role": "user", "content": "live instruction"},
]
@pytest.mark.asyncio
async def test_mid_history_cache_control_row_is_never_sent_for_compression(guardrail: HeadroomGuardrail):
wire, result = await _wire_and_result(guardrail, CACHE_MARKED_HISTORY_MESSAGES)
cached_row = CACHE_MARKED_HISTORY_MESSAGES[3]
assert cached_row not in wire
assert not any(row.get("tool_call_id") == "old_1" for row in wire)
assert result["structured_messages"][3] == cached_row
# ---------------------------------------------------------------------------
# #38558: a client that runs its own tool loop (e.g. Claude Code via the MCP
# gateway) executes headroom_retrieve and echoes the recovered original content