fix(agent): propagate task cancellation (#6284)

This commit is contained in:
InfinityPacer
2026-08-13 08:12:00 +08:00
committed by GitHub
parent a6e8ba8b57
commit 4d11a38496
6 changed files with 333 additions and 22 deletions
+50 -22
View File
@@ -2091,7 +2091,7 @@ class MoviePilotAgent:
logger.info(f"Agent执行被取消: session_id={self.session_id}") logger.info(f"Agent执行被取消: session_id={self.session_id}")
self._compiled_agent_bundle = None self._compiled_agent_bundle = None
execution_error = "任务已取消" execution_error = "任务已取消"
return "任务已取消", {} raise
except Exception as e: except Exception as e:
self._compiled_agent_bundle = None self._compiled_agent_bundle = None
execution_error = str(e) execution_error = str(e)
@@ -2436,9 +2436,9 @@ class AgentManager:
result = await self._process_message_internal(task) result = await self._process_message_internal(task)
if task.completion_future and not task.completion_future.done(): if task.completion_future and not task.completion_future.done():
task.completion_future.set_result(result) task.completion_future.set_result(result)
except asyncio.CancelledError as err: except asyncio.CancelledError:
if task.completion_future and not task.completion_future.done(): if task.completion_future and not task.completion_future.done():
task.completion_future.set_exception(err) task.completion_future.cancel()
raise raise
except Exception as e: except Exception as e:
logger.error(f"处理会话 {session_id} 的消息失败: {e}") logger.error(f"处理会话 {session_id} 的消息失败: {e}")
@@ -2452,14 +2452,28 @@ class AgentManager:
logger.info(f"会话 {session_id} 的worker被取消") logger.info(f"会话 {session_id} 的worker被取消")
finally: finally:
# 清理已完成的worker记录 # 清理已完成的worker记录
self._session_workers.pop(session_id, None) # noqa current_worker = asyncio.current_task()
if self._session_workers.get(session_id) is current_worker:
self._session_workers.pop(session_id, None) # noqa
# 如果队列为空,清理队列 # 如果队列为空,清理队列
if ( if (
session_id in self._session_queues self._session_queues.get(session_id) is queue
and self._session_queues[session_id].empty() and queue.empty()
): ):
self._session_queues.pop(session_id, None) self._session_queues.pop(session_id, None)
@staticmethod
def _discard_queued_messages(queue: asyncio.Queue) -> None:
"""丢弃会话队列时同步结束等待任务完成的调用方。"""
while not queue.empty():
try:
task = queue.get_nowait()
except asyncio.QueueEmpty:
break
if task.completion_future and not task.completion_future.done():
task.completion_future.cancel()
queue.task_done()
@staticmethod @staticmethod
async def _start_task_processing_status(task: _MessageTask) -> None: async def _start_task_processing_status(task: _MessageTask) -> None:
""" """
@@ -2550,27 +2564,37 @@ class AgentManager:
""" """
stopped = False stopped = False
# 取消该会话的worker(会触发 _execute_agent 中的 CancelledError worker = self._session_workers.get(session_id)
if session_id in self._session_workers: queue = self._session_queues.get(session_id)
self._session_workers[session_id].cancel() if queue and self._session_queues.get(session_id) is queue:
self._session_queues.pop(session_id, None)
# 先摘下旧队列;清理期间的新消息进入新队列,但等待旧 worker 完全退出后再执行。
if worker:
worker.cancel()
if queue:
self._discard_queued_messages(queue)
if worker:
try: try:
await self._session_workers[session_id] await worker
except asyncio.CancelledError: except asyncio.CancelledError:
pass pass
self._session_workers.pop(session_id, None) # noqa if self._session_workers.get(session_id) is worker:
self._session_workers.pop(session_id, None) # noqa
stopped = True
if queue:
stopped = True stopped = True
# 清空队列中待处理的消息 new_queue = self._session_queues.get(session_id)
if session_id in self._session_queues: current_worker = self._session_workers.get(session_id)
queue = self._session_queues[session_id] if (
while not queue.empty(): new_queue
try: and not new_queue.empty()
queue.get_nowait() and (not current_worker or current_worker.done())
queue.task_done() ):
except asyncio.QueueEmpty: self._session_workers[session_id] = asyncio.create_task(
break self._session_worker(session_id)
self._session_queues.pop(session_id, None) )
stopped = True
if stopped: if stopped:
logger.info(f"会话 {session_id} 的Agent推理已应急停止") logger.info(f"会话 {session_id} 的Agent推理已应急停止")
@@ -2683,6 +2707,10 @@ class AgentManager:
success = not result_text.startswith( success = not result_text.startswith(
(AGENT_EXECUTION_ERROR_PREFIX, "处理消息时发生错误") (AGENT_EXECUTION_ERROR_PREFIX, "处理消息时发生错误")
) )
except asyncio.CancelledError:
success = False
result = "Agent 定时任务已取消"
raise
except Exception as err: except Exception as err:
success = False success = False
result = f"Agent 定时任务执行失败:{str(err)}" result = f"Agent 定时任务执行失败:{str(err)}"
+4
View File
@@ -965,6 +965,10 @@ class Scheduler(ConfigReloadMixin, metaclass=SingletonClass):
result = await coro result = await coro
error = self.__get_result_error(result) error = self.__get_result_error(result)
success = error is None success = error is None
except asyncio.CancelledError:
success = False
error = "任务已取消"
raise
except Exception as err: except Exception as err:
success = False success = False
error = str(err) error = str(err)
+153
View File
@@ -0,0 +1,153 @@
import asyncio
from unittest.mock import AsyncMock
import pytest
from app.agent import AgentManager, MoviePilotAgent
def test_execute_agent_propagates_task_cancellation():
"""取消 Agent 执行时应终止调用方任务,不能转换成普通完成结果。"""
started = asyncio.Event()
class _BlockingAgent:
"""等待取消的最小 LangGraph 替身。"""
async def ainvoke(self, _payload, config=None): # noqa: ARG002
"""阻塞到外层任务取消。"""
started.set()
await asyncio.Event().wait()
async def _run_scenario():
agent = MoviePilotAgent(session_id="session-1", user_id="10001")
agent._should_stream = lambda: False
agent._create_agent = AsyncMock(return_value=_BlockingAgent())
agent.stream_handler.stop_streaming = AsyncMock(return_value=(False, ""))
execution = asyncio.create_task(agent._execute_agent([]))
await asyncio.wait_for(started.wait(), timeout=1)
execution.cancel()
with pytest.raises(asyncio.CancelledError):
await execution
agent.stream_handler.stop_streaming.assert_awaited_once()
asyncio.run(_run_scenario())
def test_stop_current_task_cancels_waiters_and_allows_next_message():
"""停止会话应结束当前及排队请求,并允许同一会话继续处理消息。"""
async def _run_scenario():
manager = AgentManager()
started = asyncio.Event()
async def _block_current_task(_task):
started.set()
await asyncio.Event().wait()
manager._process_message_internal = _block_current_task
first_waiter = asyncio.create_task(
manager.process_message(
session_id="session-1",
user_id="10001",
message="first",
wait_for_completion=True,
)
)
await asyncio.wait_for(started.wait(), timeout=1)
second_waiter = asyncio.create_task(
manager.process_message(
session_id="session-1",
user_id="10001",
message="second",
wait_for_completion=True,
)
)
await asyncio.sleep(0)
try:
assert await asyncio.wait_for(
manager.stop_current_task("session-1"), timeout=1
) is True
with pytest.raises(asyncio.CancelledError):
await first_waiter
with pytest.raises(asyncio.CancelledError):
await asyncio.wait_for(second_waiter, timeout=1)
manager._process_message_internal = AsyncMock(return_value="resumed")
result = await asyncio.wait_for(
manager.process_message(
session_id="session-1",
user_id="10001",
message="next",
wait_for_completion=True,
),
timeout=1,
)
assert result == "resumed"
finally:
await manager.stop_current_task("session-1")
for waiter in (first_waiter, second_waiter):
if not waiter.done():
waiter.cancel()
await asyncio.gather(
first_waiter,
second_waiter,
return_exceptions=True,
)
asyncio.run(_run_scenario())
def test_stop_queues_new_message_until_cancellation_cleanup_finishes():
"""旧 worker 清理期间到达的新消息应保留,并在清理完成后执行。"""
async def _run_scenario():
manager = AgentManager()
current_started = asyncio.Event()
cancellation_cleanup_started = asyncio.Event()
release_cleanup = asyncio.Event()
async def _process(task):
if task.message == "current":
current_started.set()
await asyncio.Event().wait()
return "next-completed"
async def _finish_status(_task):
cancellation_cleanup_started.set()
await release_cleanup.wait()
manager._process_message_internal = _process
manager._finish_task_processing_status = _finish_status
current_waiter = asyncio.create_task(
manager.process_message(
session_id="session-1",
user_id="10001",
message="current",
wait_for_completion=True,
)
)
await asyncio.wait_for(current_started.wait(), timeout=1)
stop_task = asyncio.create_task(manager.stop_current_task("session-1"))
await asyncio.wait_for(cancellation_cleanup_started.wait(), timeout=1)
next_waiter = asyncio.create_task(
manager.process_message(
session_id="session-1",
user_id="10001",
message="next",
wait_for_completion=True,
)
)
await asyncio.sleep(0)
assert not next_waiter.done()
release_cleanup.set()
assert await asyncio.wait_for(stop_task, timeout=1) is True
assert await asyncio.wait_for(next_waiter, timeout=1) == "next-completed"
with pytest.raises(asyncio.CancelledError):
await current_waiter
asyncio.run(_run_scenario())
+30
View File
@@ -1,3 +1,4 @@
import asyncio
import json import json
import threading import threading
from datetime import datetime, timedelta from datetime import datetime, timedelta
@@ -591,6 +592,35 @@ async def test_agent_manager_runs_contextless_task_in_broadcast_mode(
post_message.assert_not_awaited() post_message.assert_not_awaited()
@pytest.mark.anyio
async def test_agent_manager_records_cancelled_scheduled_task_as_failed() -> None:
"""被用户停止的定时 Agent 任务不得记录为成功完成。"""
user_id = f"cancel-{uuid4().hex}"
task = AgentTaskOper().add(
name="取消中的后台检查",
content="检查资源并报告",
trigger_type="cron",
cron_expression="0 * * * *",
run_at=None,
user_id=user_id,
username="admin",
session_id=f"session-{user_id}",
channel=None,
source="api",
original_chat_id=None,
)
manager = AgentManager()
manager.process_message = AsyncMock(side_effect=asyncio.CancelledError)
with pytest.raises(asyncio.CancelledError):
await manager.execute_scheduled_task(task.id)
completed = AgentTaskOper().get(task.id)
assert completed.last_status == "failed"
assert completed.last_result == "Agent 定时任务已取消"
assert completed.run_count == 1
@pytest.mark.anyio @pytest.mark.anyio
async def test_cached_agent_clears_channel_for_background_task() -> None: async def test_cached_agent_clears_channel_for_background_task() -> None:
"""复用会话 Agent 时,后台任务必须覆盖上一轮保留的渠道信息。""" """复用会话 Agent 时,后台任务必须覆盖上一轮保留的渠道信息。"""
+25
View File
@@ -2,6 +2,8 @@ import asyncio
import threading import threading
from uuid import uuid4 from uuid import uuid4
import pytest
from app.core.config import global_vars from app.core.config import global_vars
from app.scheduler import Scheduler from app.scheduler import Scheduler
@@ -144,6 +146,29 @@ def test_scheduler_runs_async_job_from_current_event_loop(monkeypatch):
assert progress.status == "success" assert progress.status == "success"
def test_scheduler_records_cancelled_async_job_as_failed():
"""协程任务取消时运行时进度不得收敛为成功。"""
job_id = f"test-cancelled-{uuid4()}"
async def task():
"""模拟传播到调度器外壳的取消。"""
raise asyncio.CancelledError
async def run_task():
job = scheduler._Scheduler__prepare_job(job_id)
with pytest.raises(asyncio.CancelledError):
await scheduler._Scheduler__run_coro_job(task(), job_id, job)
scheduler = _build_scheduler(job_id, task)
asyncio.run(run_task())
progress = scheduler.get_progress(job_id)
assert progress.enable is False
assert progress.status == "failed"
assert progress.success is False
assert progress.error == "任务已取消"
def test_scheduler_returns_none_for_unknown_job(): def test_scheduler_returns_none_for_unknown_job():
"""未注册且无历史进度的定时服务应返回空。""" """未注册且无历史进度的定时服务应返回空。"""
job_id = f"test-unknown-{uuid4()}" job_id = f"test-unknown-{uuid4()}"
+71
View File
@@ -1053,6 +1053,77 @@ def test_web_agent_stream_emits_heartbeat_during_idle_tool_wait():
assert '"type": "done"' in body assert '"type": "done"' in body
def test_web_agent_stop_finishes_stream_without_error():
"""停止运行中的 Web Agent 后应正常结束 SSE,不能继续等待或报执行错误。"""
payload = schemas.AgentWebChatRequest(
text="执行长任务",
session_id="browser-stop",
)
request = SimpleNamespace(is_disconnected=AsyncMock(return_value=False))
user = SimpleNamespace(id=1, name="admin", is_superuser=True)
session_id = "web-agent:stop"
class BlockingWebAgent:
"""阻塞到会话 worker 被停止的 Web Agent 替身。"""
started = None
def __init__(self, **kwargs):
self.__dict__.update(kwargs)
async def process(self, _message, **_kwargs):
"""等待外层 worker 取消。"""
self.started.set()
await asyncio.Event().wait()
async def cleanup(self):
"""模拟 Agent 资源清理。"""
return None
async def scenario():
BlockingWebAgent.started = asyncio.Event()
response = await web_agent_stream(payload, request, user)
iterator = response.body_iterator.__aiter__()
received = [await asyncio.wait_for(anext(iterator), timeout=1)]
await asyncio.wait_for(BlockingWebAgent.started.wait(), timeout=1)
assert await asyncio.wait_for(
agent_manager.stop_current_task(session_id), timeout=1
) is True
while '"type": "done"' not in "".join(received):
received.append(await asyncio.wait_for(anext(iterator), timeout=1))
await iterator.aclose()
return "".join(received)
try:
with patch("app.api.endpoints.agent.settings.AI_AGENT_ENABLE", True), patch(
"app.api.endpoints.agent._is_web_agent_traditional_message",
return_value=False,
), patch(
"app.api.endpoints.agent._has_web_agent_traditional_interaction",
return_value=False,
), patch(
"app.api.endpoints.agent._build_web_agent_session_id",
return_value=session_id,
), patch.object(
MessageChain,
"bind_user_session",
), patch(
"app.api.endpoints.agent._WebAgentMoviePilotAgent",
BlockingWebAgent,
), patch(
"app.api.endpoints.agent._save_web_agent_display_snapshot",
):
body = asyncio.run(scenario())
finally:
agent_manager._session_queues.pop(session_id, None)
agent_manager._session_workers.pop(session_id, None)
agent_manager.active_agents.pop(session_id, None)
assert '"type": "done"' in body
assert '"type": "error"' not in body
def test_web_agent_traditional_stream_keeps_alive_and_saves_after_done(): def test_web_agent_traditional_stream_keeps_alive_and_saves_after_done():
"""传统消息等待期间应保活,且展示快照不能阻塞终态。""" """传统消息等待期间应保活,且展示快照不能阻塞终态。"""
payload = schemas.AgentWebChatRequest(text="/状态", session_id="traditional-heartbeat") payload = schemas.AgentWebChatRequest(text="/状态", session_id="traditional-heartbeat")