fix(agent): close file edit activity on cancellation

This commit is contained in:
chengyongru 2026-07-08 23:04:22 +08:00 committed by Xubin Ren
parent 6851a6ebd4
commit 0ae28571e7
2 changed files with 99 additions and 3 deletions

View File

@ -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:

View File

@ -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()