fix(core): interrupt subagent tool waits on stop (#5850)

* fix(core): interrupt subagent tool waits on stop

* test: relax subagent handoff timeout

* test: cover stop-aware tool interruption

* refactor: unify runner stop state

* refactor: simplify tool executor interruption

* fix: preserve tool interruption propagation

* refactor: tighten interruption helpers

---------

Co-authored-by: idiotsj <idiotsj@users.noreply.github.com>
This commit is contained in:
SJ
2026-03-21 00:59:52 +08:00
committed by GitHub
co-authored by idiotsj
parent d2e0bc778a
commit b816f26fe6
2 changed files with 306 additions and 65 deletions
@@ -4,6 +4,8 @@ import sys
import time
import traceback
import typing as T
from collections.abc import AsyncIterator
from contextlib import suppress
from dataclasses import dataclass, field
from mcp.types import (
@@ -80,6 +82,18 @@ class FollowUpTicket:
resolved: asyncio.Event = field(default_factory=asyncio.Event)
class _ToolExecutionInterrupted(Exception):
"""Raised when a running tool call is interrupted by a stop request."""
ToolExecutorResultT = T.TypeVar("ToolExecutorResultT")
USER_INTERRUPTION_MESSAGE = (
"[SYSTEM: User actively interrupted the response generation. "
"Partial output before interruption is preserved.]"
)
class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
def _get_persona_custom_error_message(self) -> str | None:
"""Read persona-level custom error message from event extras when available."""
@@ -154,8 +168,8 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
self.tool_executor = tool_executor
self.agent_hooks = agent_hooks
self.run_context = run_context
self._stop_requested = False
self._aborted = False
self._abort_signal = asyncio.Event()
self._pending_follow_ups: list[FollowUpTicket] = []
self._follow_up_seq = 0
@@ -208,6 +222,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
"func_tool": self.req.func_tool,
"session_id": self.req.session_id,
"extra_user_content_parts": self.req.extra_user_content_parts, # list[ContentPart]
"abort_signal": self._abort_signal,
}
if include_model:
# For primary provider we keep explicit model selection if provided.
@@ -398,10 +413,10 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
),
),
)
if self._stop_requested:
if self._is_stop_requested():
llm_resp_result = LLMResponse(
role="assistant",
completion_text="[SYSTEM: User actively interrupted the response generation. Partial output before interruption is preserved.]",
completion_text=USER_INTERRUPTION_MESSAGE,
reasoning_content=llm_response.reasoning_content,
reasoning_signature=llm_response.reasoning_signature,
)
@@ -417,49 +432,13 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
break # got final response
if not llm_resp_result:
if self._stop_requested:
if self._is_stop_requested():
llm_resp_result = LLMResponse(role="assistant", completion_text="")
else:
return
if self._stop_requested:
logger.info("Agent execution was requested to stop by user.")
llm_resp = llm_resp_result
if llm_resp.role != "assistant":
llm_resp = LLMResponse(
role="assistant",
completion_text="[SYSTEM: User actively interrupted the response generation. Partial output before interruption is preserved.]",
)
self.final_llm_resp = llm_resp
self._aborted = True
self._transition_state(AgentState.DONE)
self.stats.end_time = time.time()
parts = []
if llm_resp.reasoning_content or llm_resp.reasoning_signature:
parts.append(
ThinkPart(
think=llm_resp.reasoning_content,
encrypted=llm_resp.reasoning_signature,
)
)
if llm_resp.completion_text:
parts.append(TextPart(text=llm_resp.completion_text))
if parts:
self.run_context.messages.append(
Message(role="assistant", content=parts)
)
try:
await self.agent_hooks.on_agent_done(self.run_context, llm_resp)
except Exception as e:
logger.error(f"Error in on_agent_done hook: {e}", exc_info=True)
yield AgentResponse(
type="aborted",
data=AgentResponseData(chain=MessageChain(type="aborted")),
)
self._resolve_unconsumed_follow_ups()
if self._is_stop_requested():
yield await self._finalize_aborted_step(llm_resp_result)
return
# 处理 LLM 响应
@@ -534,27 +513,31 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
tool_call_result_blocks = []
cached_images = [] # Collect cached images for LLM visibility
async for result in self._handle_function_tools(self.req, llm_resp):
if result.kind == "tool_call_result_blocks":
if result.tool_call_result_blocks is not None:
tool_call_result_blocks = result.tool_call_result_blocks
elif result.kind == "cached_image":
if result.cached_image is not None:
# Collect cached image info
cached_images.append(result.cached_image)
elif result.kind == "message_chain":
chain = result.message_chain
if chain is None or chain.type is None:
# should not happen
continue
if chain.type == "tool_direct_result":
ar_type = "tool_call_result"
else:
ar_type = chain.type
yield AgentResponse(
type=ar_type,
data=AgentResponseData(chain=chain),
)
try:
async for result in self._handle_function_tools(self.req, llm_resp):
if result.kind == "tool_call_result_blocks":
if result.tool_call_result_blocks is not None:
tool_call_result_blocks = result.tool_call_result_blocks
elif result.kind == "cached_image":
if result.cached_image is not None:
# Collect cached image info
cached_images.append(result.cached_image)
elif result.kind == "message_chain":
chain = result.message_chain
if chain is None or chain.type is None:
# should not happen
continue
if chain.type == "tool_direct_result":
ar_type = "tool_call_result"
else:
ar_type = chain.type
yield AgentResponse(
type=ar_type,
data=AgentResponseData(chain=chain),
)
except _ToolExecutionInterrupted:
yield await self._finalize_aborted_step(llm_resp)
return
# 将结果添加到上下文中
parts = []
@@ -754,7 +737,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
)
_final_resp: CallToolResult | None = None
async for resp in executor: # type: ignore
async for resp in self._iter_tool_executor_results(executor): # type: ignore
if isinstance(resp, CallToolResult):
res = resp
_final_resp = resp
@@ -855,6 +838,8 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
except Exception as e:
logger.error(f"Error in on_tool_end hook: {e}", exc_info=True)
except Exception as e:
if isinstance(e, _ToolExecutionInterrupted):
raise
logger.warning(traceback.format_exc())
_append_tool_call_result(
func_tool_id,
@@ -945,6 +930,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
func_tool=param_subset,
model=self.req.model,
session_id=self.req.session_id,
abort_signal=self._abort_signal,
)
if requery_resp:
llm_resp = requery_resp
@@ -956,10 +942,102 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
return self._state in (AgentState.DONE, AgentState.ERROR)
def request_stop(self) -> None:
self._stop_requested = True
self._abort_signal.set()
def _is_stop_requested(self) -> bool:
return self._abort_signal.is_set()
def was_aborted(self) -> bool:
return self._aborted
def get_final_llm_resp(self) -> LLMResponse | None:
return self.final_llm_resp
async def _finalize_aborted_step(
self,
llm_resp: LLMResponse | None = None,
) -> AgentResponse:
logger.info("Agent execution was requested to stop by user.")
if llm_resp is None:
llm_resp = LLMResponse(role="assistant", completion_text="")
if llm_resp.role != "assistant":
llm_resp = LLMResponse(
role="assistant",
completion_text=USER_INTERRUPTION_MESSAGE,
)
self.final_llm_resp = llm_resp
self._aborted = True
self._transition_state(AgentState.DONE)
self.stats.end_time = time.time()
parts = []
if llm_resp.reasoning_content or llm_resp.reasoning_signature:
parts.append(
ThinkPart(
think=llm_resp.reasoning_content,
encrypted=llm_resp.reasoning_signature,
)
)
if llm_resp.completion_text:
parts.append(TextPart(text=llm_resp.completion_text))
if parts:
self.run_context.messages.append(Message(role="assistant", content=parts))
try:
await self.agent_hooks.on_agent_done(self.run_context, llm_resp)
except Exception as e:
logger.error(f"Error in on_agent_done hook: {e}", exc_info=True)
self._resolve_unconsumed_follow_ups()
return AgentResponse(
type="aborted",
data=AgentResponseData(chain=MessageChain(type="aborted")),
)
async def _close_executor(self, executor: T.Any) -> None:
close_executor = getattr(executor, "aclose", None)
if close_executor is None:
return
with suppress(asyncio.CancelledError, RuntimeError, StopAsyncIteration):
await close_executor()
async def _iter_tool_executor_results(
self,
executor: AsyncIterator[ToolExecutorResultT],
) -> T.AsyncGenerator[ToolExecutorResultT, None]:
while True:
if self._is_stop_requested():
await self._close_executor(executor)
raise _ToolExecutionInterrupted(
"Tool execution interrupted before reading the next tool result."
)
next_result_task = asyncio.create_task(anext(executor))
abort_task = asyncio.create_task(self._abort_signal.wait())
try:
done, _ = await asyncio.wait(
{next_result_task, abort_task},
return_when=asyncio.FIRST_COMPLETED,
)
if abort_task in done:
if not next_result_task.done():
next_result_task.cancel()
with suppress(asyncio.CancelledError, StopAsyncIteration):
await next_result_task
await self._close_executor(executor)
raise _ToolExecutionInterrupted(
"Tool execution interrupted by a stop request."
)
try:
yield next_result_task.result()
except StopAsyncIteration:
return
finally:
if not abort_task.done():
abort_task.cancel()
with suppress(asyncio.CancelledError):
await abort_task
+163
View File
@@ -1,5 +1,7 @@
import asyncio
import os
import sys
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
@@ -7,10 +9,13 @@ import pytest
# 将项目根目录添加到 sys.path
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..")))
from astrbot.core.agent.agent import Agent
from astrbot.core.agent.hooks import BaseAgentRunHooks
from astrbot.core.agent.handoff import HandoffTool
from astrbot.core.agent.run_context import ContextWrapper
from astrbot.core.agent.runners.tool_loop_agent_runner import ToolLoopAgentRunner
from astrbot.core.agent.tool import FunctionTool, ToolSet
from astrbot.core.astr_agent_tool_exec import FunctionToolExecutor
from astrbot.core.provider.entities import LLMResponse, ProviderRequest, TokenUsage
from astrbot.core.provider.provider import Provider
@@ -127,6 +132,31 @@ class MockAbortableStreamProvider(MockProvider):
)
class MockToolCallProvider(MockProvider):
def __init__(self, tool_name: str, tool_args: dict[str, str] | None = None):
super().__init__()
self.tool_name = tool_name
self.tool_args = tool_args or {}
self.abort_signal = None
async def text_chat(self, **kwargs) -> LLMResponse:
self.call_count += 1
self.abort_signal = kwargs.get("abort_signal")
return LLMResponse(
role="assistant",
completion_text="",
tools_call_name=[self.tool_name],
tools_call_args=[self.tool_args],
tools_call_ids=[f"call_{self.tool_name}"],
usage=TokenUsage(input_other=10, output=5),
)
class MockHandoffProvider(MockToolCallProvider):
def __init__(self, handoff_tool_name: str):
super().__init__(handoff_tool_name, {"input": "delegate this task"})
class MockHooks(BaseAgentRunHooks):
"""模拟钩子函数"""
@@ -163,6 +193,41 @@ class MockAgentContext:
self.event = event
class BlockingSubagentContext:
def __init__(self):
self.started = asyncio.Event()
self.cancelled = False
async def get_current_chat_provider_id(self, _umo: str) -> str:
return "provider-id"
def get_config(self, **_kwargs):
return {"provider_settings": {}}
async def tool_loop_agent(self, **_kwargs):
self.started.set()
try:
await asyncio.Future()
except asyncio.CancelledError:
self.cancelled = True
raise
class BlockingToolState:
def __init__(self):
self.started = asyncio.Event()
self.cancelled = False
async def handler(self, event, query: str = ""):
del event, query
self.started.set()
try:
await asyncio.Future()
except asyncio.CancelledError:
self.cancelled = True
raise
@pytest.fixture
def mock_provider():
return MockProvider()
@@ -466,6 +531,104 @@ async def test_stop_signal_returns_aborted_and_persists_partial_message(
assert runner.run_context.messages[-1].role == "assistant"
@pytest.mark.asyncio
async def test_stop_interrupts_pending_subagent_handoff(mock_hooks):
subagent_context = BlockingSubagentContext()
event = MockEvent("webchat:FriendMessage:webchat!user!session", "user")
handoff_tool = HandoffTool(
Agent(name="subagent", instructions="subagent-instructions", tools=[]),
tool_description="Delegate tasks to the subagent.",
)
provider = MockHandoffProvider(handoff_tool.name)
request = ProviderRequest(
prompt="delegate",
func_tool=ToolSet(tools=[handoff_tool]),
contexts=[],
)
runner = ToolLoopAgentRunner()
await runner.reset(
provider=provider,
request=request,
run_context=ContextWrapper(
context=SimpleNamespace(event=event, context=subagent_context)
),
tool_executor=FunctionToolExecutor(),
agent_hooks=mock_hooks,
streaming=False,
)
step_iter = runner.step()
first_resp = await step_iter.__anext__()
assert first_resp.type == "tool_call"
assert provider.abort_signal is not None
assert provider.abort_signal.is_set() is False
pending_resp = asyncio.create_task(step_iter.__anext__())
await asyncio.wait_for(subagent_context.started.wait(), timeout=5)
runner.request_stop()
assert provider.abort_signal.is_set() is True
aborted_resp = await asyncio.wait_for(pending_resp, timeout=1)
assert aborted_resp.type == "aborted"
assert runner.was_aborted() is True
assert subagent_context.cancelled is True
with pytest.raises(StopAsyncIteration):
await step_iter.__anext__()
@pytest.mark.asyncio
async def test_stop_interrupts_pending_regular_tool(mock_hooks):
tool_state = BlockingToolState()
event = MockEvent("webchat:FriendMessage:webchat!user!session", "user")
tool = FunctionTool(
name="long_tool",
description="A long-running test tool",
parameters={"type": "object", "properties": {"query": {"type": "string"}}},
handler=tool_state.handler,
)
provider = MockToolCallProvider(tool.name, {"query": "slow"})
request = ProviderRequest(
prompt="run a slow tool",
func_tool=ToolSet(tools=[tool]),
contexts=[],
)
runner = ToolLoopAgentRunner()
await runner.reset(
provider=provider,
request=request,
run_context=ContextWrapper(
context=SimpleNamespace(event=event, context=SimpleNamespace())
),
tool_executor=FunctionToolExecutor(),
agent_hooks=mock_hooks,
streaming=False,
)
step_iter = runner.step()
first_resp = await step_iter.__anext__()
assert first_resp.type == "tool_call"
assert provider.abort_signal is not None
assert provider.abort_signal.is_set() is False
pending_resp = asyncio.create_task(step_iter.__anext__())
await asyncio.wait_for(tool_state.started.wait(), timeout=5)
runner.request_stop()
assert provider.abort_signal.is_set() is True
aborted_resp = await asyncio.wait_for(pending_resp, timeout=5)
assert aborted_resp.type == "aborted"
assert runner.was_aborted() is True
assert tool_state.cancelled is True
with pytest.raises(StopAsyncIteration):
await step_iter.__anext__()
@pytest.mark.asyncio
async def test_tool_result_injects_follow_up_notice(
runner, mock_provider, provider_request, mock_tool_executor, mock_hooks