refactor(bus): type outbound runtime events

This commit is contained in:
chengyongru
2026-07-01 20:17:00 +08:00
committed by Xubin Ren
parent f6d1dba32a
commit 5f4cfbcb16
45 changed files with 1206 additions and 741 deletions
+77 -43
View File
@@ -9,6 +9,15 @@ import pytest
import nanobot.agent.runner as runner_module
from nanobot.agent.loop import AgentLoop
from nanobot.bus.events import InboundMessage
from nanobot.bus.outbound_events import (
GoalStatusEvent,
ProgressEvent,
SessionUpdatedEvent,
StreamDeltaEvent,
StreamedResponseEvent,
StreamEndEvent,
TurnEndEvent,
)
from nanobot.bus.queue import MessageBus
from nanobot.providers.base import LLMResponse, ToolCallRequest
from nanobot.session.webui_turns import WebuiTurnCoordinator
@@ -260,25 +269,45 @@ class TestToolEventProgress:
)
await loop._dispatch(msg)
# Drain all outbound messages and find the one carrying _tool_events
# Drain all outbound messages and find the one carrying tool events.
outbound = []
while bus.outbound_size > 0:
outbound.append(await bus.consume_outbound())
tool_event_msgs = [m for m in outbound if m.metadata and m.metadata.get("_tool_events")]
assert tool_event_msgs, "expected at least one outbound message with _tool_events"
tool_event_msgs = [
m
for m in outbound
if isinstance(m.event, ProgressEvent) and m.event.tool_events
]
assert tool_event_msgs, "expected at least one outbound message with tool events"
start_msgs = [m for m in tool_event_msgs if m.metadata["_tool_events"][0]["phase"] == "start"]
finish_msgs = [m for m in tool_event_msgs if m.metadata["_tool_events"][0]["phase"] in ("end", "error")]
start_msgs = [
m
for m in tool_event_msgs
if isinstance(m.event, ProgressEvent)
and m.event.tool_events
and m.event.tool_events[0]["phase"] == "start"
]
finish_msgs = [
m
for m in tool_event_msgs
if isinstance(m.event, ProgressEvent)
and m.event.tool_events
and m.event.tool_events[0]["phase"] in ("end", "error")
]
assert start_msgs, "expected a start-phase tool event"
assert finish_msgs, "expected a finish-phase tool event"
start = start_msgs[0].metadata["_tool_events"][0]
assert isinstance(start_msgs[0].event, ProgressEvent)
assert start_msgs[0].event.tool_events is not None
start = start_msgs[0].event.tool_events[0]
assert start["name"] == "exec"
assert start["call_id"] == "tc1"
assert start["result"] is None
finish = finish_msgs[0].metadata["_tool_events"][0]
assert isinstance(finish_msgs[0].event, ProgressEvent)
assert finish_msgs[0].event.tool_events is not None
finish = finish_msgs[0].event.tool_events[0]
assert finish["phase"] == "end"
assert finish["result"] == "file.txt"
@@ -309,7 +338,8 @@ class TestToolEventProgress:
await invoke_file_edit_progress(progress, edit_events)
outbound = await bus.consume_outbound()
assert outbound.channel == "telegram"
assert outbound.metadata["_file_edit_events"] == edit_events
assert isinstance(outbound.event, ProgressEvent)
assert outbound.event.file_edit_events == edit_events
@pytest.mark.asyncio
async def test_goal_turn_keeps_live_file_edit_progress_for_webui(self, tmp_path: Path) -> None:
@@ -389,7 +419,8 @@ class TestToolEventProgress:
edit_events = [
event
for msg in outbound
for event in msg.metadata.get("_file_edit_events", [])
if isinstance(msg.event, ProgressEvent)
for event in msg.event.file_edit_events or []
]
assert any(
event["status"] == "editing"
@@ -433,8 +464,8 @@ class TestToolEventProgress:
outbound.append(await bus.consume_outbound())
assert [m.content for m in outbound] == ["Hello"]
assert not any(m.metadata.get("_progress") for m in outbound)
assert not any(m.metadata.get("_streamed") for m in outbound)
assert not any(isinstance(m.event, ProgressEvent) for m in outbound)
assert not any(isinstance(m.event, StreamedResponseEvent) for m in outbound)
provider.chat_stream_with_retry.assert_not_awaited()
provider.chat_with_retry.assert_awaited_once()
@@ -443,7 +474,7 @@ class TestToolEventProgress:
self,
tmp_path: Path,
) -> None:
"""Streaming channels still receive provider deltas through _stream_delta messages."""
"""Streaming channels still receive provider deltas through stream events."""
bus = MessageBus()
provider = MagicMock()
provider.supports_progress_deltas = True
@@ -473,21 +504,19 @@ class TestToolEventProgress:
while bus.outbound_size > 0:
outbound.append(await bus.consume_outbound())
deltas = [m for m in outbound if m.metadata.get("_stream_delta")]
stream_end = [m for m in outbound if m.metadata.get("_stream_end")]
deltas = [m for m in outbound if isinstance(m.event, StreamDeltaEvent)]
stream_end = [m for m in outbound if isinstance(m.event, StreamEndEvent)]
final = [
m for m in outbound
if not m.metadata.get("_stream_delta")
and not m.metadata.get("_stream_end")
and not m.metadata.get("_turn_end")
and not m.metadata.get("_goal_status")
if not isinstance(m.event, StreamDeltaEvent | StreamEndEvent)
and not isinstance(m.event, TurnEndEvent | GoalStatusEvent)
]
assert [m.content for m in deltas] == ["Hel", "lo"]
assert len(stream_end) == 1
assert final[-1].content == "Hello"
assert final[-1].metadata.get("_streamed") is True
turn_end_msgs = [m for m in outbound if m.metadata.get("_turn_end")]
assert isinstance(final[-1].event, StreamedResponseEvent)
turn_end_msgs = [m for m in outbound if isinstance(m.event, TurnEndEvent)]
assert len(turn_end_msgs) == 1
assert turn_end_msgs[0].content == ""
provider.chat_with_retry.assert_not_awaited()
@@ -528,23 +557,28 @@ class TestToolEventProgress:
while bus.outbound_size > 0:
outbound.append(await bus.consume_outbound())
deltas = [m for m in outbound if m.metadata.get("_stream_delta")]
stream_end = [m for m in outbound if m.metadata.get("_stream_end")]
deltas = [m for m in outbound if isinstance(m.event, StreamDeltaEvent)]
stream_end = [m for m in outbound if isinstance(m.event, StreamEndEvent)]
final = [
m for m in outbound
if not m.metadata.get("_stream_delta")
and not m.metadata.get("_stream_end")
and not m.metadata.get("_turn_end")
and not m.metadata.get("_goal_status")
if not isinstance(m.event, StreamDeltaEvent | StreamEndEvent)
and not isinstance(m.event, TurnEndEvent | GoalStatusEvent)
]
assert [m.content for m in deltas] == ["partial", "full retry response"]
assert [m.metadata.get("_resuming") for m in stream_end] == [True, False]
assert deltas[0].metadata.get("_stream_id") == stream_end[0].metadata.get("_stream_id")
assert deltas[1].metadata.get("_stream_id") == stream_end[1].metadata.get("_stream_id")
assert deltas[0].metadata.get("_stream_id") != deltas[1].metadata.get("_stream_id")
assert [m.event.resuming for m in stream_end if isinstance(m.event, StreamEndEvent)] == [
True,
False,
]
assert isinstance(deltas[0].event, StreamDeltaEvent)
assert isinstance(deltas[1].event, StreamDeltaEvent)
assert isinstance(stream_end[0].event, StreamEndEvent)
assert isinstance(stream_end[1].event, StreamEndEvent)
assert deltas[0].event.stream_id == stream_end[0].event.stream_id
assert deltas[1].event.stream_id == stream_end[1].event.stream_id
assert deltas[0].event.stream_id != deltas[1].event.stream_id
assert final[-1].content == "full retry response"
assert final[-1].metadata.get("_streamed") is True
assert isinstance(final[-1].event, StreamedResponseEvent)
provider.chat_with_retry.assert_not_awaited()
@pytest.mark.asyncio
@@ -623,9 +657,9 @@ class TestToolEventProgress:
done_msgs = [m for m in outbound if m.content == "Done"]
assert len(done_msgs) == 1
assert not done_msgs[0].metadata.get("_turn_end")
assert not isinstance(done_msgs[0].event, TurnEndEvent)
turn_end_msgs = [m for m in outbound if m.metadata.get("_turn_end")]
turn_end_msgs = [m for m in outbound if isinstance(m.event, TurnEndEvent)]
assert len(turn_end_msgs) == 1
assert turn_end_msgs[0].content == ""
assert turn_end_msgs[0].chat_id == "chat1"
@@ -659,14 +693,14 @@ class TestToolEventProgress:
outbound.append(await bus.consume_outbound())
error_msgs = [m for m in outbound if m.content == "Sorry, I encountered an error."]
turn_end_msgs = [m for m in outbound if m.metadata.get("_turn_end")]
statuses = [m for m in outbound if m.metadata.get("_goal_status")]
turn_end_msgs = [m for m in outbound if isinstance(m.event, TurnEndEvent)]
statuses = [m for m in outbound if isinstance(m.event, GoalStatusEvent)]
assert len(error_msgs) == 1
assert len(turn_end_msgs) == 1
assert turn_end_msgs[0].content == ""
assert turn_end_msgs[0].chat_id == "chat1"
assert [m.metadata["goal_status"] for m in statuses] == ["idle"]
assert [m.event.status for m in statuses if isinstance(m.event, GoalStatusEvent)] == ["idle"]
assert outbound.index(error_msgs[0]) < outbound.index(turn_end_msgs[0])
assert outbound.index(turn_end_msgs[0]) < outbound.index(statuses[-1])
@@ -705,27 +739,27 @@ class TestToolEventProgress:
outbound: list = []
for _ in range(12):
outbound.append(await asyncio.wait_for(bus.consume_outbound(), timeout=0.5))
if outbound[-1].metadata.get("_turn_end"):
if isinstance(outbound[-1].event, TurnEndEvent):
break
else:
raise AssertionError("_turn_end message not found")
raise AssertionError("turn-end event not found")
done_with_body = [m for m in outbound if m.content == "Done"]
assert len(done_with_body) == 1
assert outbound[-1].metadata.get("_turn_end") is True
assert isinstance(outbound[-1].event, TurnEndEvent)
await asyncio.wait_for(title_started.wait(), timeout=0.5)
release_title.set()
session_updated = None
for _ in range(10):
candidate = await asyncio.wait_for(bus.consume_outbound(), timeout=0.5)
if (candidate.metadata or {}).get("_session_updated"):
if isinstance(candidate.event, SessionUpdatedEvent):
session_updated = candidate
break
assert session_updated is not None
assert (session_updated.metadata or {}).get("_session_updated") is True
assert (session_updated.metadata or {}).get("_session_update_scope") == "metadata"
assert isinstance(session_updated.event, SessionUpdatedEvent)
assert session_updated.event.scope == "metadata"
assert provider.chat_with_retry.await_count == 2
@pytest.mark.asyncio
@@ -837,4 +871,4 @@ class TestToolEventProgress:
assert len(outbound) == 1
assert outbound[0].content == "Done"
assert (outbound[0].metadata or {}).get("_turn_end") is not True
assert not isinstance(outbound[0].event, TurnEndEvent)