fix(heartbeat): fail closed on internal checks

This commit is contained in:
Xubin Ren 2026-05-31 01:07:04 +08:00
parent 0cc58a80a4
commit 362f9629e2
6 changed files with 232 additions and 19 deletions

View File

@ -83,6 +83,10 @@ class MessageTool(Tool, ContextAware):
"message_record_channel_delivery",
default=False,
)
self._suppress_delivery_var: ContextVar[bool] = ContextVar(
"message_suppress_delivery",
default=False,
)
@classmethod
def create(cls, ctx: Any) -> Tool:
@ -121,6 +125,14 @@ class MessageTool(Tool, ContextAware):
"""Restore previous proactive delivery recording state."""
self._record_channel_delivery_var.reset(token)
def set_suppress_delivery(self, active: bool):
"""Temporarily suppress real channel delivery for internal checks."""
return self._suppress_delivery_var.set(active)
def reset_suppress_delivery(self, token) -> None:
"""Restore previous channel delivery suppression state."""
self._suppress_delivery_var.reset(token)
@property
def _sent_in_turn(self) -> bool:
return self._sent_in_turn_var.get()
@ -217,6 +229,9 @@ class MessageTool(Tool, ContextAware):
if not channel or not chat_id:
return "Error: No target channel/chat specified"
if self._suppress_delivery_var.get():
return "Message suppressed during internal check"
if not self._send_callback:
return "Error: Message sending not configured"

View File

@ -100,8 +100,8 @@ _HEARTBEAT_PREAMBLE = (
"[Your response will be delivered directly to the user's messaging app. "
"Output ONLY the final user-facing message. Never reference internal "
"files (HEARTBEAT.md, AWARENESS.md, etc.), your instructions, or your "
"decision process. If nothing needs reporting, respond with just "
"'All clear.' and nothing else.]\n\n"
"decision process. If nothing needs reporting, respond with a brief "
"no-op status and nothing else.]\n\n"
)
@ -991,13 +991,21 @@ def _run_gateway(
+ f"Review the following HEARTBEAT.md and report any active tasks:\n\n{content}"
)
resp = await agent.process_direct(
prompt,
session_key="heartbeat",
channel=channel,
chat_id=chat_id,
on_progress=_silent,
)
message_suppress_token = None
if isinstance(message_tool, MessageTool):
message_suppress_token = message_tool.set_suppress_delivery(True)
try:
resp = await agent.process_direct(
prompt,
session_key="heartbeat",
channel=channel,
chat_id=chat_id,
on_progress=_silent,
)
finally:
if isinstance(message_tool, MessageTool) and message_suppress_token is not None:
message_tool.reset_suppress_delivery(message_suppress_token)
response = resp.content if resp else ""
# Keep a small tail of heartbeat history so the loop stays bounded.
@ -1009,7 +1017,7 @@ def _run_gateway(
return None
should_notify = await evaluate_response(
response, prompt, agent.provider, agent.model,
response, prompt, agent.provider, agent.model, default_notify=False,
)
if should_notify:
logger.info("Heartbeat: completed, delivering response")

View File

@ -44,12 +44,15 @@ async def evaluate_response(
task_context: str,
provider: LLMProvider,
model: str,
*,
default_notify: bool = True,
) -> bool:
"""Decide whether a background-task result should be delivered to the user.
Uses a lightweight tool-call LLM request (same pattern as heartbeat
``_decide()``). Falls back to ``True`` (notify) on any failure so
that important messages are never silently dropped.
Uses a lightweight tool-call LLM request. ``default_notify`` controls
the fallback path when the evaluator cannot produce a valid decision:
user-scheduled reminders stay fail-open, while internal checks such as
heartbeat can fail closed.
"""
try:
llm_response = await provider.chat_with_retry(
@ -71,19 +74,23 @@ async def evaluate_response(
if not llm_response.should_execute_tools:
if llm_response.has_tool_calls:
logger.warning(
"evaluate_response: ignoring tool calls under finish_reason='{}', defaulting to notify",
"evaluate_response: ignoring tool calls under finish_reason='{}', defaulting to notify={}",
llm_response.finish_reason,
default_notify,
)
else:
logger.warning("evaluate_response: no tool call returned, defaulting to notify")
return True
logger.warning(
"evaluate_response: no tool call returned, defaulting to notify={}",
default_notify,
)
return default_notify
args = llm_response.tool_calls[0].arguments
should_notify = args.get("should_notify", True)
should_notify = args.get("should_notify", default_notify)
reason = args.get("reason", "")
logger.info("evaluate_response: should_notify={}, reason={}", should_notify, reason)
return bool(should_notify)
except Exception:
logger.exception("evaluate_response failed, defaulting to notify")
return True
logger.exception("evaluate_response failed, defaulting to notify={}", default_notify)
return default_notify

View File

@ -56,8 +56,38 @@ async def test_fallback_on_error() -> None:
assert result is True
@pytest.mark.asyncio
async def test_fallback_can_fail_closed() -> None:
class FailingProvider(DummyProvider):
async def chat(self, *args, **kwargs) -> LLMResponse:
raise RuntimeError("provider down")
provider = FailingProvider([])
result = await evaluate_response(
"some response",
"some task",
provider,
"m",
default_notify=False,
)
assert result is False
@pytest.mark.asyncio
async def test_no_tool_call_fallback() -> None:
provider = DummyProvider([LLMResponse(content="I think you should notify", tool_calls=[])])
result = await evaluate_response("some response", "some task", provider, "m")
assert result is True
@pytest.mark.asyncio
async def test_no_tool_call_can_fail_closed() -> None:
provider = DummyProvider([LLMResponse(content="I think you should notify", tool_calls=[])])
result = await evaluate_response(
"some response",
"some task",
provider,
"m",
default_notify=False,
)
assert result is False

View File

@ -1394,6 +1394,138 @@ def test_gateway_cron_job_suppresses_intermediate_progress(
bus.publish_outbound.assert_not_awaited()
def test_gateway_heartbeat_fails_closed_and_suppresses_message_tool(
monkeypatch, tmp_path: Path
) -> None:
"""Heartbeat only delivers after an explicit positive evaluation, and
internal checks cannot bypass that gate with the proactive message tool."""
from nanobot.agent.tools.message import MessageTool
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")
config.workspace_path.mkdir(parents=True)
(config.workspace_path / "HEARTBEAT.md").write_text(
"Check whether anything needs attention.",
encoding="utf-8",
)
bus = MagicMock()
bus.publish_outbound = AsyncMock()
seen: dict[str, object] = {}
class _FakeSession:
def retain_recent_legal_suffix(self, _keep: int) -> None:
seen["retained"] = True
class _FakeSessionManager:
def __init__(self, _workspace: Path) -> None:
self.session = _FakeSession()
def list_sessions(self) -> list[dict[str, object]]:
return [{"key": "lark:chat-1", "updated_at": "2026-05-30T00:00:00"}]
def get_or_create(self, key: str) -> _FakeSession:
seen["session_key"] = key
return self.session
def save(self, session: _FakeSession) -> None:
seen["saved"] = session
class _FakeCron:
def __init__(self, _store_path: Path) -> None:
self.on_job = None
seen["cron"] = self
def status(self) -> dict[str, int]:
return {"jobs": 0}
def register_system_job(self, job: CronJob) -> CronJob:
if job.name == "heartbeat":
seen["heartbeat_job"] = job
raise _StopGatewayError("stop")
return job
class _FakeDream:
model = None
max_batch_size = 0
max_iterations = 0
annotate_line_ages = False
async def run(self) -> None:
return None
class _FakeAgentLoop:
@classmethod
def from_config(cls, config, bus=None, **extra):
return cls(bus=bus, **extra)
def __init__(self, bus=None, **kwargs) -> None:
self.model = "test-model"
self.provider = object()
self.sessions = kwargs["session_manager"]
self.dream = _FakeDream()
self.tools = {
"message": MessageTool(send_callback=bus.publish_outbound),
}
async def process_direct(self, *_args, **_kwargs):
result = await self.tools["message"].execute(
content="All clear.",
channel="lark",
chat_id="chat-1",
)
seen["message_tool_result"] = result
return OutboundMessage(
channel="lark",
chat_id="chat-1",
content="All clear.",
)
async def close_mcp(self) -> None:
return None
def stop(self) -> None:
return None
class _FakeChannels:
enabled_channels = ["lark"]
async def _capture_evaluate(*_args, **kwargs) -> bool:
seen["default_notify"] = kwargs.get("default_notify")
return False
_patch_cli_command_runtime(
monkeypatch,
config,
message_bus=lambda: bus,
session_manager=_FakeSessionManager,
cron_service=_FakeCron,
)
monkeypatch.setattr("nanobot.cli.commands.AgentLoop", _FakeAgentLoop)
monkeypatch.setattr(
"nanobot.channels.manager.ChannelManager",
lambda *_args, **_kwargs: _FakeChannels(),
)
monkeypatch.setattr("nanobot.cli.commands.evaluate_response", _capture_evaluate)
result = runner.invoke(app, ["gateway", "--config", str(config_file)])
assert isinstance(result.exception, _StopGatewayError)
cron = seen["cron"]
response = asyncio.run(cron.on_job(seen["heartbeat_job"]))
assert response == "All clear."
assert seen["message_tool_result"] == "Message suppressed during internal check"
assert seen["default_notify"] is False
assert seen["session_key"] == "heartbeat"
assert seen["retained"] is True
bus.publish_outbound.assert_not_awaited()
def test_gateway_workspace_override_does_not_migrate_legacy_cron(
monkeypatch, tmp_path: Path
) -> None:

View File

@ -58,6 +58,27 @@ async def test_message_tool_marks_channel_delivery_only_when_enabled() -> None:
assert sent[1].metadata == {"_record_channel_delivery": True}
@pytest.mark.asyncio
async def test_message_tool_can_suppress_delivery_for_internal_checks() -> None:
sent: list[OutboundMessage] = []
async def _send(msg: OutboundMessage) -> None:
sent.append(msg)
tool = MessageTool(send_callback=_send)
token = tool.set_suppress_delivery(True)
try:
result = await tool.execute(content="All clear.", channel="lark", chat_id="chat-1")
finally:
tool.reset_suppress_delivery(token)
assert result == "Message suppressed during internal check"
assert sent == []
await tool.execute(content="real update", channel="lark", chat_id="chat-1")
assert [msg.content for msg in sent] == ["real update"]
@pytest.mark.asyncio
async def test_message_tool_records_media_deliveries() -> None:
sent: list[OutboundMessage] = []