mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-09-04 23:17:20 +08:00
fix(runtime): 收敛跨线程任务提交终态 (#6446)
This commit is contained in:
@@ -3,7 +3,7 @@ from concurrent.futures import Future
|
||||
from dataclasses import replace
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
from app.agent import MoviePilotAgent
|
||||
from app.agent.orchestrator import MoviePilotAgent
|
||||
from app.agent.orchestrator import AgentManagerQueueFullError
|
||||
from app.agent.tools.impl.ask_user_choice import (
|
||||
AskUserChoiceTool,
|
||||
@@ -12,13 +12,14 @@ from app.agent.tools.impl.ask_user_choice import (
|
||||
from app.agent.tools.impl.send_message import SendMessageTool
|
||||
from app.chain.message import MessageChain
|
||||
from app.runtime.config import global_vars, settings
|
||||
from app.db import SessionFactory
|
||||
from app.db.session import SessionFactory
|
||||
from app.db.oper.message import MessageOper
|
||||
from app.db.models.message import Message
|
||||
from app.application.messaging.agent import AgentInteractionOption, agent_interaction_manager
|
||||
from app.application.messaging.interaction import InteractionContext
|
||||
from app.application.messaging.media import media_interaction_manager
|
||||
from app.schemas.types import NotificationChannel, MessageType
|
||||
from app.runtime.tasks import TaskRegistry
|
||||
|
||||
|
||||
def _clear_messages() -> None:
|
||||
@@ -61,6 +62,23 @@ def test_agent_session_clear_uses_owned_threadsafe_submission():
|
||||
}
|
||||
|
||||
|
||||
def test_agent_session_clear_handles_closed_task_registry():
|
||||
"""宿主已停止接收任务时,同步消息清理链应记录拒绝而不是泄漏协程。"""
|
||||
registry = TaskRegistry()
|
||||
asyncio.run(registry.shutdown(timeout_seconds=0.01))
|
||||
manager = Mock(clear_session=AsyncMock())
|
||||
|
||||
with patch.object(global_vars, "CURRENT_EVENT_LOOP", _running_loop_stub()), patch(
|
||||
"app.chain.message.get_running_agent_manager", return_value=manager
|
||||
), patch("app.chain.message.get_task_registry", return_value=registry), patch(
|
||||
"app.chain.message.logger"
|
||||
) as logger:
|
||||
MessageChain._schedule_agent_session_clear("session-1", "10001")
|
||||
|
||||
logger.warning.assert_called_once()
|
||||
assert "正在关闭" in logger.warning.call_args.args[0]
|
||||
|
||||
|
||||
def test_remote_session_clear_reuses_owned_clear_scheduler():
|
||||
"""远程清理命令应复用唯一 Agent 会话清理入口。"""
|
||||
chain = MessageChain()
|
||||
|
||||
@@ -1,9 +1,12 @@
|
||||
"""服务端统计兼容入口的后台任务生命周期回归。"""
|
||||
|
||||
import asyncio
|
||||
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
from app.adapters.external.server import MoviePilotServerHelper
|
||||
from app.runtime.config import global_vars
|
||||
from app.runtime.tasks import TaskRegistry
|
||||
|
||||
|
||||
def test_legacy_subscription_reports_use_owned_threadsafe_tasks() -> None:
|
||||
@@ -40,3 +43,18 @@ def test_legacy_subscription_report_rejects_without_runtime_loop() -> None:
|
||||
"""宿主生命周期不可用时应拒绝提交,并保持布尔返回合同。"""
|
||||
with patch.object(global_vars, "CURRENT_EVENT_LOOP", None):
|
||||
assert MoviePilotServerHelper.sub_done_async({"media_id": "1"}) is False
|
||||
|
||||
|
||||
def test_legacy_subscription_report_handles_closed_task_registry() -> None:
|
||||
"""运行 loop 尚在但宿主已封口时,兼容入口应保持布尔失败合同。"""
|
||||
registry = TaskRegistry()
|
||||
asyncio.run(registry.shutdown(timeout_seconds=0.01))
|
||||
loop = Mock(**{"is_running.return_value": True, "is_closed.return_value": False})
|
||||
|
||||
with patch.object(global_vars, "CURRENT_EVENT_LOOP", loop), patch(
|
||||
"app.adapters.external.server.get_task_registry", return_value=registry
|
||||
), patch("app.adapters.external.server.logger") as logger:
|
||||
assert MoviePilotServerHelper.sub_done_async({"media_id": "1"}) is False
|
||||
|
||||
logger.warning.assert_called_once()
|
||||
assert "正在关闭" in logger.warning.call_args.args[0]
|
||||
|
||||
+212
-14
@@ -1,13 +1,39 @@
|
||||
"""进程内后台任务登记与关停语义测试。"""
|
||||
|
||||
import asyncio
|
||||
import concurrent.futures
|
||||
import inspect
|
||||
import threading
|
||||
import weakref
|
||||
from collections.abc import Coroutine
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from app.runtime.tasks import TaskRegistry
|
||||
|
||||
|
||||
class _CloseCountingCoroutine(Coroutine[Any, Any, None]):
|
||||
"""记录 pending submission 的协程关闭所有权。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.close_count = 0
|
||||
|
||||
def __await__(self):
|
||||
return self
|
||||
|
||||
def send(self, _value):
|
||||
raise StopIteration
|
||||
|
||||
def throw(self, typ, val=None, tb=None):
|
||||
if val is None:
|
||||
raise typ
|
||||
raise val.with_traceback(tb)
|
||||
|
||||
def close(self) -> None:
|
||||
self.close_count += 1
|
||||
|
||||
|
||||
def test_task_registry_removes_completed_task() -> None:
|
||||
"""正常完成的任务应自动退出登记表,避免长期持有请求对象。"""
|
||||
|
||||
@@ -124,37 +150,209 @@ def test_task_registry_owns_threadsafe_submission_until_shutdown() -> None:
|
||||
|
||||
|
||||
def test_task_registry_rejects_threadsafe_submission_after_shutdown() -> None:
|
||||
"""关停先赢得竞态时应关闭协程并通过 completion 报告拒绝原因。"""
|
||||
"""关停先赢得竞态时应关闭协程并同步拒绝提交。"""
|
||||
|
||||
async def scenario() -> None:
|
||||
registry = TaskRegistry()
|
||||
loop = asyncio.get_running_loop()
|
||||
reports: list[dict[str, object]] = []
|
||||
previous_handler = loop.get_exception_handler()
|
||||
loop.set_exception_handler(lambda _, context: reports.append(context))
|
||||
|
||||
async def late_worker() -> None:
|
||||
"""模拟关停完成后从宿主线程到达的晚任务。"""
|
||||
|
||||
try:
|
||||
assert await registry.shutdown(timeout_seconds=1.0) is True
|
||||
completion = await asyncio.to_thread(
|
||||
assert await registry.shutdown(timeout_seconds=1.0) is True
|
||||
coroutine = late_worker()
|
||||
with pytest.raises(RuntimeError, match="正在关闭"):
|
||||
await asyncio.to_thread(
|
||||
registry.submit_threadsafe,
|
||||
late_worker(),
|
||||
coroutine,
|
||||
loop=loop,
|
||||
owner="test.threadsafe-late",
|
||||
)
|
||||
with pytest.raises(RuntimeError, match="正在关闭"):
|
||||
await asyncio.wrap_future(completion)
|
||||
|
||||
assert registry.records == ()
|
||||
assert reports[-1]["owner"] == "test.threadsafe-late"
|
||||
finally:
|
||||
loop.set_exception_handler(previous_handler)
|
||||
assert inspect.getcoroutinestate(coroutine) == inspect.CORO_CLOSED
|
||||
assert registry.records == ()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_task_registry_shutdown_closes_submission_before_loop_dispatch() -> None:
|
||||
"""回调尚未执行时关停也必须关闭协程并结束 completion。"""
|
||||
registry = TaskRegistry()
|
||||
loop = asyncio.new_event_loop()
|
||||
|
||||
async def worker() -> None:
|
||||
"""不应被迟到的 loop callback 启动。"""
|
||||
|
||||
coroutine = worker()
|
||||
completion = registry.submit_threadsafe(
|
||||
coroutine,
|
||||
loop=loop,
|
||||
owner="test.threadsafe-pending",
|
||||
)
|
||||
|
||||
assert asyncio.run(registry.shutdown(timeout_seconds=0.01)) is True
|
||||
assert completion.cancelled()
|
||||
assert inspect.getcoroutinestate(coroutine) == inspect.CORO_CLOSED
|
||||
|
||||
loop.run_until_complete(asyncio.sleep(0))
|
||||
assert registry.records == ()
|
||||
loop.close()
|
||||
|
||||
|
||||
def test_task_registry_caller_cancels_before_loop_dispatch() -> None:
|
||||
"""调用方在 callback 前取消时应关闭协程,迟到 callback 不得创建 Task。"""
|
||||
registry = TaskRegistry()
|
||||
loop = asyncio.new_event_loop()
|
||||
|
||||
async def worker() -> None:
|
||||
"""不应在 completion 取消后启动。"""
|
||||
|
||||
coroutine = worker()
|
||||
completion = registry.submit_threadsafe(
|
||||
coroutine,
|
||||
loop=loop,
|
||||
owner="test.threadsafe-cancel-before-dispatch",
|
||||
)
|
||||
|
||||
assert completion.cancel()
|
||||
assert inspect.getcoroutinestate(coroutine) == inspect.CORO_CLOSED
|
||||
loop.run_until_complete(asyncio.sleep(0))
|
||||
|
||||
assert registry.records == ()
|
||||
loop.close()
|
||||
|
||||
|
||||
def test_task_registry_closes_submission_when_loop_rejects_callback() -> None:
|
||||
"""目标 loop 已关闭时同步失败,并且不遗留未等待协程。"""
|
||||
registry = TaskRegistry()
|
||||
loop = asyncio.new_event_loop()
|
||||
loop.close()
|
||||
|
||||
async def worker() -> None:
|
||||
"""模拟无法进入目标 loop 的提交。"""
|
||||
|
||||
coroutine = worker()
|
||||
with pytest.raises(RuntimeError, match="closed"):
|
||||
registry.submit_threadsafe(
|
||||
coroutine,
|
||||
loop=loop,
|
||||
owner="test.threadsafe-closed-loop",
|
||||
)
|
||||
|
||||
assert inspect.getcoroutinestate(coroutine) == inspect.CORO_CLOSED
|
||||
|
||||
|
||||
def test_task_registry_shutdown_owns_close_before_loop_rejection() -> None:
|
||||
"""shutdown 已取走 pending 后,迟到的投递失败不得再次关闭同一协程。"""
|
||||
registry = TaskRegistry()
|
||||
dispatch_entered = threading.Event()
|
||||
release_dispatch = threading.Event()
|
||||
coroutine = _CloseCountingCoroutine()
|
||||
|
||||
class BlockingClosedLoop:
|
||||
"""让 shutdown 稳定发生在投递抛错之前。"""
|
||||
|
||||
@staticmethod
|
||||
def call_soon_threadsafe(_callback) -> None:
|
||||
dispatch_entered.set()
|
||||
release_dispatch.wait()
|
||||
raise RuntimeError("loop closed")
|
||||
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor:
|
||||
submission = executor.submit(
|
||||
registry.submit_threadsafe,
|
||||
coroutine,
|
||||
loop=BlockingClosedLoop(),
|
||||
owner="test.threadsafe-shutdown-before-rejection",
|
||||
)
|
||||
assert dispatch_entered.wait(1.0)
|
||||
assert asyncio.run(registry.shutdown(timeout_seconds=0.01)) is True
|
||||
release_dispatch.set()
|
||||
with pytest.raises(RuntimeError, match="loop closed"):
|
||||
submission.result(timeout=1.0)
|
||||
|
||||
assert coroutine.close_count == 1
|
||||
|
||||
|
||||
def test_task_registry_shutdown_rejects_late_successful_dispatch() -> None:
|
||||
"""shutdown 取走 pending 后,迟到的成功投递也不得发布真实 Task。"""
|
||||
registry = TaskRegistry()
|
||||
dispatch_entered = threading.Event()
|
||||
release_dispatch = threading.Event()
|
||||
callbacks = []
|
||||
coroutine = _CloseCountingCoroutine()
|
||||
|
||||
class BlockingLoop:
|
||||
"""让 shutdown 稳定发生在 loop 接受 callback 之前。"""
|
||||
|
||||
@staticmethod
|
||||
def call_soon_threadsafe(callback) -> None:
|
||||
dispatch_entered.set()
|
||||
release_dispatch.wait()
|
||||
callbacks.append(callback)
|
||||
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor:
|
||||
submission = executor.submit(
|
||||
registry.submit_threadsafe,
|
||||
coroutine,
|
||||
loop=BlockingLoop(),
|
||||
owner="test.threadsafe-shutdown-before-dispatch",
|
||||
)
|
||||
assert dispatch_entered.wait(1.0)
|
||||
assert asyncio.run(registry.shutdown(timeout_seconds=0.01)) is True
|
||||
release_dispatch.set()
|
||||
completion = submission.result(timeout=1.0)
|
||||
|
||||
assert completion.cancelled()
|
||||
callbacks[0]()
|
||||
assert coroutine.close_count == 1
|
||||
assert registry.records == ()
|
||||
|
||||
|
||||
def test_task_registry_completion_does_not_retain_itself_after_cancellation() -> None:
|
||||
"""终态回调不得让已取消 completion 依赖循环 GC 才能释放。"""
|
||||
registry = TaskRegistry()
|
||||
loop = asyncio.new_event_loop()
|
||||
completion = registry.submit_threadsafe(
|
||||
_CloseCountingCoroutine(),
|
||||
loop=loop,
|
||||
owner="test.threadsafe-completion-release",
|
||||
)
|
||||
completion_ref = weakref.ref(completion)
|
||||
|
||||
assert completion.cancel()
|
||||
loop.run_until_complete(asyncio.sleep(0))
|
||||
del completion
|
||||
|
||||
assert completion_ref() is None
|
||||
loop.close()
|
||||
|
||||
|
||||
def test_task_registry_pending_non_cancellable_work_does_not_start_on_shutdown() -> None:
|
||||
"""尚未发布成 Task 的 non-cancellable 工作仍应在关停时直接关闭。"""
|
||||
registry = TaskRegistry()
|
||||
loop = asyncio.new_event_loop()
|
||||
|
||||
async def worker() -> None:
|
||||
"""未开始的同步兼容工作不应拖延关停。"""
|
||||
|
||||
coroutine = worker()
|
||||
completion = registry.submit_threadsafe(
|
||||
coroutine,
|
||||
loop=loop,
|
||||
owner="test.threadsafe-pending-non-cancellable",
|
||||
cancel_on_shutdown=False,
|
||||
)
|
||||
|
||||
assert asyncio.run(registry.shutdown(timeout_seconds=0.01)) is True
|
||||
assert completion.cancelled()
|
||||
assert inspect.getcoroutinestate(coroutine) == inspect.CORO_CLOSED
|
||||
|
||||
loop.run_until_complete(asyncio.sleep(0))
|
||||
assert registry.records == ()
|
||||
loop.close()
|
||||
|
||||
|
||||
def test_task_registry_keeps_timed_out_sync_owner_until_real_completion() -> None:
|
||||
"""同步线程超过关停预算后仍应保留 owner,不能把包装任务取消成伪完成。"""
|
||||
|
||||
|
||||
@@ -16,6 +16,7 @@ from app.chain.message import MessageChain
|
||||
from app.chain.transfer import TransferChain
|
||||
from app.application.messaging.interaction import InteractionContext
|
||||
from app.runtime.config import global_vars, settings
|
||||
from app.runtime.tasks import TaskRegistry
|
||||
from app.schemas.types import NotificationChannel
|
||||
|
||||
|
||||
@@ -128,10 +129,18 @@ class TestTransferFailedRetryButtons(unittest.TestCase):
|
||||
errmsg="未识别到媒体信息",
|
||||
)
|
||||
|
||||
def _close_pending_coro(coro, *args, **kwargs):
|
||||
"""关闭被调度的协程:测试中事件循环未运行,不关闭会残留 never-awaited 警告。"""
|
||||
coro.close()
|
||||
async_messages = []
|
||||
|
||||
def _run_pending_coro(coro, *args, **kwargs):
|
||||
asyncio.run(coro)
|
||||
|
||||
async def _capture_message(message):
|
||||
async_messages.append(message)
|
||||
|
||||
async def _finish_immediately(**kwargs):
|
||||
kwargs["output_callback"]("ok")
|
||||
|
||||
manager = SimpleNamespace(run_background_prompt=_finish_immediately)
|
||||
loop = Mock(**{"is_running.return_value": True, "is_closed.return_value": False})
|
||||
with patch.object(global_vars, "CURRENT_EVENT_LOOP", loop), patch.object(
|
||||
settings, "AI_AGENT_ENABLE", True
|
||||
@@ -141,12 +150,14 @@ class TestTransferFailedRetryButtons(unittest.TestCase):
|
||||
) as history_oper_cls, patch(
|
||||
"app.chain._transfer.build_manual_redo_prompt",
|
||||
return_value="retry transfer prompt",
|
||||
), patch(
|
||||
"app.chain._transfer.get_running_agent_manager", return_value=manager
|
||||
), patch("app.chain._transfer.get_task_registry") as get_registry:
|
||||
get_registry.return_value.submit_threadsafe.side_effect = (
|
||||
_close_pending_coro
|
||||
_run_pending_coro
|
||||
)
|
||||
history_oper_cls.return_value.get.return_value = history
|
||||
with patch.object(chain, "post_message") as post_message:
|
||||
with patch.object(chain, "async_post_message", side_effect=_capture_message):
|
||||
chain.handle_failed_transfer_callback(
|
||||
callback_data="transfer_ai_retry_34",
|
||||
channel=NotificationChannel.Telegram,
|
||||
@@ -160,10 +171,90 @@ class TestTransferFailedRetryButtons(unittest.TestCase):
|
||||
get_registry.return_value.submit_threadsafe.call_args.kwargs["owner"],
|
||||
"chain.transfer.ai_takeover",
|
||||
)
|
||||
self.assertEqual(len(async_messages), 2)
|
||||
self.assertEqual(
|
||||
async_messages[0].title,
|
||||
"已将整理记录 #34 交给智能助手处理",
|
||||
)
|
||||
self.assertEqual(async_messages[1].title, "智能助手整理完成")
|
||||
|
||||
def test_transfer_ai_retry_callback_reports_closed_task_registry(self):
|
||||
"""宿主停止接收任务时,不得向用户报告智能助手已接管。"""
|
||||
chain = TransferChain()
|
||||
chain.runtime_config = replace(chain.runtime_config, ai_agent_enable=True)
|
||||
history = SimpleNamespace(id=34)
|
||||
registry = TaskRegistry()
|
||||
asyncio.run(registry.shutdown(timeout_seconds=0.01))
|
||||
loop = Mock(**{"is_running.return_value": True, "is_closed.return_value": False})
|
||||
|
||||
with patch.object(global_vars, "CURRENT_EVENT_LOOP", loop), patch.object(
|
||||
settings, "AI_AGENT_ENABLE", True
|
||||
), patch(
|
||||
"app.chain._transfer.get_chain_transfer_history_port"
|
||||
) as history_port, patch(
|
||||
"app.chain._transfer.build_manual_redo_prompt",
|
||||
return_value="retry transfer prompt",
|
||||
), patch(
|
||||
"app.chain._transfer.get_task_registry", return_value=registry
|
||||
), patch(
|
||||
"app.chain._transfer.logger"
|
||||
) as logger, patch.object(
|
||||
chain, "post_message"
|
||||
) as post_message:
|
||||
history_port.return_value.get.return_value = history
|
||||
|
||||
chain.handle_failed_transfer_callback(
|
||||
callback_data="transfer_ai_retry_34",
|
||||
channel=NotificationChannel.Telegram,
|
||||
source="telegram-test",
|
||||
userid="10001",
|
||||
username="tester",
|
||||
)
|
||||
|
||||
logger.warning.assert_called_once()
|
||||
self.assertEqual(post_message.call_count, 1)
|
||||
self.assertEqual(
|
||||
post_message.call_args_list[0].args[0].title,
|
||||
"已将整理记录 #34 交给智能助手处理",
|
||||
post_message.call_args.args[0].title,
|
||||
"智能助手整理失败",
|
||||
)
|
||||
self.assertNotIn("已将", post_message.call_args.args[0].title)
|
||||
|
||||
def test_transfer_ai_retry_callback_reports_unavailable_event_loop(self):
|
||||
"""主循环不可用时,应在创建后台协程前返回明确失败提示。"""
|
||||
chain = TransferChain()
|
||||
chain.runtime_config = replace(chain.runtime_config, ai_agent_enable=True)
|
||||
history = SimpleNamespace(id=34)
|
||||
|
||||
with patch.object(global_vars, "CURRENT_EVENT_LOOP", None), patch.object(
|
||||
settings, "AI_AGENT_ENABLE", True
|
||||
), patch(
|
||||
"app.chain._transfer.get_chain_transfer_history_port"
|
||||
) as history_port, patch(
|
||||
"app.chain._transfer.build_manual_redo_prompt",
|
||||
return_value="retry transfer prompt",
|
||||
), patch(
|
||||
"app.chain._transfer.get_task_registry"
|
||||
) as get_registry, patch(
|
||||
"app.chain._transfer.logger"
|
||||
) as logger, patch.object(
|
||||
chain, "post_message"
|
||||
) as post_message:
|
||||
history_port.return_value.get.return_value = history
|
||||
|
||||
chain.handle_failed_transfer_callback(
|
||||
callback_data="transfer_ai_retry_34",
|
||||
channel=NotificationChannel.Telegram,
|
||||
source="telegram-test",
|
||||
userid="10001",
|
||||
username="tester",
|
||||
)
|
||||
|
||||
get_registry.return_value.submit_threadsafe.assert_not_called()
|
||||
logger.warning.assert_called_once()
|
||||
self.assertEqual(post_message.call_count, 1)
|
||||
self.assertEqual(
|
||||
post_message.call_args.args[0].title,
|
||||
"智能助手整理失败",
|
||||
)
|
||||
|
||||
def test_transfer_ai_retry_callback_uses_successful_move_dest_as_source(self):
|
||||
|
||||
Reference in New Issue
Block a user