fix(anthropic): keep pinging during agentic hooks and fail instead of replaying server-fulfilled tool_use

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
mateo 2026-08-08 07:30:24 +00:00
parent e4c2ad4627
commit f994068a73
4 changed files with 190 additions and 59 deletions

View file

@ -9,11 +9,13 @@ follow-up response is chained as Phase 2 of the same iterator.
In hold-back mode (``hold_back=True``), chunks are buffered instead of
yielded live, with SSE ping events emitted while the upstream message is
in flight. On exhaustion the hooks run first: if a follow-up response
replaces the message, only the follow-up is yielded and the buffered
message is dropped; otherwise the buffer is replayed verbatim. This is
required for server-fulfilled tools (e.g. ``headroom_retrieve``), whose
tool_use blocks must never reach a client that cannot execute them.
in flight and while the agentic hooks run. On exhaustion the hooks run
first: if a follow-up response replaces the message, only the follow-up
is yielded and the buffered message is dropped; otherwise the buffer is
replayed verbatim, unless it holds a tool_use for a server-fulfilled tool
(e.g. ``headroom_retrieve``), in which case an SSE ``error`` event is
emitted because such a block must never reach a client that cannot
execute it.
"""
import asyncio
@ -26,6 +28,11 @@ from litellm._logging import verbose_logger
PING_SSE_BYTES: Final = b'event: ping\ndata: {"type": "ping"}\n\n'
HOLD_BACK_PING_INTERVAL_SECONDS: Final = 15.0
SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES: Final = (
b"event: error\n"
b'data: {"type": "error", "error": {"type": "api_error", "message": '
b'"Server-side tool retrieval failed, so this turn could not be completed. Please retry."}}\n\n'
)
# ---------------------------------------------------------------------------
# SSE parsing helpers (module-level to keep the class lean)
@ -170,6 +177,7 @@ class AgenticAnthropicStreamingIterator:
custom_llm_provider: str,
kwargs: dict,
hold_back: bool = False,
server_fulfilled_tool_names: frozenset[str] = frozenset(),
ping_interval_seconds: float = HOLD_BACK_PING_INTERVAL_SECONDS,
):
self._inner = completion_stream.__aiter__()
@ -182,6 +190,7 @@ class AgenticAnthropicStreamingIterator:
self._custom_llm_provider = custom_llm_provider
self._kwargs = kwargs
self._hold_back = hold_back
self._server_fulfilled_tool_names = server_fulfilled_tool_names
self._ping_interval_seconds = ping_interval_seconds
self._collected_bytes: list[bytes] = []
@ -189,7 +198,9 @@ class AgenticAnthropicStreamingIterator:
self._hook_processing_done = False
self._follow_up_iterator: AsyncIterator | None = None
self._drain_task: asyncio.Task | None = None
self._hook_task: asyncio.Task | None = None
self._replay_index = 0
self._error_emitted = False
def __aiter__(self):
return self
@ -223,22 +234,42 @@ class AgenticAnthropicStreamingIterator:
except StopAsyncIteration:
return
async def _completed_within_ping_interval(self, task: asyncio.Task) -> bool:
try:
await asyncio.wait_for(asyncio.shield(task), timeout=self._ping_interval_seconds)
except asyncio.TimeoutError:
return False
return True
async def _anext_held_back(self) -> bytes:
if self._drain_task is None:
self._drain_task = asyncio.create_task(self._drain_upstream())
return PING_SSE_BYTES
while not self._stream_exhausted:
try:
await asyncio.wait_for(asyncio.shield(self._drain_task), timeout=self._ping_interval_seconds)
except asyncio.TimeoutError:
if not self._stream_exhausted:
if not await self._completed_within_ping_interval(self._drain_task):
return PING_SSE_BYTES
self._stream_exhausted = True
await self._process_agentic_hooks()
if self._hook_task is None:
self._hook_task = asyncio.create_task(self._process_agentic_hooks())
if not await self._completed_within_ping_interval(self._hook_task):
return PING_SSE_BYTES
if self._follow_up_iterator is not None:
return await self._follow_up_iterator.__anext__()
if self._buffer_holds_server_fulfilled_tool_use():
if self._error_emitted:
raise StopAsyncIteration
self._error_emitted = True
verbose_logger.error(
"AgenticStreamingIterator: hooks did not replace a message containing a server-fulfilled "
"tool_use [model=%s]; emitting an SSE error instead of leaking the tool call to the client",
self._model,
)
return SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES
if self._replay_index < len(self._collected_bytes):
chunk: Final = self._collected_bytes[self._replay_index]
self._replay_index += 1
@ -246,18 +277,40 @@ class AgenticAnthropicStreamingIterator:
raise StopAsyncIteration
def _buffer_holds_server_fulfilled_tool_use(self) -> bool:
if not self._server_fulfilled_tool_names:
return False
started_blocks: Final = (
data.get("content_block")
for event_type, data in _parse_sse_events(b"".join(self._collected_bytes))
if event_type == "content_block_start"
)
return any(
isinstance(block, dict)
and block.get("type") == "tool_use"
and block.get("name") in self._server_fulfilled_tool_names
for block in started_blocks
)
@staticmethod
async def _settle_task(task: asyncio.Task | None) -> None:
if task is None:
return
if task.done():
if not task.cancelled():
task.exception()
return
task.cancel()
with contextlib.suppress(asyncio.CancelledError):
await task
async def aclose(self) -> None:
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
aclose_if_supported,
)
if self._drain_task is not None and self._drain_task.done():
if not self._drain_task.cancelled():
self._drain_task.exception()
elif self._drain_task is not None:
self._drain_task.cancel()
with contextlib.suppress(asyncio.CancelledError):
await self._drain_task
await self._settle_task(self._drain_task)
await self._settle_task(self._hook_task)
await aclose_if_supported(self._inner)
await aclose_if_supported(self._follow_up_iterator)

View file

@ -2179,6 +2179,10 @@ class BaseLLMHTTPHandler:
AgenticAnthropicStreamingIterator,
)
held_back_tool_names: Final = self._server_fulfilled_tools_in_request(
logging_obj=logging_obj,
tools=anthropic_messages_optional_request_params.get("tools"),
)
initial_response = AgenticAnthropicStreamingIterator(
completion_stream=completion_stream,
http_handler=self,
@ -2189,10 +2193,8 @@ class BaseLLMHTTPHandler:
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
kwargs={**kwargs, "api_key": api_key} if api_key else kwargs,
hold_back=self._should_hold_back_stream(
logging_obj=logging_obj,
tools=anthropic_messages_optional_request_params.get("tools"),
),
hold_back=bool(held_back_tool_names),
server_fulfilled_tool_names=held_back_tool_names,
)
return AnthropicMessagesStreamingResponse(
completion_stream=initial_response,
@ -5038,22 +5040,23 @@ class BaseLLMHTTPHandler:
return False
@staticmethod
def _should_hold_back_stream(logging_obj: LiteLLMLoggingObj, tools: object) -> bool:
def _server_fulfilled_tools_in_request(logging_obj: LiteLLMLoggingObj, tools: object) -> frozenset[str]:
"""
True when the request carries a tool that a registered callback fulfills
server-side (e.g. ``headroom_retrieve``). The model's tool_use for such a
tool must never reach the client, which cannot execute it: the agentic
loop replaces the whole message with a follow-up response, so the stream
is buffered (with ping keepalives) instead of forwarded live.
The request's tools that a registered callback fulfills server-side (e.g.
``headroom_retrieve``). The model's tool_use for such a tool must never
reach the client, which cannot execute it: the agentic loop replaces the
whole message with a follow-up response, so a stream carrying any of
these is buffered (with ping keepalives) instead of forwarded live.
"""
if not isinstance(tools, list) or not tools:
return False
return frozenset()
from litellm.litellm_core_utils.prompt_templates.factory import has_tool_with_name
return any(
has_tool_with_name(tools, name)
return frozenset(
name
for cb in _custom_logger_callbacks(logging_obj)
for name in getattr(cb, "server_fulfilled_tool_names", frozenset())
if has_tool_with_name(tools, name)
)
@staticmethod

View file

@ -15,6 +15,7 @@ sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
PING_SSE_BYTES,
SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES,
AgenticAnthropicStreamingIterator,
_handle_content_block_delta,
_handle_content_block_start,
@ -261,6 +262,7 @@ def _build_hold_back_iterator(
stream: MockAsyncStream,
mock_handler: MagicMock,
ping_interval_seconds: float = 15.0,
server_fulfilled_tool_names: frozenset = frozenset({"litellm_content_retrieve"}),
) -> AgenticAnthropicStreamingIterator:
return AgenticAnthropicStreamingIterator(
completion_stream=stream,
@ -273,6 +275,7 @@ def _build_hold_back_iterator(
custom_llm_provider="anthropic",
kwargs={},
hold_back=True,
server_fulfilled_tool_names=server_fulfilled_tool_names,
ping_interval_seconds=ping_interval_seconds,
)
@ -920,27 +923,76 @@ class TestAgenticStreamingIteratorHoldBack:
mock_handler._call_agentic_completion_hooks.assert_not_awaited()
@pytest.mark.asyncio
async def test_should_replay_buffer_when_hook_processing_errors(self):
"""A hook crash degrades to replaying the original message rather than dropping it."""
async def test_should_emit_pings_while_hooks_are_slow(self):
"""Retrieval and follow-up generation can outlast a client's idle timeout, so hooks get keepalives too."""
chunks = _build_tool_use_stream()
phase2_chunks = [b"follow-up-chunk"]
async def slow_hooks(**_kwargs):
await asyncio.sleep(0.12)
return MockAsyncStream(phase2_chunks)
mock_handler = MagicMock()
mock_handler._call_agentic_completion_hooks = AsyncMock(side_effect=slow_hooks)
iterator = _build_hold_back_iterator(
MockAsyncStream(chunks),
mock_handler,
ping_interval_seconds=0.02,
)
collected = []
async for chunk in iterator:
collected.append(chunk)
assert collected.count(PING_SSE_BYTES) >= 4
assert [c for c in collected if c != PING_SSE_BYTES] == phase2_chunks
@pytest.mark.asyncio
async def test_should_error_instead_of_replaying_server_fulfilled_tool_use_when_hook_crashes(self):
"""A hook crash must not replay the buffered retrieval tool_use: that is the unknown-tool bug."""
chunks = _build_tool_use_stream()
mock_handler = MagicMock()
mock_handler._call_agentic_completion_hooks = AsyncMock(side_effect=RuntimeError("hook exploded"))
mock_logging = MagicMock()
mock_logging.litellm_call_id = "test_call_holdback"
iterator = _build_hold_back_iterator(MockAsyncStream(chunks), mock_handler)
iterator = AgenticAnthropicStreamingIterator(
completion_stream=MockAsyncStream(chunks),
http_handler=mock_handler,
model="claude-sonnet-4-20250514",
messages=[],
anthropic_messages_provider_config=MagicMock(),
anthropic_messages_optional_request_params={},
logging_obj=mock_logging,
custom_llm_provider="anthropic",
kwargs={},
hold_back=True,
collected = []
async for chunk in iterator:
collected.append(chunk)
assert [c for c in collected if c != PING_SSE_BYTES] == [SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES]
assert b"litellm_content_retrieve" not in b"".join(collected)
@pytest.mark.asyncio
async def test_should_error_instead_of_replaying_when_no_hook_fires_on_tool_use(self):
"""Hooks returning None on a retrieval tool_use is still a leak, so the turn fails loudly."""
chunks = _build_tool_use_stream()
mock_handler = MagicMock()
mock_handler._call_agentic_completion_hooks = AsyncMock(return_value=None)
iterator = _build_hold_back_iterator(MockAsyncStream(chunks), mock_handler)
collected = []
async for chunk in iterator:
collected.append(chunk)
assert [c for c in collected if c != PING_SSE_BYTES] == [SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES]
@pytest.mark.asyncio
async def test_should_replay_client_owned_tool_use_verbatim(self):
"""Only server-fulfilled tools are withheld: a client's own tool_use still reaches it byte-identical."""
chunks = _build_tool_use_stream()
mock_handler = MagicMock()
mock_handler._call_agentic_completion_hooks = AsyncMock(return_value=None)
iterator = _build_hold_back_iterator(
MockAsyncStream(chunks),
mock_handler,
server_fulfilled_tool_names=frozenset({"headroom_retrieve"}),
)
collected = []
@ -968,3 +1020,26 @@ class TestAgenticStreamingIteratorHoldBack:
await iterator.aclose()
assert iterator._drain_task.cancelled()
@pytest.mark.asyncio
async def test_aclose_cancels_in_flight_hook_task(self):
"""Closing while hooks are running must not leave the retrieval follow-up task orphaned."""
chunks = _build_tool_use_stream()
async def never_finishing_hooks(**_kwargs):
await asyncio.sleep(5.0)
mock_handler = MagicMock()
mock_handler._call_agentic_completion_hooks = AsyncMock(side_effect=never_finishing_hooks)
iterator = _build_hold_back_iterator(
MockAsyncStream(chunks),
mock_handler,
ping_interval_seconds=0.02,
)
while iterator._hook_task is None:
await iterator.__anext__()
await iterator.aclose()
assert iterator._hook_task.cancelled()

View file

@ -2073,9 +2073,9 @@ async def test_anthropic_invalid_thinking_signature_retry_resigns_bedrock_reques
assert retry_authorization != first_attempt_headers["Authorization"]
class TestShouldHoldBackStream:
"""_should_hold_back_stream gates the buffered (non-leaking) streaming mode
for server-fulfilled tools like headroom_retrieve."""
class TestServerFulfilledToolsInRequest:
"""_server_fulfilled_tools_in_request gates the buffered (non-leaking) streaming
mode for server-fulfilled tools like headroom_retrieve."""
@staticmethod
def _logging_obj_with(callbacks):
@ -2093,12 +2093,9 @@ class TestShouldHoldBackStream:
{"name": "Bash", "input_schema": {"type": "object"}},
{"name": "headroom_retrieve", "input_schema": {"type": "object"}},
]
assert (
BaseLLMHTTPHandler._should_hold_back_stream(
logging_obj=self._logging_obj_with([RetrievalCallback()]), tools=tools
)
is True
)
assert BaseLLMHTTPHandler._server_fulfilled_tools_in_request(
logging_obj=self._logging_obj_with([RetrievalCallback()]), tools=tools
) == frozenset({"headroom_retrieve"})
def test_should_stream_live_when_tool_absent_from_request(self):
from litellm.integrations.custom_logger import CustomLogger
@ -2108,10 +2105,10 @@ class TestShouldHoldBackStream:
tools = [{"name": "Bash", "input_schema": {"type": "object"}}]
assert (
BaseLLMHTTPHandler._should_hold_back_stream(
BaseLLMHTTPHandler._server_fulfilled_tools_in_request(
logging_obj=self._logging_obj_with([RetrievalCallback()]), tools=tools
)
is False
== frozenset()
)
def test_should_stream_live_when_no_callback_declares_tool_names(self):
@ -2119,14 +2116,17 @@ class TestShouldHoldBackStream:
tools = [{"name": "headroom_retrieve", "input_schema": {"type": "object"}}]
assert (
BaseLLMHTTPHandler._should_hold_back_stream(
BaseLLMHTTPHandler._server_fulfilled_tools_in_request(
logging_obj=self._logging_obj_with([CustomLogger()]), tools=tools
)
is False
== frozenset()
)
def test_should_stream_live_without_tools(self):
assert BaseLLMHTTPHandler._should_hold_back_stream(logging_obj=self._logging_obj_with([]), tools=None) is False
assert (
BaseLLMHTTPHandler._server_fulfilled_tools_in_request(logging_obj=self._logging_obj_with([]), tools=None)
== frozenset()
)
def test_interception_callbacks_declare_their_retrieval_tools(self):
from litellm.integrations.compression_interception.handler import (