fix(agent): close async chat persistence lifecycle

This commit is contained in:
InfinityPacer
2026-08-23 11:44:18 +08:00
parent 6f5ee96152
commit 0ba4a7e5e3
14 changed files with 448 additions and 72 deletions
+9
View File
@@ -4,6 +4,7 @@ from fastapi import Depends
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.api.context import ( from app.api.context import (
get_agent_chat_runtime,
get_agent_chat_repository, get_agent_chat_repository,
get_agent_chat_transaction, get_agent_chat_transaction,
get_async_session, get_async_session,
@@ -11,6 +12,7 @@ from app.api.context import (
) )
from app.application.messaging.chat import ( from app.application.messaging.chat import (
AgentChatService, AgentChatService,
AgentChatPersistenceService,
AsyncAgentChatRepository, AsyncAgentChatRepository,
AsyncUnitOfWork, AsyncUnitOfWork,
) )
@@ -26,6 +28,13 @@ def get_agent_chat_service(
return AgentChatService(chat_repository, unit_of_work) return AgentChatService(chat_repository, unit_of_work)
def get_agent_chat_persistence(
runtime: HostRuntime = Depends(get_agent_chat_runtime),
) -> AgentChatPersistenceService:
"""从类型化 Agent 运行时获取有界会话写入端口。"""
return runtime.persistence
def get_message_query_service( def get_message_query_service(
db: AsyncSession = Depends(get_async_session), db: AsyncSession = Depends(get_async_session),
runtime: HostRuntime = Depends(get_host_runtime), runtime: HostRuntime = Depends(get_host_runtime),
+2
View File
@@ -5,6 +5,7 @@
""" """
from app.api.dependencies.agent import ( from app.api.dependencies.agent import (
get_agent_chat_persistence,
get_agent_chat_service, get_agent_chat_service,
get_message_query_service, get_message_query_service,
) )
@@ -52,6 +53,7 @@ from app.api.dependencies.workflow import (
# 兼容聚合入口只显式列出既有 FastAPI 依赖,不向插件制造新的动态导出规则。 # 兼容聚合入口只显式列出既有 FastAPI 依赖,不向插件制造新的动态导出规则。
__all__ = [ __all__ = [
"get_agent_chat_persistence",
"get_agent_chat_service", "get_agent_chat_service",
"get_auth_service", "get_auth_service",
"get_current_active_manage_user", "get_current_active_manage_user",
+91 -45
View File
@@ -47,17 +47,25 @@ from app.command import Command
from app.runtime.config import global_vars from app.runtime.config import global_vars
from app.runtime.events import Event, EventManager from app.runtime.events import Event, EventManager
from app.api.principal import ApiPrincipal from app.api.principal import ApiPrincipal
from app.api.dependencies.agent import get_agent_chat_service from app.api.dependencies.agent import (
get_agent_chat_persistence,
get_agent_chat_service,
)
from app.api.dependencies.auth import get_current_active_user from app.api.dependencies.auth import get_current_active_user
from app.application.messaging.chat import ( from app.application.messaging.chat import (
AgentChatRecord, AgentChatRecord,
AgentChatPersistenceService,
AgentChatService, AgentChatService,
get_configured_agent_chat_service, get_configured_agent_chat_service,
get_configured_agent_chat_persistence, get_configured_agent_chat_persistence,
) )
from app.application.security.user import get_configured_user_id_lookup from app.application.security.user import get_configured_user_id_lookup
from app.application.configuration import get_api_runtime_config_snapshot from app.application.configuration import get_api_runtime_config_snapshot
from app.application.messaging.agent import attach_web_agent_edit_queue, detach_web_agent_edit_queue from app.application.messaging.agent import (
attach_web_agent_edit_queue,
create_web_agent_background_task,
detach_web_agent_edit_queue,
)
from app.application.messaging.agent import agent_interaction_manager from app.application.messaging.agent import agent_interaction_manager
from app.application.messaging.agent import ( from app.application.messaging.agent import (
build_agent_choice_button_rows, build_agent_choice_button_rows,
@@ -88,7 +96,6 @@ _WEB_AGENT_FILE_REGISTRY: dict[str, dict[str, Any]] = {}
_WEB_AGENT_MESSAGE_QUEUES: dict[str, list[Queue[_SchemaMessage]]] = {} _WEB_AGENT_MESSAGE_QUEUES: dict[str, list[Queue[_SchemaMessage]]] = {}
_WEB_AGENT_MESSAGE_LOCK = Lock() _WEB_AGENT_MESSAGE_LOCK = Lock()
_WEB_AGENT_MESSAGE_LISTENER_REGISTERED = False _WEB_AGENT_MESSAGE_LISTENER_REGISTERED = False
_WEB_AGENT_BACKGROUND_TASKS: set[asyncio.Task] = set()
class _WebAgentEventPublisher: class _WebAgentEventPublisher:
@@ -517,13 +524,16 @@ def _build_web_agent_session_id(user: ApiPrincipal, session_id: Optional[str]) -
async def _build_web_agent_session_id_async( async def _build_web_agent_session_id_async(
user: ApiPrincipal, user: ApiPrincipal,
session_id: Optional[str], session_id: Optional[str],
service: Optional[AgentChatService] = None,
) -> str: ) -> str:
"""异步解析 Web Agent 会话 ID,并复用异步会话查询端口。""" """异步解析 Web Agent 会话 ID,并复用异步会话查询端口。"""
seed = str(session_id or "").strip() or uuid.uuid4().hex seed = str(session_id or "").strip() or uuid.uuid4().hex
if seed.startswith(WEB_AGENT_SESSION_PREFIX): if seed.startswith(WEB_AGENT_SESSION_PREFIX):
return seed return seed
try: try:
existing_chat = await get_configured_agent_chat_service().get(seed) if service is None:
service = get_configured_agent_chat_service()
existing_chat = await service.get(seed)
if existing_chat and AgentChatService.can_access(existing_chat, user): if existing_chat and AgentChatService.can_access(existing_chat, user):
return seed return seed
except Exception as e: except Exception as e:
@@ -658,36 +668,41 @@ async def _save_web_agent_display_snapshot(
current_user: ApiPrincipal, current_user: ApiPrincipal,
messages: list[dict], messages: list[dict],
client_session_id: Optional[str] = None, client_session_id: Optional[str] = None,
service: Optional[AgentChatService] = None,
persistence: Optional[AgentChatPersistenceService] = None,
) -> None: ) -> None:
""" """
保存 WebAgent 当前展示消息快照。 保存 WebAgent 当前展示消息快照。
""" """
try: if service is None:
existing_chat = await get_configured_agent_chat_service().get(session_id) # 直接调用该内部 helper 时没有 FastAPI 依赖注入上下文。
await get_configured_agent_chat_persistence().async_save_display_messages( service = get_configured_agent_chat_service()
session_id=session_id, existing_chat = await service.get(session_id)
user_id=(existing_chat.user_id if existing_chat else str(current_user.id)), if persistence is None:
username=(existing_chat.username if existing_chat else current_user.name), # 直接调用该内部 helper 时没有 FastAPI 依赖注入上下文。
channel=( persistence = get_configured_agent_chat_persistence()
existing_chat.channel await persistence.async_save_display_messages(
if existing_chat and existing_chat.channel session_id=session_id,
else NotificationChannel.WebAgent user_id=(existing_chat.user_id if existing_chat else str(current_user.id)),
), username=(existing_chat.username if existing_chat else current_user.name),
source=( channel=(
existing_chat.source existing_chat.channel
if existing_chat and existing_chat.source if existing_chat and existing_chat.channel
else WEB_AGENT_SOURCE else NotificationChannel.WebAgent
), ),
original_chat_id=existing_chat.original_chat_id if existing_chat else None, source=(
client_session_id=( existing_chat.source
existing_chat.client_session_id if existing_chat and existing_chat.source
if existing_chat and existing_chat.client_session_id else WEB_AGENT_SOURCE
else client_session_id ),
), original_chat_id=existing_chat.original_chat_id if existing_chat else None,
messages=messages, client_session_id=(
) existing_chat.client_session_id
except Exception as e: if existing_chat and existing_chat.client_session_id
logger.debug(f"保存WebAgent展示历史失败: {e}") else client_session_id
),
messages=messages,
)
def _build_web_agent_sse( def _build_web_agent_sse(
@@ -757,6 +772,7 @@ def _sanitize_web_agent_upload_name(
async def _get_web_agent_upload_dir( async def _get_web_agent_upload_dir(
user: ApiPrincipal, user: ApiPrincipal,
session_id: Optional[str], session_id: Optional[str],
service: Optional[AgentChatService] = None,
) -> Path: ) -> Path:
""" """
计算当前 Web Agent 会话的临时附件目录。 计算当前 Web Agent 会话的临时附件目录。
@@ -765,7 +781,11 @@ async def _get_web_agent_upload_dir(
:param session_id: 前端会话标识 :param session_id: 前端会话标识
:return: 已创建的临时附件目录 :return: 已创建的临时附件目录
""" """
server_session_id = await _build_web_agent_session_id_async(user, session_id) server_session_id = await _build_web_agent_session_id_async(
user,
session_id,
service,
)
safe_session_id = server_session_id.replace(":", "_") safe_session_id = server_session_id.replace(":", "_")
upload_dir = get_api_runtime_config_snapshot().temp_path / "agent_uploads" / safe_session_id upload_dir = get_api_runtime_config_snapshot().temp_path / "agent_uploads" / safe_session_id
upload_dir.mkdir(parents=True, exist_ok=True) upload_dir.mkdir(parents=True, exist_ok=True)
@@ -1701,6 +1721,7 @@ async def upload_web_agent_file(
file: UploadFile = File(...), file: UploadFile = File(...),
session_id: Optional[str] = Form(None), session_id: Optional[str] = Form(None),
current_user: ApiPrincipal = Depends(get_current_active_user), current_user: ApiPrincipal = Depends(get_current_active_user),
service: AgentChatService = Depends(get_agent_chat_service),
) -> _SchemaResponse: ) -> _SchemaResponse:
""" """
上传 Web 智能助手对话附件。 上传 Web 智能助手对话附件。
@@ -1712,7 +1733,7 @@ async def upload_web_agent_file(
""" """
mime_type = file.content_type or mimetypes.guess_type(file.filename or "")[0] mime_type = file.content_type or mimetypes.guess_type(file.filename or "")[0]
safe_name = _sanitize_web_agent_upload_name(file.filename, mime_type) safe_name = _sanitize_web_agent_upload_name(file.filename, mime_type)
upload_dir = await _get_web_agent_upload_dir(current_user, session_id) upload_dir = await _get_web_agent_upload_dir(current_user, session_id, service)
target_path = upload_dir / f"{uuid.uuid4().hex[:8]}_{safe_name}" target_path = upload_dir / f"{uuid.uuid4().hex[:8]}_{safe_name}"
size = await _save_web_agent_upload(file, target_path) size = await _save_web_agent_upload(file, target_path)
attachment = _register_web_agent_file( attachment = _register_web_agent_file(
@@ -1846,6 +1867,7 @@ async def get_agent_chat_session(
server_session_id = await _build_web_agent_session_id_async( server_session_id = await _build_web_agent_session_id_async(
current_user, current_user,
session_id, session_id,
service,
) )
if server_session_id != session_id: if server_session_id != session_id:
chat = await _get_accessible_agent_chat( chat = await _get_accessible_agent_chat(
@@ -1884,6 +1906,7 @@ async def save_agent_chat_display(
payload: _SchemaAgentChatDisplaySaveRequest, payload: _SchemaAgentChatDisplaySaveRequest,
current_user: ApiPrincipal = Depends(get_current_active_user), current_user: ApiPrincipal = Depends(get_current_active_user),
service: AgentChatService = Depends(get_agent_chat_service), service: AgentChatService = Depends(get_agent_chat_service),
persistence: AgentChatPersistenceService = Depends(get_agent_chat_persistence),
) -> _SchemaResponse: ) -> _SchemaResponse:
""" """
保存前端聚合后的 Agent 展示消息。 保存前端聚合后的 Agent 展示消息。
@@ -1911,6 +1934,8 @@ async def save_agent_chat_display(
current_user=current_user, current_user=current_user,
messages=messages, messages=messages,
client_session_id=existing_chat.client_session_id if existing_chat else session_id, client_session_id=existing_chat.client_session_id if existing_chat else session_id,
service=service,
persistence=persistence,
) )
chat = await service.get_accessible(session_id, current_user) chat = await service.get_accessible(session_id, current_user)
if not chat: if not chat:
@@ -1964,6 +1989,7 @@ async def stop_web_agent_session_task(
server_session_id = await _build_web_agent_session_id_async( server_session_id = await _build_web_agent_session_id_async(
current_user, current_user,
session_id, session_id,
service,
) )
chat = await _get_accessible_agent_chat( chat = await _get_accessible_agent_chat(
service, service,
@@ -1988,6 +2014,8 @@ async def _web_agent_stream_impl(
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: Optional[AgentChatService] = None,
persistence: Optional[AgentChatPersistenceService] = None,
) -> StreamingResponse: ) -> StreamingResponse:
""" """
Web 智能助手流式对话。 Web 智能助手流式对话。
@@ -1998,11 +2026,18 @@ async def _web_agent_stream_impl(
:return: SSE 流式响应 :return: SSE 流式响应
""" """
prompt = payload.text.strip() prompt = payload.text.strip()
if not isinstance(service, AgentChatService):
# 直接调用公开函数时不经过 FastAPI 依赖解析;生产路由总是传入运行时服务。
service = get_configured_agent_chat_service()
if not isinstance(persistence, AgentChatPersistenceService):
# 直接调用公开函数时不经过 FastAPI 依赖解析;生产路由总是传入运行时端口。
persistence = get_configured_agent_chat_persistence()
locale = LocaleHelper.get_locale_from_request(request) locale = LocaleHelper.get_locale_from_request(request)
display_prompt = (payload.display_text or payload.text).strip() display_prompt = (payload.display_text or payload.text).strip()
session_id = await _build_web_agent_session_id_async( session_id = await _build_web_agent_session_id_async(
current_user, current_user,
payload.session_id, payload.session_id,
service,
) )
is_secret_confirmation_candidate = ( is_secret_confirmation_candidate = (
prompt in {"确认", "取消"} prompt in {"确认", "取消"}
@@ -2110,13 +2145,13 @@ async def _web_agent_stream_impl(
current_user=current_user, current_user=current_user,
messages=display_messages, messages=display_messages,
client_session_id=payload.session_id or session_id, client_session_id=payload.session_id or session_id,
service=service,
persistence=persistence,
) )
except Exception as err: except Exception as err:
logger.error(f"保存WebAgent传统消息快照失败: {str(err)}") logger.error(f"保存WebAgent传统消息快照失败: {str(err)}")
snapshot_task = asyncio.create_task(save_display_snapshot()) snapshot_task = create_web_agent_background_task(save_display_snapshot())
_WEB_AGENT_BACKGROUND_TASKS.add(snapshot_task)
snapshot_task.add_done_callback(_WEB_AGENT_BACKGROUND_TASKS.discard)
await asyncio.sleep(0) await asyncio.sleep(0)
for event in events: for event in events:
event_payload = copy.deepcopy(event) event_payload = copy.deepcopy(event)
@@ -2268,16 +2303,19 @@ async def _web_agent_stream_impl(
# 终态先进入 SSE 队列,避免展示快照落库延迟前端结束动画。 # 终态先进入 SSE 队列,避免展示快照落库延迟前端结束动画。
event_publisher.publish(done_event) event_publisher.publish(done_event)
if not is_secret_confirmation_control: if not is_secret_confirmation_control:
await _save_web_agent_display_snapshot( try:
session_id=session_id, await _save_web_agent_display_snapshot(
current_user=current_user, session_id=session_id,
messages=display_messages, current_user=current_user,
client_session_id=payload.session_id or session_id, messages=display_messages,
) client_session_id=payload.session_id or session_id,
service=service,
persistence=persistence,
)
except Exception as err:
logger.error(f"保存WebAgent展示历史失败:{err}")
task = asyncio.create_task(run_agent()) task = create_web_agent_background_task(run_agent())
_WEB_AGENT_BACKGROUND_TASKS.add(task)
task.add_done_callback(_WEB_AGENT_BACKGROUND_TASKS.discard)
disconnected = False disconnected = False
terminal_sent = False terminal_sent = False
try: try:
@@ -2349,6 +2387,14 @@ 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),
) -> StreamingResponse: ) -> StreamingResponse:
"""Web 智能助手流式对话的稳定公开路由入口。""" """Web 智能助手流式对话的稳定公开路由入口。"""
return await _web_agent_stream_impl(payload, request, current_user) return await _web_agent_stream_impl(
payload,
request,
current_user,
service,
persistence,
)
+22 -1
View File
@@ -1,9 +1,10 @@
import asyncio
import uuid import uuid
from dataclasses import dataclass, field from dataclasses import dataclass, field
from datetime import datetime, timedelta from datetime import datetime, timedelta
from queue import Queue from queue import Queue
from threading import Lock from threading import Lock
from typing import Callable, Dict, Iterable, List, Optional, Tuple, Union from typing import Awaitable, Callable, Dict, Iterable, List, Optional, Tuple, Union
from app.schemas.types import NotificationChannel from app.schemas.types import NotificationChannel
@@ -177,6 +178,26 @@ _WEB_AGENT_EDIT_QUEUES: dict[str, list[Queue[dict]]] = {}
_WEB_AGENT_EDIT_LOCK = Lock() _WEB_AGENT_EDIT_LOCK = Lock()
_ChannelAdminResolver = Callable[[Optional[dict]], Iterable[Union[str, int]]] _ChannelAdminResolver = Callable[[Optional[dict]], Iterable[Union[str, int]]]
_CHANNEL_ADMIN_RESOLVERS: dict[str, _ChannelAdminResolver] = {} _CHANNEL_ADMIN_RESOLVERS: dict[str, _ChannelAdminResolver] = {}
_WEB_AGENT_BACKGROUND_TASKS: set[asyncio.Task[object]] = set()
def create_web_agent_background_task(
coroutine: Awaitable[object],
) -> asyncio.Task[object]:
"""登记 Web Agent 后台任务,使应用关闭时可以统一收口。"""
task = asyncio.create_task(coroutine)
_WEB_AGENT_BACKGROUND_TASKS.add(task)
task.add_done_callback(_WEB_AGENT_BACKGROUND_TASKS.discard)
return task
async def shutdown_web_agent_background_tasks() -> None:
"""取消并等待 Web Agent 后台任务,避免关闭数据库后仍提交快照。"""
tasks = tuple(_WEB_AGENT_BACKGROUND_TASKS)
for task in tasks:
task.cancel()
if tasks:
await asyncio.gather(*tasks, return_exceptions=True)
def register_channel_admin_resolver( def register_channel_admin_resolver(
+49 -9
View File
@@ -10,6 +10,7 @@ from weakref import WeakValueDictionary
from app.application.database import ( from app.application.database import (
AsyncDatabaseExecutor, AsyncDatabaseExecutor,
DatabaseWorkerClosedError,
DatabaseWorkerOverloadedError, DatabaseWorkerOverloadedError,
) )
from app.schemas.agent import AgentChatSessionDetail, AgentChatSessionSummary from app.schemas.agent import AgentChatSessionDetail, AgentChatSessionSummary
@@ -17,6 +18,7 @@ from app.runtime.observability import record_metric
DEFAULT_AGENT_CHAT_WRITE_CAPACITY = 32 DEFAULT_AGENT_CHAT_WRITE_CAPACITY = 32
DEFAULT_AGENT_CHAT_SESSION_CAPACITY = 4
def has_custom_agent_chat_title(value: Optional[str]) -> bool: def has_custom_agent_chat_title(value: Optional[str]) -> bool:
@@ -143,7 +145,8 @@ class SyncAgentChatRepository(Protocol):
... ...
SyncAgentChatRepositoryFactory = Callable[[], SyncAgentChatRepository] SyncAgentChatRepositoryFactory = Callable[[object], SyncAgentChatRepository]
SyncAgentChatTransaction = Callable[[Callable[[object], object]], object]
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
@@ -355,15 +358,24 @@ class AgentChatPersistenceService:
self, self,
repository: SyncAgentChatRepositoryFactory, repository: SyncAgentChatRepositoryFactory,
async_executor: AsyncDatabaseExecutor, async_executor: AsyncDatabaseExecutor,
sync_transaction: SyncAgentChatTransaction,
capacity: int = DEFAULT_AGENT_CHAT_WRITE_CAPACITY, capacity: int = DEFAULT_AGENT_CHAT_WRITE_CAPACITY,
session_capacity: int = DEFAULT_AGENT_CHAT_SESSION_CAPACITY,
) -> None: ) -> None:
"""保存同步仓储工厂和异步执行端口""" """保存同步仓储工厂、事务端口和两级写入容量"""
if capacity < 1: if capacity < 1:
raise ValueError("AgentChat 写入容量必须大于 0") raise ValueError("AgentChat 写入容量必须大于 0")
if session_capacity < 1:
raise ValueError("AgentChat 单会话写入容量必须大于 0")
self._repository = repository self._repository = repository
self._async_executor = async_executor self._async_executor = async_executor
self._sync_transaction = sync_transaction
self._capacity = capacity self._capacity = capacity
self._session_capacity = session_capacity
self._pending_writes = 0 self._pending_writes = 0
self._pending_by_session: dict[str, int] = {}
self._active_tasks: set[asyncio.Task[object]] = set()
self._closing = False
# append_display_messages 属于读取旧快照后整列写回的复合操作;按会话串行化, # append_display_messages 属于读取旧快照后整列写回的复合操作;按会话串行化,
# 才能在 worker 并发下保持首次建行和既有会话追加的完整性。弱引用避免长期运行 # 才能在 worker 并发下保持首次建行和既有会话追加的完整性。弱引用避免长期运行
# 中为一次性会话永久保留锁对象,不限制不同会话之间的 worker 并行度。 # 中为一次性会话永久保留锁对象,不限制不同会话之间的 worker 并行度。
@@ -383,24 +395,52 @@ class AgentChatPersistenceService:
operation: Callable[[SyncAgentChatRepository], object], operation: Callable[[SyncAgentChatRepository], object],
) -> None: ) -> None:
"""在线程 worker 内完成同步写入并丢弃仓储对象返回值。""" """在线程 worker 内完成同步写入并丢弃仓储对象返回值。"""
# 会话锁前的等待也纳入固定总量,避免公开展示保存入口形成无界应用层队列 # 同时限制全局和单会话等待量,避免一个热点会话占满总 admission 后饿死其他会话
if self._pending_writes >= self._capacity: if self._closing:
raise DatabaseWorkerClosedError("AgentChat 持久化服务当前不可接收任务")
session_pending = self._pending_by_session.get(session_id, 0)
if (
self._pending_writes >= self._capacity
or session_pending >= self._session_capacity
):
record_metric("agent.chat.persistence.rejected") record_metric("agent.chat.persistence.rejected")
raise DatabaseWorkerOverloadedError( raise DatabaseWorkerOverloadedError(
f"AgentChat 写入容量已用尽(上限 {self._capacity}" f"AgentChat 写入容量已用尽(全局上限 {self._capacity}"
f"单会话上限 {self._session_capacity}"
) )
self._pending_writes += 1 self._pending_writes += 1
record_metric("agent.chat.persistence.pending", self._pending_writes) self._pending_by_session[session_id] = session_pending + 1
current = asyncio.current_task()
if current is not None:
self._active_tasks.add(current)
record_metric("agent.chat.persistence.pending", 1)
try: try:
async with self._session_lock(session_id): async with self._session_lock(session_id):
def execute() -> None: def execute() -> None:
"""执行同步写入,不让 ORM 对象越过 worker 边界。""" """在单一同步事务中执行写入,不让 ORM 对象越过 worker 边界。"""
operation(self._repository()) self._sync_transaction(
lambda session: operation(self._repository(session))
)
await self._async_executor.run(execute) await self._async_executor.run(execute)
finally: finally:
self._pending_writes -= 1 self._pending_writes -= 1
record_metric("agent.chat.persistence.pending", self._pending_writes) remaining = self._pending_by_session.get(session_id, 1) - 1
if remaining:
self._pending_by_session[session_id] = remaining
else:
self._pending_by_session.pop(session_id, None)
if current is not None:
self._active_tasks.discard(current)
record_metric("agent.chat.persistence.pending", -1)
async def shutdown(self) -> None:
"""拒绝新写入并等待当前会话锁和 worker 操作取得终态。"""
self._closing = True
current = asyncio.current_task()
tasks = tuple(task for task in self._active_tasks if task is not current)
if tasks:
await asyncio.gather(*tasks, return_exceptions=True)
async def async_append_display_messages( async def async_append_display_messages(
self, self,
+2 -1
View File
@@ -269,7 +269,8 @@ class AgentChatOper(DbOper):
) )
if not chat: if not chat:
return None return None
display_messages = self._normalize_messages(chat.display_messages) # JSON 列不是 MutableList;必须复制旧列表,原地 extend 会让 SQLAlchemy 误认为字段未变化。
display_messages = list(self._normalize_messages(chat.display_messages))
display_messages.extend(self._normalize_messages(messages)) display_messages.extend(self._normalize_messages(messages))
title = chat.title if self.has_custom_title(chat.title) else None title = chat.title if self.has_custom_title(chat.title) else None
return self.save_display_messages( return self.save_display_messages(
+2
View File
@@ -6,6 +6,7 @@ from typing import Protocol
from app.application.messaging.chat import ( from app.application.messaging.chat import (
AsyncAgentChatRepository, AsyncAgentChatRepository,
AgentChatPersistenceService,
AsyncUnitOfWork, AsyncUnitOfWork,
) )
from app.application.outbox import AsyncOutboxTransaction from app.application.outbox import AsyncOutboxTransaction
@@ -113,6 +114,7 @@ class AgentChatRuntime:
async_session: AsyncSessionProvider async_session: AsyncSessionProvider
repository: AgentChatRepositoryFactory repository: AgentChatRepositoryFactory
transaction: AsyncUnitOfWorkFactory transaction: AsyncUnitOfWorkFactory
persistence: AgentChatPersistenceService
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
+15 -7
View File
@@ -67,7 +67,9 @@ from app.application.messaging.chat import (
AgentChatService, AgentChatService,
configure_agent_chat_persistence, configure_agent_chat_persistence,
configure_agent_chat_service, configure_agent_chat_service,
get_configured_agent_chat_persistence,
) )
from app.application.messaging.agent import shutdown_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
@@ -576,6 +578,11 @@ 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())
await run_step("Web Agent后台任务", shutdown_web_agent_background_tasks)
await run_step(
"Agent会话持久化",
lambda: get_configured_agent_chat_persistence().shutdown(),
)
await run_step("数据库任务", stop_database_worker) await run_step("数据库任务", stop_database_worker)
if _database_worker is None: if _database_worker is None:
await run_step("数据库连接", close_database) await run_step("数据库连接", close_database)
@@ -642,11 +649,18 @@ async def init_modules() -> HostRuntime:
chain=lambda: build_chain_runtime_config(settings), chain=lambda: build_chain_runtime_config(settings),
) )
runtime_settings = _build_runtime_settings_service() runtime_settings = _build_runtime_settings_service()
agent_chat_persistence = AgentChatPersistenceService(
repository=lambda session: AgentChatOper(session),
async_executor=database_worker,
sync_transaction=transaction_runner.sync,
capacity=database_worker.snapshot().capacity,
)
host_runtime = HostRuntime( host_runtime = HostRuntime(
agent_chat=AgentChatRuntime( agent_chat=AgentChatRuntime(
async_session=get_async_db, async_session=get_async_db,
repository=AgentChatOper, repository=AgentChatOper,
transaction=SqlAlchemyAsyncUnitOfWork, transaction=SqlAlchemyAsyncUnitOfWork,
persistence=agent_chat_persistence,
), ),
persistence=PersistenceRuntime( persistence=PersistenceRuntime(
sync_session=get_db, sync_session=get_db,
@@ -714,13 +728,7 @@ async def init_modules() -> HostRuntime:
) )
configure_database_governance(build_database_governance()) configure_database_governance(build_database_governance())
configure_agent_chat_service(AgentChatService(repository=AgentChatOper())) configure_agent_chat_service(AgentChatService(repository=AgentChatOper()))
configure_agent_chat_persistence( configure_agent_chat_persistence(agent_chat_persistence)
AgentChatPersistenceService(
repository=AgentChatOper,
async_executor=database_worker,
capacity=database_worker.snapshot().capacity,
)
)
configure_user_lookups( configure_user_lookups(
by_id=lambda user_id: UserOper().get_by_id(user_id), by_id=lambda user_id: UserOper().get_by_id(user_id),
by_name=lambda username: UserOper().get_by_name(username), by_name=lambda username: UserOper().get_by_name(username),
+2 -1
View File
@@ -263,8 +263,9 @@ def configure_plugin_system_services():
) )
configure_agent_chat_persistence( configure_agent_chat_persistence(
AgentChatPersistenceService( AgentChatPersistenceService(
repository=AgentChatOper, repository=lambda session: AgentChatOper(session),
async_executor=database_executor, async_executor=database_executor,
sync_transaction=transaction_runner.sync,
) )
) )
configure_agent_chat_service(AgentChatService(repository=AgentChatOper())) configure_agent_chat_service(AgentChatService(repository=AgentChatOper()))
+1
View File
@@ -6097,6 +6097,7 @@
"app.startup.modules_initializer -> app.application.history", "app.startup.modules_initializer -> app.application.history",
"app.startup.modules_initializer -> app.application.image", "app.startup.modules_initializer -> app.application.image",
"app.startup.modules_initializer -> app.application.messaging", "app.startup.modules_initializer -> app.application.messaging",
"app.startup.modules_initializer -> app.application.messaging.agent",
"app.startup.modules_initializer -> app.application.messaging.chat", "app.startup.modules_initializer -> app.application.messaging.chat",
"app.startup.modules_initializer -> app.application.messaging.message", "app.startup.modules_initializer -> app.application.messaging.message",
"app.startup.modules_initializer -> app.application.module", "app.startup.modules_initializer -> app.application.module",
+186 -6
View File
@@ -4,16 +4,24 @@ from __future__ import annotations
import asyncio import asyncio
import threading import threading
from types import SimpleNamespace
from unittest.mock import AsyncMock, call, patch
from uuid import uuid4 from uuid import uuid4
import pytest import pytest
from sqlalchemy import delete, select from sqlalchemy import delete, select
from app.application.database import DatabaseWorkerOverloadedError from app.application.database import (
DatabaseWorkerClosedError,
DatabaseWorkerOverloadedError,
)
from app.application.messaging.chat import AgentChatPersistenceService, AgentChatService from app.application.messaging.chat import AgentChatPersistenceService, AgentChatService
from app.api.endpoints.agent import save_agent_chat_display
from app.db.models.agentchat import AgentChat from app.db.models.agentchat import AgentChat
from app.db.oper.agentchat import AgentChatOper from app.db.oper.agentchat import AgentChatOper
from app.db.session import SessionFactory, async_session_scope from app.db.session import SessionFactory, async_session_scope
from app.db.uow import run_sync_transaction
from app.schemas.agent import AgentChatDisplaySaveRequest
from app.db.worker import DatabaseWorker from app.db.worker import DatabaseWorker
@@ -65,8 +73,9 @@ async def test_agent_chat_persistence_runs_sync_repository_inside_worker() -> No
executor = _Executor() executor = _Executor()
repository = _Repository() repository = _Repository()
service = AgentChatPersistenceService( service = AgentChatPersistenceService(
repository=lambda: repository, repository=lambda _session: repository,
async_executor=executor, async_executor=executor,
sync_transaction=lambda operation: operation(object()),
) )
caller_thread_id = threading.get_ident() caller_thread_id = threading.get_ident()
@@ -111,8 +120,9 @@ async def test_agent_chat_persistence_propagates_worker_failure() -> None:
raise RuntimeError("worker failed") raise RuntimeError("worker failed")
service = AgentChatPersistenceService( service = AgentChatPersistenceService(
repository=_Repository, repository=lambda _session: _Repository(),
async_executor=FailingExecutor(), async_executor=FailingExecutor(),
sync_transaction=lambda operation: operation(object()),
) )
with pytest.raises(RuntimeError, match="worker failed"): with pytest.raises(RuntimeError, match="worker failed"):
@@ -123,6 +133,83 @@ async def test_agent_chat_persistence_propagates_worker_failure() -> None:
) )
@pytest.mark.asyncio
async def test_agent_chat_persistence_pending_metric_uses_deltas() -> None:
"""pending 是 UpDownCounter,准入和释放必须分别记录增减量。"""
service = AgentChatPersistenceService(
repository=lambda _session: _Repository(),
async_executor=_Executor(),
sync_transaction=lambda operation: operation(object()),
)
with patch("app.application.messaging.chat.record_metric") as record_metric:
await service.async_save_agent_messages(
session_id="metric-session",
user_id="1",
messages=[],
)
record_metric.assert_has_calls(
[
call("agent.chat.persistence.pending", 1),
call("agent.chat.persistence.pending", -1),
]
)
@pytest.mark.asyncio
async def test_authoritative_display_save_propagates_worker_overload() -> None:
"""权威 PUT 保存不能把 worker 背压吞成成功或普通业务失败。"""
repository = AsyncMock()
repository.async_get.return_value = None
service = AgentChatService(repository=repository)
class OverloadedPersistence:
async def async_save_display_messages(self, **_kwargs):
raise DatabaseWorkerOverloadedError("busy")
with pytest.raises(DatabaseWorkerOverloadedError, match="busy"):
await save_agent_chat_display(
session_id="overloaded-session",
payload=AgentChatDisplaySaveRequest(messages=[]),
current_user=SimpleNamespace(id=1, name="admin", is_superuser=True),
service=service,
persistence=OverloadedPersistence(),
)
@pytest.mark.asyncio
async def test_agent_chat_persistence_rolls_back_compound_write(monkeypatch) -> None:
"""复合写入中途失败时,创建或更新不能留下半成品。"""
worker = DatabaseWorker(max_workers=1, capacity=4)
await worker.start()
session_id = f"worker-rollback-{uuid4().hex}"
persistence = AgentChatPersistenceService(
repository=lambda session: AgentChatOper(session),
async_executor=worker,
sync_transaction=run_sync_transaction,
)
original = AgentChatOper.save_display_messages
def fail_after_stage(self, *args, **kwargs):
original(self, *args, **kwargs)
raise RuntimeError("display snapshot failed")
monkeypatch.setattr(AgentChatOper, "save_display_messages", fail_after_stage)
try:
with pytest.raises(RuntimeError, match="display snapshot failed"):
await persistence.async_append_display_messages(
session_id=session_id,
user_id="rollback-user",
messages=[{"role": "user", "content": "not committed"}],
)
async with async_session_scope() as session:
result = await session.execute(
select(AgentChat).where(AgentChat.session_id == session_id)
)
assert result.scalars().first() is None
finally:
await worker.shutdown()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_agent_chat_persistence_bounds_session_waiters_and_releases_cancelled() -> None: async def test_agent_chat_persistence_bounds_session_waiters_and_releases_cancelled() -> None:
"""同会话锁等待受总量限制,取消等待不会遗留 admission。""" """同会话锁等待受总量限制,取消等待不会遗留 admission。"""
@@ -139,9 +226,11 @@ async def test_agent_chat_persistence_bounds_session_waiters_and_releases_cancel
executor = BlockingExecutor() executor = BlockingExecutor()
service = AgentChatPersistenceService( service = AgentChatPersistenceService(
repository=_Repository, repository=lambda _session: _Repository(),
async_executor=executor, async_executor=executor,
sync_transaction=lambda operation: operation(object()),
capacity=2, capacity=2,
session_capacity=2,
) )
first = asyncio.create_task( first = asyncio.create_task(
service.async_save_agent_messages( service.async_save_agent_messages(
@@ -177,6 +266,95 @@ async def test_agent_chat_persistence_bounds_session_waiters_and_releases_cancel
assert service._pending_writes == 0 assert service._pending_writes == 0
@pytest.mark.asyncio
async def test_agent_chat_persistence_session_admission_is_fair() -> None:
"""热点会话的锁等待不能占满全局容量并拒绝其他会话。"""
class BlockingExecutor:
def __init__(self) -> None:
self.started = asyncio.Event()
self.release = asyncio.Event()
async def run(self, operation):
self.started.set()
await self.release.wait()
return operation()
executor = BlockingExecutor()
service = AgentChatPersistenceService(
repository=lambda _session: _Repository(),
async_executor=executor,
sync_transaction=lambda operation: operation(object()),
capacity=4,
session_capacity=2,
)
first = asyncio.create_task(
service.async_save_agent_messages(
session_id="hot-session", user_id="1", messages=[]
)
)
await executor.started.wait()
second = asyncio.create_task(
service.async_save_agent_messages(
session_id="hot-session", user_id="1", messages=[]
)
)
await asyncio.sleep(0)
with pytest.raises(DatabaseWorkerOverloadedError):
await service.async_save_agent_messages(
session_id="hot-session", user_id="1", messages=[]
)
other = asyncio.create_task(
service.async_save_agent_messages(
session_id="other-session", user_id="1", messages=[]
)
)
await asyncio.sleep(0)
assert not other.done()
executor.release.set()
await first
await second
await other
@pytest.mark.asyncio
async def test_agent_chat_persistence_shutdown_drains_active_writes() -> None:
"""关闭持久化端口时拒绝新写入并等待现有会话写入收口。"""
class BlockingExecutor:
def __init__(self) -> None:
self.started = asyncio.Event()
self.release = asyncio.Event()
async def run(self, operation):
self.started.set()
await self.release.wait()
return operation()
executor = BlockingExecutor()
service = AgentChatPersistenceService(
repository=lambda _session: _Repository(),
async_executor=executor,
sync_transaction=lambda operation: operation(object()),
)
write = asyncio.create_task(
service.async_save_agent_messages(
session_id="shutdown-session", user_id="1", messages=[]
)
)
await executor.started.wait()
shutdown = asyncio.create_task(service.shutdown())
await asyncio.sleep(0)
assert not shutdown.done()
with pytest.raises(DatabaseWorkerClosedError):
await service.async_save_agent_messages(
session_id="new-session", user_id="1", messages=[]
)
executor.release.set()
await write
await shutdown
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_agent_chat_persistence_uses_real_worker_and_sqlite_transaction() -> None: async def test_agent_chat_persistence_uses_real_worker_and_sqlite_transaction() -> None:
"""真实 AgentChat Oper 经 worker 写入后可被 native async 查询恢复。""" """真实 AgentChat Oper 经 worker 写入后可被 native async 查询恢复。"""
@@ -184,8 +362,9 @@ async def test_agent_chat_persistence_uses_real_worker_and_sqlite_transaction()
await worker.start() await worker.start()
session_id = f"worker-{uuid4().hex}" session_id = f"worker-{uuid4().hex}"
persistence = AgentChatPersistenceService( persistence = AgentChatPersistenceService(
repository=AgentChatOper, repository=lambda session: AgentChatOper(session),
async_executor=worker, async_executor=worker,
sync_transaction=run_sync_transaction,
) )
query = AgentChatService(repository=AgentChatOper()) query = AgentChatService(repository=AgentChatOper())
@@ -221,8 +400,9 @@ async def test_agent_chat_persistence_serializes_same_session_writes() -> None:
session_id = f"worker-race-{uuid4().hex}" session_id = f"worker-race-{uuid4().hex}"
existing_session_id = f"worker-race-existing-{uuid4().hex}" existing_session_id = f"worker-race-existing-{uuid4().hex}"
persistence = AgentChatPersistenceService( persistence = AgentChatPersistenceService(
repository=AgentChatOper, repository=lambda session: AgentChatOper(session),
async_executor=worker, async_executor=worker,
sync_transaction=run_sync_transaction,
) )
async def append(content: str) -> None: async def append(content: str) -> None:
+26
View File
@@ -4,6 +4,10 @@ from unittest.mock import AsyncMock, MagicMock
import pytest import pytest
import app.agent.orchestrator as agent_module import app.agent.orchestrator as agent_module
from app.application.messaging.agent import (
create_web_agent_background_task,
shutdown_web_agent_background_tasks,
)
from app.agent.orchestrator import ( from app.agent.orchestrator import (
AGENT_SESSION_QUEUE_MAX_SIZE, AGENT_SESSION_QUEUE_MAX_SIZE,
AgentManager, AgentManager,
@@ -14,6 +18,28 @@ from app.agent.memory import MemoryManager
from app.startup import agent_initializer, modules_initializer from app.startup import agent_initializer, modules_initializer
@pytest.mark.anyio
async def test_web_agent_background_tasks_are_cancelled_and_drained() -> None:
"""Web Agent 任务关闭后不得继续占用循环或提交晚到的快照。"""
started = asyncio.Event()
finished = asyncio.Event()
async def blocked_task() -> None:
started.set()
try:
await asyncio.Event().wait()
finally:
finished.set()
task = create_web_agent_background_task(blocked_task())
await started.wait()
await shutdown_web_agent_background_tasks()
assert task.done()
assert task.cancelled()
assert finished.is_set()
@pytest.mark.anyio @pytest.mark.anyio
async def test_agent_entrypoint_initializes_on_calling_loop(monkeypatch) -> None: async def test_agent_entrypoint_initializes_on_calling_loop(monkeypatch) -> None:
"""Agent 启动入口必须在应用主循环完成初始化。""" """Agent 启动入口必须在应用主循环完成初始化。"""
+12 -2
View File
@@ -13,6 +13,7 @@ from app.api.context import (
get_agent_chat_repository, get_agent_chat_repository,
get_agent_chat_transaction, get_agent_chat_transaction,
) )
from app.api.dependencies.agent import get_agent_chat_persistence
from app.startup import lifecycle from app.startup import lifecycle
from app.startup.context import ( from app.startup.context import (
AgentChatRuntime, AgentChatRuntime,
@@ -59,6 +60,10 @@ class _UnitOfWork:
"""模拟回滚。""" """模拟回滚。"""
class _AgentChatPersistence:
"""提供 AgentChat 运行时所需的最小写端口。"""
class _SyncUnitOfWork: class _SyncUnitOfWork:
"""记录绑定会话的同步事务替身。""" """记录绑定会话的同步事务替身。"""
@@ -119,6 +124,7 @@ def _runtime() -> HostRuntime:
async_session=async_session, async_session=async_session,
repository=_Repository, repository=_Repository,
transaction=_UnitOfWork, transaction=_UnitOfWork,
persistence=_AgentChatPersistence(),
), ),
persistence=PersistenceRuntime( persistence=PersistenceRuntime(
sync_session=sync_session, sync_session=sync_session,
@@ -186,15 +192,19 @@ def test_fastapi_dependencies_use_fake_runtime_without_real_services() -> None:
async def probe( async def probe(
repository=Depends(get_agent_chat_repository), repository=Depends(get_agent_chat_repository),
unit_of_work=Depends(get_agent_chat_transaction), unit_of_work=Depends(get_agent_chat_transaction),
persistence=Depends(get_agent_chat_persistence),
) -> dict[str, bool]: ) -> dict[str, bool]:
"""返回两个类型化能力是否绑定同一请求会话。""" """返回两个类型化能力是否绑定同一请求会话。"""
return {"same_session": repository.session is unit_of_work.session} return {
"same_session": repository.session is unit_of_work.session,
"has_persistence": persistence is app.state.host_runtime.agent_chat.persistence,
}
with TestClient(app) as client: with TestClient(app) as client:
response = client.get("/probe") response = client.get("/probe")
assert response.status_code == 200 assert response.status_code == 200
assert response.json() == {"same_session": True} assert response.json() == {"same_session": True, "has_persistence": True}
def test_official_api_dependencies_do_not_use_string_data_locator() -> None: def test_official_api_dependencies_do_not_use_string_data_locator() -> None:
+29
View File
@@ -570,6 +570,35 @@ def test_stop_modules_continues_after_internal_owner_failures(monkeypatch):
_assert_completed_once(dependency) _assert_completed_once(dependency)
def test_stop_modules_drains_web_agent_tasks_before_persistence(monkeypatch):
"""关闭时必须先收口 Web Agent 后台任务,再关闭会话持久化端口。"""
order = []
monkeypatch.setattr(modules_initializer, "stop_agent", AsyncMock())
dependencies = _patch_module_shutdown_dependencies(monkeypatch)
monkeypatch.setattr(
modules_initializer,
"shutdown_web_agent_background_tasks",
AsyncMock(side_effect=lambda: order.append("web-agent")),
)
persistence = MagicMock()
persistence.shutdown = AsyncMock(side_effect=lambda: order.append("persistence"))
monkeypatch.setattr(
modules_initializer,
"get_configured_agent_chat_persistence",
MagicMock(return_value=persistence),
)
monkeypatch.setattr(
modules_initializer,
"stop_database_worker",
AsyncMock(side_effect=lambda: order.append("database")),
)
monkeypatch.setattr(modules_initializer, "_database_worker", object())
asyncio.run(modules_initializer.stop_modules())
assert order == ["web-agent", "persistence", "database"]
def _patch_module_shutdown_dependencies(monkeypatch) -> dict: def _patch_module_shutdown_dependencies(monkeypatch) -> dict:
"""替换 stop_modules 的资源所有者,避免测试启动真实后台服务""" """替换 stop_modules 的资源所有者,避免测试启动真实后台服务"""
dependencies = {} dependencies = {}