mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-09-04 02:01:48 +03:00
refactor(agent): make checkpoint recovery ownership explicit
This commit is contained in:
@@ -18,6 +18,7 @@ import pytest
|
||||
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.session.recovery import RUNTIME_CHECKPOINT_KEY
|
||||
|
||||
|
||||
def _make_provider():
|
||||
@@ -41,51 +42,6 @@ def _make_loop(tmp_path: Path) -> AgentLoop:
|
||||
return AgentLoop(bus=bus, provider=provider, workspace=tmp_path)
|
||||
|
||||
|
||||
class TestStopPreservesContext:
|
||||
"""Verify that /stop restores partial context via checkpoint."""
|
||||
|
||||
def test_restore_checkpoint_method_exists(self, tmp_path):
|
||||
"""AgentLoop should have _restore_runtime_checkpoint."""
|
||||
loop = _make_loop(tmp_path)
|
||||
assert hasattr(loop, "_restore_runtime_checkpoint")
|
||||
|
||||
def test_checkpoint_key_constant(self, tmp_path):
|
||||
"""The runtime checkpoint key should be defined."""
|
||||
loop = _make_loop(tmp_path)
|
||||
assert loop._RUNTIME_CHECKPOINT_KEY == "runtime_checkpoint"
|
||||
|
||||
def test_cancel_dispatch_restores_checkpoint(self, tmp_path):
|
||||
"""When a task is cancelled, the checkpoint should be restored."""
|
||||
loop = _make_loop(tmp_path)
|
||||
session = MagicMock()
|
||||
session.metadata = {
|
||||
"runtime_checkpoint": {
|
||||
"phase": "awaiting_tools",
|
||||
"iteration": 0,
|
||||
"assistant_message": {
|
||||
"role": "assistant",
|
||||
"content": "Let me search for that.",
|
||||
"tool_calls": [{"id": "tc_1", "type": "function",
|
||||
"function": {"name": "web_search", "arguments": "{}"}}],
|
||||
},
|
||||
"completed_tool_results": [
|
||||
{"role": "tool", "tool_call_id": "tc_1",
|
||||
"content": "Search results: ..."},
|
||||
],
|
||||
"pending_tool_calls": [],
|
||||
}
|
||||
}
|
||||
session.messages = [
|
||||
{"role": "user", "content": "Search for something"},
|
||||
]
|
||||
loop.sessions.get_or_create.return_value = session
|
||||
|
||||
restored = loop._restore_runtime_checkpoint(session)
|
||||
assert restored is True
|
||||
assert len(session.messages) > 1
|
||||
assert "runtime_checkpoint" not in session.metadata
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatch_cancellation_restores_checkpoint():
|
||||
"""Regression for #2966: /stop interrupting _dispatch must materialize the
|
||||
@@ -93,9 +49,8 @@ async def test_dispatch_cancellation_restores_checkpoint():
|
||||
unwinds, so the next turn can see the partial work.
|
||||
|
||||
This exercises the real _dispatch path (locks, pending queues, the
|
||||
CancelledError handler) rather than poking _restore_runtime_checkpoint in
|
||||
isolation, so a future refactor that drops the cancel-time restore is
|
||||
caught by CI instead of silently regressing.
|
||||
CancelledError handler), so a future refactor that drops the cancel-time
|
||||
restore is caught by CI instead of silently regressing.
|
||||
"""
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.bus.queue import MessageBus
|
||||
@@ -112,7 +67,7 @@ async def test_dispatch_cancellation_restores_checkpoint():
|
||||
mock_subagent_manager.return_value.cancel_by_session = AsyncMock(return_value=0)
|
||||
loop = AgentLoop(bus=bus, provider=provider, workspace=workspace)
|
||||
|
||||
checkpoint_key = loop._RUNTIME_CHECKPOINT_KEY
|
||||
checkpoint_key = RUNTIME_CHECKPOINT_KEY
|
||||
session = SimpleNamespace(
|
||||
key="test:c1",
|
||||
metadata={
|
||||
@@ -168,7 +123,19 @@ async def test_dispatch_cancellation_keeps_checkpoint_for_gateway_shutdown(tmp_p
|
||||
"""Gateway shutdown preserves the checkpoint; an explicit stop restores it."""
|
||||
loop = _make_loop(tmp_path)
|
||||
loop.preserve_inflight_turns_on_shutdown()
|
||||
loop._restore_runtime_checkpoint = MagicMock() # type: ignore[method-assign]
|
||||
checkpoint_key = RUNTIME_CHECKPOINT_KEY
|
||||
checkpoint = {
|
||||
"phase": "final_response",
|
||||
"assistant_message": {"role": "assistant", "content": "finished"},
|
||||
"completed_tool_results": [],
|
||||
"pending_tool_calls": [],
|
||||
}
|
||||
session = SimpleNamespace(
|
||||
metadata={checkpoint_key: checkpoint},
|
||||
messages=[],
|
||||
provider_state=None,
|
||||
)
|
||||
loop.sessions.get_or_create.return_value = session
|
||||
|
||||
async def _cancel(*_args: object, **_kwargs: object) -> None:
|
||||
raise asyncio.CancelledError()
|
||||
@@ -182,4 +149,5 @@ async def test_dispatch_cancellation_keeps_checkpoint_for_gateway_shutdown(tmp_p
|
||||
InboundMessage(channel="test", sender_id="u1", chat_id="c1", content="work")
|
||||
)
|
||||
|
||||
loop._restore_runtime_checkpoint.assert_not_called()
|
||||
assert session.metadata[checkpoint_key] == checkpoint
|
||||
assert session.messages == []
|
||||
|
||||
Reference in New Issue
Block a user