fix(backends,markdown): stream turn events live, no blank lines for tool blocks

This commit is contained in:
hh
2026-08-28 17:27:35 +02:00
parent fd8cbacd7b
commit 7ae87aeeb8
4 changed files with 121 additions and 15 deletions
+27 -12
View File
@@ -296,15 +296,21 @@ class ClaudeSdkBackend:
live = await self._acquire(key, session_id=session_id, history=prior, spec=spec) live = await self._acquire(key, session_id=session_id, history=prior, spec=spec)
message_id = f"msg_{uuid.uuid4().hex}" message_id = f"msg_{uuid.uuid4().hex}"
yield build_message_start(message_id=message_id, model=self._agent.model) yield build_message_start(message_id=message_id, model=self._agent.model)
turn = _Turn()
async with live.lock: async with live.lock:
live.running_turn = turn_id or message_id live.running_turn = turn_id or message_id
live.last_used = time.monotonic() live.last_used = time.monotonic()
try: try:
turn = await self._run_turn(live, prompt, observer=observer) # Events go out as the CLI produces them: the frontends
# stream text and thinking live, the turn is not buffered.
async for event in self._run_turn(live, prompt, turn, observer):
yield event
except Exception: except Exception:
live.running_turn = None live.running_turn = None
await self._pool.close(key) await self._pool.close(key)
if not (live.resumed and live.turns == 0): # A dead resume can be reseeded from history, but only
# while nothing of this turn has reached the caller yet.
if not (live.resumed and live.turns == 0) or turn.events:
raise raise
_log.exception( _log.exception(
"resume of %s failed, reseeding from history", live.session_id "resume of %s failed, reseeding from history", live.session_id
@@ -312,10 +318,10 @@ class ClaudeSdkBackend:
live = await self._acquire( live = await self._acquire(
key, session_id=None, history=prior, spec=spec key, session_id=None, history=prior, spec=spec
) )
turn = _Turn()
async with live.lock: async with live.lock:
live.running_turn = turn_id or message_id live.running_turn = turn_id or message_id
turn = await self._run_turn(live, prompt, observer=observer) async for event in self._run_turn(live, prompt, turn, observer):
for event in turn.events:
yield event yield event
live.turns += 1 live.turns += 1
live.last_used = time.monotonic() live.last_used = time.monotonic()
@@ -373,11 +379,16 @@ class ClaudeSdkBackend:
self, self,
live: Session, live: Session,
prompt: str, prompt: str,
*, turn: _Turn,
observer: Callable[[Any], None] | None = None, observer: Callable[[Any], None] | None = None,
) -> _Turn: ) -> AsyncIterator[MessageStreamEvent]:
"""Run one prompt, yielding wire events as they arrive.
``turn`` is filled in place (result, synthesized history, count of
events already yielded) so the caller can finish bookkeeping - and
decide whether a retry is still possible - after a failure.
"""
streaming = self._agent.options.include_partial_messages streaming = self._agent.options.include_partial_messages
turn = _Turn()
raw: list[Any] = [] raw: list[Any] = []
next_index = 0 next_index = 0
offset = 0 offset = 0
@@ -404,12 +415,16 @@ class ClaudeSdkBackend:
if isinstance(index, int): if isinstance(index, int):
next_index = max(next_index, offset + index + 1) next_index = max(next_index, offset + index + 1)
if streaming: if streaming:
turn.events.extend(_emit_stream_event(event, offset + index)) for out in _emit_stream_event(event, offset + index):
turn.events += 1
yield out
elif isinstance(message, AssistantMessage): elif isinstance(message, AssistantMessage):
raw.append(message) raw.append(message)
if not streaming: if not streaming:
for block in message.content: for block in message.content:
turn.events.extend(_emit_block(block, next_index)) for out in _emit_block(block, next_index):
turn.events += 1
yield out
next_index += 1 next_index += 1
elif isinstance(message, UserMessage): elif isinstance(message, UserMessage):
raw.append(message) raw.append(message)
@@ -423,11 +438,10 @@ class ClaudeSdkBackend:
"turn: agent=%s session=%s events=%d synthesized=%d stop=%s", "turn: agent=%s session=%s events=%d synthesized=%d stop=%s",
self._agent.name, self._agent.name,
live.session_id, live.session_id,
len(turn.events), turn.events,
len(turn.synthesized), len(turn.synthesized),
turn.stop_reason, turn.stop_reason,
) )
return turn
async def _acquire( async def _acquire(
self, self,
@@ -606,7 +620,8 @@ class _SessionSpec:
@dataclass @dataclass
class _Turn: class _Turn:
events: list[Any] = field(default_factory=list) events: int = 0
"""Wire events already yielded to the caller."""
synthesized: list[dict[str, Any]] = field(default_factory=list) synthesized: list[dict[str, Any]] = field(default_factory=list)
result: ResultMessage | None = None result: ResultMessage | None = None
stop_reason: StopReason = "end_turn" stop_reason: StopReason = "end_turn"
@@ -33,7 +33,9 @@ __all__ = [
# Empty ``### User:`` block appended after each assistant reply so the # Empty ``### User:`` block appended after each assistant reply so the
# human has an obvious place to type the next turn. Parser drops empty # human has an obvious place to type the next turn. Parser drops empty
# user blocks, so this doesn't re-trigger dispatch on its own. # user blocks, so this doesn't re-trigger dispatch on its own.
USER_SCAFFOLD = "### User:\n" # Blank line after the header, like every rendered turn - the file stays
# symmetric whether the human or the gateway wrote the marker.
USER_SCAFFOLD = "### User:\n\n"
# Default 4-backtick fence so tool results that contain literal ```` ``` ```` # Default 4-backtick fence so tool results that contain literal ```` ``` ````
@@ -73,7 +75,12 @@ def render_assistant_message(message: Message) -> str:
""" """
parts: list[str] = ["### Assistant:", ""] parts: list[str] = ["### Assistant:", ""]
for block in message.content: for block in message.content:
parts.extend(_render_block(block)) lines = list(_render_block(block))
if not lines:
# Tool calls render to nothing - no separator for them either,
# or every tool leaves a blank line behind.
continue
parts.extend(lines)
parts.append("") parts.append("")
return "\n".join(parts).rstrip() + "\n" return "\n".join(parts).rstrip() + "\n"
+39
View File
@@ -1,3 +1,4 @@
import asyncio
import tempfile import tempfile
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any
@@ -419,3 +420,41 @@ async def test_close_disconnects(cwd: Path) -> None:
) )
assert FakeClient.instances[0].connected is False assert FakeClient.instances[0].connected is False
assert backend.sessions == {} assert backend.sessions == {}
class GatedClient(FakeClient):
"""Streams one delta, then waits for ``gate`` before finishing the turn."""
gate: asyncio.Event
async def receive_response(self):
yield StreamEvent(
uuid="u", session_id="s", event={"type": "message_start", "message": {}}
)
for e in _stream(0, "first"):
yield e
await GatedClient.gate.wait()
yield AssistantMessage(content=[TextBlock(text="first")], model="m")
yield _result(self.session_id)
async def test_deltas_reach_the_caller_before_the_turn_ends(cwd: Path) -> None:
"""The turn is not buffered: a delta is observable while the CLI still runs."""
GatedClient.gate = asyncio.Event()
backend = _backend(cwd, InMemorySessionStore())
backend._factory = GatedClient
events = backend.complete(
agent=backend.agent,
messages=[{"role": "user", "content": "hi"}],
conversation_id="conv-live",
)
seen: list[Any] = []
async for ev in events:
seen.append(ev)
if isinstance(ev, RawContentBlockDeltaEvent):
break
assert seen[-1].delta.text == "first"
assert not GatedClient.gate.is_set()
GatedClient.gate.set()
rest = [e async for e in events]
assert isinstance(rest[-1], RawMessageStopEvent)
+45
View File
@@ -0,0 +1,45 @@
from anthropic.types import Message, TextBlock, ThinkingBlock, ToolUseBlock, Usage
from beaver_gateway.frontends.markdown import renderer
def _message(*blocks: object) -> Message:
return Message(
id="m",
type="message",
role="assistant",
model="x",
content=list(blocks), # type: ignore[arg-type]
stop_reason="end_turn",
stop_sequence=None,
usage=Usage(input_tokens=1, output_tokens=1),
)
def test_tool_calls_leave_no_blank_lines() -> None:
message = _message(
TextBlock(type="text", text="first"),
ToolUseBlock(type="tool_use", id="t1", name="Bash", input={"command": "ls"}),
ToolUseBlock(type="tool_use", id="t2", name="Read", input={"file_path": "x"}),
ToolUseBlock(type="tool_use", id="t3", name="Grep", input={"pattern": "y"}),
TextBlock(type="text", text="second"),
)
assert renderer.render_assistant_message(message) == (
"### Assistant:\n\nfirst\n\nsecond\n"
)
def test_thinking_renders_as_collapsed_callout() -> None:
message = _message(
ThinkingBlock(type="thinking", thinking="a\nb", signature="s"),
TextBlock(type="text", text="answer"),
)
assert renderer.render_assistant_message(message) == (
"### Assistant:\n\n> [!thinking]-\n> a\n> b\n\nanswer\n"
)
def test_user_scaffold_matches_rendered_user_turn_shape() -> None:
body = renderer.append_to_body("### Assistant:\n\nhi\n", renderer.USER_SCAFFOLD)
assert body.endswith("\n\n---\n\n### User:\n\n")
assert renderer.render_user_text("q").startswith("### User:\n\n")