From e6baecafcd78a9eacaa4fe4c62528332c3b26c97 Mon Sep 17 00:00:00 2001 From: chengyongru Date: Fri, 24 Jul 2026 10:20:43 +0800 Subject: [PATCH] fix(agent): close length recovery lifecycle gaps --- nanobot/agent/loop.py | 15 +++- nanobot/agent/runner.py | 22 +++-- nanobot/agent/turn_delivery.py | 8 ++ nanobot/channels/mattermost/runtime.py | 6 +- .../tests/test_mattermost_channel.py | 27 ++++++ tests/agent/test_loop_progress.py | 90 +++++++++++++++++++ tests/agent/test_runner_core.py | 34 +++++++ 7 files changed, 195 insertions(+), 7 deletions(-) diff --git a/nanobot/agent/loop.py b/nanobot/agent/loop.py index e972b6be..eab6f92e 100644 --- a/nanobot/agent/loop.py +++ b/nanobot/agent/loop.py @@ -1070,7 +1070,12 @@ class AgentLoop: # Push final content through stream so streaming channels (e.g. Feishu) # update the card instead of leaving it empty. if on_stream and on_stream_end and should_stream: - await on_stream(result.final_content or "") + stream_content = ( + result.pending_stream_content + if result.pending_stream_content is not None + else result.final_content or "" + ) + await on_stream(stream_content) await on_stream_end(resuming=False) elif result.stop_reason == "error": logger.error("LLM returned error: {}", (result.final_content or "")[:200]) @@ -1217,6 +1222,14 @@ class AgentLoop: for _, coordinator in self._automation_turn_coordinators: coordinator.complete(msg, error=asyncio.CancelledError()) logger.info("Task cancelled for session {}", session_key) + try: + await delivery.abort_stream() + except Exception: + logger.debug( + "Could not close stream for cancelled session {}", + session_key, + exc_info=True, + ) # Preserve partial context from the interrupted turn so # the user does not lose tool results and assistant # messages accumulated before /stop. The checkpoint was diff --git a/nanobot/agent/runner.py b/nanobot/agent/runner.py index 3d217822..acc54abb 100644 --- a/nanobot/agent/runner.py +++ b/nanobot/agent/runner.py @@ -113,6 +113,8 @@ class AgentRunResult: error: str | None = None tool_events: list[dict[str, str]] = field(default_factory=list) had_injections: bool = False + # Terminal tail to emit when the preceding final-content prefix was already streamed. + pending_stream_content: str | None = None class AgentRunner: @@ -398,6 +400,7 @@ class AgentRunner: had_injections = False injection_cycles = 0 compacted_tool_call_ids: set[str] = set() + pending_stream_content: str | None = None governance_config = ContextGovernanceConfig( provider=spec.runtime.provider, model=spec.runtime.model, @@ -718,17 +721,25 @@ class AgentRunner: ) if drained_after_max_iterations: had_injections = True - final_content = None + terminal_content = None if spec.finalize_on_max_iterations: - final_content = await self._try_finalize_after_max_iterations( + terminal_content = await self._try_finalize_after_max_iterations( spec, hook, messages, usage, ) - if final_content is None: - final_content = self._max_iterations_fallback(spec) - self._append_final_message(messages, final_content) + if terminal_content is None: + terminal_content = self._max_iterations_fallback(spec) + if length_recovery_parts: + terminal_tail = f"\n\n{terminal_content.lstrip()}" + final_content = ( + "".join(length_recovery_parts).rstrip() + terminal_tail + ).strip() + pending_stream_content = terminal_tail + else: + final_content = terminal_content + self._append_final_message(messages, terminal_content) return AgentRunResult( final_content=final_content, @@ -739,6 +750,7 @@ class AgentRunner: error=error, tool_events=tool_events, had_injections=had_injections, + pending_stream_content=pending_stream_content, ) def _build_request_kwargs( diff --git a/nanobot/agent/turn_delivery.py b/nanobot/agent/turn_delivery.py index a415551d..9f2aba0a 100644 --- a/nanobot/agent/turn_delivery.py +++ b/nanobot/agent/turn_delivery.py @@ -126,6 +126,7 @@ class TurnDelivery: lifecycle_message: InboundMessage = field(init=False) _stream_base_id: str | None = field(init=False, default=None) _stream_segment: int = field(init=False, default=0) + _stream_open: bool = field(init=False, default=False) def __post_init__(self) -> None: self.delivery_message = dataclasses.replace( @@ -284,6 +285,7 @@ class TurnDelivery: metadata=self.delivery_message.metadata, ) ) + self._stream_open = True async def _publish_stream_end( self, @@ -303,5 +305,11 @@ class TurnDelivery: metadata=self.delivery_message.metadata, ) ) + self._stream_open = merge_next if not merge_next: self._stream_segment += 1 + + async def abort_stream(self) -> None: + """Close an interrupted stream so stateful channels can release its buffer.""" + if self._stream_open: + await self._publish_stream_end() diff --git a/nanobot/channels/mattermost/runtime.py b/nanobot/channels/mattermost/runtime.py index d15bcd3f..473dac31 100644 --- a/nanobot/channels/mattermost/runtime.py +++ b/nanobot/channels/mattermost/runtime.py @@ -533,7 +533,11 @@ class MattermostChannel(BaseChannel): final += delta if resuming: - self._clear_stream_state(stream_id) + if merge_next: + self._stream_buffers[stream_id] = final + self._stream_committed[stream_id] = final + else: + self._clear_stream_state(stream_id) return if final and not meta.get("_progress"): diff --git a/nanobot/channels/mattermost/tests/test_mattermost_channel.py b/nanobot/channels/mattermost/tests/test_mattermost_channel.py index 3b27ccf1..f9afcdd1 100644 --- a/nanobot/channels/mattermost/tests/test_mattermost_channel.py +++ b/nanobot/channels/mattermost/tests/test_mattermost_channel.py @@ -582,6 +582,33 @@ async def test_stream_end_keyword_resuming_does_not_post_or_mark_done(): assert "s1" not in channel._stream_buffers +@pytest.mark.asyncio +async def test_stream_end_merge_next_preserves_buffer_until_final_end(): + channel, fake = _make_channel() + channel._self_id = "bot_id" + fake.set_post_response("/api/v4/posts", {"id": "stream_post_1"}) + await channel.send_delta("chan_1", "first ", stream_id="s1") + + await channel.send_delta( + "chan_1", + "boundary ", + stream_id="s1", + stream_end=True, + resuming=True, + merge_next=True, + ) + + assert channel._stream_buffers["s1"] == "first boundary " + + await channel.send_delta("chan_1", "second", stream_id="s1") + await channel.send_delta("chan_1", "", stream_id="s1", stream_end=True) + + posts = [call for call in fake.post_calls if call["path"] == "/api/v4/posts"] + assert len(posts) == 1 + assert posts[0]["json"]["message"] == "first boundary second" + assert "s1" not in channel._stream_buffers + + @pytest.mark.asyncio async def test_stream_end_failure_keeps_buffer_for_retry(): channel, fake = _make_channel() diff --git a/tests/agent/test_loop_progress.py b/tests/agent/test_loop_progress.py index aca35927..39ecb727 100644 --- a/tests/agent/test_loop_progress.py +++ b/tests/agent/test_loop_progress.py @@ -580,6 +580,96 @@ class TestToolEventProgress: assert [event.merge_next for event in endings] == [True, False] assert {event.stream_id for event in [*deltas, *endings]} == {deltas[0].stream_id} + @pytest.mark.asyncio + async def test_length_recovery_at_max_iterations_streams_only_missing_tail( + self, + tmp_path: Path, + ) -> None: + bus = MessageBus() + provider = MagicMock() + provider.supports_progress_deltas = True + provider.get_default_model.return_value = "test-model" + + async def chat_stream_with_retry(*, on_content_delta, **kwargs): + await on_content_delta("partial") + return LLMResponse(content="partial", finish_reason="length") + + provider.chat_stream_with_retry = chat_stream_with_retry + provider.chat_with_retry = AsyncMock( + return_value=LLMResponse(content="summary", finish_reason="stop") + ) + loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") + _attach_webui_runtime_events(loop, bus) + loop.max_iterations = 1 + loop.tools.get_definitions = MagicMock(return_value=[]) + loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign] + + await loop._dispatch(InboundMessage( + channel="websocket", + sender_id="u1", + chat_id="chat1", + content="give a long answer", + metadata={"_wants_stream": True}, + )) + + outbound = [] + while bus.outbound_size > 0: + outbound.append(await bus.consume_outbound()) + + deltas = [m.event for m in outbound if isinstance(m.event, StreamDeltaEvent)] + endings = [m.event for m in outbound if isinstance(m.event, StreamEndEvent)] + final = [m for m in outbound if isinstance(m.event, StreamedResponseEvent)] + + assert [event.content for event in deltas] == ["partial", "\n\nsummary"] + assert [event.merge_next for event in endings] == [True, False] + assert {event.stream_id for event in [*deltas, *endings]} == {deltas[0].stream_id} + assert [message.content for message in final] == ["partial\n\nsummary"] + + @pytest.mark.asyncio + async def test_cancelled_length_recovery_closes_merged_stream( + self, + tmp_path: Path, + ) -> None: + bus = MessageBus() + provider = MagicMock() + provider.get_default_model.return_value = "test-model" + loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") + + async def cancel_after_merge( + _msg: InboundMessage, + *, + on_stream, + on_stream_end, + **_kwargs, + ): + assert on_stream is not None + assert on_stream_end is not None + await on_stream("partial") + await on_stream_end(resuming=True, merge_next=True) + raise asyncio.CancelledError + + loop._process_message = cancel_after_merge # type: ignore[method-assign] + + with pytest.raises(asyncio.CancelledError): + await loop._dispatch(InboundMessage( + channel="websocket", + sender_id="u1", + chat_id="chat1", + content="give a long answer", + metadata={"_wants_stream": True}, + )) + + outbound = [] + while bus.outbound_size > 0: + outbound.append(await bus.consume_outbound()) + + endings = [m.event for m in outbound if isinstance(m.event, StreamEndEvent)] + assert [(event.resuming, event.merge_next) for event in endings] == [ + (True, True), + (False, False), + ] + assert endings[0].stream_id == endings[1].stream_id + @pytest.mark.asyncio async def test_non_streamed_finalization_is_delivered_as_regular_message( self, diff --git a/tests/agent/test_runner_core.py b/tests/agent/test_runner_core.py index 584df875..883fe24d 100644 --- a/tests/agent/test_runner_core.py +++ b/tests/agent/test_runner_core.py @@ -482,6 +482,40 @@ async def test_runner_length_recovery_returns_all_segments(): assert provider.chat_with_retry.await_count == 3 +@pytest.mark.asyncio +async def test_runner_length_recovery_preserves_prefix_at_max_iterations(): + """Budget exhaustion must not replace output already produced by recovery.""" + from nanobot.agent.runner import AgentRunner + + provider = MagicMock(spec=LLMProvider) + provider.chat_with_retry = AsyncMock( + return_value=LLMResponse(content="partial answer", finish_reason="length") + ) + tools = MagicMock() + tools.get_definitions.return_value = [] + + runner = AgentRunner() + result = await runner.run(make_run_spec( + provider, + initial_messages=[{"role": "user", "content": "give a long answer"}], + tools=tools, + model="test-model", + max_iterations=1, + max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, + finalize_on_max_iterations=False, + max_iterations_message="limit reached", + )) + + assert result.stop_reason == "max_iterations" + assert result.final_content == "partial answer\n\nlimit reached" + assert result.pending_stream_content == "\n\nlimit reached" + assert [ + message["content"] + for message in result.messages + if message.get("role") == "assistant" + ] == ["partial answer", "limit reached"] + + @pytest.mark.asyncio async def test_runner_length_recovery_does_not_leak_across_tool_calls(): """A recovered prefix belongs only to its contiguous response chain."""