From b5db9fcd52de67ea181bfb2411328e55423e6b92 Mon Sep 17 00:00:00 2001 From: chengyongru Date: Wed, 3 Jun 2026 16:49:30 +0800 Subject: [PATCH] fix(cron): gate buffered streaming delivery maintainer edit: buffer cron stream chunks until evaluate_response approves notification, so streaming channels do not leak suppressed reminders. Limit cron turn_end markers to delivered WebSocket messages. --- nanobot/cli/commands.py | 52 ++++++++++----- tests/cli/test_commands.py | 132 +++++++++++++++++++++++++++++++++++++ 2 files changed, 166 insertions(+), 18 deletions(-) diff --git a/nanobot/cli/commands.py b/nanobot/cli/commands.py index eac065994..b06149898 100644 --- a/nanobot/cli/commands.py +++ b/nanobot/cli/commands.py @@ -754,24 +754,34 @@ def _run_gateway( target_channel = channels.channels.get(channel_name) except NameError: target_channel = None - wants_stream = target_channel is not None and target_channel.supports_streaming + wants_stream = bool( + job.payload.deliver + and job.payload.to + and target_channel is not None + and target_channel.supports_streaming + ) stream_base_id = None stream_segment = 0 + stream_had_delta = False + stream_events: list[OutboundMessage] = [] def _current_stream_id() -> str: return f"{stream_base_id}:{stream_segment}" async def _on_stream(delta: str) -> None: + nonlocal stream_had_delta meta = dict(job.payload.channel_meta) meta["_stream_delta"] = True meta["_stream_id"] = _current_stream_id() - await bus.publish_outbound(OutboundMessage( + stream_events.append(OutboundMessage( channel=channel_name, chat_id=chat_id, content=delta, metadata=meta, )) + if delta: + stream_had_delta = True async def _on_stream_end(*, resuming: bool = False) -> None: nonlocal stream_segment @@ -779,7 +789,7 @@ def _run_gateway( meta["_stream_end"] = True meta["_resuming"] = resuming meta["_stream_id"] = _current_stream_id() - await bus.publish_outbound(OutboundMessage( + stream_events.append(OutboundMessage( channel=channel_name, chat_id=chat_id, content="", @@ -790,6 +800,20 @@ def _run_gateway( if wants_stream: stream_base_id = f"cron:{job.id}:{time.time_ns()}" + async def _publish_buffered_stream() -> None: + for event in stream_events: + await bus.publish_outbound(event) + + async def _publish_turn_end_if_needed() -> None: + if channel_name != "websocket" or not job.payload.to: + return + await bus.publish_outbound(OutboundMessage( + channel=channel_name, + chat_id=chat_id, + content="", + metadata={**job.payload.channel_meta, "_turn_end": True}, + )) + try: resp = await agent.process_direct( reminder_note, @@ -809,22 +833,18 @@ def _run_gateway( response = resp.content if resp else "" if job.payload.deliver and isinstance(message_tool, MessageTool) and message_tool._sent_in_turn: - if wants_stream: - await bus.publish_outbound(OutboundMessage( - channel=channel_name, - chat_id=chat_id, - content="", - metadata={**job.payload.channel_meta, "_turn_end": True}, - )) + await _publish_turn_end_if_needed() return response + delivered = False if job.payload.deliver and job.payload.to and response: should_notify = await evaluate_response( response, reminder_note, agent.provider, agent.model, ) if should_notify: meta = dict(job.payload.channel_meta) - if wants_stream: + if wants_stream and stream_had_delta: + await _publish_buffered_stream() meta["_streamed"] = True await _deliver_to_channel( OutboundMessage( @@ -836,13 +856,9 @@ def _run_gateway( record=True, session_key=job.payload.session_key, ) - if wants_stream: - await bus.publish_outbound(OutboundMessage( - channel=channel_name, - chat_id=chat_id, - content="", - metadata={**job.payload.channel_meta, "_turn_end": True}, - )) + delivered = True + if delivered: + await _publish_turn_end_if_needed() return response cron.on_job = on_cron_job diff --git a/tests/cli/test_commands.py b/tests/cli/test_commands.py index b1f4ae81b..aa8f4232f 100644 --- a/tests/cli/test_commands.py +++ b/tests/cli/test_commands.py @@ -1369,6 +1369,9 @@ def test_gateway_cron_job_streams_when_channel_supports_it( monkeypatch.setattr("nanobot.bus.queue.MessageBus", lambda: bus) monkeypatch.setattr("nanobot.session.manager.SessionManager", lambda _workspace: object()) + async def _always_notify(*_args, **_kwargs) -> bool: + return True + class _FakeStreamingChannel: supports_streaming = True @@ -1434,6 +1437,10 @@ def test_gateway_cron_job_streams_when_channel_supports_it( monkeypatch.setattr("nanobot.cron.service.CronService", _FakeCron) monkeypatch.setattr("nanobot.cli.commands.AgentLoop", _FakeAgentLoop) monkeypatch.setattr("nanobot.channels.manager.ChannelManager", _FakeChannelManager) + monkeypatch.setattr( + "nanobot.utils.evaluator.evaluate_response", + _always_notify, + ) result = runner.invoke(app, ["gateway", "--config", str(config_file)]) assert result.exit_code == 0 @@ -1473,6 +1480,131 @@ def test_gateway_cron_job_streams_when_channel_supports_it( assert calls[4].args[0].metadata.get("_turn_end") is True +def test_gateway_cron_job_streaming_respects_evaluator_rejection( + monkeypatch, tmp_path: Path +) -> None: + """Streaming cron output must not reach the channel before evaluator approval.""" + config_file = tmp_path / "instance" / "config.json" + config_file.parent.mkdir(parents=True) + config_file.write_text("{}") + + config = Config() + config.agents.defaults.workspace = str(tmp_path / "config-workspace") + bus = MagicMock() + bus.publish_outbound = AsyncMock() + seen: dict[str, object] = {} + + monkeypatch.setattr("nanobot.config.loader.set_config_path", lambda _path: None) + monkeypatch.setattr("nanobot.config.loader.load_config", lambda _path=None: config) + monkeypatch.setattr("nanobot.cli.commands.sync_workspace_templates", lambda _path: None) + monkeypatch.setattr("nanobot.providers.factory.make_provider", lambda _config: _fake_provider()) + monkeypatch.setattr( + "nanobot.providers.factory.build_provider_snapshot", + lambda _config: _test_provider_snapshot(object(), _config), + ) + monkeypatch.setattr( + "nanobot.providers.factory.load_provider_snapshot", + lambda _config_path=None: _test_provider_snapshot(object(), config), + ) + monkeypatch.setattr("nanobot.bus.queue.MessageBus", lambda: bus) + monkeypatch.setattr("nanobot.session.manager.SessionManager", lambda _workspace: object()) + + async def _reject(*_args, **_kwargs) -> bool: + seen["evaluated"] = True + return False + + class _FakeStreamingChannel: + supports_streaming = True + + class _FakeChannelManager: + def __init__(self, *_args, **_kwargs) -> None: + self.channels = {"websocket": _FakeStreamingChannel()} + self.enabled_channels = ["websocket"] + + async def start_all(self): + pass + + async def stop_all(self): + pass + + class _FakeCron: + def __init__(self, _store_path: Path) -> None: + self.on_job = None + seen["cron"] = self + + def status(self): + return {"enabled": True, "jobs": 0, "next_wake_at_ms": None} + + def register_system_job(self, job): + pass + + def stop(self): + pass + + class _FakeAgentLoop: + @classmethod + def from_config(cls, config, bus=None, **extra): + return cls(**extra) + def __init__(self, *args, **kwargs) -> None: + self.model = "test-model" + self.provider = object() + self.tools = {} + self.dream = MagicMock() + self.sessions = MagicMock() + + async def process_direct(self, *_args, on_stream=None, on_stream_end=None, **_kwargs): + seen["on_stream"] = on_stream + seen["on_stream_end"] = on_stream_end + if on_stream: + await on_stream("This should not leak") + if on_stream_end: + await on_stream_end(resuming=False) + return OutboundMessage( + channel="websocket", + chat_id="user-1", + content="This should not leak", + ) + + async def close_mcp(self) -> None: + return None + + async def run(self) -> None: + return None + + def stop(self) -> None: + return None + + monkeypatch.setattr("nanobot.cron.service.CronService", _FakeCron) + monkeypatch.setattr("nanobot.cli.commands.AgentLoop", _FakeAgentLoop) + monkeypatch.setattr("nanobot.channels.manager.ChannelManager", _FakeChannelManager) + monkeypatch.setattr( + "nanobot.utils.evaluator.evaluate_response", + _reject, + ) + + result = runner.invoke(app, ["gateway", "--config", str(config_file)]) + assert result.exit_code == 0 + + cron = seen["cron"] + job = CronJob( + id="cron-stream-rejected-test", + name="test-stream-rejected", + payload=CronPayload( + message="Say something optional.", + deliver=True, + channel="websocket", + to="user-1", + ), + ) + response = asyncio.run(cron.on_job(job)) + + assert response == "This should not leak" + assert seen["on_stream"] is not None + assert seen["on_stream_end"] is not None + assert seen["evaluated"] is True + bus.publish_outbound.assert_not_awaited() + + def test_gateway_workspace_override_does_not_migrate_legacy_cron( monkeypatch, tmp_path: Path ) -> None: