mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
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:
parent
44d84360fb
commit
2a998c2f38
2 changed files with 274 additions and 4 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue