mirror of
https://github.com/AstrBotDevs/AstrBot.git
synced 2026-09-24 16:39:52 +08:00
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:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user