mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(guardrails): walk nested tool results iteratively, with a depth bound
CI flagged _collect_content as recursive. It was, and worse, it was unbounded: a tool_result nests its own content, the nesting is caller controlled, and the descent had nothing to stop it. That is a JSON bomb, not a style issue. Now an explicit queue with a depth bound of 8. Real payloads nest one or two deep. The queue is walked in document order because the shield maps its replies back by position, so collection order is part of the contract.
This commit is contained in:
parent
403b06ea76
commit
ce921f49e4
2 changed files with 46 additions and 12 deletions
|
|
@ -71,6 +71,10 @@ JsonBody: TypeAlias = dict
|
|||
|
||||
# One redactable span: the text as it stands, and the write that puts the
|
||||
# replacement back where it came from.
|
||||
# How far a tool_result chain is followed. Real payloads nest one or two deep; the
|
||||
# bound is what stops a crafted one from becoming an unbounded walk.
|
||||
_MAX_CONTENT_DEPTH: Final = 8
|
||||
|
||||
_Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's param list.
|
||||
|
||||
# The accumulator the collectors below append into. It never escapes
|
||||
|
|
@ -113,19 +117,32 @@ def _collect_prompt(data: MutableRequest, slots: _SlotSink) -> None:
|
|||
|
||||
|
||||
def _collect_content(container: MutableRequest, slots: _SlotSink) -> None:
|
||||
"""`content` is either a string or the multimodal list of typed parts."""
|
||||
content: Final = container.get("content")
|
||||
if isinstance(content, str):
|
||||
_collect(container, "content", slots)
|
||||
return
|
||||
for part in content if isinstance(content, list) else ():
|
||||
if not isinstance(part, dict):
|
||||
"""Collects `content`, a string or a list of typed parts.
|
||||
|
||||
An Anthropic tool_result nests its own content, so this has to descend. It walks
|
||||
with an explicit stack and a depth bound rather than by recursion: the nesting is
|
||||
caller controlled, and an unbounded descent is a JSON bomb.
|
||||
"""
|
||||
# Walked in document order: the shield maps its replies back by position, so the
|
||||
# order spans are collected in is part of the contract.
|
||||
pending: Final[list] = [(container, 0)] # mutable-ok: local queue, never escapes.
|
||||
cursor = 0 # rebind-ok: advances through the queue.
|
||||
while cursor < len(pending):
|
||||
node, depth = pending[cursor]
|
||||
cursor += 1
|
||||
content = node.get("content")
|
||||
if isinstance(content, str):
|
||||
_collect(node, "content", slots)
|
||||
continue
|
||||
_collect(part, "text", slots)
|
||||
# An Anthropic tool_result carries its own content, as a string or as more
|
||||
# blocks. Image and audio parts have no text and fall through untouched.
|
||||
if "content" in part:
|
||||
_collect_content(part, slots)
|
||||
if depth >= _MAX_CONTENT_DEPTH:
|
||||
continue
|
||||
for part in content if isinstance(content, list) else ():
|
||||
if not isinstance(part, dict):
|
||||
continue
|
||||
# Image and audio parts have no text and fall through untouched.
|
||||
_collect(part, "text", slots)
|
||||
if "content" in part:
|
||||
pending.append((part, depth + 1))
|
||||
|
||||
|
||||
def _collect_participant_name(message: MutableRequest, slots: _SlotSink) -> None:
|
||||
|
|
|
|||
|
|
@ -384,6 +384,23 @@ class TestRequestCoverage:
|
|||
assert data["messages"][0]["content"][0]["content"] == "[EMAIL_1]"
|
||||
assert data["messages"][0]["content"][1]["content"][0]["text"] == "[EMAIL_2]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deeply_nested_tool_results_are_bounded(self):
|
||||
"""Nesting is caller controlled, so the descent has to stop somewhere.
|
||||
|
||||
The walk must terminate on a payload built to be pathological, rather than
|
||||
following it as far as it goes.
|
||||
"""
|
||||
guardrail = _guardrail()
|
||||
_mock_post(guardrail, {"texts": ["ok"] * 64})
|
||||
|
||||
deep: dict = {"type": "tool_result", "content": "jane.doe@example.com"}
|
||||
for _ in range(200):
|
||||
deep = {"type": "tool_result", "content": [deep]}
|
||||
data = {"messages": [{"role": "user", "content": [deep]}]}
|
||||
|
||||
await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_completions_suffix_is_redacted(self):
|
||||
"""LiteLLM forwards the legacy `suffix` to providers that support it."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue