fix(headroom): stop re-compressing retrieved CCR content in client tool loops

When the headroom_retrieve tool is exposed to a client that runs its own
tool-execution loop (the LiteLLM MCP gateway path), the client executes the
retrieve call and sends the recovered original content back as a tool result
on the next turn. The guardrail then compressed that row again, and because
CCR is content-addressed it collapsed back to the exact same hash it was just
retrieved from. The model never saw the expansion and the agent looped.

Hold tool-result rows that carry headroom_retrieve output back from the
compression service, the same way the live turn and trailing tool exchange are
already protected, so the expansion survives. Retrieve calls are matched by the
direct headroom_retrieve name and the mcp__<server>__headroom_retrieve gateway
name. Because a long gateway name is truncated past 64 chars in the
OpenAI-translated view the guardrail scans, the pairing also falls back to the
tool-call id read from the request's own untranslated messages, which is never
truncated.

Fixes #38558
This commit is contained in:
QuantumBreakz 2026-08-28 04:49:36 +05:00
parent 44d84360fb
commit 2a998c2f38
2 changed files with 274 additions and 4 deletions

View file

@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, TypeGuard
import httpx
from fastapi import HTTPException
from httpx import Response as HttpxResponse
from pydantic import TypeAdapter
import litellm
from litellm._logging import verbose_proxy_logger
@ -48,6 +49,10 @@ BYPASS_HEADER: Final = "x-headroom-bypass"
HEADROOM_RETRIEVE_TOOL_NAME: Final = "headroom_retrieve"
_HASH_PATTERN: Final = re.compile(r"hash=([a-f0-9]{24})")
_HASH_CACHE_TTL_SECONDS: Final = 15 * 60
# Narrows the base class's bare-dict ``request_data`` at the boundary so its
# untranslated messages can be read with concrete types (values pass through by
# reference, so this is a shallow top-level reconstruction).
_REQUEST_DATA_ADAPTER: Final = TypeAdapter(dict[str, object])
def _is_str_object_dict(value: object) -> TypeGuard[dict[str, object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip
@ -112,16 +117,119 @@ def _restore_content_shapes(
return restored
def _protected_indices(messages: Sequence[Mapping[str, object]]) -> frozenset[int]:
def _tool_call_name(tool_call: Mapping[str, object]) -> str | None:
function: Final = tool_call.get("function")
if not _is_str_object_dict(function):
return None
name: Final = function.get("name")
return name if isinstance(name, str) else None
def _is_retrieve_tool_name(name: str | None) -> bool:
"""Match the retrieve tool whether called directly or via the MCP gateway.
Server-side the tool is ``headroom_retrieve``; exposed through LiteLLM's MCP
gateway a client calls it as ``mcp__<server>__headroom_retrieve``.
"""
return name is not None and (
name == HEADROOM_RETRIEVE_TOOL_NAME or name.endswith(f"__{HEADROOM_RETRIEVE_TOOL_NAME}")
)
def _retrieve_call_ids_in_message(message: Mapping[str, object]) -> frozenset[str]:
if message.get("role") != "assistant":
return frozenset()
tool_calls: Final = message.get("tool_calls")
if not _is_object_list(tool_calls):
return frozenset()
return frozenset(
str(tool_call["id"])
for tool_call in tool_calls
if _is_str_object_dict(tool_call) and tool_call.get("id") and _is_retrieve_tool_name(_tool_call_name(tool_call))
)
def _anthropic_tool_use_retrieve_id(block: object) -> str | None:
if not _is_str_object_dict(block) or block.get("type") != "tool_use":
return None
name: Final = block.get("name")
call_id: Final = block.get("id")
if isinstance(name, str) and call_id is not None and _is_retrieve_tool_name(name):
return str(call_id)
return None
def _anthropic_retrieve_ids_in_message(message: Mapping[str, object]) -> frozenset[str]:
content: Final = message.get("content")
if not _is_object_list(content):
return frozenset()
return frozenset(call_id for block in content if (call_id := _anthropic_tool_use_retrieve_id(block)) is not None)
def _raw_retrieve_call_ids(messages: object) -> frozenset[str]:
"""Retrieve-tool call ids read from the request's own, untranslated messages.
The guardrail otherwise scans an OpenAI-translated view where a tool name
over 64 chars is truncated to ``{prefix}_{hash}``, which drops the
``__headroom_retrieve`` suffix a long ``mcp__<server>__`` prefix pushes past
the limit. Tool-call ids are never truncated, so pairing the tool result to
an id read from the original request keeps the match intact. Both wire
shapes are handled: OpenAI ``tool_calls`` and Anthropic ``tool_use`` blocks.
"""
if not _is_object_list(messages):
return frozenset()
return frozenset(
call_id
for message in messages
if _is_str_object_dict(message)
for call_id in _retrieve_call_ids_in_message(message) | _anthropic_retrieve_ids_in_message(message)
)
def _retrieval_result_indices(
messages: Sequence[Mapping[str, object]], extra_retrieve_call_ids: frozenset[str] = frozenset()
) -> frozenset[int]:
"""Indices of tool-result rows that carry ``headroom_retrieve`` output.
When the retrieve tool is exposed to a client that runs its own tool loop
(the LiteLLM MCP gateway path), the client executes the call and sends the
recovered original content back as a tool result on the next turn. That
content is exactly what a prior compression stubbed, so compressing it again
re-derives the identical content hash: a no-op that strands the model on the
marker and loops the agent. Hold those rows back so the expansion survives.
``extra_retrieve_call_ids`` carries ids recovered from the untruncated
request so the pairing survives tool-name truncation (see
``_raw_retrieve_call_ids``).
"""
retrieve_call_ids: Final = extra_retrieve_call_ids | frozenset(
call_id for message in messages for call_id in _retrieve_call_ids_in_message(message)
)
if not retrieve_call_ids:
return frozenset()
return frozenset(
index
for index, message in enumerate(messages)
if message.get("role") in ("tool", "function") and str(message.get("tool_call_id")) in retrieve_call_ids
)
def _protected_indices(
messages: Sequence[Mapping[str, object]], extra_retrieve_call_ids: frozenset[str] = frozenset()
) -> frozenset[int]:
"""Indices headroom must not send to the compression service.
``get_protected_indices`` is litellm's own compression policy: the system
rows, the last user row, the last assistant row. It is expanded over whole
rows, the last user row, the last assistant row. Rows carrying just-retrieved
``headroom_retrieve`` output are added so re-compression can't collapse them
back to the marker they were expanded from. The union is expanded over whole
tool exchanges the way ``compress()`` expands it, so a protected assistant
tool call cannot end up answered by a marker standing in for the result the
model just asked for.
"""
protected: Final = frozenset(get_protected_indices(messages))
protected: Final = frozenset(get_protected_indices(messages)) | _retrieval_result_indices(
messages, extra_retrieve_call_ids
)
return protected | frozenset(
index
for group in group_tool_exchanges(messages)
@ -640,7 +748,11 @@ class HeadroomGuardrail(CustomGuardrail):
# /v1/compress grows a field for sending the live turn as the retrieval
# query without compressing it: query-aware compression reads the newest
# user message, so it is withheld here at some cost to history ranking.
protected_indices: Final = _protected_indices(messages)
# request_data is a bare dict on the base signature; narrow it before
# reading the untranslated messages so long tool names can be recovered.
raw_messages: Final = _REQUEST_DATA_ADAPTER.validate_python(request_data).get("messages")
raw_retrieve_call_ids: Final = _raw_retrieve_call_ids(raw_messages)
protected_indices: Final = _protected_indices(messages, raw_retrieve_call_ids)
compressible: Final = [m for i, m in enumerate(messages) if i not in protected_indices]
if not compressible:
return inputs

View file

@ -1998,6 +1998,164 @@ async def test_history_is_still_compressed(guardrail: HeadroomGuardrail):
assert has_headroom_retrieve_tool(result.get("tools") or [])
# ---------------------------------------------------------------------------
# #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
# back as a tool result. Compressing that row re-derives the same content hash
# it was just retrieved from -- the marker returns and the agent loops. The
# retrieved row must be held back from the compression service.
# ---------------------------------------------------------------------------
RETRIEVE_ECHO_MESSAGES = [
{"role": "system", "content": "You are Claude Code. " + "S" * 5000},
{"role": "user", "content": "H" * 5000},
{
"role": "assistant",
"content": "Expanding the marker.",
"tool_calls": [
{
"id": "hr_1",
"type": "function",
"function": {
"name": "mcp__headroom__headroom_retrieve",
"arguments": '{"hash": "b573993006976af767214fac"}',
},
}
],
},
{"role": "tool", "tool_call_id": "hr_1", "content": "RETRIEVED BODY " + "R" * 5000},
{"role": "assistant", "content": "Older answer. " + "O" * 5000},
{"role": "user", "content": "now summarize the description"},
]
@pytest.mark.asyncio
async def test_retrieved_content_is_never_recompressed(guardrail: HeadroomGuardrail):
"""The tool result carrying headroom_retrieve output is held back, so it can
never collapse back to the hash it was just retrieved from."""
wire, result = await _wire_and_result(guardrail, RETRIEVE_ECHO_MESSAGES)
assert not any(row.get("tool_call_id") == "hr_1" for row in wire)
assert not any("RETRIEVED BODY" in json.dumps(row) for row in wire)
# Reaches the model byte-identical, so no marker stands in for the expansion.
assert result["structured_messages"][3] == RETRIEVE_ECHO_MESSAGES[3]
# Negative control: unrelated history is still compressed, not a no-op.
assert any(row.get("content") == "H" * 5000 for row in wire)
@pytest.mark.asyncio
async def test_retrieved_content_guard_matches_direct_tool_name(guardrail: HeadroomGuardrail):
"""Server-side the tool is named headroom_retrieve (no MCP prefix); its
result must be protected the same way."""
messages = [
{"role": "system", "content": "sys " + "S" * 5000},
{"role": "user", "content": "H" * 5000},
{
"role": "assistant",
"content": "",
"tool_calls": [
{
"id": "hr_direct",
"type": "function",
"function": {"name": HEADROOM_RETRIEVE_TOOL_NAME, "arguments": "{}"},
}
],
},
{"role": "tool", "tool_call_id": "hr_direct", "content": "RETRIEVED BODY " + "R" * 5000},
{"role": "assistant", "content": "Older. " + "O" * 5000},
{"role": "user", "content": "summarize"},
]
wire, result = await _wire_and_result(guardrail, messages)
assert not any(row.get("tool_call_id") == "hr_direct" for row in wire)
assert result["structured_messages"][3] == messages[3]
@pytest.mark.asyncio
async def test_retrieved_content_protected_when_mcp_tool_name_is_truncated(guardrail: HeadroomGuardrail):
"""A long mcp__<server>__headroom_retrieve name is truncated past 64 chars in
the OpenAI-translated view the guardrail scans, dropping the suffix. The call
id read from the request's own Anthropic tool_use (never truncated) still
pairs the retrieved row so it is held back."""
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
truncate_tool_name,
)
long_name = "mcp__" + "s" * 45 + "__" + HEADROOM_RETRIEVE_TOOL_NAME
assert len(long_name) > 64
truncated = truncate_tool_name(long_name)
assert not truncated.endswith(HEADROOM_RETRIEVE_TOOL_NAME)
# What the guardrail scans: OpenAI-translated messages with the truncated name.
structured = [
{"role": "system", "content": "sys " + "S" * 5000},
{"role": "user", "content": "H" * 5000},
{
"role": "assistant",
"content": "",
"tool_calls": [{"id": "hr_long", "type": "function", "function": {"name": truncated, "arguments": "{}"}}],
},
{"role": "tool", "tool_call_id": "hr_long", "content": "RETRIEVED BODY " + "R" * 5000},
{"role": "assistant", "content": "Older. " + "O" * 5000},
{"role": "user", "content": "summarize"},
]
# The request's own messages, untranslated: Anthropic tool_use carries the full name.
raw_messages = [
{"role": "assistant", "content": [{"type": "tool_use", "id": "hr_long", "name": long_name, "input": {}}]},
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "hr_long", "content": "RETRIEVED BODY"}]},
]
inputs = GenericGuardrailAPIInputs(texts=["x"], structured_messages=json.loads(json.dumps(structured)))
sent: dict = {}
def _echo(**kwargs):
sent["messages"] = kwargs["json"]["messages"]
return _make_compress_response(json.loads(json.dumps(kwargs["json"]["messages"])))
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock, side_effect=_echo):
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data={"model": "claude-sonnet-4-5-20250929", "messages": raw_messages},
input_type="request",
)
assert not any(row.get("tool_call_id") == "hr_long" for row in sent["messages"])
assert result["structured_messages"][3] == structured[3]
assert any(row.get("content") == "H" * 5000 for row in sent["messages"])
def test_raw_retrieve_call_ids_covers_both_shapes_and_ignores_others():
"""Retrieve ids are read from OpenAI tool_calls and Anthropic tool_use blocks;
non-retrieve calls, non-tool_use blocks, string content, and non-list inputs
yield nothing."""
from litellm.proxy.guardrails.guardrail_hooks.headroom.headroom import _raw_retrieve_call_ids
messages = [
{
"role": "assistant",
"tool_calls": [
{"id": "oa1", "function": {"name": HEADROOM_RETRIEVE_TOOL_NAME}},
{"id": "other", "function": {"name": "get_weather"}},
{"id": "malformed", "function": {"name": 123}},
{"id": "nofunc"},
],
},
{
"role": "assistant",
"content": [
{"type": "tool_use", "id": "an1", "name": "mcp__hr__headroom_retrieve", "input": {}},
{"type": "tool_use", "id": "an2", "name": "jira_get_issue", "input": {}},
{"type": "text", "text": "noise"},
],
},
{"role": "user", "content": "plain string content, not a list"},
]
assert _raw_retrieve_call_ids(messages) == frozenset({"oa1", "an1"})
assert _raw_retrieve_call_ids("not a list") == frozenset()
assert _raw_retrieve_call_ids(None) == frozenset()
@pytest.mark.asyncio
async def test_nothing_compressible_returns_inputs_untouched(guardrail: HeadroomGuardrail):
"""A single-turn request is all protected, so there is nothing to send and