fix(agent): close file edit activity on cancellation
This commit is contained in:
@@ -6,7 +6,12 @@ from collections.abc import Awaitable, Callable
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from nanobot.agent.hook import AgentHook, AgentHookContext, AgentTurnHookContext
|
from nanobot.agent.hook import (
|
||||||
|
AgentHook,
|
||||||
|
AgentHookContext,
|
||||||
|
AgentRunHookContext,
|
||||||
|
AgentTurnHookContext,
|
||||||
|
)
|
||||||
from nanobot.providers.base import ToolCallRequest
|
from nanobot.providers.base import ToolCallRequest
|
||||||
from nanobot.utils.file_edit_events import (
|
from nanobot.utils.file_edit_events import (
|
||||||
FileEditTracker,
|
FileEditTracker,
|
||||||
@@ -71,9 +76,11 @@ class FileEditActivityHook(AgentHook):
|
|||||||
params: Any,
|
params: Any,
|
||||||
result: Any,
|
result: Any,
|
||||||
) -> None:
|
) -> None:
|
||||||
trackers = self._trackers_by_call.pop(self._tool_call_key(tool_call), [])
|
key = self._tool_call_key(tool_call)
|
||||||
|
trackers = self._trackers_by_call.get(key, [])
|
||||||
if trackers:
|
if trackers:
|
||||||
await self._emit([build_file_edit_end_event(tracker) for tracker in trackers])
|
await self._emit([build_file_edit_end_event(tracker) for tracker in trackers])
|
||||||
|
self._trackers_by_call.pop(key, None)
|
||||||
|
|
||||||
async def on_execute_tool_error(
|
async def on_execute_tool_error(
|
||||||
self,
|
self,
|
||||||
@@ -83,11 +90,30 @@ class FileEditActivityHook(AgentHook):
|
|||||||
params: Any,
|
params: Any,
|
||||||
error: Any,
|
error: Any,
|
||||||
) -> None:
|
) -> None:
|
||||||
trackers = self._trackers_by_call.pop(self._tool_call_key(tool_call), [])
|
key = self._tool_call_key(tool_call)
|
||||||
|
trackers = self._trackers_by_call.get(key, [])
|
||||||
if trackers:
|
if trackers:
|
||||||
await self._emit([
|
await self._emit([
|
||||||
build_file_edit_error_event(tracker, str(error)) for tracker in trackers
|
build_file_edit_error_event(tracker, str(error)) for tracker in trackers
|
||||||
])
|
])
|
||||||
|
self._trackers_by_call.pop(key, None)
|
||||||
|
|
||||||
|
async def on_finally(self, context: AgentRunHookContext) -> None:
|
||||||
|
if context.stop_reason != "cancelled" or not self._trackers_by_call:
|
||||||
|
return
|
||||||
|
trackers = [
|
||||||
|
tracker
|
||||||
|
for trackers in self._trackers_by_call.values()
|
||||||
|
for tracker in trackers
|
||||||
|
]
|
||||||
|
self._trackers_by_call.clear()
|
||||||
|
await self._emit([
|
||||||
|
build_file_edit_error_event(
|
||||||
|
tracker,
|
||||||
|
"Task interrupted before this tool finished.",
|
||||||
|
)
|
||||||
|
for tracker in trackers
|
||||||
|
])
|
||||||
|
|
||||||
async def _emit(self, events: list[dict[str, Any]]) -> None:
|
async def _emit(self, events: list[dict[str, Any]]) -> None:
|
||||||
if self._on_progress is not None:
|
if self._on_progress is not None:
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
"""Tests for provider progress delta routing in the shared runner."""
|
"""Tests for provider progress delta routing in the shared runner."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
from unittest.mock import AsyncMock, MagicMock
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -280,3 +281,72 @@ async def test_runner_marks_file_edit_activity_failed_when_tool_errors(tmp_path)
|
|||||||
assert progress_events[-1]["phase"] == "error"
|
assert progress_events[-1]["phase"] == "error"
|
||||||
assert progress_events[-1]["status"] == "error"
|
assert progress_events[-1]["status"] == "error"
|
||||||
provider.chat_with_retry.assert_not_awaited()
|
provider.chat_with_retry.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_runner_marks_file_edit_activity_failed_when_cancelled(tmp_path):
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
|
progress_events: list[dict] = []
|
||||||
|
executing = asyncio.Event()
|
||||||
|
target = tmp_path / "cancelled.txt"
|
||||||
|
target.write_text("old\n", encoding="utf-8")
|
||||||
|
|
||||||
|
async def progress_cb(content, *, file_edit_events=None, **kwargs):
|
||||||
|
if file_edit_events:
|
||||||
|
progress_events.extend(file_edit_events)
|
||||||
|
|
||||||
|
class SlowWriteTool(WriteFileTool):
|
||||||
|
async def execute(self, path=None, content=None, **kwargs):
|
||||||
|
executing.set()
|
||||||
|
await asyncio.sleep(60)
|
||||||
|
return "ok"
|
||||||
|
|
||||||
|
tool = SlowWriteTool(workspace=tmp_path)
|
||||||
|
|
||||||
|
class Tools:
|
||||||
|
def get_definitions(self):
|
||||||
|
return [{"type": "function", "function": {"name": "write_file"}}]
|
||||||
|
|
||||||
|
def prepare_call(self, name, params):
|
||||||
|
return tool, params, None
|
||||||
|
|
||||||
|
async def chat_stream_with_retry(**kwargs):
|
||||||
|
return LLMResponse(
|
||||||
|
content=None,
|
||||||
|
tool_calls=[
|
||||||
|
ToolCallRequest(
|
||||||
|
id="call-write",
|
||||||
|
name="write_file",
|
||||||
|
arguments={"path": "cancelled.txt", "content": "new\n"},
|
||||||
|
)
|
||||||
|
],
|
||||||
|
usage={},
|
||||||
|
)
|
||||||
|
|
||||||
|
provider.chat_stream_with_retry = chat_stream_with_retry
|
||||||
|
provider.chat_with_retry = AsyncMock()
|
||||||
|
tools = Tools()
|
||||||
|
|
||||||
|
runner = AgentRunner(provider)
|
||||||
|
task = asyncio.create_task(runner.run(AgentRunSpec(
|
||||||
|
initial_messages=[{"role": "user", "content": "write a file"}],
|
||||||
|
tools=tools,
|
||||||
|
model="test-model",
|
||||||
|
max_iterations=2,
|
||||||
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
|
progress_callback=progress_cb,
|
||||||
|
workspace=tmp_path,
|
||||||
|
hook=FileEditActivityHook(on_progress=progress_cb, workspace=tmp_path),
|
||||||
|
)))
|
||||||
|
await asyncio.wait_for(executing.wait(), timeout=1)
|
||||||
|
|
||||||
|
task.cancel()
|
||||||
|
with pytest.raises(asyncio.CancelledError):
|
||||||
|
await task
|
||||||
|
|
||||||
|
assert [event["phase"] for event in progress_events] == ["start", "error"]
|
||||||
|
assert progress_events[-1]["path"] == "cancelled.txt"
|
||||||
|
assert progress_events[-1]["status"] == "error"
|
||||||
|
assert progress_events[-1]["error"] == "Task interrupted before this tool finished."
|
||||||
|
provider.chat_with_retry.assert_not_awaited()
|
||||||
|
|||||||
Reference in New Issue
Block a user