fix(anthropic): require message stop for streams
(cherry picked from commit 8bca0d97c602624f6acaaa3073a8fbb43f718e3c)
This commit is contained in:
@@ -26,6 +26,7 @@ from agent.anthropic_endpoints import (
|
||||
from agent.anthropic_message_convert import (
|
||||
convert_messages_to_anthropic, convert_tools_to_anthropic, normalize_model_name,
|
||||
)
|
||||
from agent.errors import EmptyStreamError
|
||||
|
||||
from hermes_cli import __version__ as _HERMES_VERSION
|
||||
|
||||
@@ -734,9 +735,13 @@ def _stream_final_message(stream_fn, api_kwargs, log_prefix, on_stream_event, on
|
||||
# has given up, so abandon the stream (``with`` closes it) instead of streaming an answer
|
||||
# nobody reads.
|
||||
# Some SDK versions drop optional message_delta metadata from the final snapshot.
|
||||
# Non-iterable shims (get_final_message-only) skip straight to the snapshot.
|
||||
# An iterable stream must end in message_stop; a get_final_message-only shim cannot
|
||||
# prove completion and is treated as a retryable incomplete response.
|
||||
stop_details = None
|
||||
saw_message_stop = False
|
||||
for event in (stream if isinstance(stream, Iterable) else ()):
|
||||
if getattr(event, "type", None) == "message_stop":
|
||||
saw_message_stop = True
|
||||
if getattr(event, "type", None) == "message_delta":
|
||||
details = getattr(getattr(event, "delta", None), "stop_details", None)
|
||||
if details is not None:
|
||||
@@ -752,6 +757,10 @@ def _stream_final_message(stream_fn, api_kwargs, log_prefix, on_stream_event, on
|
||||
raise
|
||||
except Exception:
|
||||
logger.debug("%son_stream_event callback failed", log_prefix, exc_info=True)
|
||||
if not saw_message_stop:
|
||||
raise EmptyStreamError(
|
||||
"Anthropic Messages stream ended before message_stop (possible upstream stream drop)."
|
||||
)
|
||||
message = stream.get_final_message()
|
||||
if stop_details is not None:
|
||||
message.stop_details = stop_details
|
||||
|
||||
@@ -3300,6 +3300,8 @@ class _StreamingCall(StreamingWaitMonitor):
|
||||
# Eventless stream: the SDK's get_final_message() raises AssertionError (no
|
||||
# message_start); shims may fabricate a contentless Message. All -> EmptyStreamError.
|
||||
saw_stream_event = False
|
||||
saw_message_stop = False
|
||||
pending_deltas = []
|
||||
self.last_chunk_time["t"] = time.time()
|
||||
_diag = self._new_diag()
|
||||
self._writer_token = self._attempt_stream_response = None
|
||||
@@ -3340,6 +3342,8 @@ class _StreamingCall(StreamingWaitMonitor):
|
||||
if self.agent._interrupt_requested:
|
||||
break
|
||||
event_type = getattr(event, "type", None)
|
||||
if event_type == "message_stop":
|
||||
saw_message_stop = True
|
||||
if event_type == "content_block_start":
|
||||
block = getattr(event, "content_block", None)
|
||||
if block and getattr(block, "type", None) == "tool_use":
|
||||
@@ -3355,11 +3359,15 @@ class _StreamingCall(StreamingWaitMonitor):
|
||||
if delta_type == "text_delta":
|
||||
text = getattr(delta, "text", "")
|
||||
if text and not has_tool_use:
|
||||
self._emit_text(text)
|
||||
pending_deltas.append(("text", text))
|
||||
elif delta_type == "thinking_delta" and getattr(delta, "thinking", ""):
|
||||
self._emit_reasoning(delta.thinking)
|
||||
pending_deltas.append(("reasoning", delta.thinking))
|
||||
raw_stream = _stream_context["stream"]
|
||||
if not self.agent._interrupt_requested and raw_stream is not None:
|
||||
if not saw_message_stop:
|
||||
raise EmptyStreamError(
|
||||
"Anthropic Messages stream ended before message_stop (possible upstream stream drop)."
|
||||
)
|
||||
try:
|
||||
base_final_message = raw_stream.get_final_message()
|
||||
# The SDK snapshot keeps only stop_reason/stop_sequence from message_delta; the
|
||||
@@ -3382,11 +3390,21 @@ class _StreamingCall(StreamingWaitMonitor):
|
||||
|
||||
if self.agent._interrupt_requested:
|
||||
return None
|
||||
def _flush_completed_deltas():
|
||||
for kind, text in pending_deltas:
|
||||
if kind == "text":
|
||||
self._emit_text(text)
|
||||
else:
|
||||
self._emit_reasoning(text)
|
||||
if base_final_message is not None:
|
||||
self._check_anthropic_message(base_final_message, tool_drop=False)
|
||||
if not stream.output_modified:
|
||||
return self._check_anthropic_message(base_final_message)
|
||||
return self._check_anthropic_message(accumulator.response(base_final_message))
|
||||
response = self._check_anthropic_message(base_final_message)
|
||||
_flush_completed_deltas()
|
||||
return response
|
||||
response = self._check_anthropic_message(accumulator.response(base_final_message))
|
||||
_flush_completed_deltas()
|
||||
return response
|
||||
|
||||
# ── retry loop ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
106
tests/agent/test_anthropic_message_stop.py
Normal file
106
tests/agent/test_anthropic_message_stop.py
Normal file
@@ -0,0 +1,106 @@
|
||||
"""Regression coverage for incomplete native Anthropic Messages streams (#121320)."""
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _stream_cm(final_message, events):
|
||||
stream = MagicMock()
|
||||
stream.__iter__ = MagicMock(return_value=iter(events))
|
||||
stream.get_final_message = MagicMock(return_value=final_message)
|
||||
cm = MagicMock()
|
||||
cm.__enter__ = MagicMock(return_value=stream)
|
||||
cm.__exit__ = MagicMock(return_value=False)
|
||||
return cm
|
||||
|
||||
|
||||
def _event(event_type, **fields):
|
||||
return SimpleNamespace(type=event_type, **fields)
|
||||
|
||||
|
||||
def _text_delta(text):
|
||||
return _event(
|
||||
"content_block_delta",
|
||||
delta=SimpleNamespace(type="text_delta", text=text),
|
||||
)
|
||||
|
||||
|
||||
def _thinking_delta(thinking):
|
||||
return _event(
|
||||
"content_block_delta",
|
||||
delta=SimpleNamespace(type="thinking_delta", thinking=thinking),
|
||||
)
|
||||
|
||||
|
||||
def _agent():
|
||||
from run_agent import AIAgent
|
||||
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://api.anthropic.com",
|
||||
model="claude-test",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
agent.api_mode = "anthropic_messages"
|
||||
agent._interrupt_requested = False
|
||||
agent._anthropic_client = MagicMock()
|
||||
agent._anthropic_api_key = "test-key"
|
||||
agent._create_request_anthropic_client = lambda *args, **kwargs: agent._anthropic_client
|
||||
return agent
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("dropped_event", "expected_callback"),
|
||||
[
|
||||
(_text_delta("HALF-ANSWER"), "RECOVERED"),
|
||||
(_thinking_delta("PARTIAL-THOUGHT"), "RECOVERED"),
|
||||
],
|
||||
)
|
||||
def test_anthropic_eof_before_message_stop_retries_without_delivering_partial_output(
|
||||
monkeypatch, dropped_event, expected_callback,
|
||||
):
|
||||
"""A final SDK snapshot cannot turn an unterminated SSE response into success."""
|
||||
monkeypatch.setenv("HERMES_STREAM_RETRIES", "1")
|
||||
agent = _agent()
|
||||
delivered = []
|
||||
agent.stream_delta_callback = delivered.append
|
||||
dropped = SimpleNamespace(content=[SimpleNamespace(type="text", text="HALF-ANSWER")], stop_reason="end_turn")
|
||||
recovered = SimpleNamespace(content=[SimpleNamespace(type="text", text="RECOVERED")], stop_reason="end_turn")
|
||||
agent._anthropic_client.messages.stream.side_effect = [
|
||||
_stream_cm(dropped, [_event("message_start"), dropped_event]),
|
||||
_stream_cm(recovered, [_event("message_start"), _text_delta("RECOVERED"), _event("message_stop")]),
|
||||
]
|
||||
|
||||
response = agent._interruptible_streaming_api_call({"model": "claude-test"})
|
||||
|
||||
assert response is recovered
|
||||
assert agent._anthropic_client.messages.stream.call_count == 2
|
||||
assert delivered == [expected_callback]
|
||||
|
||||
|
||||
def test_anthropic_message_stop_accepts_completed_stream_without_retry():
|
||||
agent = _agent()
|
||||
completed = SimpleNamespace(content=[SimpleNamespace(type="text", text="done")], stop_reason="end_turn")
|
||||
agent._anthropic_client.messages.stream.return_value = _stream_cm(
|
||||
completed,
|
||||
[_event("message_start"), _text_delta("done"), _event("message_stop")],
|
||||
)
|
||||
|
||||
assert agent._interruptible_streaming_api_call({"model": "claude-test"}) is completed
|
||||
assert agent._anthropic_client.messages.stream.call_count == 1
|
||||
|
||||
|
||||
def test_auxiliary_anthropic_eof_before_message_stop_raises_empty_stream():
|
||||
from agent.anthropic_adapter import _stream_final_message
|
||||
from agent.errors import EmptyStreamError
|
||||
|
||||
partial = SimpleNamespace(content=[SimpleNamespace(type="text", text="HALF")], stop_reason="end_turn")
|
||||
with pytest.raises(EmptyStreamError, match="message_stop"):
|
||||
_stream_final_message(
|
||||
lambda **kwargs: _stream_cm(partial, [_event("message_start"), _text_delta("HALF")]),
|
||||
{"model": "claude-test"}, "", None, None,
|
||||
)
|
||||
@@ -1120,6 +1120,7 @@ class TestAnthropicStreamCallbacks:
|
||||
type="content_block_start",
|
||||
content_block=SimpleNamespace(type="tool_use", name="terminal"),
|
||||
),
|
||||
SimpleNamespace(type="message_stop"),
|
||||
]
|
||||
|
||||
final_message = SimpleNamespace(
|
||||
@@ -1176,7 +1177,7 @@ class TestAnthropicStreamCallbacks:
|
||||
good_stream = MagicMock()
|
||||
good_stream.__enter__ = MagicMock(return_value=good_stream)
|
||||
good_stream.__exit__ = MagicMock(return_value=False)
|
||||
good_stream.__iter__ = MagicMock(return_value=iter([]))
|
||||
good_stream.__iter__ = MagicMock(return_value=iter([SimpleNamespace(type="message_stop")]))
|
||||
good_stream.get_final_message.return_value = final_message
|
||||
|
||||
agent._anthropic_client = MagicMock()
|
||||
@@ -1231,7 +1232,7 @@ class TestAnthropicStreamCallbacks:
|
||||
good_stream = MagicMock()
|
||||
good_stream.__enter__ = MagicMock(return_value=good_stream)
|
||||
good_stream.__exit__ = MagicMock(return_value=False)
|
||||
good_stream.__iter__ = MagicMock(return_value=iter([]))
|
||||
good_stream.__iter__ = MagicMock(return_value=iter([SimpleNamespace(type="message_stop")]))
|
||||
good_stream.get_final_message.return_value = repaired_message
|
||||
|
||||
seen_tools = []
|
||||
@@ -1254,9 +1255,7 @@ class TestAnthropicStreamCallbacks:
|
||||
assert seen_tools[1][0]["eager_input_streaming"] is False
|
||||
|
||||
def test_anthropic_partial_tool_names_do_not_survive_into_next_attempt(self):
|
||||
"""A tool name from an attempt that died before any text is attempt-local: when the
|
||||
retry streams plain text and then drops, the partial stub must not blame ``old_tool``
|
||||
(that would also make the third attempt look mid-tool-call and thus retryable)."""
|
||||
"""A tool name from a failed attempt cannot leak into a later native-stream retry."""
|
||||
from run_agent import AIAgent
|
||||
|
||||
agent = AIAgent(
|
||||
@@ -1270,6 +1269,17 @@ class TestAnthropicStreamCallbacks:
|
||||
agent.api_mode = "anthropic_messages"
|
||||
agent._interrupt_requested = False
|
||||
|
||||
recovered_stream = MagicMock()
|
||||
recovered_stream.__enter__ = MagicMock(return_value=recovered_stream)
|
||||
recovered_stream.__exit__ = MagicMock(return_value=False)
|
||||
recovered_stream.__iter__ = MagicMock(return_value=iter([
|
||||
SimpleNamespace(type="content_block_delta",
|
||||
delta=SimpleNamespace(type="text_delta", text="Recovered")),
|
||||
SimpleNamespace(type="message_stop"),
|
||||
]))
|
||||
recovered_stream.get_final_message.return_value = SimpleNamespace(
|
||||
content=[SimpleNamespace(type="text", text="Recovered")], stop_reason="end_turn"
|
||||
)
|
||||
attempts = [
|
||||
_AnthropicEventStream([SimpleNamespace(type="content_block_start",
|
||||
content_block=SimpleNamespace(type="tool_use", name="old_tool"))],
|
||||
@@ -1277,21 +1287,22 @@ class TestAnthropicStreamCallbacks:
|
||||
_AnthropicEventStream([SimpleNamespace(type="content_block_delta",
|
||||
delta=SimpleNamespace(type="text_delta", text="Plain answer."))],
|
||||
ConnectionError("connection dropped")),
|
||||
recovered_stream,
|
||||
]
|
||||
agent._anthropic_client = MagicMock()
|
||||
agent._anthropic_client.messages.stream.side_effect = lambda **kwargs: attempts.pop(0)
|
||||
agent._create_request_anthropic_client = lambda *a, **k: agent._anthropic_client
|
||||
emitted = []
|
||||
# A real consumer: delivered text is recorded, so the second attempt counts as
|
||||
# partial delivery (a bare _fire_stream_delta override records nothing).
|
||||
# Native Anthropic text is held until message_stop, so the dropped second
|
||||
# attempt remains retryable and cannot be persisted as a partial answer.
|
||||
agent.stream_delta_callback = emitted.append
|
||||
|
||||
response = agent._interruptible_streaming_api_call(
|
||||
{"model": agent.model, "tools": [{"name": "old_tool", "input_schema": {"type": "object"}}]})
|
||||
|
||||
assert agent._anthropic_client.messages.stream.call_count == 2
|
||||
assert "old_tool" not in (response.choices[0].message.content or "")
|
||||
assert not any("old_tool" in t for t in emitted)
|
||||
assert agent._anthropic_client.messages.stream.call_count == 3
|
||||
assert response.content[0].text == "Recovered"
|
||||
assert emitted == ["Recovered"]
|
||||
|
||||
@patch("run_agent.AIAgent._replace_primary_openai_client")
|
||||
def test_generic_anthropic_valueerror_still_propagates_without_stream_retry(
|
||||
|
||||
@@ -92,10 +92,11 @@ class TestAnthropicTransport:
|
||||
|
||||
details = {"type": "refusal", "category": "general_harms", "explanation": "classifier halt"}
|
||||
delta_event = SimpleNamespace(type="message_delta", delta=SimpleNamespace(stop_reason="refusal", stop_details=details))
|
||||
message_stop_event = SimpleNamespace(type="message_stop")
|
||||
|
||||
def _stream_cm(final):
|
||||
stream = MagicMock()
|
||||
stream.__iter__ = MagicMock(return_value=iter([delta_event]))
|
||||
stream.__iter__ = MagicMock(return_value=iter([delta_event, message_stop_event]))
|
||||
stream.get_final_message = MagicMock(return_value=final)
|
||||
cm = MagicMock()
|
||||
cm.__enter__, cm.__exit__ = MagicMock(return_value=stream), MagicMock(return_value=False)
|
||||
|
||||
Reference in New Issue
Block a user