mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
e4c2ad4627
commit
f994068a73
4 changed files with 190 additions and 59 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue