mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-06 17:38:35 +00:00
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.
This commit is contained in:
parent
ca17292768
commit
b5db9fcd52
@ -754,24 +754,34 @@ def _run_gateway(
|
|||||||
target_channel = channels.channels.get(channel_name)
|
target_channel = channels.channels.get(channel_name)
|
||||||
except NameError:
|
except NameError:
|
||||||
target_channel = None
|
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_base_id = None
|
||||||
stream_segment = 0
|
stream_segment = 0
|
||||||
|
stream_had_delta = False
|
||||||
|
stream_events: list[OutboundMessage] = []
|
||||||
|
|
||||||
def _current_stream_id() -> str:
|
def _current_stream_id() -> str:
|
||||||
return f"{stream_base_id}:{stream_segment}"
|
return f"{stream_base_id}:{stream_segment}"
|
||||||
|
|
||||||
async def _on_stream(delta: str) -> None:
|
async def _on_stream(delta: str) -> None:
|
||||||
|
nonlocal stream_had_delta
|
||||||
meta = dict(job.payload.channel_meta)
|
meta = dict(job.payload.channel_meta)
|
||||||
meta["_stream_delta"] = True
|
meta["_stream_delta"] = True
|
||||||
meta["_stream_id"] = _current_stream_id()
|
meta["_stream_id"] = _current_stream_id()
|
||||||
await bus.publish_outbound(OutboundMessage(
|
stream_events.append(OutboundMessage(
|
||||||
channel=channel_name,
|
channel=channel_name,
|
||||||
chat_id=chat_id,
|
chat_id=chat_id,
|
||||||
content=delta,
|
content=delta,
|
||||||
metadata=meta,
|
metadata=meta,
|
||||||
))
|
))
|
||||||
|
if delta:
|
||||||
|
stream_had_delta = True
|
||||||
|
|
||||||
async def _on_stream_end(*, resuming: bool = False) -> None:
|
async def _on_stream_end(*, resuming: bool = False) -> None:
|
||||||
nonlocal stream_segment
|
nonlocal stream_segment
|
||||||
@ -779,7 +789,7 @@ def _run_gateway(
|
|||||||
meta["_stream_end"] = True
|
meta["_stream_end"] = True
|
||||||
meta["_resuming"] = resuming
|
meta["_resuming"] = resuming
|
||||||
meta["_stream_id"] = _current_stream_id()
|
meta["_stream_id"] = _current_stream_id()
|
||||||
await bus.publish_outbound(OutboundMessage(
|
stream_events.append(OutboundMessage(
|
||||||
channel=channel_name,
|
channel=channel_name,
|
||||||
chat_id=chat_id,
|
chat_id=chat_id,
|
||||||
content="",
|
content="",
|
||||||
@ -790,6 +800,20 @@ def _run_gateway(
|
|||||||
if wants_stream:
|
if wants_stream:
|
||||||
stream_base_id = f"cron:{job.id}:{time.time_ns()}"
|
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:
|
try:
|
||||||
resp = await agent.process_direct(
|
resp = await agent.process_direct(
|
||||||
reminder_note,
|
reminder_note,
|
||||||
@ -809,22 +833,18 @@ def _run_gateway(
|
|||||||
response = resp.content if resp else ""
|
response = resp.content if resp else ""
|
||||||
|
|
||||||
if job.payload.deliver and isinstance(message_tool, MessageTool) and message_tool._sent_in_turn:
|
if job.payload.deliver and isinstance(message_tool, MessageTool) and message_tool._sent_in_turn:
|
||||||
if wants_stream:
|
await _publish_turn_end_if_needed()
|
||||||
await bus.publish_outbound(OutboundMessage(
|
|
||||||
channel=channel_name,
|
|
||||||
chat_id=chat_id,
|
|
||||||
content="",
|
|
||||||
metadata={**job.payload.channel_meta, "_turn_end": True},
|
|
||||||
))
|
|
||||||
return response
|
return response
|
||||||
|
|
||||||
|
delivered = False
|
||||||
if job.payload.deliver and job.payload.to and response:
|
if job.payload.deliver and job.payload.to and response:
|
||||||
should_notify = await evaluate_response(
|
should_notify = await evaluate_response(
|
||||||
response, reminder_note, agent.provider, agent.model,
|
response, reminder_note, agent.provider, agent.model,
|
||||||
)
|
)
|
||||||
if should_notify:
|
if should_notify:
|
||||||
meta = dict(job.payload.channel_meta)
|
meta = dict(job.payload.channel_meta)
|
||||||
if wants_stream:
|
if wants_stream and stream_had_delta:
|
||||||
|
await _publish_buffered_stream()
|
||||||
meta["_streamed"] = True
|
meta["_streamed"] = True
|
||||||
await _deliver_to_channel(
|
await _deliver_to_channel(
|
||||||
OutboundMessage(
|
OutboundMessage(
|
||||||
@ -836,13 +856,9 @@ def _run_gateway(
|
|||||||
record=True,
|
record=True,
|
||||||
session_key=job.payload.session_key,
|
session_key=job.payload.session_key,
|
||||||
)
|
)
|
||||||
if wants_stream:
|
delivered = True
|
||||||
await bus.publish_outbound(OutboundMessage(
|
if delivered:
|
||||||
channel=channel_name,
|
await _publish_turn_end_if_needed()
|
||||||
chat_id=chat_id,
|
|
||||||
content="",
|
|
||||||
metadata={**job.payload.channel_meta, "_turn_end": True},
|
|
||||||
))
|
|
||||||
return response
|
return response
|
||||||
|
|
||||||
cron.on_job = on_cron_job
|
cron.on_job = on_cron_job
|
||||||
|
|||||||
@ -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.bus.queue.MessageBus", lambda: bus)
|
||||||
monkeypatch.setattr("nanobot.session.manager.SessionManager", lambda _workspace: object())
|
monkeypatch.setattr("nanobot.session.manager.SessionManager", lambda _workspace: object())
|
||||||
|
|
||||||
|
async def _always_notify(*_args, **_kwargs) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
class _FakeStreamingChannel:
|
class _FakeStreamingChannel:
|
||||||
supports_streaming = True
|
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.cron.service.CronService", _FakeCron)
|
||||||
monkeypatch.setattr("nanobot.cli.commands.AgentLoop", _FakeAgentLoop)
|
monkeypatch.setattr("nanobot.cli.commands.AgentLoop", _FakeAgentLoop)
|
||||||
monkeypatch.setattr("nanobot.channels.manager.ChannelManager", _FakeChannelManager)
|
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)])
|
result = runner.invoke(app, ["gateway", "--config", str(config_file)])
|
||||||
assert result.exit_code == 0
|
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
|
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(
|
def test_gateway_workspace_override_does_not_migrate_legacy_cron(
|
||||||
monkeypatch, tmp_path: Path
|
monkeypatch, tmp_path: Path
|
||||||
) -> None:
|
) -> None:
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user