mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-09-07 00:16:57 +08:00
fix(agent): close chat lifecycle safely
This commit is contained in:
@@ -17,7 +17,7 @@ from app.application.messaging.chat import (
|
|||||||
AsyncUnitOfWork,
|
AsyncUnitOfWork,
|
||||||
)
|
)
|
||||||
from app.application.messaging.message import MessageQueryService
|
from app.application.messaging.message import MessageQueryService
|
||||||
from app.startup.context import HostRuntime
|
from app.startup.context import AgentChatRuntime, HostRuntime
|
||||||
|
|
||||||
|
|
||||||
def get_agent_chat_service(
|
def get_agent_chat_service(
|
||||||
@@ -29,7 +29,7 @@ def get_agent_chat_service(
|
|||||||
|
|
||||||
|
|
||||||
def get_agent_chat_persistence(
|
def get_agent_chat_persistence(
|
||||||
runtime: HostRuntime = Depends(get_agent_chat_runtime),
|
runtime: AgentChatRuntime = Depends(get_agent_chat_runtime),
|
||||||
) -> AgentChatPersistenceService:
|
) -> AgentChatPersistenceService:
|
||||||
"""从类型化 Agent 运行时获取有界会话写入端口。"""
|
"""从类型化 Agent 运行时获取有界会话写入端口。"""
|
||||||
return runtime.persistence
|
return runtime.persistence
|
||||||
|
|||||||
@@ -1937,10 +1937,13 @@ async def save_agent_chat_display(
|
|||||||
service=service,
|
service=service,
|
||||||
persistence=persistence,
|
persistence=persistence,
|
||||||
)
|
)
|
||||||
chat = await service.get_accessible(session_id, current_user)
|
# 写入由独立 worker 事务完成,使用组合根登记的短会话服务复读,避免请求会话
|
||||||
|
# 的 identity map 返回写入前的 ORM 快照。
|
||||||
|
chat_service = get_configured_agent_chat_service()
|
||||||
|
chat = await chat_service.get_accessible(session_id, current_user)
|
||||||
if not chat:
|
if not chat:
|
||||||
return _SchemaResponse(success=False, message="会话保存失败")
|
return _SchemaResponse(success=False, message="会话保存失败")
|
||||||
return _SchemaResponse(success=True, data=service.to_summary(chat))
|
return _SchemaResponse(success=True, data=chat_service.to_summary(chat))
|
||||||
|
|
||||||
|
|
||||||
@router.delete(
|
@router.delete(
|
||||||
@@ -2027,7 +2030,7 @@ async def _web_agent_stream_impl(
|
|||||||
"""
|
"""
|
||||||
prompt = payload.text.strip()
|
prompt = payload.text.strip()
|
||||||
if not isinstance(service, AgentChatService):
|
if not isinstance(service, AgentChatService):
|
||||||
# 直接调用公开函数时不经过 FastAPI 依赖解析;生产路由总是传入运行时服务。
|
# SSE 后台任务可能在请求依赖释放后继续运行,查询服务必须自行取得短会话。
|
||||||
service = get_configured_agent_chat_service()
|
service = get_configured_agent_chat_service()
|
||||||
if not isinstance(persistence, AgentChatPersistenceService):
|
if not isinstance(persistence, AgentChatPersistenceService):
|
||||||
# 直接调用公开函数时不经过 FastAPI 依赖解析;生产路由总是传入运行时端口。
|
# 直接调用公开函数时不经过 FastAPI 依赖解析;生产路由总是传入运行时端口。
|
||||||
@@ -2387,7 +2390,6 @@ async def web_agent_stream(
|
|||||||
payload: _SchemaAgentWebChatRequest,
|
payload: _SchemaAgentWebChatRequest,
|
||||||
request: Request,
|
request: Request,
|
||||||
current_user: ApiPrincipal = Depends(get_current_active_user),
|
current_user: ApiPrincipal = Depends(get_current_active_user),
|
||||||
service: AgentChatService = Depends(get_agent_chat_service),
|
|
||||||
persistence: AgentChatPersistenceService = Depends(get_agent_chat_persistence),
|
persistence: AgentChatPersistenceService = Depends(get_agent_chat_persistence),
|
||||||
) -> StreamingResponse:
|
) -> StreamingResponse:
|
||||||
"""Web 智能助手流式对话的稳定公开路由入口。"""
|
"""Web 智能助手流式对话的稳定公开路由入口。"""
|
||||||
@@ -2395,6 +2397,5 @@ async def web_agent_stream(
|
|||||||
payload,
|
payload,
|
||||||
request,
|
request,
|
||||||
current_user,
|
current_user,
|
||||||
service,
|
persistence=persistence,
|
||||||
persistence,
|
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -202,6 +202,13 @@ async def shutdown_web_agent_background_tasks() -> None:
|
|||||||
await asyncio.wait(tasks)
|
await asyncio.wait(tasks)
|
||||||
|
|
||||||
|
|
||||||
|
async def wait_web_agent_background_tasks() -> None:
|
||||||
|
"""等待已登记的 Web Agent 任务完成取消后的最终收尾。"""
|
||||||
|
tasks = tuple(_WEB_AGENT_BACKGROUND_TASKS)
|
||||||
|
if tasks:
|
||||||
|
await asyncio.gather(*tasks, return_exceptions=True)
|
||||||
|
|
||||||
|
|
||||||
def register_channel_admin_resolver(
|
def register_channel_admin_resolver(
|
||||||
channel: Union[NotificationChannel, str],
|
channel: Union[NotificationChannel, str],
|
||||||
resolver: _ChannelAdminResolver,
|
resolver: _ChannelAdminResolver,
|
||||||
|
|||||||
@@ -103,10 +103,21 @@ async def run_shutdown_step(
|
|||||||
try:
|
try:
|
||||||
result = callback()
|
result = callback()
|
||||||
if inspect.isawaitable(result):
|
if inspect.isawaitable(result):
|
||||||
|
task = asyncio.ensure_future(result)
|
||||||
if timeout_seconds:
|
if timeout_seconds:
|
||||||
await asyncio.wait_for(result, timeout=timeout_seconds)
|
try:
|
||||||
|
await asyncio.wait_for(
|
||||||
|
asyncio.shield(task), timeout=timeout_seconds
|
||||||
|
)
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
logger.error("关闭%s超时,等待其取消收口", name)
|
||||||
|
task.cancel()
|
||||||
|
try:
|
||||||
|
await task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
else:
|
else:
|
||||||
await result
|
await task
|
||||||
except Exception as err:
|
except Exception as err:
|
||||||
logger.error(f"关闭{name}失败:{err}")
|
logger.error(f"关闭{name}失败:{err}")
|
||||||
|
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
import asyncio
|
||||||
import inspect
|
import inspect
|
||||||
import sys
|
import sys
|
||||||
from typing import Callable
|
from typing import Callable
|
||||||
@@ -69,7 +70,10 @@ from app.application.messaging.chat import (
|
|||||||
configure_agent_chat_service,
|
configure_agent_chat_service,
|
||||||
get_configured_agent_chat_persistence,
|
get_configured_agent_chat_persistence,
|
||||||
)
|
)
|
||||||
from app.application.messaging.agent import shutdown_web_agent_background_tasks
|
from app.application.messaging.agent import (
|
||||||
|
shutdown_web_agent_background_tasks,
|
||||||
|
wait_web_agent_background_tasks,
|
||||||
|
)
|
||||||
from app.application.security.user import configure_user_lookups
|
from app.application.security.user import configure_user_lookups
|
||||||
from app.application.security.auth import AuthService, configure_auth_service
|
from app.application.security.auth import AuthService, configure_auth_service
|
||||||
from app.application.security.passkeys import PasskeyService, configure_passkey_service
|
from app.application.security.passkeys import PasskeyService, configure_passkey_service
|
||||||
@@ -559,14 +563,19 @@ async def stop_modules():
|
|||||||
"""
|
"""
|
||||||
服务关闭
|
服务关闭
|
||||||
"""
|
"""
|
||||||
async def run_step(name: str, callback: Callable[[], object]) -> None:
|
async def run_step(name: str, callback: Callable[[], object]) -> bool:
|
||||||
"""单个模块资源关闭失败时继续执行后续阶段"""
|
"""单个模块资源关闭失败时继续执行后续阶段"""
|
||||||
try:
|
try:
|
||||||
result = callback()
|
result = callback()
|
||||||
if inspect.isawaitable(result):
|
if inspect.isawaitable(result):
|
||||||
await result
|
await result
|
||||||
|
return True
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
logger.warning("关闭%s时收到取消请求,继续执行资源收口", name)
|
||||||
|
return False
|
||||||
except Exception as err:
|
except Exception as err:
|
||||||
logger.error(f"关闭{name}失败:{err}")
|
logger.error(f"关闭{name}失败:{err}")
|
||||||
|
return True
|
||||||
|
|
||||||
await run_step("AI智能体", stop_agent)
|
await run_step("AI智能体", stop_agent)
|
||||||
await run_step("模块", lambda: ModuleManager().shutdown())
|
await run_step("模块", lambda: ModuleManager().shutdown())
|
||||||
@@ -578,21 +587,32 @@ async def stop_modules():
|
|||||||
await run_step("消息服务", stop_message)
|
await run_step("消息服务", stop_message)
|
||||||
await run_step("Redis缓存连接", lambda: RedisHelper().close())
|
await run_step("Redis缓存连接", lambda: RedisHelper().close())
|
||||||
await run_step("异步Redis缓存连接", lambda: AsyncRedisHelper().close())
|
await run_step("异步Redis缓存连接", lambda: AsyncRedisHelper().close())
|
||||||
# 先关闭持久化准入,取消中的 Web Agent finally 才会快速拒绝晚到的快照写入。
|
# Web Agent 的取消 finally 可能还要写入最终展示快照,必须先完成任务收尾,再关闭写入准入。
|
||||||
await run_step(
|
web_agent_drained = await run_step(
|
||||||
"Agent会话持久化准入",
|
"Web Agent后台任务", shutdown_web_agent_background_tasks
|
||||||
lambda: get_configured_agent_chat_persistence().begin_shutdown(),
|
|
||||||
)
|
)
|
||||||
await run_step("Web Agent后台任务", shutdown_web_agent_background_tasks)
|
if not web_agent_drained:
|
||||||
await run_step(
|
web_agent_drained = await run_step(
|
||||||
"Agent会话持久化",
|
"Web Agent后台任务收尾", wait_web_agent_background_tasks
|
||||||
lambda: get_configured_agent_chat_persistence().shutdown(),
|
)
|
||||||
)
|
if web_agent_drained:
|
||||||
await run_step("数据库任务", stop_database_worker)
|
await run_step(
|
||||||
if _database_worker is None:
|
"Agent会话持久化准入",
|
||||||
await run_step("数据库连接", close_database)
|
lambda: get_configured_agent_chat_persistence().begin_shutdown(),
|
||||||
|
)
|
||||||
|
persistence_drained = await run_step(
|
||||||
|
"Agent会话持久化",
|
||||||
|
lambda: get_configured_agent_chat_persistence().shutdown(),
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
logger.error("数据库任务未收敛,跳过数据库连接关闭以避免运行中事务使用已释放连接")
|
persistence_drained = False
|
||||||
|
logger.error("Web Agent任务未完成收尾,跳过持久化和数据库关闭以保护活动事务")
|
||||||
|
if persistence_drained:
|
||||||
|
await run_step("数据库任务", stop_database_worker)
|
||||||
|
if _database_worker is None:
|
||||||
|
await run_step("数据库连接", close_database)
|
||||||
|
else:
|
||||||
|
logger.error("数据库任务未收敛,跳过数据库连接关闭以避免运行中事务使用已释放连接")
|
||||||
await run_step("前端服务", stop_frontend)
|
await run_step("前端服务", stop_frontend)
|
||||||
await run_step("临时文件", clear_temp)
|
await run_step("临时文件", clear_temp)
|
||||||
|
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ from __future__ import annotations
|
|||||||
import asyncio
|
import asyncio
|
||||||
import threading
|
import threading
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import AsyncMock, call, patch
|
from unittest.mock import AsyncMock, MagicMock, call, patch
|
||||||
from uuid import uuid4
|
from uuid import uuid4
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -176,6 +176,53 @@ async def test_authoritative_display_save_propagates_worker_overload() -> None:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_authoritative_display_save_reads_fresh_projection_after_worker_write(
|
||||||
|
monkeypatch,
|
||||||
|
) -> None:
|
||||||
|
"""权威展示保存的响应必须读取 worker 提交后的最新投影。"""
|
||||||
|
existing_chat = SimpleNamespace(
|
||||||
|
user_id="1",
|
||||||
|
username="admin",
|
||||||
|
channel="WebAgent",
|
||||||
|
source="web-agent",
|
||||||
|
original_chat_id=None,
|
||||||
|
client_session_id="client-1",
|
||||||
|
)
|
||||||
|
updated_chat = SimpleNamespace(
|
||||||
|
session_id="fresh-session",
|
||||||
|
message_count=2,
|
||||||
|
)
|
||||||
|
request_service = SimpleNamespace(
|
||||||
|
get_accessible=AsyncMock(return_value=existing_chat),
|
||||||
|
get=AsyncMock(return_value=existing_chat),
|
||||||
|
)
|
||||||
|
canonical_service = SimpleNamespace(
|
||||||
|
get_accessible=AsyncMock(return_value=updated_chat),
|
||||||
|
to_summary=MagicMock(return_value="fresh-summary"),
|
||||||
|
)
|
||||||
|
persistence = SimpleNamespace(async_save_display_messages=AsyncMock())
|
||||||
|
current_user = SimpleNamespace(id=1, name="admin", is_superuser=True)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"app.api.endpoints.agent.get_configured_agent_chat_service",
|
||||||
|
MagicMock(return_value=canonical_service),
|
||||||
|
)
|
||||||
|
|
||||||
|
response = await save_agent_chat_display(
|
||||||
|
session_id="fresh-session",
|
||||||
|
payload=AgentChatDisplaySaveRequest(messages=[]),
|
||||||
|
current_user=current_user,
|
||||||
|
service=request_service,
|
||||||
|
persistence=persistence,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.success is True
|
||||||
|
assert response.data == "fresh-summary"
|
||||||
|
canonical_service.get_accessible.assert_awaited_once_with(
|
||||||
|
"fresh-session", current_user
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_agent_chat_persistence_rolls_back_compound_write(monkeypatch) -> None:
|
async def test_agent_chat_persistence_rolls_back_compound_write(monkeypatch) -> None:
|
||||||
"""复合写入中途失败时,创建或更新不能留下半成品。"""
|
"""复合写入中途失败时,创建或更新不能留下半成品。"""
|
||||||
|
|||||||
@@ -571,7 +571,7 @@ def test_stop_modules_continues_after_internal_owner_failures(monkeypatch):
|
|||||||
|
|
||||||
|
|
||||||
def test_stop_modules_drains_web_agent_tasks_before_persistence(monkeypatch):
|
def test_stop_modules_drains_web_agent_tasks_before_persistence(monkeypatch):
|
||||||
"""关闭时先关闭持久化准入,再收口 Web Agent 和已有写入。"""
|
"""关闭时先收口 Web Agent,再关闭持久化准入和数据库任务。"""
|
||||||
order = []
|
order = []
|
||||||
monkeypatch.setattr(modules_initializer, "stop_agent", AsyncMock())
|
monkeypatch.setattr(modules_initializer, "stop_agent", AsyncMock())
|
||||||
dependencies = _patch_module_shutdown_dependencies(monkeypatch)
|
dependencies = _patch_module_shutdown_dependencies(monkeypatch)
|
||||||
@@ -599,12 +599,53 @@ def test_stop_modules_drains_web_agent_tasks_before_persistence(monkeypatch):
|
|||||||
|
|
||||||
asyncio.run(modules_initializer.stop_modules())
|
asyncio.run(modules_initializer.stop_modules())
|
||||||
|
|
||||||
assert order == [
|
assert order == ["web-agent", "persistence-admission", "persistence", "database"]
|
||||||
"persistence-admission",
|
|
||||||
"web-agent",
|
|
||||||
"persistence",
|
@pytest.mark.asyncio
|
||||||
"database",
|
async def test_shutdown_timeout_does_not_skip_database_worker_cleanup(monkeypatch):
|
||||||
]
|
"""模块关闭超时取消当前步骤后仍应继续收口数据库 worker。"""
|
||||||
|
started = asyncio.Event()
|
||||||
|
|
||||||
|
async def blocked_web_agent_shutdown():
|
||||||
|
started.set()
|
||||||
|
await asyncio.Event().wait()
|
||||||
|
|
||||||
|
monkeypatch.setattr(modules_initializer, "stop_agent", AsyncMock())
|
||||||
|
_patch_module_shutdown_dependencies(monkeypatch)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
modules_initializer,
|
||||||
|
"shutdown_web_agent_background_tasks",
|
||||||
|
blocked_web_agent_shutdown,
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
modules_initializer,
|
||||||
|
"wait_web_agent_background_tasks",
|
||||||
|
AsyncMock(),
|
||||||
|
)
|
||||||
|
persistence = MagicMock()
|
||||||
|
persistence.begin_shutdown = MagicMock()
|
||||||
|
persistence.shutdown = AsyncMock()
|
||||||
|
monkeypatch.setattr(
|
||||||
|
modules_initializer,
|
||||||
|
"get_configured_agent_chat_persistence",
|
||||||
|
MagicMock(return_value=persistence),
|
||||||
|
)
|
||||||
|
stop_database_worker = AsyncMock()
|
||||||
|
monkeypatch.setattr(modules_initializer, "stop_database_worker", stop_database_worker)
|
||||||
|
monkeypatch.setattr(modules_initializer, "_database_worker", object())
|
||||||
|
|
||||||
|
shutdown = asyncio.create_task(
|
||||||
|
lifecycle.run_shutdown_step(
|
||||||
|
"模块服务",
|
||||||
|
modules_initializer.stop_modules,
|
||||||
|
timeout_seconds=0.01,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await started.wait()
|
||||||
|
await shutdown
|
||||||
|
|
||||||
|
stop_database_worker.assert_awaited_once_with()
|
||||||
|
|
||||||
|
|
||||||
def _patch_module_shutdown_dependencies(monkeypatch) -> dict:
|
def _patch_module_shutdown_dependencies(monkeypatch) -> dict:
|
||||||
|
|||||||
@@ -355,6 +355,23 @@ def test_web_agent_stream_returns_error_for_unknown_command():
|
|||||||
handle_message.assert_not_called()
|
handle_message.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
def test_web_agent_stream_does_not_bind_request_scoped_chat_service():
|
||||||
|
"""流式路由不能把请求级 Agent 会话服务捕获到后台任务。"""
|
||||||
|
from app.api.dependencies.agent import get_agent_chat_service
|
||||||
|
from app.api.endpoints import agent as agent_endpoint
|
||||||
|
|
||||||
|
route = next(
|
||||||
|
route
|
||||||
|
for route in agent_endpoint.router.routes
|
||||||
|
if getattr(route, "name", None) == "web_agent_stream"
|
||||||
|
)
|
||||||
|
|
||||||
|
assert all(
|
||||||
|
dependency.call is not get_agent_chat_service
|
||||||
|
for dependency in route.dependant.dependencies
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_build_web_agent_message_update_event_converts_buttons():
|
def test_build_web_agent_message_update_event_converts_buttons():
|
||||||
"""WebAgent 编辑消息应转换为可原地更新卡片的事件。"""
|
"""WebAgent 编辑消息应转换为可原地更新卡片的事件。"""
|
||||||
event = build_web_agent_message_update_event(
|
event = build_web_agent_message_update_event(
|
||||||
|
|||||||
Reference in New Issue
Block a user