From 6e69258e3ce56a75a8b4f8543fab930082e77700 Mon Sep 17 00:00:00 2001 From: jxxghp Date: Sun, 23 Aug 2026 23:07:56 +0800 Subject: [PATCH] refactor: make model sessions explicit --- app/db/base.py | 66 ++-- app/db/decorators.py | 153 +-------- app/db/models/agentchat.py | 5 - app/db/models/agenttask.py | 42 +-- app/db/models/agenttaskrun.py | 3 - app/db/models/downloadhistory.py | 17 - app/db/models/mediaserver.py | 8 - app/db/models/message.py | 118 +++---- app/db/models/passkey.py | 60 +--- app/db/models/plugindata.py | 43 +-- app/db/models/site.py | 127 ++----- app/db/models/siteicon.py | 17 +- app/db/models/sitestatistic.py | 19 +- app/db/models/siteuserdata.py | 6 - app/db/models/subscribe.py | 317 ++++++----------- app/db/models/subscribehistory.py | 10 +- app/db/models/systemconfig.py | 7 +- app/db/models/transferhistory.py | 109 ++---- app/db/models/transferpending.py | 2 - app/db/models/user.py | 75 +---- app/db/models/workflow.py | 9 - app/db/oper/agentchat.py | 4 + app/db/oper/downloadfailure.py | 24 +- app/db/oper/downloadhistory.py | 2 +- app/db/oper/plugindata.py | 46 ++- app/db/oper/site.py | 95 +++--- app/db/oper/subscribe.py | 253 +++++--------- app/db/oper/systemconfig.py | 11 +- app/db/oper/transferhistory.py | 2 +- app/db/oper/user.py | 16 +- app/db/oper/userconfig.py | 11 +- app/db/oper/workflow.py | 11 + app/db/uow.py | 12 +- app/sdk/_legacy/history.py | 7 +- app/sdk/_legacy/subscribe.py | 3 +- app/sdk/_legacy/user.py | 2 - app/testing/bootstrap.py | 6 +- docs/architecture-overview.md | 9 +- .../backend-architecture-governance.md | 4 +- .../backend-architecture-next-stage.md | 94 ++---- docs/rules/10-data-and-persistent.md | 30 +- tests/conftest.py | 35 +- .../architecture/dependency-baseline.json | 31 +- .../startup-performance-baseline.json | 318 +++++++++--------- tests/test_agent_task_runs.py | 23 +- tests/test_architecture_dependencies.py | 88 +++-- tests/test_db_base_crud.py | 39 ++- tests/test_db_config_user_queries.py | 121 +++---- tests/test_db_downloadhistory_queries.py | 22 +- tests/test_db_media_identity_normalizer.py | 1 + tests/test_db_mediaserver_queries.py | 55 +-- tests/test_db_plugin_message_agent_queries.py | 83 +++-- tests/test_db_public_api.py | 4 +- tests/test_db_site_queries.py | 67 ++-- tests/test_db_subscribe_queries.py | 100 +++--- tests/test_db_transferhistory_queries.py | 25 +- tests/test_db_transferpending_queries.py | 29 +- tests/test_db_workflow_queries.py | 53 ++- tests/test_mediaserver_sync_incremental.py | 64 ++-- tests/test_subscribe_create_command.py | 12 +- tests/test_subscribe_oper.py | 99 +++--- tests/test_systemconfig_oper.py | 17 +- ...transfer_download_history_oper_sessions.py | 51 +++ ..._transfer_download_history_query_compat.py | 111 ------ tests/test_web_agent_stream.py | 6 +- 65 files changed, 1299 insertions(+), 2010 deletions(-) create mode 100644 tests/test_transfer_download_history_oper_sessions.py delete mode 100644 tests/test_transfer_download_history_query_compat.py diff --git a/app/db/base.py b/app/db/base.py index f28043057..388d16703 100644 --- a/app/db/base.py +++ b/app/db/base.py @@ -1,7 +1,7 @@ """ ORM 基类与数据访问基类。 -Base 提供声明式基类与兼容行为(字典转换、旧增删改查便利方法); +Base 提供声明式基类与显式会话增删改查原语; DbOper 是各业务 Oper 的基类,持有一个可注入的会话。 """ from collections.abc import Awaitable, Callable @@ -12,12 +12,6 @@ from sqlalchemy import (CursorResult, Executable, Identity, Integer, Sequence, from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import DeclarativeBase, Mapped, Session, declared_attr, mapped_column -from app.db.decorators import ( - legacy_async_db_query, - legacy_async_db_update, - legacy_db_query, - legacy_db_update, -) from app.db.uow import run_async_transaction, run_sync_transaction from app.runtime.config import settings @@ -70,98 +64,87 @@ class Base(DeclarativeBase): # type: ignore[misc] # SQLAlchemy 无 py.typed 继承本类的模型一律使用 mapped_column() + Mapped[] 注解;确需非映射的类级属性时 用 ClassVar 显式声明,而不是把这个标志加回来。 - create/get/update/delete/list/truncate 及其异步版本仅保留旧插件 ABI。宿主新代码应由 - Application Command 定义事务边界,经显式 Session 调用 Oper,不得新增对这些方法的依赖。 + create/get/update/delete/list/truncate 及其异步版本都是显式会话原语:只在调用方 + Session 中暂存或查询,不自行创建、提交、回滚或关闭事务。宿主业务代码应通过 Oper + 或 Application Command 使用这些能力,插件不得直接依赖宿主模型。 """ # 由 get_id_column() 在各模型中提供实际的列定义,这里只声明类型供 IDE 使用 id: Mapped[int] - @legacy_db_update def create(self, db: Session) -> None: - """兼容旧插件调用:新增当前模型并提交。""" + """在调用方同步事务中暂存当前模型。""" db.add(self) - @legacy_async_db_update async def async_create(self, db: AsyncSession) -> Self: - """兼容旧插件调用:异步新增当前模型、刷新主键并提交。""" + """在调用方异步事务中暂存当前模型并刷新主键。""" db.add(self) await db.flush() return self @classmethod - @legacy_db_query def get(cls, db: Session, rid: int) -> Optional[Self]: - """兼容旧插件调用:按主键查询当前模型。""" + """在调用方同步会话中按主键查询当前模型。""" return cast( Optional[Self], db.execute(select(cls).where(and_(cls.id == rid))).scalars().first(), ) @classmethod - @legacy_async_db_query async def async_get(cls, db: AsyncSession, rid: int) -> Optional[Self]: - """兼容旧插件调用:异步按主键查询当前模型。""" + """在调用方异步会话中按主键查询当前模型。""" result = await db.execute(select(cls).where(and_(cls.id == rid))) return cast(Optional[Self], result.scalars().first()) - @legacy_db_update def update(self, db: Session, payload: dict[str, Any]) -> None: - """兼容旧插件调用:更新当前模型字段并提交。""" + """在调用方同步事务中更新当前模型字段。""" for key, value in payload.items(): setattr(self, key, value) if inspect(self).detached: db.add(self) - @legacy_async_db_update async def async_update( self, db: AsyncSession, payload: dict[str, Any], ) -> None: - """兼容旧插件调用:异步更新当前模型字段并提交。""" + """在调用方异步事务中更新当前模型字段。""" for key, value in payload.items(): setattr(self, key, value) if inspect(self).detached: db.add(self) @classmethod - @legacy_db_update def delete(cls, db: Session, rid: Any) -> None: - """兼容旧插件调用:按主键删除当前模型并提交。""" + """在调用方同步事务中按主键删除当前模型。""" db.execute(delete(cls).where(and_(cls.id == rid))) @classmethod - @legacy_async_db_update async def async_delete(cls, db: AsyncSession, rid: Any) -> None: - """兼容旧插件调用:异步按主键删除当前模型并提交。""" + """在调用方异步事务中按主键删除当前模型。""" result = await db.execute(select(cls).where(and_(cls.id == rid))) user = result.scalars().first() if user: await db.delete(user) @classmethod - @legacy_db_update def truncate(cls, db: Session) -> None: - """兼容旧插件调用:清空当前模型表并提交。""" + """在调用方同步事务中清空当前模型表。""" db.execute(delete(cls)) @classmethod - @legacy_async_db_update async def async_truncate(cls, db: AsyncSession) -> None: - """兼容旧插件调用:异步清空当前模型表并提交。""" + """在调用方异步事务中清空当前模型表。""" await db.execute(delete(cls)) @classmethod - @legacy_db_query def list(cls, db: Session) -> List[Self]: - """兼容旧插件调用:查询当前模型的全部记录。""" + """在调用方同步会话中查询当前模型的全部记录。""" return list(db.execute(select(cls)).scalars().all()) @classmethod - @legacy_async_db_query async def async_list(cls, db: AsyncSession) -> List[Self]: - """兼容旧插件调用:异步查询当前模型的全部记录。""" + """在调用方异步会话中查询当前模型的全部记录。""" result = await db.execute(select(cls)) return list(result.scalars().all()) @@ -183,19 +166,19 @@ class DbOper: """ def __init__(self, db: Optional[Union[Session, AsyncSession]] = None): - """保存调用方会话;无会话写入由组合根兼容事务执行器承接。""" + """保存调用方会话;无会话调用由组合根事务执行器承接。""" self._db = db def _execute_sync_write(self, operation: Callable[[Session], T]) -> T: - """在当前同步会话暂存,或委托组合根创建兼容事务。""" + """在当前同步会话暂存,或委托组合根创建事务。""" if self._db is None or isinstance(self._db, AsyncSession): # 旧调用可能在同一 Oper 上混用同步/异步方法;跨会话类型时使用匹配的 - # 兼容事务,不能把 AsyncSession 交给同步 SQLAlchemy API。 + # 独立事务,不能把 AsyncSession 交给同步 SQLAlchemy API。 return run_sync_transaction(operation) return operation(self._db) def _execute_sync_query(self, operation: Callable[[Session], T]) -> T: - """在当前同步会话查询,或委托组合根创建一次性兼容会话。""" + """在当前同步会话查询,或委托组合根创建一次性会话。""" if self._db is None or isinstance(self._db, AsyncSession): return run_sync_transaction(operation) return operation(self._db) @@ -204,10 +187,9 @@ class DbOper: self, operation: Callable[[AsyncSession], Awaitable[T]], ) -> T: - """在当前异步会话暂存,或委托组合根创建兼容事务。""" + """在当前异步会话暂存,或委托组合根创建事务。""" if self._db is None or isinstance(self._db, Session): - # 与查询装饰器的历史行为一致:同步会话不会被错误传入异步模型写入, - # 而是由组合根另开匹配的异步事务。 + # 同步会话不会被错误传入异步模型写入,而是由组合根另开匹配的异步事务。 return await run_async_transaction(operation) return await operation(self._db) @@ -215,13 +197,13 @@ class DbOper: self, operation: Callable[[AsyncSession], Awaitable[T]], ) -> T: - """在当前异步会话查询,或委托组合根创建一次性兼容会话。""" + """在当前异步会话查询,或委托组合根创建一次性会话。""" if self._db is None or isinstance(self._db, Session): return await run_async_transaction(operation) return await operation(self._db) def _stage_create(self, model: TModel) -> TModel: - """在显式同步事务中暂存新模型,不触发 Base 的兼容提交装饰器。""" + """在调用方或组合根持有的同步事务中暂存新模型。""" def stage(session: Session) -> TModel: """把模型加入当前同步会话。""" session.add(model) diff --git a/app/db/decorators.py b/app/db/decorators.py index 3a858418e..33f3bd0c7 100644 --- a/app/db/decorators.py +++ b/app/db/decorators.py @@ -5,8 +5,8 @@ 未显式传入会话时自动创建,并在结束时归还——异步路径经 async_session_scope 收口, 连接池与配额都在那里生效。 -收尾故障(rollback / close / __aexit__ 自身抛异常)一律只记日志、不上抛,正式装饰器 -和 legacy 兼容壳的处理一致。理由与代价都要写明,别当成漏写的 raise: +收尾故障(rollback / close / __aexit__ 自身抛异常)一律只记日志、不上抛。理由与代价 +都要写明,别当成漏写的 raise: - 连接断开、事务已失效这类故障恰恰最容易发生在「出错之后」的收尾阶段。裸写收尾语句时 它一抛错就顶替掉原始异常,调用方看到的只剩「connection reset」,业务异常连类型都被 @@ -16,8 +16,6 @@ SQLAlchemy 归还连接时已在池层吞掉异常并 invalidate 坏连接,再把释放故障升级成调用方 的异常,只会让一次已经落库的写入看起来像失败,诱发重复提交。 """ -from functools import wraps -from inspect import Parameter, signature from typing import Any, Awaitable, Callable, Optional, TypeVar from sqlalchemy.ext.asyncio import AsyncSession @@ -279,150 +277,3 @@ def async_db_query(func: Callable[..., Awaitable[_R]]) -> Callable[..., Awaitabl return result return wrapper - - -def legacy_db_query(func: Callable[..., _R]) -> Callable[..., _R]: - """保留旧 Model 查询 ABI,同时让新调用方复用显式 Session。 - - 旧插件通常省略 ``db``,直接把业务参数放在第一个位置;通用 ``db_query`` - 装饰器只适用于固定的 ``(db, ...)`` 形状,不能把这类位置参数直接套进去。 - 这里按签名插入会话,避免丢失旧插件传入的第一个业务参数。 - """ - - @wraps(func) - def wrapper(*args: Any, **kwargs: Any) -> _R: - db = _get_args_db(args, kwargs) - if db is not None: - return func(*args, **kwargs) - - session = ScopedSession() - call_args, call_kwargs = _inject_legacy_db(func, args, kwargs, session) - try: - return func(*call_args, **call_kwargs) - finally: - try: - session.close() - except Exception as close_err: # noqa: BLE001 释放故障不得改变旧 ABI 返回值 - logger.error(f"释放数据库会话失败:{close_err}") - - return wrapper - - -def legacy_async_db_query( - func: Callable[..., Awaitable[_R]], -) -> Callable[..., Awaitable[_R]]: - """保留旧 Model 异步查询 ABI,同时让新调用方复用显式 AsyncSession。""" - - @wraps(func) - async def wrapper(*args: Any, **kwargs: Any) -> _R: - db = _get_args_async_db(args, kwargs) - if db is not None: - return await func(*args, **kwargs) - - async with async_session_scope() as session: - call_args, call_kwargs = _inject_legacy_db(func, args, kwargs, session) - return await func(*call_args, **call_kwargs) - - return wrapper - - -def legacy_db_update(func: Callable[..., _R]) -> Callable[..., _R]: - """保留旧 Model 同步写 ABI,并维持历史自动提交语义。 - - 该装饰器只供已经公开的 Model/Base 方法兼容仓外插件。宿主新写路径必须 - 通过 Application Command、显式 Session 和 UnitOfWork 完成事务收口。 - """ - - @wraps(func) - def wrapper(*args: Any, **kwargs: Any) -> _R: - db = _get_args_db(args, kwargs) - owns_session = db is None - if db is None: - db = ScopedSession() - args, kwargs = _inject_legacy_db(func, args, kwargs, db) - try: - result = func(*args, **kwargs) - db.commit() - return result - except Exception: - try: - db.rollback() - except Exception as rollback_err: # noqa: BLE001 回滚失败不能掩盖原始异常 - logger.error(f"事务回滚失败,原始异常将原样上抛:{rollback_err}") - raise - finally: - if owns_session: - try: - db.close() - except Exception as close_err: # noqa: BLE001 释放故障不得改变旧 ABI 结果 - logger.error(f"释放数据库会话失败:{close_err}") - - return wrapper - - -def legacy_async_db_update( - func: Callable[..., Awaitable[_R]], -) -> Callable[..., Awaitable[_R]]: - """保留旧 Model 异步写 ABI,并维持历史自动提交语义。 - - 该装饰器只承接既有兼容面;新宿主代码不得用它创建隐式事务。 - """ - - @wraps(func) - async def wrapper(*args: Any, **kwargs: Any) -> _R: - db = _get_args_async_db(args, kwargs) - owns_session = db is None - scope = None - if db is None: - scope = async_session_scope() - db = await scope.__aenter__() - args, kwargs = _inject_legacy_db(func, args, kwargs, db) - try: - result = await func(*args, **kwargs) - await db.commit() - return result - except Exception: - try: - await db.rollback() - except Exception as rollback_err: # noqa: BLE001 回滚失败不能掩盖原始异常 - logger.error(f"事务回滚失败,原始异常将原样上抛:{rollback_err}") - raise - finally: - if owns_session and scope is not None: - try: - await scope.__aexit__(None, None, None) - except Exception as close_err: # noqa: BLE001 释放故障不得改变旧 ABI 结果 - logger.error(f"释放数据库会话失败:{close_err}") - - return wrapper - - -def _inject_legacy_db( - func: Callable[..., _R], - args: tuple[Any, ...], - kwargs: dict[str, Any], - db: Any, -) -> tuple[tuple[Any, ...], dict[str, Any]]: - """按旧 Model 方法签名注入兼容会话,不吞掉位置业务参数。""" - call_args = list(args) - call_kwargs = dict(kwargs) - parameters = list(signature(func).parameters.values()) - db_index = next( - (index for index, parameter in enumerate(parameters) if parameter.name == "db"), - None, - ) - if "db" in call_kwargs: - call_kwargs["db"] = db - return tuple(call_args), call_kwargs - if db_index is None: - # 兼容没有显式 db 参数的极旧函数,保持调用失败方式与普通 Python 一致。 - return tuple(call_args), {"db": db, **call_kwargs} - if db_index < len(call_args) and call_args[db_index] is None: - call_args[db_index] = db - elif db_index < len(parameters) and parameters[db_index].kind is Parameter.POSITIONAL_ONLY: - call_args.insert(db_index, db) - elif db_index <= len(call_args): - call_args.insert(db_index, db) - else: - call_kwargs["db"] = db - return tuple(call_args), call_kwargs diff --git a/app/db/models/agentchat.py b/app/db/models/agentchat.py index 1537d5a9b..3ecdc7071 100644 --- a/app/db/models/agentchat.py +++ b/app/db/models/agentchat.py @@ -5,7 +5,6 @@ from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Mapped, Session, mapped_column from app.db.base import Base, get_id_column -from app.db.decorators import legacy_async_db_query, legacy_db_query class AgentChat(Base): @@ -50,7 +49,6 @@ class AgentChat(Base): ) @classmethod - @legacy_db_query def get_by_session( cls, db: Session, session_id: str, user_id: Optional[str] = None ) -> Optional["AgentChat"]: @@ -63,7 +61,6 @@ class AgentChat(Base): return db.execute(statement.order_by(cls.id.desc())).scalars().first() @classmethod - @legacy_async_db_query async def async_get_by_session( cls, db: AsyncSession, session_id: str, user_id: Optional[str] = None ) -> Optional["AgentChat"]: @@ -77,7 +74,6 @@ class AgentChat(Base): return result.scalars().first() @classmethod - @legacy_db_query def list_by_page( cls, db: Session, @@ -103,7 +99,6 @@ class AgentChat(Base): ).scalars().all()) @classmethod - @legacy_async_db_query async def async_list_by_page( cls, db: AsyncSession, diff --git a/app/db/models/agenttask.py b/app/db/models/agenttask.py index 0073dbccb..49364b477 100644 --- a/app/db/models/agenttask.py +++ b/app/db/models/agenttask.py @@ -4,7 +4,6 @@ from sqlalchemy import Boolean, Index, Integer, String, Text, select, update from sqlalchemy.orm import Mapped, Session, mapped_column from app.db.base import Base, execute_dml, get_id_column -from app.db.decorators import legacy_db_query def _get_for_user_statement( @@ -85,47 +84,28 @@ class AgentTask(Base): return task.id @classmethod - @legacy_db_query def get_for_user( cls, - db: Session | int | None = None, - task_id: int | None = None, + db: Session, + task_id: int, user_id: Optional[str] = None, ) -> Optional["AgentTask"]: - """ - 按任务 ID 和可选用户 ID 查询,并保留无 Session 的旧插件调用方式。 - """ - if task_id is None and isinstance(db, int): - task_id, db = db, None - if task_id is None: - raise TypeError("task_id is required") - - def query(session: Session) -> Optional["AgentTask"]: - """在给定会话中读取单个 Agent 任务。""" - return session.execute( - _get_for_user_statement(cls, task_id=task_id, user_id=user_id) - ).scalars().first() - - return query(db) + """在调用方会话中按任务 ID 和可选用户 ID 查询。""" + return db.execute( + _get_for_user_statement(cls, task_id=task_id, user_id=user_id) + ).scalars().first() @classmethod - @legacy_db_query def list_for_user( cls, - db: Session | None = None, + db: Session, user_id: Optional[str] = None, enabled: Optional[bool] = None, ) -> list["AgentTask"]: - """ - 按用户和启用状态查询,并保留无 Session 的旧插件调用方式。 - """ - def query(session: Session) -> list["AgentTask"]: - """在给定会话中读取 Agent 任务列表。""" - return list(session.execute( - _list_for_user_statement(cls, user_id=user_id, enabled=enabled) - ).scalars().all()) - - return query(db) + """在调用方会话中按用户和启用状态查询。""" + return list(db.execute( + _list_for_user_statement(cls, user_id=user_id, enabled=enabled) + ).scalars().all()) @classmethod def update_task( diff --git a/app/db/models/agenttaskrun.py b/app/db/models/agenttaskrun.py index 1b4d7fb12..5cd2803f0 100644 --- a/app/db/models/agenttaskrun.py +++ b/app/db/models/agenttaskrun.py @@ -4,7 +4,6 @@ from sqlalchemy import Index, Integer, String, Text, delete, select, update from sqlalchemy.orm import Mapped, Session, mapped_column from app.db.base import Base, execute_dml, get_id_column -from app.db.decorators import legacy_db_query from app.db.models.agenttask import AgentTask @@ -249,7 +248,6 @@ class AgentTaskRun(Base): return True @classmethod - @legacy_db_query def get_by_run_id( cls, db: Session, @@ -261,7 +259,6 @@ class AgentTaskRun(Base): ).scalars().first() @classmethod - @legacy_db_query def list_for_task( cls, db: Session, diff --git a/app/db/models/downloadhistory.py b/app/db/models/downloadhistory.py index 5e5fdf14f..8a1a9de02 100644 --- a/app/db/models/downloadhistory.py +++ b/app/db/models/downloadhistory.py @@ -6,7 +6,6 @@ from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Mapped, Session, mapped_column from app.db.base import Base, execute_dml, get_id_column -from app.db.decorators import legacy_async_db_query, legacy_db_query from app.db.models._constraints import media_identity_constraint from app.schemas.types import MediaSource @@ -77,7 +76,6 @@ class DownloadHistory(Base): ) @classmethod - @legacy_db_query def get_by_hash(cls, db: Session, download_hash: str): return db.execute( select(DownloadHistory) @@ -86,7 +84,6 @@ class DownloadHistory(Base): ).scalars().first() @classmethod - @legacy_db_query def get_by_hashes(cls, db: Session, download_hashes: List[str]): """ 批量查询多个下载任务的最新历史记录,避免在上层形成 N+1 查询。 @@ -119,7 +116,6 @@ class DownloadHistory(Base): ] @classmethod - @legacy_db_query def get_by_media_identity( cls, db: Session, media_source: MediaSource, media_id: str, music_type: Optional[str] = None, @@ -136,7 +132,6 @@ class DownloadHistory(Base): return list(db.execute(statement).scalars().all()) @classmethod - @legacy_db_query def list_by_page( cls, db: Session, page: int = 1, count: int = 30 ): @@ -148,7 +143,6 @@ class DownloadHistory(Base): ).scalars().all()) @classmethod - @legacy_async_db_query async def async_list_by_page( cls, db: AsyncSession, page: int = 1, count: int = 30 ): @@ -161,7 +155,6 @@ class DownloadHistory(Base): return list(result.scalars().all()) @classmethod - @legacy_async_db_query async def async_list_by_title( cls, db: AsyncSession, @@ -177,13 +170,11 @@ class DownloadHistory(Base): return list(result.scalars().all()) @classmethod - @legacy_async_db_query async def async_count(cls, db: AsyncSession): result = await db.execute(select(func.count(cls.id))) return result.scalar() @classmethod - @legacy_async_db_query async def async_count_by_title(cls, db: AsyncSession, title: str): result = await db.execute( select(func.count(cls.id)).filter(_title_like(cls.title, title)) @@ -191,14 +182,12 @@ class DownloadHistory(Base): return result.scalar() @classmethod - @legacy_db_query def get_by_path(cls, db: Session, path: str): return db.execute( select(DownloadHistory).where(DownloadHistory.path == path) ).scalars().first() @classmethod - @legacy_db_query def get_last_by( cls, db: Session, @@ -237,7 +226,6 @@ class DownloadHistory(Base): @classmethod - @legacy_db_query def list_by_user_date(cls, db: Session, date: str, username: Optional[str] = None): """ 查询某用户某时间之前的下载历史。 @@ -256,7 +244,6 @@ class DownloadHistory(Base): ).scalars().all()) @classmethod - @legacy_db_query def list_by_date( cls, db: Session, @@ -282,7 +269,6 @@ class DownloadHistory(Base): ).scalars().all()) @classmethod - @legacy_db_query def list_by_type(cls, db: Session, mtype: str, days: int): return list(db.execute( select(DownloadHistory).where( @@ -345,7 +331,6 @@ class DownloadFiles(Base): ) @classmethod - @legacy_db_query def get_by_hash(cls, db: Session, download_hash: str, state: Optional[int] = None): statement = select(cls).where(cls.download_hash == download_hash) if state is not None: @@ -353,7 +338,6 @@ class DownloadFiles(Base): return list(db.execute(statement).scalars().all()) @classmethod - @legacy_db_query def get_by_fullpath(cls, db: Session, fullpath: str, all_files: bool = False): result = db.execute( select(cls).where(cls.fullpath == fullpath).order_by(cls.id.desc()) @@ -361,7 +345,6 @@ class DownloadFiles(Base): return list(result.all()) if all_files else result.first() @classmethod - @legacy_db_query def get_by_savepath(cls, db: Session, savepath: str): return list(db.execute(select(cls).where(cls.savepath == savepath)).scalars().all()) diff --git a/app/db/models/mediaserver.py b/app/db/models/mediaserver.py index 26e8965aa..f4ce13d92 100644 --- a/app/db/models/mediaserver.py +++ b/app/db/models/mediaserver.py @@ -7,7 +7,6 @@ from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Mapped, Session, mapped_column from app.db.base import Base, execute_dml, get_id_column -from app.db.decorators import legacy_async_db_query, legacy_db_query from app.db.models._constraints import media_identity_constraint from app.schemas.types import MediaSource @@ -53,12 +52,10 @@ class MediaServerItem(Base): ) @classmethod - @legacy_db_query def get_by_itemid(cls, db: Session, item_id: str): return db.execute(select(cls).where(cls.item_id == item_id)).scalars().first() @classmethod - @legacy_db_query def get_by_server_itemid(cls, db: Session, server: str, item_id: str): return db.execute( select(cls).where(cls.server == server, cls.item_id == item_id) @@ -97,7 +94,6 @@ class MediaServerItem(Base): ) @classmethod - @legacy_db_query def exist_by_media_identity( cls, db: Session, media_source: MediaSource, media_id: str, mtype: str, ): @@ -109,7 +105,6 @@ class MediaServerItem(Base): )).scalars().first() @classmethod - @legacy_db_query def exists_by_title(cls, db: Session, title: str, mtype: str, year: str): statement = select(cls).where(cls.title == title) if mtype: @@ -119,13 +114,11 @@ class MediaServerItem(Base): return db.execute(statement).scalars().first() @classmethod - @legacy_async_db_query async def async_get_by_itemid(cls, db: AsyncSession, item_id: str): result = await db.execute(select(cls).filter(cls.item_id == item_id)) return result.scalars().first() @classmethod - @legacy_async_db_query async def async_exist_by_media_identity( cls, db: AsyncSession, media_source: MediaSource, media_id: str, mtype: str, ): @@ -138,7 +131,6 @@ class MediaServerItem(Base): return result.scalars().first() @classmethod - @legacy_async_db_query async def async_exists_by_title(cls, db: AsyncSession, title: str, mtype: str, year: str): if not mtype and not year: result = await db.execute(select(cls).filter(cls.title == title)) diff --git a/app/db/models/message.py b/app/db/models/message.py index 034adf974..fe3db76c6 100644 --- a/app/db/models/message.py +++ b/app/db/models/message.py @@ -5,7 +5,6 @@ from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Mapped, Session, mapped_column from app.db.base import Base, execute_dml, get_id_column -from app.db.decorators import legacy_async_db_query, legacy_db_query class Message(Base): @@ -49,33 +48,25 @@ class Message(Base): return self.to_dict() @classmethod - @legacy_db_query def list_by_page( cls, - db: Session | None = None, + db: Session, page: int = 1, count: int = 30, ) -> List["Message"]: - """ - 分页获取消息记录,兼容显式会话和旧插件无会话调用。 - """ - def query(session: Session) -> List["Message"]: - """在给定同步会话中执行消息分页查询。""" - return list(session.execute( - select(cls) - .order_by(cls.reg_time.desc(), cls.id.desc()) - .offset((page - 1) * count) - .limit(count) - ).scalars().all()) - - return query(db) + """在调用方同步会话中分页获取消息记录。""" + return list(db.execute( + select(cls) + .order_by(cls.reg_time.desc(), cls.id.desc()) + .offset((page - 1) * count) + .limit(count) + ).scalars().all()) @classmethod - @legacy_db_query def exists_by_source( cls, - db: Session | str | None = None, - source: str | None = None, + db: Session, + source: str, ) -> bool: """ 判断指定来源标识的消息记录是否存在。 @@ -84,44 +75,29 @@ class Message(Base): :param source: 消息来源唯一标识 :return: 是否存在匹配记录 """ - if source is None and isinstance(db, str): - source, db = db, None - if source is None: - raise TypeError("source is required") - - def query(session: Session) -> bool: - """在给定同步会话中执行来源存在性查询。""" - return session.execute( - select(cls.id).where(cls.source == source).limit(1) - ).scalars().first() is not None - - return query(db) + return db.execute( + select(cls.id).where(cls.source == source).limit(1) + ).scalars().first() is not None @classmethod - @legacy_async_db_query async def async_list_by_page( - cls, db: AsyncSession | None = None, page: int = 1, count: int = 30 + cls, db: AsyncSession, page: int = 1, count: int = 30 ) -> List["Message"]: """ 异步分页获取消息记录。 """ - async def query(session: AsyncSession) -> List["Message"]: - """在给定异步会话中执行消息分页查询。""" - result = await session.execute( - select(cls) - .order_by(cls.reg_time.desc(), cls.id.desc()) - .offset((page - 1) * count) - .limit(count) - ) - return list(result.scalars().all()) - - return await query(db) + result = await db.execute( + select(cls) + .order_by(cls.reg_time.desc(), cls.id.desc()) + .offset((page - 1) * count) + .limit(count) + ) + return list(result.scalars().all()) @classmethod - @legacy_async_db_query async def async_list_sent_by_page( cls, - db: AsyncSession | None = None, + db: AsyncSession, page: int = 1, count: int = 30, all_clear_before: Optional[str] = None, @@ -131,35 +107,31 @@ class Message(Base): """ 分页获取系统发送的通知消息。 """ - async def query(session: AsyncSession) -> List["Message"]: - """在给定异步会话中执行通知消息分页查询。""" - statement = select(cls).where(cls.action == 1) - if all_clear_before: - statement = statement.where(cls.reg_time > all_clear_before) - if system_clear_before: - statement = statement.where( - or_( - and_(cls.image.isnot(None), cls.image != ""), - cls.reg_time > system_clear_before, - ) + statement = select(cls).where(cls.action == 1) + if all_clear_before: + statement = statement.where(cls.reg_time > all_clear_before) + if system_clear_before: + statement = statement.where( + or_( + and_(cls.image.isnot(None), cls.image != ""), + cls.reg_time > system_clear_before, ) - if media_clear_before: - statement = statement.where( - or_( - cls.image.is_(None), - cls.image == "", - cls.reg_time > media_clear_before, - ) - ) - result = await session.execute( - statement - .order_by(cls.reg_time.desc(), cls.id.desc()) - .offset((page - 1) * count) - .limit(count) ) - return list(result.scalars().all()) - - return await query(db) + if media_clear_before: + statement = statement.where( + or_( + cls.image.is_(None), + cls.image == "", + cls.reg_time > media_clear_before, + ) + ) + result = await db.execute( + statement + .order_by(cls.reg_time.desc(), cls.id.desc()) + .offset((page - 1) * count) + .limit(count) + ) + return list(result.scalars().all()) @classmethod def delete_before( diff --git a/app/db/models/passkey.py b/app/db/models/passkey.py index 06179467e..a26b7f0d5 100644 --- a/app/db/models/passkey.py +++ b/app/db/models/passkey.py @@ -5,10 +5,6 @@ from sqlalchemy.orm import Mapped, Session, mapped_column from datetime import datetime from app.db.base import Base, get_id_column -from app.db.decorators import ( - legacy_async_db_query, - legacy_db_query, -) def _get_by_user_id_statement(model: type["PassKey"], user_id: int): @@ -54,75 +50,51 @@ class PassKey(Base): transports: Mapped[Optional[str]] = mapped_column(String, nullable=True) @classmethod - @legacy_db_query def get_by_user_id( cls, - db: Session | int | None = None, - user_id: int | None = None, + db: Session, + user_id: int, ): - """获取用户的所有 PassKey,并保留无 Session 的旧插件调用方式。""" - if user_id is None and isinstance(db, int): - user_id, db = db, None - if user_id is None: - raise TypeError("user_id is required") - - def query(session: Session): - """在给定会话中执行启用凭证查询。""" - return list(session.execute( - _get_by_user_id_statement(cls, user_id) - ).scalars().all()) - - return query(db) + """在调用方 Session 中获取用户的所有启用 PassKey。""" + return list(db.execute( + _get_by_user_id_statement(cls, user_id) + ).scalars().all()) @classmethod - @legacy_async_db_query async def async_get_by_user_id(cls, db: AsyncSession, user_id: int): - """异步获取用户的所有 PassKey,并保留旧插件无 Session 调用。""" + """在调用方 AsyncSession 中获取用户的所有启用 PassKey。""" result = await db.execute( _get_by_user_id_statement(cls, user_id) ) return list(result.scalars().all()) @classmethod - @legacy_db_query def get_by_credential_id( cls, - db: Session | str | None = None, - credential_id: str | None = None, + db: Session, + credential_id: str, ): - """按凭证 ID 获取 PassKey,并保留无 Session 的旧插件调用方式。""" - if credential_id is None and isinstance(db, str): - credential_id, db = db, None - if credential_id is None: - raise TypeError("credential_id is required") - - def query(session: Session): - """在给定会话中执行启用凭证查询。""" - return session.execute( - _get_by_credential_id_statement(cls, credential_id) - ).scalars().first() - - return query(db) + """在调用方 Session 中按凭证 ID 获取启用 PassKey。""" + return db.execute( + _get_by_credential_id_statement(cls, credential_id) + ).scalars().first() @classmethod - @legacy_async_db_query async def async_get_by_credential_id(cls, db: AsyncSession, credential_id: str): - """异步根据凭证 ID 获取 PassKey,并保留旧插件无 Session 调用。""" + """在调用方 AsyncSession 中根据凭证 ID 获取启用 PassKey。""" result = await db.execute( _get_by_credential_id_statement(cls, credential_id) ) return result.scalars().first() @classmethod - @legacy_db_query def get_by_id(cls, db: Session, passkey_id: int): - """根据 ID 获取 PassKey,并保留旧插件无 Session 调用。""" + """在调用方 Session 中根据 ID 获取 PassKey。""" return db.execute(select(cls).where(cls.id == passkey_id)).scalars().first() @classmethod - @legacy_async_db_query async def async_get_by_id(cls, db: AsyncSession, passkey_id: int): - """异步根据 ID 获取 PassKey,并保留旧插件无 Session 调用。""" + """在调用方 AsyncSession 中根据 ID 获取 PassKey。""" result = await db.execute( select(cls).filter(cls.id == passkey_id) ) diff --git a/app/db/models/plugindata.py b/app/db/models/plugindata.py index 4aa9c9985..51744131f 100644 --- a/app/db/models/plugindata.py +++ b/app/db/models/plugindata.py @@ -4,7 +4,6 @@ from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Mapped, Session, mapped_column from app.db.base import get_id_column, Base -from app.db.decorators import legacy_async_db_query, legacy_db_query class PluginData(Base): @@ -21,44 +20,32 @@ class PluginData(Base): ) @classmethod - @legacy_db_query - def get_plugin_data(cls, db: Session | None = None, plugin_id: str | None = None): - """在调用方 Session 中读取插件全部数据,并兼容旧无会话入口。""" - if plugin_id is None: - raise TypeError("plugin_id is required") + def get_plugin_data(cls, db: Session, plugin_id: str): + """在调用方 Session 中读取插件全部数据。""" return list(db.execute(select(cls).where(cls.plugin_id == plugin_id)).scalars().all()) @classmethod - @legacy_async_db_query async def async_get_plugin_data( - cls, db: AsyncSession | None = None, plugin_id: str | None = None + cls, db: AsyncSession, plugin_id: str ): - """在调用方 AsyncSession 中读取插件全部数据,并兼容旧无会话入口。""" - if plugin_id is None: - raise TypeError("plugin_id is required") + """在调用方 AsyncSession 中读取插件全部数据。""" result = await db.execute(select(cls).where(cls.plugin_id == plugin_id)) return list(result.scalars().all()) @classmethod - @legacy_db_query def get_plugin_data_by_key( - cls, db: Session | None = None, plugin_id: str | None = None, key: str | None = None + cls, db: Session, plugin_id: str, key: str ): - """在调用方 Session 中按键读取插件数据,并兼容旧无会话入口。""" - if plugin_id is None or key is None: - raise TypeError("plugin_id and key are required") + """在调用方 Session 中按键读取插件数据。""" return db.execute( select(cls).where(cls.plugin_id == plugin_id, cls.key == key) ).scalars().first() @classmethod - @legacy_async_db_query async def async_get_plugin_data_by_key( - cls, db: AsyncSession | None = None, plugin_id: str | None = None, key: str | None = None + cls, db: AsyncSession, plugin_id: str, key: str ): - """在调用方 AsyncSession 中按键读取插件数据,并兼容旧无会话入口。""" - if plugin_id is None or key is None: - raise TypeError("plugin_id and key are required") + """在调用方 AsyncSession 中按键读取插件数据。""" result = await db.execute( select(cls).where(cls.plugin_id == plugin_id, cls.key == key) ) @@ -75,22 +62,16 @@ class PluginData(Base): db.execute(delete(cls).where(cls.plugin_id == plugin_id)) @classmethod - @legacy_db_query def get_plugin_data_by_plugin_id( - cls, db: Session | None = None, plugin_id: str | None = None + cls, db: Session, plugin_id: str ): - """在调用方 Session 中按插件 ID 读取数据,并兼容旧无会话入口。""" - if plugin_id is None: - raise TypeError("plugin_id is required") + """在调用方 Session 中按插件 ID 读取数据。""" return list(db.execute(select(cls).where(cls.plugin_id == plugin_id)).scalars().all()) @classmethod - @legacy_async_db_query async def async_get_plugin_data_by_plugin_id( - cls, db: AsyncSession | None = None, plugin_id: str | None = None + cls, db: AsyncSession, plugin_id: str ): - """在调用方 AsyncSession 中按插件 ID 读取数据,并兼容旧无会话入口。""" - if plugin_id is None: - raise TypeError("plugin_id is required") + """在调用方 AsyncSession 中按插件 ID 读取数据。""" result = await db.execute(select(cls).where(cls.plugin_id == plugin_id)) return list(result.scalars().all()) diff --git a/app/db/models/site.py b/app/db/models/site.py index 206880b3e..ddd64e623 100644 --- a/app/db/models/site.py +++ b/app/db/models/site.py @@ -6,7 +6,6 @@ from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Mapped, Session, mapped_column from app.db.base import Base, get_id_column -from app.db.decorators import legacy_async_db_query, legacy_db_query class Site(Base): @@ -58,122 +57,66 @@ class Site(Base): downloader: Mapped[Optional[str]] = mapped_column(String) @classmethod - @legacy_db_query - def get_by_domain(cls, db: Session | str | None = None, domain: str | None = None): - """按域名查询站点,兼容显式会话和旧插件无会话调用。""" - if domain is None and isinstance(db, str): - domain, db = db, None - if domain is None: - raise TypeError("domain is required") - - def query(session: Session): - """在给定同步会话中执行域名查询。""" - return session.execute(select(cls).where(cls.domain == domain)).scalars().first() - - return query(db) + def get_by_domain(cls, db: Session, domain: str): + """在调用方 Session 中按域名查询站点。""" + return db.execute(select(cls).where(cls.domain == domain)).scalars().first() @classmethod - @legacy_async_db_query async def async_get_by_domain( cls, - db: AsyncSession | str | None = None, - domain: str | None = None, + db: AsyncSession, + domain: str, ): - """异步按域名查询站点,兼容显式会话和旧插件无会话调用。""" - if domain is None and isinstance(db, str): - domain, db = db, None - if domain is None: - raise TypeError("domain is required") - - async def query(session: AsyncSession): - """在给定异步会话中执行域名查询。""" - result = await session.execute(select(cls).where(cls.domain == domain)) - return result.scalar_one_or_none() - - return await query(db) + """在调用方 AsyncSession 中按域名查询站点。""" + result = await db.execute(select(cls).where(cls.domain == domain)) + return result.scalar_one_or_none() @classmethod - @legacy_async_db_query async def async_get_by_name( cls, - db: AsyncSession | str | None = None, - name: str | None = None, + db: AsyncSession, + name: str, ): - """异步按站点名称查询,兼容显式会话和旧插件无会话调用。""" - if name is None and isinstance(db, str): - name, db = db, None - if name is None: - raise TypeError("name is required") - - async def query(session: AsyncSession): - """在给定异步会话中执行名称查询。""" - result = await session.execute(select(cls).where(cls.name == name)) - return result.scalar_one_or_none() - - return await query(db) + """在调用方 AsyncSession 中按站点名称查询。""" + result = await db.execute(select(cls).where(cls.name == name)) + return result.scalar_one_or_none() @classmethod - @legacy_db_query - def get_actives(cls, db: Session | None = None): - """查询启用站点,兼容显式会话和旧插件无会话调用。""" - def query(session: Session): - """在给定同步会话中执行启用站点查询。""" - return list(session.execute(select(cls).where(cls.is_active.is_(True))).scalars().all()) - - return query(db) + def get_actives(cls, db: Session): + """在调用方 Session 中查询启用站点。""" + return list(db.execute( + select(cls).where(cls.is_active.is_(True)) + ).scalars().all()) @classmethod - @legacy_async_db_query - async def async_get_actives(cls, db: AsyncSession | None = None): - """异步查询启用站点,兼容显式会话和旧插件无会话调用。""" - async def query(session: AsyncSession): - """在给定异步会话中执行启用站点查询。""" - result = await session.execute(select(cls).where(cls.is_active.is_(True))) - return list(result.scalars().all()) - - return await query(db) + async def async_get_actives(cls, db: AsyncSession): + """在调用方 AsyncSession 中查询启用站点。""" + result = await db.execute(select(cls).where(cls.is_active.is_(True))) + return list(result.scalars().all()) @classmethod - @legacy_db_query - def list_order_by_pri(cls, db: Session | None = None): - """按优先级升序查询站点,兼容显式会话和旧插件无会话调用。""" - def query(session: Session): - """在给定同步会话中执行优先级查询。""" - return list(session.execute(select(cls).order_by(cls.pri)).scalars().all()) - - return query(db) + def list_order_by_pri(cls, db: Session): + """在调用方 Session 中按优先级升序查询站点。""" + return list(db.execute(select(cls).order_by(cls.pri)).scalars().all()) @classmethod - @legacy_async_db_query - async def async_list_order_by_pri(cls, db: AsyncSession | None = None): - """异步按优先级升序查询站点,兼容显式会话和旧插件无会话调用。""" - async def query(session: AsyncSession): - """在给定异步会话中执行优先级查询。""" - result = await session.execute(select(cls).order_by(cls.pri)) - return list(result.scalars().all()) - - return await query(db) + async def async_list_order_by_pri(cls, db: AsyncSession): + """在调用方 AsyncSession 中按优先级升序查询站点。""" + result = await db.execute(select(cls).order_by(cls.pri)) + return list(result.scalars().all()) @classmethod - @legacy_db_query def get_domains_by_ids( cls, - db: Session | list[int] | None = None, - ids: list[int] | None = None, + db: Session, + ids: list[int], ): - """按 ID 查询域名,兼容显式会话和旧插件无会话调用。""" - if ids is None and isinstance(db, list): - ids, db = db, None - if ids is None: - raise TypeError("ids is required") + """在调用方 Session 中按 ID 查询域名。""" if not ids: return [] - - def query(session: Session): - """在给定同步会话中执行域名投影查询。""" - return list(session.execute(select(cls.domain).where(cls.id.in_(ids))).scalars().all()) - - return query(db) + return list(db.execute( + select(cls.domain).where(cls.id.in_(ids)) + ).scalars().all()) @classmethod def reset(cls, db: Session): diff --git a/app/db/models/siteicon.py b/app/db/models/siteicon.py index 2882c5cbd..176ede035 100644 --- a/app/db/models/siteicon.py +++ b/app/db/models/siteicon.py @@ -4,7 +4,6 @@ from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Mapped, Session, mapped_column from app.db.base import Base, get_id_column -from app.db.decorators import legacy_async_db_query class SiteIcon(Base): @@ -27,19 +26,11 @@ class SiteIcon(Base): return db.execute(select(cls).where(cls.domain == domain)).scalars().first() @classmethod - @legacy_async_db_query async def async_get_by_domain( cls, - db: AsyncSession | None = None, - domain: str | None = None, + db: AsyncSession, + domain: str, ): """在调用方 AsyncSession 中查询站点图标。""" - if domain is None: - raise TypeError("domain is required") - - async def query(session: AsyncSession): - """在给定异步会话中执行站点图标查询。""" - result = await session.execute(select(cls).where(cls.domain == domain)) - return result.scalar_one_or_none() - - return await query(db) + result = await db.execute(select(cls).where(cls.domain == domain)) + return result.scalar_one_or_none() diff --git a/app/db/models/sitestatistic.py b/app/db/models/sitestatistic.py index cf187c1bb..ba0fe4e2b 100644 --- a/app/db/models/sitestatistic.py +++ b/app/db/models/sitestatistic.py @@ -6,7 +6,6 @@ from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Mapped, Session, mapped_column from app.db.base import get_id_column, Base -from app.db.decorators import legacy_async_db_query class SiteStatistic(Base): @@ -35,22 +34,14 @@ class SiteStatistic(Base): return db.execute(select(cls).where(cls.domain == domain)).scalars().first() @classmethod - @legacy_async_db_query async def async_get_by_domain( cls, - db: AsyncSession | None = None, - domain: str | None = None, + db: AsyncSession, + domain: str, ): - """在调用方 AsyncSession 中查询站点统计,并兼容旧无会话调用。""" - if domain is None: - raise TypeError("domain is required") - - async def query(session: AsyncSession): - """在给定异步会话中执行站点统计查询。""" - result = await session.execute(select(cls).where(cls.domain == domain)) - return result.scalar_one_or_none() - - return await query(db) + """在调用方 AsyncSession 中查询站点统计。""" + result = await db.execute(select(cls).where(cls.domain == domain)) + return result.scalar_one_or_none() @classmethod def reset(cls, db: Session): diff --git a/app/db/models/siteuserdata.py b/app/db/models/siteuserdata.py index 452e8d4c3..7970955ff 100644 --- a/app/db/models/siteuserdata.py +++ b/app/db/models/siteuserdata.py @@ -6,7 +6,6 @@ from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Mapped, Session, mapped_column from app.db.base import Base, execute_dml, get_id_column -from app.db.decorators import legacy_async_db_query, legacy_db_query class SiteUserData(Base): @@ -61,7 +60,6 @@ class SiteUserData(Base): ) @classmethod - @legacy_db_query def get_by_domain(cls, db: Session, domain: str, workdate: Optional[str] = None, worktime: Optional[str] = None): statement = select(cls).where(cls.domain == domain) if workdate and worktime: @@ -72,7 +70,6 @@ class SiteUserData(Base): return list(db.execute(statement).scalars().all()) @classmethod - @legacy_async_db_query async def async_get_by_domain(cls, db: AsyncSession, domain: str, workdate: Optional[str] = None, worktime: Optional[str] = None): query = select(cls).filter(cls.domain == domain) if workdate and worktime: @@ -83,12 +80,10 @@ class SiteUserData(Base): return list(result.scalars().all()) @classmethod - @legacy_db_query def get_by_date(cls, db: Session, date: str): return list(db.execute(select(cls).where(cls.updated_day == date)).scalars().all()) @classmethod - @legacy_db_query def get_latest(cls, db: Session): """ 获取各站点最新一天的数据 @@ -113,7 +108,6 @@ class SiteUserData(Base): ).scalars().all()) @classmethod - @legacy_async_db_query async def async_get_latest(cls, db: AsyncSession): """ 异步获取各站点最新一天的数据 diff --git a/app/db/models/subscribe.py b/app/db/models/subscribe.py index 92c125d60..e9a94a1b9 100644 --- a/app/db/models/subscribe.py +++ b/app/db/models/subscribe.py @@ -6,7 +6,6 @@ from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Mapped, Session, mapped_column from app.db.base import get_id_column, Base -from app.db.decorators import legacy_async_db_query, legacy_db_query from app.db.models._constraints import media_identity_constraint from app.schemas.types import MUSIC_ENTITY_RECORDING, MediaSource @@ -140,9 +139,8 @@ class Subscribe(Base): return condition @classmethod - @legacy_db_query def exists( - cls, db: Session | MediaSource | None = None, + cls, db: Session, media_source: MediaSource | str | None = None, media_id: str | None = None, season: Optional[int] = None, @@ -150,27 +148,21 @@ class Subscribe(Base): music_type: Optional[str] = None, ): """按媒体身份、季号与剧集组查询已有订阅。""" - if db is not None and not isinstance(db, Session): - media_source, media_id, db = db, media_source, None condition = cls._identity_condition( media_source, media_id, music_type ) if condition is None: return None - def query(session: Session): - """在给定会话中执行订阅身份查询。""" - statement = select(cls).where(condition) - if season is not None: - statement = statement.where(cls.season == season) - return session.execute( - statement.where(cls.episode_group == episode_group) - ).scalars().first() - return query(db) + statement = select(cls).where(condition) + if season is not None: + statement = statement.where(cls.season == season) + return db.execute( + statement.where(cls.episode_group == episode_group) + ).scalars().first() @classmethod - @legacy_async_db_query async def async_exists( - cls, db: AsyncSession | MediaSource | None = None, + cls, db: AsyncSession, media_source: MediaSource | str | None = None, media_id: str | None = None, season: Optional[int] = None, @@ -178,28 +170,22 @@ class Subscribe(Base): music_type: Optional[str] = None, ): """异步按媒体身份、季号与剧集组查询已有订阅。""" - if db is not None and not isinstance(db, AsyncSession): - media_source, media_id, db = db, media_source, None condition = cls._identity_condition( media_source, media_id, music_type ) if condition is None: return None - async def query(session: AsyncSession): - """在给定异步会话中执行订阅身份查询。""" - statement = select(cls).where(condition) - if season is not None: - statement = statement.where(cls.season == season) - result = await session.execute( - statement.where(cls.episode_group == episode_group) - ) - return result.scalars().first() - return await query(db) + statement = select(cls).where(condition) + if season is not None: + statement = statement.where(cls.season == season) + result = await db.execute( + statement.where(cls.episode_group == episode_group) + ) + return result.scalars().first() @classmethod - @legacy_db_query def exists_by_username( - cls, db: Session | str | None = None, + cls, db: Session, username: str | MediaSource | None = None, media_source: MediaSource | str | None = None, media_id: str | None = None, @@ -210,8 +196,6 @@ class Subscribe(Base): """ 按订阅 owner、媒体身份、季号与剧集组查询订阅行。 """ - if db is not None and not isinstance(db, Session): - username, media_source, media_id, db = db, username, media_source, None if not username: return None condition = cls._identity_condition( @@ -219,20 +203,16 @@ class Subscribe(Base): ) if condition is None: return None - def query(session: Session): - """在给定会话中执行订阅 owner 查询。""" - statement = select(cls).where(cls.username == username, condition) - if season is not None: - statement = statement.where(cls.season == season) - return session.execute( - statement.where(cls.episode_group == episode_group) - ).scalars().first() - return query(db) + statement = select(cls).where(cls.username == username, condition) + if season is not None: + statement = statement.where(cls.season == season) + return db.execute( + statement.where(cls.episode_group == episode_group) + ).scalars().first() @classmethod - @legacy_async_db_query async def async_exists_by_username( - cls, db: AsyncSession | str | None = None, + cls, db: AsyncSession, username: str | MediaSource | None = None, media_source: MediaSource | str | None = None, media_id: str | None = None, season: Optional[int] = None, @@ -242,8 +222,6 @@ class Subscribe(Base): """ 异步按订阅 owner、媒体身份、季号与剧集组查询订阅行。 """ - if db is not None and not isinstance(db, AsyncSession): - username, media_source, media_id, db = db, username, media_source, None if not username: return None condition = cls._identity_condition( @@ -251,112 +229,76 @@ class Subscribe(Base): ) if condition is None: return None - async def query(session: AsyncSession): - """在给定异步会话中执行订阅 owner 查询。""" - statement = select(cls).where(cls.username == username, condition) - if season is not None: - statement = statement.where(cls.season == season) - result = await session.execute( - statement.where(cls.episode_group == episode_group) - ) - return result.scalars().first() - return await query(db) + statement = select(cls).where(cls.username == username, condition) + if season is not None: + statement = statement.where(cls.season == season) + result = await db.execute( + statement.where(cls.episode_group == episode_group) + ) + return result.scalars().first() @classmethod - @legacy_db_query - def get_by_state(cls, db: Session | str | None = None, state: str | None = None): - """按状态列表查询订阅,兼容显式会话和旧插件无会话调用。""" - if not isinstance(db, Session): - state, db = db if state is None else state, None - def query(session: Session): - """在给定会话中执行状态查询。""" - statement = select(cls) - if state: - statement = statement.where(cls.state.in_(state.split(','))) - return list(session.execute(statement).scalars().all()) - return query(db) + def get_by_state(cls, db: Session, state: str | None = None): + """在调用方 Session 中按状态列表查询订阅。""" + statement = select(cls) + if state: + statement = statement.where(cls.state.in_(state.split(','))) + return list(db.execute(statement).scalars().all()) @classmethod - @legacy_async_db_query async def async_get_by_state( - cls, db: AsyncSession | str | None = None, state: str | None = None + cls, db: AsyncSession, state: str | None = None ): - """异步按状态列表查询订阅,兼容显式会话和旧插件无会话调用。""" - if not isinstance(db, AsyncSession): - state, db = db if state is None else state, None - async def query(session: AsyncSession): - """在给定异步会话中执行状态查询。""" - statement = select(cls) - if state: - statement = statement.where(cls.state.in_(state.split(','))) - result = await session.execute(statement) - return list(result.scalars().all()) - return await query(db) + """在调用方 AsyncSession 中按状态列表查询订阅。""" + statement = select(cls) + if state: + statement = statement.where(cls.state.in_(state.split(','))) + result = await db.execute(statement) + return list(result.scalars().all()) @classmethod - @legacy_db_query def get_by_title( - cls, db: Session | str | None = None, title: str | None = None, + cls, db: Session, title: str, season: Optional[int] = None, ): - """按标题查询订阅,兼容显式会话和旧插件无会话调用。""" - if not isinstance(db, Session): - title, db = db if title is None else title, None - def query(session: Session): - """在给定会话中执行标题查询。""" - statement = select(cls).where(cls.name == title) - if season is not None: - statement = statement.where(cls.season == season) - return session.execute(statement).scalars().first() - return query(db) + """在调用方 Session 中按标题查询订阅。""" + statement = select(cls).where(cls.name == title) + if season is not None: + statement = statement.where(cls.season == season) + return db.execute(statement).scalars().first() @classmethod - @legacy_async_db_query async def async_get_by_title( - cls, db: AsyncSession | str | None = None, title: str | None = None, + cls, db: AsyncSession, title: str, season: Optional[int] = None, ): - """异步按标题查询订阅,兼容显式会话和旧插件无会话调用。""" - if not isinstance(db, AsyncSession): - title, db = db if title is None else title, None - async def query(session: AsyncSession): - """在给定异步会话中执行标题查询。""" - statement = select(cls).where(cls.name == title) - if season is not None: - statement = statement.where(cls.season == season) - result = await session.execute(statement) - return result.scalars().first() - return await query(db) + """在调用方 AsyncSession 中按标题查询订阅。""" + statement = select(cls).where(cls.name == title) + if season is not None: + statement = statement.where(cls.season == season) + result = await db.execute(statement) + return result.scalars().first() @classmethod - @legacy_async_db_query async def async_list_by_title( - cls, db: AsyncSession | str | None = None, title: str | None = None, + cls, db: AsyncSession, title: str, season: Optional[int] = None, ): - """异步按标题查询候选订阅列表,兼容显式会话和旧插件无会话调用。""" - if not isinstance(db, AsyncSession): - title, db = db if title is None else title, None - async def query(session: AsyncSession): - """在给定异步会话中执行标题列表查询。""" - statement = select(cls).where(cls.name == title) - if season is not None: - statement = statement.where(cls.season == season) - result = await session.execute(statement) - return list(result.scalars().all()) - return await query(db) + """在调用方 AsyncSession 中按标题查询候选订阅列表。""" + statement = select(cls).where(cls.name == title) + if season is not None: + statement = statement.where(cls.season == season) + result = await db.execute(statement) + return list(result.scalars().all()) @classmethod - @legacy_db_query def list_by_media_identity( - cls, db: Session | MediaSource | None = None, + cls, db: Session, media_source: MediaSource | str | None = None, media_id: str | None = None, music_type: Optional[str] = None, ): """同步按统一媒体身份查询候选订阅列表。""" - if db is not None and not isinstance(db, Session): - media_source, media_id, db = db, media_source, None condition = cls._identity_condition( media_source=media_source, media_id=media_id, @@ -364,22 +306,16 @@ class Subscribe(Base): ) if condition is None: return [] - def query(session: Session): - """在给定会话中执行媒体身份列表查询。""" - return list(session.execute(select(cls).where(condition)).scalars().all()) - return query(db) + return list(db.execute(select(cls).where(condition)).scalars().all()) @classmethod - @legacy_async_db_query async def async_list_by_media_identity( - cls, db: AsyncSession | MediaSource | None = None, + cls, db: AsyncSession, media_source: MediaSource | str | None = None, media_id: str | None = None, music_type: Optional[str] = None, ): """异步按统一媒体身份查询候选订阅列表。""" - if db is not None and not isinstance(db, AsyncSession): - media_source, media_id, db = db, media_source, None condition = cls._identity_condition( media_source=media_source, media_id=media_id, @@ -387,16 +323,12 @@ class Subscribe(Base): ) if condition is None: return [] - async def query(session: AsyncSession): - """在给定异步会话中执行媒体身份列表查询。""" - result = await session.execute(select(cls).where(condition)) - return list(result.scalars().all()) - return await query(db) + result = await db.execute(select(cls).where(condition)) + return list(result.scalars().all()) @classmethod - @legacy_db_query def get_by( - cls, db: Session | str | None = None, + cls, db: Session, type: str | MediaSource | None = None, media_source: MediaSource | str | None = None, media_id: str | None = None, @@ -406,8 +338,6 @@ class Subscribe(Base): """ 根据条件查询订阅 """ - if db is not None and not isinstance(db, Session): - type, media_source, media_id, db = db, type, media_source, None condition = cls._identity_condition( media_source, media_id, music_type ) @@ -416,15 +346,11 @@ class Subscribe(Base): statement = select(cls).where(condition, cls.type == type) if season is not None: statement = statement.where(cls.season == season) - def query(session: Session): - """在给定会话中执行类型媒体查询。""" - return session.execute(statement).scalars().first() - return query(db) + return db.execute(statement).scalars().first() @classmethod - @legacy_async_db_query async def async_get_by( - cls, db: AsyncSession | str | None = None, + cls, db: AsyncSession, type: str | MediaSource | None = None, media_source: MediaSource | str | None = None, media_id: str | None = None, @@ -434,8 +360,6 @@ class Subscribe(Base): """ 根据条件查询订阅 """ - if db is not None and not isinstance(db, AsyncSession): - type, media_source, media_id, db = db, type, media_source, None condition = cls._identity_condition( media_source, media_id, music_type ) @@ -444,76 +368,49 @@ class Subscribe(Base): query = select(cls).filter(condition, cls.type == type) if season is not None: query = query.filter(cls.season == season) - async def execute_query(session: AsyncSession): - """在给定异步会话中执行类型媒体查询。""" - result = await session.execute(query) - return result.scalars().first() - return await execute_query(db) + result = await db.execute(query) + return result.scalars().first() @classmethod - @legacy_db_query - def list_by_username(cls, db: Session | str | None = None, username: str | None = None, + def list_by_username(cls, db: Session, username: str, state: Optional[str] = None, mtype: Optional[str] = None): - """按用户筛选订阅,兼容显式会话和旧插件无会话调用。""" - if not isinstance(db, Session): - username, db = db if username is None else username, None - def query(session: Session): - """在给定会话中执行用户筛选查询。""" - statement = select(cls).where(cls.username == username) - if state: - statement = statement.where(cls.state == state) - if mtype: - statement = statement.where(cls.type == mtype) - return list(session.execute(statement).scalars().all()) - return query(db) + """在调用方 Session 中按用户筛选订阅。""" + statement = select(cls).where(cls.username == username) + if state: + statement = statement.where(cls.state == state) + if mtype: + statement = statement.where(cls.type == mtype) + return list(db.execute(statement).scalars().all()) @classmethod - @legacy_async_db_query - async def async_list_by_username(cls, db: AsyncSession | str | None = None, - username: str | None = None, state: Optional[str] = None, + async def async_list_by_username(cls, db: AsyncSession, + username: str, state: Optional[str] = None, mtype: Optional[str] = None): - """异步按用户筛选订阅,兼容显式会话和旧插件无会话调用。""" - if not isinstance(db, AsyncSession): - username, db = db if username is None else username, None - async def query(session: AsyncSession): - """在给定异步会话中执行用户筛选查询。""" - statement = select(cls).where(cls.username == username) - if state: - statement = statement.where(cls.state == state) - if mtype: - statement = statement.where(cls.type == mtype) - result = await session.execute(statement) - return list(result.scalars().all()) - return await query(db) + """在调用方 AsyncSession 中按用户筛选订阅。""" + statement = select(cls).where(cls.username == username) + if state: + statement = statement.where(cls.state == state) + if mtype: + statement = statement.where(cls.type == mtype) + result = await db.execute(statement) + return list(result.scalars().all()) @classmethod - @legacy_db_query - def list_by_type(cls, db: Session | str | None = None, mtype: str | None = None, days: int = 7): - """按类型查询最近时间窗内的订阅,兼容显式会话和旧插件无会话调用。""" - if not isinstance(db, Session): - mtype, db = db if mtype is None else mtype, None - def query(session: Session): - """在给定会话中执行时间窗订阅查询。""" - return list(session.execute(select(cls).where( - cls.type == mtype, - cls.date >= time.strftime("%Y-%m-%d %H:%M:%S", - time.localtime(time.time() - 86400 * int(days))) - )).scalars().all()) - return query(db) + def list_by_type(cls, db: Session, mtype: str, days: int = 7): + """在调用方 Session 中按类型查询最近时间窗内的订阅。""" + return list(db.execute(select(cls).where( + cls.type == mtype, + cls.date >= time.strftime("%Y-%m-%d %H:%M:%S", + time.localtime(time.time() - 86400 * int(days))) + )).scalars().all()) @classmethod - @legacy_async_db_query - async def async_list_by_type(cls, db: AsyncSession | str | None = None, - mtype: str | None = None, days: int = 7): - """异步按类型查询最近时间窗内的订阅,兼容显式会话和旧插件无会话调用。""" - if not isinstance(db, AsyncSession): - mtype, db = db if mtype is None else mtype, None - async def query(session: AsyncSession): - """在给定异步会话中执行时间窗订阅查询。""" - result = await session.execute(select(cls).where( - cls.type == mtype, - cls.date >= time.strftime("%Y-%m-%d %H:%M:%S", - time.localtime(time.time() - 86400 * int(days))) - )) - return list(result.scalars().all()) - return await query(db) + async def async_list_by_type(cls, db: AsyncSession, + mtype: str, days: int = 7): + """在调用方 AsyncSession 中按类型查询最近时间窗内的订阅。""" + result = await db.execute(select(cls).where( + cls.type == mtype, + cls.date >= time.strftime("%Y-%m-%d %H:%M:%S", + time.localtime(time.time() - 86400 * int(days))) + )) + return list(result.scalars().all()) diff --git a/app/db/models/subscribehistory.py b/app/db/models/subscribehistory.py index f095571c8..664789565 100644 --- a/app/db/models/subscribehistory.py +++ b/app/db/models/subscribehistory.py @@ -5,7 +5,6 @@ from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Mapped, Session, mapped_column from app.db.base import Base, get_id_column -from app.db.decorators import legacy_async_db_query, legacy_db_query from app.db.models._constraints import media_identity_constraint from app.schemas.types import MUSIC_ENTITY_RECORDING, MediaSource @@ -107,9 +106,8 @@ class SubscribeHistory(Base): ) @classmethod - @legacy_db_query def list_by_type(cls, db: Session, mtype: str, page: int = 1, count: int = 30): - """按媒体类型分页查询订阅历史,并保留旧插件无 Session 调用。""" + """在调用方 Session 中按媒体类型分页查询订阅历史。""" return list(db.execute( select(cls).where( cls.type == mtype @@ -119,9 +117,8 @@ class SubscribeHistory(Base): ).scalars().all()) @classmethod - @legacy_async_db_query async def async_list_by_type(cls, db: AsyncSession, mtype: str, page: int = 1, count: int = 30): - """异步按媒体类型分页查询订阅历史,并保留旧插件无 Session 调用。""" + """在调用方 AsyncSession 中按媒体类型分页查询订阅历史。""" result = await db.execute( select(cls).filter( cls.type == mtype @@ -132,7 +129,6 @@ class SubscribeHistory(Base): return list(result.scalars().all()) @classmethod - @legacy_async_db_query async def async_list_by_type_and_username( cls, db: AsyncSession, @@ -177,7 +173,6 @@ class SubscribeHistory(Base): return condition @classmethod - @legacy_db_query def exists( cls, db: Session, media_source: MediaSource, media_id: str, season: Optional[int] = None, @@ -197,7 +192,6 @@ class SubscribeHistory(Base): return db.execute(statement).scalars().first() @classmethod - @legacy_async_db_query async def async_exists( cls, db: AsyncSession, media_source: MediaSource, media_id: str, season: Optional[int] = None, diff --git a/app/db/models/systemconfig.py b/app/db/models/systemconfig.py index efd88d2d0..01e926830 100644 --- a/app/db/models/systemconfig.py +++ b/app/db/models/systemconfig.py @@ -4,7 +4,6 @@ from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Mapped, Session, mapped_column from app.db.base import Base, get_id_column -from app.db.decorators import legacy_async_db_query, legacy_db_query class SystemConfig(Base): @@ -18,15 +17,13 @@ class SystemConfig(Base): value: Mapped[Optional[Any]] = mapped_column(JSON) @classmethod - @legacy_db_query def get_by_key(cls, db: Session, key: str): - """按配置键查询系统配置,并保留旧插件无 Session 调用。""" + """在调用方 Session 中按配置键查询系统配置。""" return db.execute(select(cls).where(cls.key == key)).scalars().first() @classmethod - @legacy_async_db_query async def async_get_by_key(cls, db: AsyncSession, key: str): - """异步按配置键查询系统配置,并保留旧插件无 Session 调用。""" + """在调用方 AsyncSession 中按配置键查询系统配置。""" result = await db.execute(select(cls).where(cls.key == key)) return result.scalar_one_or_none() diff --git a/app/db/models/transferhistory.py b/app/db/models/transferhistory.py index 6d68caffd..bbea0f5ee 100644 --- a/app/db/models/transferhistory.py +++ b/app/db/models/transferhistory.py @@ -8,10 +8,6 @@ from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Mapped, Session, mapped_column from app.db.base import Base, execute_dml, get_id_column -from app.db.decorators import ( - legacy_async_db_query, - legacy_db_query, -) from app.db.models._constraints import media_identity_constraint from app.schemas.types import MUSIC_ENTITY_ALBUM, MUSIC_ENTITY_RECORDING, MediaSource, MediaType @@ -97,7 +93,6 @@ class TransferHistory(Base): ) @classmethod - @legacy_db_query def list_by_title(cls, db: Session, title: str, page: int = 1, count: int = 30, status: Optional[bool] = None, wildcard: bool = False): if wildcard: @@ -124,7 +119,6 @@ class TransferHistory(Base): return list(db.execute(statement).scalars().all()) @classmethod - @legacy_async_db_query async def async_list_by_title(cls, db: AsyncSession, title: str, page: int = 1, count: int = 30, status: Optional[bool] = None, wildcard: bool = False): if wildcard: @@ -152,7 +146,6 @@ class TransferHistory(Base): return list(result.scalars().all()) @classmethod - @legacy_db_query def list_by_page(cls, db: Session, page: int = 1, count: int = 30, status: Optional[bool] = None): statement = select(cls) if status is not None: @@ -166,7 +159,6 @@ class TransferHistory(Base): return list(db.execute(statement).scalars().all()) @classmethod - @legacy_async_db_query async def async_list_by_page(cls, db: AsyncSession, page: int = 1, count: int = 30, status: Optional[bool] = None): if status is not None: @@ -188,30 +180,19 @@ class TransferHistory(Base): return list(result.scalars().all()) @classmethod - @legacy_db_query def get_by_hash( cls, - db: Session | str | None = None, - download_hash: str | None = None, + db: Session, + download_hash: str, ): - """按下载哈希查询最新记录,兼容旧插件无会话调用。""" - if download_hash is None and isinstance(db, str): - download_hash, db = db, None - if download_hash is None: - raise TypeError("download_hash is required") - - def query(session: Session): - """在调用方提供的同步会话中执行哈希查询。""" - return session.execute( - select(cls).where(cls.download_hash == download_hash) - ).scalars().first() - - return query(db) + """在调用方 Session 中按下载哈希查询最新记录。""" + return db.execute( + select(cls).where(cls.download_hash == download_hash) + ).scalars().first() @classmethod - @legacy_db_query def get_by_src( - cls, db: Session | str | None = None, src: str | None = None, + cls, db: Session, src: str, storage: Optional[str] = None ) -> Optional["TransferHistory"]: """ @@ -222,26 +203,14 @@ class TransferHistory(Base): :param storage: 源存储类型 :return: 命中的整理记录,未命中时返回 None """ - if src is None and isinstance(db, str): - src, db = db, None - if src is None: - raise TypeError("src is required") - - def query(session: Session): - """在调用方提供的同步会话中执行源路径查询。""" - statement = select(cls).where(cls.src == src) - if storage: - statement = statement.where(cls.src_storage == storage) - return session.execute( - statement.order_by(cls.id.desc()) - ).scalars().first() - - return query(db) + statement = select(cls).where(cls.src == src) + if storage: + statement = statement.where(cls.src_storage == storage) + return db.execute(statement.order_by(cls.id.desc())).scalars().first() @classmethod - @legacy_db_query def get_success_by_src( - cls, db: Session | str | None = None, src: str | None = None, + cls, db: Session, src: str, storage: Optional[str] = None ) -> Optional["TransferHistory"]: """ @@ -254,26 +223,14 @@ class TransferHistory(Base): :param storage: 源存储类型 :return: 命中的成功整理记录,未命中时返回 None """ - if src is None and isinstance(db, str): - src, db = db, None - if src is None: - raise TypeError("src is required") - - def query(session: Session): - """在调用方提供的同步会话中执行成功源路径查询。""" - statement = select(cls).where(cls.src == src, cls.status.is_(True)) - if storage: - statement = statement.where(cls.src_storage == storage) - return session.execute( - statement.order_by(cls.id.desc()) - ).scalars().first() - - return query(db) + statement = select(cls).where(cls.src == src, cls.status.is_(True)) + if storage: + statement = statement.where(cls.src_storage == storage) + return db.execute(statement.order_by(cls.id.desc())).scalars().first() @classmethod - @legacy_db_query def get_by_dest( - cls, db: Session | str | None = None, dest: str | None = None, + cls, db: Session, dest: str, storage: Optional[str] = None ) -> Optional["TransferHistory"]: """ @@ -284,24 +241,12 @@ class TransferHistory(Base): :param storage: 目标存储类型 :return: 命中的整理记录,未命中时返回 None """ - if dest is None and isinstance(db, str): - dest, db = db, None - if dest is None: - raise TypeError("dest is required") - - def query(session: Session): - """在调用方提供的同步会话中执行目标路径查询。""" - statement = select(cls).where(cls.dest == dest) - if storage: - statement = statement.where(cls.dest_storage == storage) - return session.execute( - statement.order_by(cls.id.desc()) - ).scalars().first() - - return query(db) + statement = select(cls).where(cls.dest == dest) + if storage: + statement = statement.where(cls.dest_storage == storage) + return db.execute(statement.order_by(cls.id.desc())).scalars().first() @classmethod - @legacy_db_query def list_success_by_src( cls, db: Session, @@ -341,7 +286,6 @@ class TransferHistory(Base): return list(db.execute(statement).scalars().all()) @classmethod - @legacy_db_query def list_success_move_by_dest( cls, db: Session, @@ -384,14 +328,12 @@ class TransferHistory(Base): return list(db.execute(statement).scalars().all()) @classmethod - @legacy_db_query def list_by_hash(cls, db: Session, download_hash: str): return list(db.execute( select(cls).where(cls.download_hash == download_hash) ).scalars().all()) @classmethod - @legacy_db_query def statistic(cls, db: Session, days: int = 7): """ 统计最近days天的下载历史数量,按日期分组返回每日数量 @@ -408,7 +350,6 @@ class TransferHistory(Base): ).all()) @classmethod - @legacy_db_query def monthly_media_statistics(cls, db: Session): """ 统计当月成功整理的电影、电视剧、剧集和音乐数量。 @@ -474,7 +415,6 @@ class TransferHistory(Base): return 1 @classmethod - @legacy_async_db_query async def async_statistic(cls, db: AsyncSession, days: int = 7): """ 统计最近days天的下载历史数量,按日期分组返回每日数量 @@ -489,7 +429,6 @@ class TransferHistory(Base): return result.all() @classmethod - @legacy_db_query def count(cls, db: Session, status: Optional[bool] = None): statement = select(func.count(cls.id)) if status is not None: @@ -497,7 +436,6 @@ class TransferHistory(Base): return db.execute(statement).scalar() @classmethod - @legacy_async_db_query async def async_count(cls, db: AsyncSession, status: Optional[bool] = None): if status is not None: result = await db.execute( @@ -510,7 +448,6 @@ class TransferHistory(Base): return result.scalar() @classmethod - @legacy_db_query def count_by_title(cls, db: Session, title: str, status: Optional[bool] = None, wildcard: bool = False): if wildcard: text_filter = or_( @@ -530,7 +467,6 @@ class TransferHistory(Base): return db.execute(statement).scalar() @classmethod - @legacy_async_db_query async def async_count_by_title(cls, db: AsyncSession, title: str, status: Optional[bool] = None, wildcard: bool = False): if wildcard: text_filter = or_( @@ -551,7 +487,6 @@ class TransferHistory(Base): return result.scalar() @classmethod - @legacy_db_query def list_by(cls, db: Session, mtype: Optional[str] = None, title: Optional[str] = None, year: Optional[str] = None, season: Optional[str] = None, episode: Optional[str] = None, @@ -589,7 +524,6 @@ class TransferHistory(Base): return list(db.execute(statement).scalars().all()) @classmethod - @legacy_db_query def get_by_media_identity( cls, db: Session, media_source: MediaSource, media_id: str, mtype: Optional[str] = None, @@ -636,7 +570,6 @@ class TransferHistory(Base): return history @classmethod - @legacy_db_query def list_by_date(cls, db: Session, date: str): """ 查询某时间之后的转移历史 diff --git a/app/db/models/transferpending.py b/app/db/models/transferpending.py index c88b7be39..9cfd9211d 100644 --- a/app/db/models/transferpending.py +++ b/app/db/models/transferpending.py @@ -4,7 +4,6 @@ from sqlalchemy import Index, String, delete, select from sqlalchemy.orm import Mapped, Session, mapped_column from app.db.base import Base, execute_dml, get_id_column -from app.db.decorators import legacy_db_query class TransferPending(Base): @@ -73,7 +72,6 @@ class TransferPending(Base): ) @classmethod - @legacy_db_query def list_all(cls, db: Session, limit: Optional[int] = 5000) -> List["TransferPending"]: """ 列出全部待整理登记,供启动回放使用。 diff --git a/app/db/models/user.py b/app/db/models/user.py index e77a50aee..53a54b603 100644 --- a/app/db/models/user.py +++ b/app/db/models/user.py @@ -4,7 +4,6 @@ from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Mapped, Session, mapped_column from app.db.base import Base, get_id_column -from app.db.decorators import legacy_async_db_query, legacy_db_query class User(Base): @@ -35,78 +34,38 @@ class User(Base): settings: Mapped[Optional[Any]] = mapped_column(JSON, default=dict) @classmethod - @legacy_db_query def get_by_name( cls, - db: Session | str | None = None, - name: str | None = None, + db: Session, + name: str, ): - """按用户名查询用户,兼容显式会话和旧插件无会话调用。""" - if name is None and isinstance(db, str): - name, db = db, None - if name is None: - raise TypeError("name is required") - - def query(session: Session): - """在给定会话中执行用户名查询。""" - return session.execute(select(cls).where(cls.name == name)).scalars().first() - - return query(db) + """在调用方同步会话中按用户名查询用户。""" + return db.execute(select(cls).where(cls.name == name)).scalars().first() @classmethod - @legacy_async_db_query async def async_get_by_name( cls, - db: AsyncSession | str | None = None, - name: str | None = None, + db: AsyncSession, + name: str, ): - """异步按用户名查询,兼容显式会话和旧插件无会话调用。""" - if name is None and isinstance(db, str): - name, db = db, None - if name is None: - raise TypeError("name is required") - - async def query(session: AsyncSession): - """在给定异步会话中执行用户名查询。""" - result = await session.execute(select(cls).filter(cls.name == name)) - return result.scalars().first() - - return await query(db) + """在调用方异步会话中按用户名查询用户。""" + result = await db.execute(select(cls).filter(cls.name == name)) + return result.scalars().first() @classmethod - @legacy_db_query - def get_by_id(cls, db: Session | int | None = None, user_id: int | None = None): - """按用户 ID 查询用户,兼容显式会话和旧插件无会话调用。""" - if user_id is None and isinstance(db, int): - user_id, db = db, None - if user_id is None: - raise TypeError("user_id is required") - - def query(session: Session): - """在给定会话中执行用户 ID 查询。""" - return session.execute(select(cls).where(cls.id == user_id)).scalars().first() - - return query(db) + def get_by_id(cls, db: Session, user_id: int): + """在调用方同步会话中按用户 ID 查询用户。""" + return db.execute(select(cls).where(cls.id == user_id)).scalars().first() @classmethod - @legacy_async_db_query async def async_get_by_id( cls, - db: AsyncSession | int | None = None, - user_id: int | None = None, + db: AsyncSession, + user_id: int, ): - """异步按用户 ID 查询,兼容显式会话和旧插件无会话调用。""" - if user_id is None and isinstance(db, int): - user_id, db = db, None - if user_id is None: - raise TypeError("user_id is required") - - async def query(session: AsyncSession): - """在给定异步会话中执行用户 ID 查询。""" - result = await session.execute(select(cls).filter(cls.id == user_id)) - return result.scalars().first() - - return await query(db) + """在调用方异步会话中按用户 ID 查询用户。""" + result = await db.execute(select(cls).filter(cls.id == user_id)) + return result.scalars().first() def delete_by_name(self, db: Session, name: str): user = self.get_by_name(db, name) diff --git a/app/db/models/workflow.py b/app/db/models/workflow.py index 0aaab8c72..f0ddf3602 100644 --- a/app/db/models/workflow.py +++ b/app/db/models/workflow.py @@ -7,7 +7,6 @@ from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.ext.asyncio import AsyncSession from app.db.base import Base, get_id_column -from app.db.decorators import legacy_async_db_query, legacy_db_query class Workflow(Base): @@ -56,18 +55,15 @@ class Workflow(Base): ) @classmethod - @legacy_db_query def get_enabled_workflows(cls, db): return list(db.execute(select(cls).where(cls.state != 'P')).scalars().all()) @classmethod - @legacy_async_db_query async def async_get_enabled_workflows(cls, db: AsyncSession): result = await db.execute(select(cls).where(cls.state != 'P')) return list(result.scalars().all()) @classmethod - @legacy_db_query def get_timer_triggered_workflows(cls, db): """获取定时触发的工作流""" return list(db.execute(select(cls).where( @@ -81,7 +77,6 @@ class Workflow(Base): )).scalars().all()) @classmethod - @legacy_async_db_query async def async_get_timer_triggered_workflows(cls, db: AsyncSession): """异步获取定时触发的工作流""" result = await db.execute(select(cls).where( @@ -96,7 +91,6 @@ class Workflow(Base): return list(result.scalars().all()) @classmethod - @legacy_db_query def get_event_triggered_workflows(cls, db): """获取事件触发的工作流""" return list(db.execute(select(cls).where( @@ -107,7 +101,6 @@ class Workflow(Base): )).scalars().all()) @classmethod - @legacy_async_db_query async def async_get_event_triggered_workflows(cls, db: AsyncSession): """异步获取事件触发的工作流""" result = await db.execute(select(cls).where( @@ -119,12 +112,10 @@ class Workflow(Base): return list(result.scalars().all()) @classmethod - @legacy_db_query def get_by_name(cls, db, name: str): return db.execute(select(cls).where(cls.name == name)).scalars().first() @classmethod - @legacy_async_db_query async def async_get_by_name(cls, db: AsyncSession, name: str): result = await db.execute(select(cls).where(cls.name == name)) return result.scalars().first() diff --git a/app/db/oper/agentchat.py b/app/db/oper/agentchat.py index 66c9240d5..3622fa879 100644 --- a/app/db/oper/agentchat.py +++ b/app/db/oper/agentchat.py @@ -321,6 +321,10 @@ class AgentChatOper(DbOper): await self._stage_async_delete(AgentChat, chat.id) return True + def delete_by_id(self, chat_id: int) -> None: + """在 Oper 事务边界内按主键删除 Agent 会话。""" + self._stage_delete(AgentChat, chat_id) + async def async_stage_delete( self, session_id: str, diff --git a/app/db/oper/downloadfailure.py b/app/db/oper/downloadfailure.py index 8798f18ee..2f1d5fd12 100644 --- a/app/db/oper/downloadfailure.py +++ b/app/db/oper/downloadfailure.py @@ -17,10 +17,12 @@ class DownloadFailureOper(DbOper): """ 批量按指纹查询仍在冷却期的失败记录。 """ - failures = DownloadFailure.get_active_by_fingerprints( - self._db, - fingerprints=fingerprints, - now_time=now_time, + failures = self._execute_sync_query( + lambda session: DownloadFailure.get_active_by_fingerprints( + session, + fingerprints=fingerprints, + now_time=now_time, + ) ) return { failure.fingerprint: failure @@ -38,12 +40,14 @@ class DownloadFailureOper(DbOper): """ 新增或更新资源失败记录。 """ - return DownloadFailure.record_failure( - self._db, - fingerprint=fingerprint, - now_time=now_time, - next_retry_at=next_retry_at, - **kwargs, + return self._execute_sync_write( + lambda session: DownloadFailure.record_failure( + session, + fingerprint=fingerprint, + now_time=now_time, + next_retry_at=next_retry_at, + **kwargs, + ) ) def delete_expired( diff --git a/app/db/oper/downloadhistory.py b/app/db/oper/downloadhistory.py index 08f9902df..9c173fa6d 100644 --- a/app/db/oper/downloadhistory.py +++ b/app/db/oper/downloadhistory.py @@ -289,7 +289,7 @@ class DownloadHistoryOper(DbOper): self._stage_delete(DownloadHistory, historyid) def stage_delete_history(self, historyid: int) -> None: - """暂存下载记录删除,不由模型装饰器提交事务。""" + """暂存下载记录删除,事务由调用方统一提交。""" self._db.execute( sqlalchemy_delete(DownloadHistory).where( DownloadHistory.id == historyid diff --git a/app/db/oper/plugindata.py b/app/db/oper/plugindata.py index a86525b53..218b8cf02 100644 --- a/app/db/oper/plugindata.py +++ b/app/db/oper/plugindata.py @@ -19,7 +19,11 @@ class PluginDataOper(DbOper): :param key: 数据key :param value: 数据值 """ - plugin = PluginData.get_plugin_data_by_key(self._db, plugin_id, key) + plugin = self._execute_sync_query( + lambda session: PluginData.get_plugin_data_by_key( + session, plugin_id, key + ) + ) if plugin: self._stage_update(plugin, { "value": value @@ -35,8 +39,10 @@ class PluginDataOper(DbOper): :param key: 数据键 :param value: 数据值 """ - plugin = await PluginData.async_get_plugin_data_by_key( - self._db, plugin_id, key + plugin = await self._execute_async_query( + lambda session: PluginData.async_get_plugin_data_by_key( + session, plugin_id, key + ) ) if plugin: await self._stage_async_update(plugin, {"value": value}) @@ -52,12 +58,18 @@ class PluginDataOper(DbOper): :param key: 数据key """ if key: - data = PluginData.get_plugin_data_by_key(self._db, plugin_id, key) + data = self._execute_sync_query( + lambda session: PluginData.get_plugin_data_by_key( + session, plugin_id, key + ) + ) if not data: return None return data.value else: - return PluginData.get_plugin_data(self._db, plugin_id) + return self._execute_sync_query( + lambda session: PluginData.get_plugin_data(session, plugin_id) + ) async def async_get_data(self, plugin_id: str, key: Optional[str] = None) -> Any: """ @@ -66,13 +78,17 @@ class PluginDataOper(DbOper): :param key: 数据key """ if key: - data = await PluginData.async_get_plugin_data_by_key( - self._db, plugin_id, key + data = await self._execute_async_query( + lambda session: PluginData.async_get_plugin_data_by_key( + session, plugin_id, key + ) ) if not data: return None return data.value - return await PluginData.async_get_plugin_data(self._db, plugin_id) + return await self._execute_async_query( + lambda session: PluginData.async_get_plugin_data(session, plugin_id) + ) def del_data(self, plugin_id: str, key: Optional[str] = None) -> Any: """ @@ -81,7 +97,7 @@ class PluginDataOper(DbOper): :param key: 数据key """ def stage(session: Session) -> None: - """把兼容删除入口映射到调用方或组合根持有的事务。""" + """把删除入口映射到调用方或组合根持有的事务。""" if key: PluginData.del_plugin_data_by_key(session, plugin_id, key) else: @@ -109,11 +125,19 @@ class PluginDataOper(DbOper): 获取插件所有数据 :param plugin_id: 插件id """ - return PluginData.get_plugin_data_by_plugin_id(self._db, plugin_id) + return self._execute_sync_query( + lambda session: PluginData.get_plugin_data_by_plugin_id( + session, plugin_id + ) + ) async def async_get_data_all(self, plugin_id: str) -> Any: """ 异步获取插件所有数据。 :param plugin_id: 插件id """ - return await PluginData.async_get_plugin_data_by_plugin_id(self._db, plugin_id) + return await self._execute_async_query( + lambda session: PluginData.async_get_plugin_data_by_plugin_id( + session, plugin_id + ) + ) diff --git a/app/db/oper/site.py b/app/db/oper/site.py index e1be9263b..d6d7581f4 100644 --- a/app/db/oper/site.py +++ b/app/db/oper/site.py @@ -75,7 +75,7 @@ class SiteOper(DbOper): site_id: int, payload: Mapping[str, Any], ) -> bool: - """暂存站点字段更新,不由模型装饰器提前提交。""" + """暂存站点字段更新,事务由调用方统一提交。""" site = await self.async_get(site_id) if not site: return False @@ -338,18 +338,22 @@ class SiteOper(DbOper): async def async_get_icon_by_domain(self, domain: str) -> Optional[SiteIcon]: """异步按域名获取站点图标。""" - return await SiteIcon.async_get_by_domain(self._db, domain) + return await self._execute_async_query( + lambda session: SiteIcon.async_get_by_domain(session, domain) + ) async def async_get_statistic_by_domain( self, domain: str, ) -> Optional[SiteStatistic]: """异步按域名获取站点统计。""" - return await SiteStatistic.async_get_by_domain(self._db, domain) + return await self._execute_async_query( + lambda session: SiteStatistic.async_get_by_domain(session, domain) + ) async def async_list_statistics(self) -> List[SiteStatistic]: """异步获取所有站点统计。""" - return await SiteStatistic.async_list(self._db) + return await self._execute_async_query(SiteStatistic.async_list) def get_userdata_by_date(self, date: str) -> List[SiteUserData]: """ @@ -371,7 +375,9 @@ class SiteOper(DbOper): """ 按域名获取站点图标 """ - return SiteIcon.get_by_domain(self._db, domain) + return self._execute_sync_query( + lambda session: SiteIcon.get_by_domain(session, domain) + ) def update_icon(self, name: str, domain: str, icon_url: str, icon_base64: str) -> bool: """ @@ -467,60 +473,59 @@ class SiteOper(DbOper): """ 异步站点访问成功 """ - lst_date = datetime.now().strftime("%Y-%m-%d %H:%M:%S") - sta = await SiteStatistic.async_get_by_domain(self._db, domain) - if sta: - # 使用深复制确保 note 是全新的字典对象 - note = dict(sta.note) if sta.note else {} - avg_seconds = None - - if seconds is not None: - note[lst_date] = seconds or 1 - avg_times = len(note.keys()) - if avg_times > 10: - note = dict(sorted(note.items(), key=lambda x: x[0], reverse=True)[:10]) - avg_seconds = sum([v for v in note.values()]) // avg_times - - await self._stage_async_update(sta, { - "success": sta.success + 1, - "seconds": avg_seconds or sta.seconds, - "lst_state": 0, - "lst_mod_date": lst_date, - "note": note - }) - else: - note = {} - if seconds is not None: - note = { - lst_date: seconds or 1 - } - await self._stage_async_create(SiteStatistic( + async def write(session: AsyncSession) -> None: + """在同一异步事务中读取并更新站点成功统计。""" + lst_date = datetime.now().strftime("%Y-%m-%d %H:%M:%S") + sta = await SiteStatistic.async_get_by_domain(session, domain) + if sta: + note = dict(sta.note) if sta.note else {} + avg_seconds = None + if seconds is not None: + note[lst_date] = seconds or 1 + avg_times = len(note.keys()) + if avg_times > 10: + note = dict(sorted( + note.items(), key=lambda item: item[0], reverse=True + )[:10]) + avg_seconds = sum(note.values()) // avg_times + sta.success += 1 + sta.seconds = avg_seconds or sta.seconds + sta.lst_state = 0 + sta.lst_mod_date = lst_date + sta.note = note + return + note = {lst_date: seconds or 1} if seconds is not None else {} + session.add(SiteStatistic( domain=domain, success=1, fail=0, seconds=seconds or 1, lst_state=0, lst_mod_date=lst_date, - note=note + note=note, )) + await self._execute_async_write(write) + async def async_fail(self, domain: str): """ 异步站点访问失败 """ - lst_date = datetime.now().strftime("%Y-%m-%d %H:%M:%S") - sta = await SiteStatistic.async_get_by_domain(self._db, domain) - if sta: - await self._stage_async_update(sta, { - "fail": sta.fail + 1, - "lst_state": 1, - "lst_mod_date": lst_date - }) - else: - await self._stage_async_create(SiteStatistic( + async def write(session: AsyncSession) -> None: + """在同一异步事务中读取并更新站点失败统计。""" + lst_date = datetime.now().strftime("%Y-%m-%d %H:%M:%S") + sta = await SiteStatistic.async_get_by_domain(session, domain) + if sta: + sta.fail += 1 + sta.lst_state = 1 + sta.lst_mod_date = lst_date + return + session.add(SiteStatistic( domain=domain, success=0, fail=1, lst_state=1, - lst_mod_date=lst_date + lst_mod_date=lst_date, )) + + await self._execute_async_write(write) diff --git a/app/db/oper/subscribe.py b/app/db/oper/subscribe.py index cd55d1e26..209858171 100644 --- a/app/db/oper/subscribe.py +++ b/app/db/oper/subscribe.py @@ -11,7 +11,7 @@ import time from collections.abc import Awaitable, Callable from typing import Any, Tuple, List, Optional -from sqlalchemy import delete as sqlalchemy_delete, select +from sqlalchemy import delete as sqlalchemy_delete from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Session @@ -100,25 +100,6 @@ class SubscribeOper(DbOper): 订阅管理 """ - @staticmethod - def _identity_statement(identity: dict, username: Optional[str] = None): - """构造订阅查重语句,SQL 所有权收口在 Oper。""" - condition = Subscribe._identity_condition( # pylint: disable=protected-access - identity.get("media_source"), - identity.get("media_id"), - identity.get("music_type"), - ) - if condition is None or username == "": - return None - statement = select(Subscribe).where(condition) - if username: - statement = statement.where(Subscribe.username == username) - if identity.get("season") is not None: - statement = statement.where(Subscribe.season == identity["season"]) - return statement.where( - Subscribe.episode_group == identity.get("episode_group") - ) - def _exists(self, identity: dict, username: Optional[str]) -> Optional[Any]: """ 按身份查重。 @@ -126,19 +107,19 @@ class SubscribeOper(DbOper): :param username: 非空时只在该用户的订阅内查 :return: 命中的订阅行,未命中为 None """ - if isinstance(self._db, Session): - statement = self._identity_statement(identity, username) - if statement is None: - return None - return self._db.execute(statement).scalars().first() - # 旧 SDK 允许无会话构造 Oper;保留其自动短会话行为,但规范入口不得走这里。 + if username == "": + return None if username: - return Subscribe.exists_by_username( - self._db, - username=username, - **identity, + return self._execute_sync_query( + lambda session: Subscribe.exists_by_username( + session, + username=username, + **identity, + ) ) - return Subscribe.exists(self._db, **identity) + return self._execute_sync_query( + lambda session: Subscribe.exists(session, **identity) + ) async def _async_exists(self, identity: dict, username: Optional[str]) -> Optional[Any]: """ @@ -147,20 +128,18 @@ class SubscribeOper(DbOper): :param username: 非空时只在该用户的订阅内查 :return: 命中的订阅行,未命中为 None """ - if isinstance(self._db, AsyncSession): - statement = self._identity_statement(identity, username) - if statement is None: + async def query(session: AsyncSession) -> Optional[Subscribe]: + """在调用方或组合根异步会话中执行订阅查重。""" + if username == "": return None - result = await self._db.execute(statement) - return result.scalars().first() - # 同步路径一样只为无会话旧入口保留 Model 的自动短会话兼容。 - if username: - return await Subscribe.async_exists_by_username( - self._db, - username=username, - **identity, - ) - return await Subscribe.async_exists(self._db, **identity) + if username: + return await Subscribe.async_exists_by_username( + session, + username=username, + **identity, + ) + return await Subscribe.async_exists(session, **identity) + return await self._execute_async_query(query) def stage_add( self, @@ -297,24 +276,15 @@ class SubscribeOper(DbOper): """ 获取订阅 """ - return self._execute_sync_query( - lambda session: session.execute( - select(Subscribe).where(Subscribe.id == sid) - ).scalars().first() - ) + return self._execute_sync_query(lambda session: Subscribe.get(session, sid)) async def async_get(self, sid: int) -> Optional[Subscribe]: """ 获取订阅 """ - if self._db is not None and not isinstance(self._db, (Session, AsyncSession)): - # 保留旧测试替身与插件注入对象对 Model ABI 的兼容入口。 - return await Subscribe.async_get(self._db, rid=sid) - async def query(session: AsyncSession) -> Optional[Subscribe]: - """在调用方异步会话中执行订阅主键查询。""" - result = await session.execute(select(Subscribe).where(Subscribe.id == sid)) - return result.scalars().first() - return await self._execute_async_query(query) + return await self._execute_async_query( + lambda session: Subscribe.async_get(session, sid) + ) async def async_list_by_media_identity( self, @@ -323,18 +293,14 @@ class SubscribeOper(DbOper): music_type: Optional[str] = None, ) -> List[Subscribe]: """异步按规范媒体身份读取订阅。""" - async def query(session: AsyncSession) -> List[Subscribe]: - """在调用方异步会话中执行媒体身份列表查询。""" - condition = Subscribe._identity_condition( # pylint: disable=protected-access - media_source, media_id, music_type + return await self._execute_async_query( + lambda session: Subscribe.async_list_by_media_identity( + session, + media_source=media_source, + media_id=media_id, + music_type=music_type, ) - if condition is None: - return [] - result = await session.execute(select(Subscribe).where(condition)) - return list(result.scalars().all()) - if isinstance(self._db, AsyncSession): - return await query(self._db) - return await self._execute_async_query(query) + ) def list_by_media_identity( self, @@ -343,15 +309,14 @@ class SubscribeOper(DbOper): music_type: Optional[str] = None, ) -> List[Subscribe]: """同步按规范媒体身份读取订阅。""" - def query(session: Session) -> List[Subscribe]: - """在调用方同步会话中执行媒体身份列表查询。""" - condition = Subscribe._identity_condition( # pylint: disable=protected-access - media_source, media_id, music_type + return self._execute_sync_query( + lambda session: Subscribe.list_by_media_identity( + session, + media_source=media_source, + media_id=media_id, + music_type=music_type, ) - if condition is None: - return [] - return list(session.execute(select(Subscribe).where(condition)).scalars().all()) - return self._execute_sync_query(query) + ) async def get_candidate( self, @@ -423,18 +388,16 @@ class SubscribeOper(DbOper): """ 根据条件查询订阅 """ - def query(session: Session) -> Optional[Subscribe]: - """在调用方同步会话中执行类型媒体查询。""" - condition = Subscribe._identity_condition( # pylint: disable=protected-access - media_source, media_id, music_type + return self._execute_sync_query( + lambda session: Subscribe.get_by( + session, + type=type, + media_source=media_source, + media_id=media_id, + season=season, + music_type=music_type, ) - if condition is None: - return None - statement = select(Subscribe).where(condition, Subscribe.type == type) - if season is not None: - statement = statement.where(Subscribe.season == season) - return session.execute(statement).scalars().first() - return self._execute_sync_query(query) + ) async def async_get_by( self, type: str, media_source: MediaSource, media_id: str, @@ -444,55 +407,34 @@ class SubscribeOper(DbOper): """ 根据条件查询订阅 """ - async def query(session: AsyncSession) -> Optional[Subscribe]: - """在调用方异步会话中执行类型媒体查询。""" - condition = Subscribe._identity_condition( # pylint: disable=protected-access - media_source, media_id, music_type + return await self._execute_async_query( + lambda session: Subscribe.async_get_by( + session, + type=type, + media_source=media_source, + media_id=media_id, + season=season, + music_type=music_type, ) - if condition is None: - return None - statement = select(Subscribe).where(condition, Subscribe.type == type) - if season is not None: - statement = statement.where(Subscribe.season == season) - result = await session.execute(statement) - return result.scalars().first() - return await self._execute_async_query(query) + ) def list(self, state: Optional[str] = None) -> List[Subscribe]: """ 获取订阅列表 """ - if state: - return self._execute_sync_query( - lambda session: list(session.execute( - select(Subscribe).where(Subscribe.state.in_(state.split(','))) - ).scalars().all()) - ) return self._execute_sync_query( - lambda session: list(session.execute(select(Subscribe)).scalars().all()) + lambda session: Subscribe.get_by_state(session, state) ) async def async_list(self, state: Optional[str] = None) -> List[Subscribe]: """ 异步获取订阅列表 """ - if self._db is not None and not isinstance(self._db, (Session, AsyncSession)): - if state: - return await Subscribe.async_get_by_state(self._db, state) - return await Subscribe.async_list(self._db) if state: - async def query(session: AsyncSession) -> List[Subscribe]: - """在调用方异步会话中执行状态列表查询。""" - result = await session.execute( - select(Subscribe).where(Subscribe.state.in_(state.split(','))) - ) - return list(result.scalars().all()) - return await self._execute_async_query(query) - async def query_all(session: AsyncSession) -> List[Subscribe]: - """在调用方异步会话中执行全量订阅查询。""" - result = await session.execute(select(Subscribe)) - return list(result.scalars().all()) - return await self._execute_async_query(query_all) + return await self._execute_async_query( + lambda session: Subscribe.async_get_by_state(session, state) + ) + return await self._execute_async_query(Subscribe.async_list) async def async_list_by_username( self, @@ -501,35 +443,28 @@ class SubscribeOper(DbOper): mtype: Optional[str] = None, ) -> List[Subscribe]: """异步按用户获取订阅。""" - if self._db is not None and not isinstance(self._db, (Session, AsyncSession)): - return await Subscribe.async_list_by_username( - self._db, username=username, state=state, mtype=mtype + return await self._execute_async_query( + lambda session: Subscribe.async_list_by_username( + session, + username=username, + state=state, + mtype=mtype, ) - async def query(session: AsyncSession) -> List[Subscribe]: - """在调用方异步会话中执行用户筛选查询。""" - statement = select(Subscribe).where(Subscribe.username == username) - if state: - statement = statement.where(Subscribe.state == state) - if mtype: - statement = statement.where(Subscribe.type == mtype) - result = await session.execute(statement) - return list(result.scalars().all()) - return await self._execute_async_query(query) + ) async def async_list_by_title( self, title: str, season: Optional[int] = None, ) -> List[Subscribe]: - """异步按标题获取订阅,供旧查询测试和迁移调用兼容。""" - async def query(session: AsyncSession) -> List[Subscribe]: - """在调用方异步会话中执行标题列表查询。""" - statement = select(Subscribe).where(Subscribe.name == title) - if season is not None: - statement = statement.where(Subscribe.season == season) - result = await session.execute(statement) - return list(result.scalars().all()) - return await self._execute_async_query(query) + """在 Oper 会话边界内异步按标题获取订阅。""" + return await self._execute_async_query( + lambda session: Subscribe.async_list_by_title( + session, + title=title, + season=season, + ) + ) def delete(self, sid: int): """ @@ -598,30 +533,22 @@ class SubscribeOper(DbOper): """ 获取指定用户的订阅 """ - def query(session: Session) -> List[Subscribe]: - """在调用方同步会话中执行用户筛选查询。""" - statement = select(Subscribe).where(Subscribe.username == username) - if state: - statement = statement.where(Subscribe.state == state) - if mtype: - statement = statement.where(Subscribe.type == mtype) - return list(session.execute(statement).scalars().all()) - return self._execute_sync_query(query) + return self._execute_sync_query( + lambda session: Subscribe.list_by_username( + session, + username=username, + state=state, + mtype=mtype, + ) + ) def list_by_type(self, mtype: str, days: int = 7) -> List[Subscribe]: """ 获取指定类型的订阅 """ - def query(session: Session) -> List[Subscribe]: - """在调用方同步会话中执行时间窗订阅查询。""" - cutoff = time.strftime( - "%Y-%m-%d %H:%M:%S", - time.localtime(time.time() - 86400 * int(days)), - ) - return list(session.execute(select(Subscribe).where( - Subscribe.type == mtype, Subscribe.date >= cutoff - )).scalars().all()) - return self._execute_sync_query(query) + return self._execute_sync_query( + lambda session: Subscribe.list_by_type(session, mtype, days) + ) def add_history(self, **kwargs): """ diff --git a/app/db/oper/systemconfig.py b/app/db/oper/systemconfig.py index 63be40ce0..ad9789c56 100644 --- a/app/db/oper/systemconfig.py +++ b/app/db/oper/systemconfig.py @@ -2,6 +2,8 @@ import copy import threading from typing import Any, Optional, Union +from sqlalchemy.orm import Session + from app.db.base import DbOper from app.db.models.systemconfig import SystemConfig from app.schemas.types import SystemConfigKey @@ -20,12 +22,15 @@ class SystemConfigOper(DbOper, metaclass=Singleton): self._write_lock = threading.RLock() self._loaded = False - def load_snapshot(self) -> None: - """从数据库加载完整配置,并一次性发布新的内存快照。""" + def load_snapshot(self, db: Optional[Session] = None) -> None: + """从显式会话或 Oper 事务边界加载配置并发布内存快照。""" with self._write_lock: + items = SystemConfig.list(db) if db is not None else self._execute_sync_query( + SystemConfig.list + ) snapshot = { item.key: copy.deepcopy(item.value) - for item in SystemConfig.list(self._db) + for item in items } with self._snapshot_lock: self.__SYSTEMCONF = snapshot diff --git a/app/db/oper/transferhistory.py b/app/db/oper/transferhistory.py index 9edcfcab6..500402c5f 100644 --- a/app/db/oper/transferhistory.py +++ b/app/db/oper/transferhistory.py @@ -269,7 +269,7 @@ class TransferHistoryOper(DbOper): self._stage_delete(TransferHistory, historyid) def stage_delete(self, historyid: int) -> None: - """暂存整理记录删除,不由模型装饰器提交事务。""" + """暂存整理记录删除,事务由调用方统一提交。""" self._db.execute( sqlalchemy_delete(TransferHistory).where( TransferHistory.id == historyid diff --git a/app/db/oper/user.py b/app/db/oper/user.py index 7f419bcf2..cecfbb7a8 100644 --- a/app/db/oper/user.py +++ b/app/db/oper/user.py @@ -27,7 +27,7 @@ class UserOper(DbOper): """ 获取用户列表 """ - return User.list(self._db) + return self._execute_sync_query(User.list) def add(self, **kwargs): """ @@ -40,15 +40,19 @@ class UserOper(DbOper): """ 根据用户名获取用户 """ - return User.get_by_name(self._db, name) + return self._execute_sync_query( + lambda session: User.get_by_name(session, name) + ) def get_by_id(self, user_id: int) -> Optional[User]: """按 ID 获取用户。""" - return User.get_by_id(self._db, user_id) + return self._execute_sync_query( + lambda session: User.get_by_id(session, user_id) + ) async def async_list(self) -> List[User]: """异步获取用户列表。""" - return await User.async_list(self._db) + return await self._execute_async_query(User.async_list) async def async_create(self, payload: dict) -> Optional[User]: """异步创建用户。""" @@ -126,7 +130,7 @@ class UserOper(DbOper): """ 获取用户权限 """ - user = User.get_by_name(self._db, name) + user = self.get_by_name(name) if user: return user.permissions or {} return {} @@ -135,7 +139,7 @@ class UserOper(DbOper): """ 获取用户个性化设置,返回None表示用户不存在 """ - user = User.get_by_name(self._db, name) + user = self.get_by_name(name) if user: return user.settings or {} return None diff --git a/app/db/oper/userconfig.py b/app/db/oper/userconfig.py index 7b1656bdc..d5f22c5e9 100644 --- a/app/db/oper/userconfig.py +++ b/app/db/oper/userconfig.py @@ -2,6 +2,8 @@ import copy import threading from typing import Any, Union, Dict, Optional +from sqlalchemy.orm import Session + from app.db.base import DbOper from app.db.models.userconfig import UserConfig from app.schemas.types import UserConfigKey @@ -20,11 +22,14 @@ class UserConfigOper(DbOper, metaclass=Singleton): self._write_lock = threading.RLock() self._loaded = False - def load_snapshot(self) -> None: - """从数据库加载完整用户配置,并一次性发布新的内存快照。""" + def load_snapshot(self, db: Optional[Session] = None) -> None: + """从显式会话或 Oper 事务边界加载用户配置并发布内存快照。""" with self._write_lock: snapshot: dict[str, dict[str, Any]] = {} - for item in UserConfig.list(self._db): + items = UserConfig.list(db) if db is not None else self._execute_sync_query( + UserConfig.list + ) + for item in items: if item.username and item.key: snapshot.setdefault(item.username, {})[item.key] = copy.deepcopy( item.value diff --git a/app/db/oper/workflow.py b/app/db/oper/workflow.py index 5ee6fde14..51334a411 100644 --- a/app/db/oper/workflow.py +++ b/app/db/oper/workflow.py @@ -1,6 +1,7 @@ from typing import List, Mapping, Tuple, Optional, Any, Protocol from sqlalchemy import delete as sqlalchemy_delete +from sqlalchemy.orm import Session from app.db.base import DbOper from app.db.models.workflow import Workflow @@ -202,6 +203,8 @@ class WorkflowOper(DbOper): def stage_start(self, wid: int) -> bool: """在调用方持有的会话中暂存运行中状态。""" + if not isinstance(self._db, Session): + raise RuntimeError("工作流暂存写入需要调用方提供同步 Session") return Workflow.start(self._db, wid) def success(self, wid: int, result: Optional[str] = None) -> bool: @@ -214,6 +217,8 @@ class WorkflowOper(DbOper): def stage_success(self, wid: int, result: Optional[str] = None) -> bool: """在调用方持有的会话中暂存成功状态。""" + if not isinstance(self._db, Session): + raise RuntimeError("工作流暂存写入需要调用方提供同步 Session") return Workflow.success(self._db, wid, result) def fail(self, wid: int, result: str) -> bool: @@ -226,6 +231,8 @@ class WorkflowOper(DbOper): def stage_fail(self, wid: int, result: str) -> bool: """在调用方持有的会话中暂存失败状态。""" + if not isinstance(self._db, Session): + raise RuntimeError("工作流暂存写入需要调用方提供同步 Session") return Workflow.fail(self._db, wid, result) def step( @@ -255,6 +262,8 @@ class WorkflowOper(DbOper): execution_state: Optional[dict[str, Any]] = None, ) -> bool: """在调用方持有的会话中暂存动作进度。""" + if not isinstance(self._db, Session): + raise RuntimeError("工作流暂存写入需要调用方提供同步 Session") return Workflow.update_current_action( self._db, wid, @@ -277,4 +286,6 @@ class WorkflowOper(DbOper): reset_count: bool = False, ) -> bool: """在调用方持有的会话中暂存执行状态重置。""" + if not isinstance(self._db, Session): + raise RuntimeError("工作流暂存写入需要调用方提供同步 Session") return Workflow.reset(self._db, wid, reset_count=reset_count) diff --git a/app/db/uow.py b/app/db/uow.py index 0e6ab270d..cd9443e21 100644 --- a/app/db/uow.py +++ b/app/db/uow.py @@ -1,4 +1,4 @@ -"""SQLAlchemy 请求级事务适配器与旧 Oper 事务执行端口。""" +"""SQLAlchemy 请求级事务适配器与无会话 Oper 事务执行端口。""" from collections.abc import Awaitable, Callable from typing import Protocol, TypeVar @@ -11,7 +11,7 @@ T = TypeVar("T") class SyncTransactionRunner(Protocol): - """为无显式 Session 的兼容写入口提供独占同步事务。""" + """为无显式 Session 的 Oper 入口提供独占同步事务。""" def __call__(self, operation: Callable[[Session], T]) -> T: """在一个独占会话中执行并提交操作。""" @@ -19,7 +19,7 @@ class SyncTransactionRunner(Protocol): class AsyncTransactionRunner(Protocol): - """为无显式 Session 的兼容写入口提供独占异步事务。""" + """为无显式 Session 的 Oper 入口提供独占异步事务。""" def __call__( self, @@ -38,14 +38,14 @@ def configure_transaction_runners( sync: SyncTransactionRunner, async_: AsyncTransactionRunner, ) -> None: - """由组合根登记旧 Oper 兼容入口使用的显式事务执行器。""" + """由组合根登记无会话 Oper 入口使用的显式事务执行器。""" global _sync_transaction_runner, _async_transaction_runner _sync_transaction_runner = sync _async_transaction_runner = async_ def run_sync_transaction(operation: Callable[[Session], T]) -> T: - """委托组合根在独占同步事务中执行兼容写操作。""" + """委托组合根在独占同步事务中执行 Oper 操作。""" if _sync_transaction_runner is None: raise RuntimeError("同步事务执行器尚未配置") return _sync_transaction_runner(operation) @@ -54,7 +54,7 @@ def run_sync_transaction(operation: Callable[[Session], T]) -> T: async def run_async_transaction( operation: Callable[[AsyncSession], Awaitable[T]], ) -> T: - """委托组合根在独占异步事务中执行兼容写操作。""" + """委托组合根在独占异步事务中执行 Oper 操作。""" if _async_transaction_runner is None: raise RuntimeError("异步事务执行器尚未配置") return await _async_transaction_runner(operation) diff --git a/app/sdk/_legacy/history.py b/app/sdk/_legacy/history.py index e8c993315..ccdaf6d14 100644 --- a/app/sdk/_legacy/history.py +++ b/app/sdk/_legacy/history.py @@ -1,9 +1,8 @@ """把旧整理历史 Oper 的业务写入方法转交给应用服务。""" -from typing import Optional +from typing import Any, Optional from app.application.history import add_transfer_fail, add_transfer_success -from app.db.models.transferhistory import TransferHistory from app.db.oper.transferhistory import TransferHistoryOper as CanonicalTransferHistoryOper from app.domain.context import MediaInfo, MusicInfo from app.domain.meta.metabase import MetaBase @@ -23,7 +22,7 @@ class TransferHistoryOper(CanonicalTransferHistoryOper): transferinfo: TransferInfo, downloader: Optional[str] = None, download_hash: Optional[str] = None, - ) -> Optional[TransferHistory]: + ) -> Optional[Any]: """ 按旧签名新增整理成功历史。 @@ -49,7 +48,7 @@ class TransferHistoryOper(CanonicalTransferHistoryOper): transferinfo: Optional[TransferInfo] = None, downloader: Optional[str] = None, download_hash: Optional[str] = None, - ) -> Optional[TransferHistory]: + ) -> Optional[Any]: """ 按旧签名新增整理失败历史。 diff --git a/app/sdk/_legacy/subscribe.py b/app/sdk/_legacy/subscribe.py index ed3f80df1..1f34aad26 100644 --- a/app/sdk/_legacy/subscribe.py +++ b/app/sdk/_legacy/subscribe.py @@ -3,7 +3,6 @@ from typing import Any, Optional from app.application.subscription.write import add_subscribe, async_add_subscribe -from app.db.models.subscribe import Subscribe from app.db.oper.subscribe import SubscribeOper as CanonicalSubscribeOper from app.domain.context import MediaInfo, MusicInfo @@ -72,4 +71,4 @@ class SubscribeOper(CanonicalSubscribeOper): ) -__all__ = ["Subscribe", "SubscribeOper"] +__all__ = ["SubscribeOper"] diff --git a/app/sdk/_legacy/user.py b/app/sdk/_legacy/user.py index c8cc11b4f..f80d78472 100644 --- a/app/sdk/_legacy/user.py +++ b/app/sdk/_legacy/user.py @@ -11,12 +11,10 @@ from app.api.deps import ( get_current_user_async, ) from app.db.oper.user import UserOper -from app.db.models.user import User __all__ = [ "UserOper", - "User", "get_current_active_manage_user", "get_current_active_manage_user_async", "get_current_active_superuser", diff --git a/app/testing/bootstrap.py b/app/testing/bootstrap.py index 155b61401..4062d8198 100644 --- a/app/testing/bootstrap.py +++ b/app/testing/bootstrap.py @@ -184,9 +184,11 @@ def prepare_backend() -> None: init_db() from app.db.oper.systemconfig import SystemConfigOper from app.db.oper.userconfig import UserConfigOper + from app.db.session import SessionFactory - SystemConfigOper().load_snapshot() - UserConfigOper().load_snapshot() + with SessionFactory() as session: + SystemConfigOper().load_snapshot(session) + UserConfigOper().load_snapshot(session) # 缓存装饰器在测试模块导入时即创建后端,先装配隔离配置对应的适配器。 from app.startup.initializers.cache import configure_cache_dependencies configure_cache_dependencies() diff --git a/docs/architecture-overview.md b/docs/architecture-overview.md index 32cb65032..360d96f58 100644 --- a/docs/architecture-overview.md +++ b/docs/architecture-overview.md @@ -379,10 +379,11 @@ flowchart LR `db/adapters/subscription.py` 创建独占 Session,`startup/composition/subscription.py` 只装配回调, `application/subscription/write.py` 决定事务与 post-commit 边界,`SubscribeOper.stage_add()` 只查重、`add` 和 `flush`。旧 SDK 显式构造的无会话 Oper 暂留兼容自动短会话,不得被新代码复用。 - `transaction-debt-baseline.json` 当前要求正式只读查询装饰器保持为 0;原有同步/异步写装饰器 - 已全部移除,`db_update` 与 `async_db_update` 必须持续保持为 0。下载/整理历史的旧插件 Model - 与工作流、媒体服务器、站点用户数据、PassKey、SubscribeHistory 旧插件 Model 调用由 `legacy_*` 兼容外壳承接,宿主 Oper 必须显式传递 Session。宿主 Oper 也不得调用 Base 保留的 - `create/update/delete/truncate` 兼容包装器;AST 门禁保证显式 Session 的提交权不会被底层抢走。 + `transaction-debt-baseline.json` 要求 Model 上的查询/写装饰器持续保持为 0。Model 与 Base + 已不再导入数据库装饰器,所有 `db` 参数都要求显式 Session;这些方法只查询或 stage,不能 + 创建、提交、回滚或关闭事务。无会话入口只存在于 Oper,由 `_execute_*` 经组合根事务执行器 + 承接;内置插件必须调用 Oper,不得直接导入宿主 Model。AST 门禁同时约束装饰器、可选 Session + 和插件到 Model 的依赖,保证提交权不会被底层抢走。 - 站点、历史、工作流、Agent 会话删除和插件数据重置已经形成同构事务切片;对应 Application Command/Service 持有 UoW,Oper 的 `stage_*` 方法只修改当前会话。插件数据重置从 `startup/initializers/plugins.py` 注入事务能力,插件直接使用 `PluginDataOper` 的旧 ABI 仅作兼容。 diff --git a/docs/refactor/backend-architecture-governance.md b/docs/refactor/backend-architecture-governance.md index b22cbbecd..f10b4fd55 100644 --- a/docs/refactor/backend-architecture-governance.md +++ b/docs/refactor/backend-architecture-governance.md @@ -377,12 +377,12 @@ app/chain/transfer.py # 保持 TransferChain 兼容门面 - `app/api/endpoints/subscribe.py` 直接持有 Session、模型和 Oper 是治理前证据;当前 endpoint→Session/Model 目标边已清零。 - Chain、Scheduler、Application 的模型直连属于治理前扫描结果;当前目标 Application/Chain/Runtime→DB 边均为零。 -- `app/db/models/subscribe.py:121` 起在 ORM 模型上定义查询方法,并通过 `@db_query` 等装饰器执行数据库访问。 +- ORM Model 仍保留贴近表结构的查询原语,但已全部要求调用方显式传入 Session;Model/Base 的查询、写入和 legacy 事务装饰器均已清零。 - `app/db/__init__.py` 的根入口和模型回流曾参与 DB SCC;该自有 SCC 已消除,旧根入口仅作为兼容边界保留。 #### 问题本质 -当前同时存在三种数据访问风格: +治理前同时存在三种数据访问风格: 1. `db/oper` 服务。 2. ORM 模型类方法。 diff --git a/docs/refactor/backend-architecture-next-stage.md b/docs/refactor/backend-architecture-next-stage.md index 76a54974e..c15a0783d 100644 --- a/docs/refactor/backend-architecture-next-stage.md +++ b/docs/refactor/backend-architecture-next-stage.md @@ -6,7 +6,7 @@ > 审计范围:宿主后端;排除 `app/plugins/**` 运行时插件副本 > 规范优先级:`AGENTS.md` 与 `docs/rules/` 高于本文 > 相关文档:`docs/architecture-overview.md`、`docs/refactor/backend-architecture-governance.md`、`docs/refactor/backend-module-refactor-compatibility.md` -> 实施进度:阶段 0~6 的宿主架构能力已完成收口;API/Application 公共复杂度基线已清零,启动组合根的 SystemConfigOper 构造点已由 14 降至 1;API 进程内后台任务已完成首批统一登记,插件仓适配、Outbox 外围扩展和 Model 查询兼容面仍按风险切片推进。2026-08-23 的长期整改阶段 0 已恢复宿主、启动性能、官方插件和 SDK 契约门禁的可信基线;阶段 1a 已补齐 TaskRegistry owner 零债务门禁和诚实的关停超时语义;阶段 1b1 已收口整理 worker、pending 回放、失败通知、进程内 AI 重试、插件监控与事件投递的生命周期所有权。 +> 实施进度:阶段 0~6 的宿主架构能力已完成收口;API/Application 公共复杂度基线已清零,启动组合根的 SystemConfigOper 构造点已由 14 降至 1;API 进程内后台任务已完成首批统一登记,插件仓适配和 Outbox 外围扩展仍按风险切片推进。Model/Base 查询与写装饰器、legacy 隐式会话外壳均已清零,插件 SDK 也不再导出宿主 Model。2026-08-23 的长期整改阶段 0 已恢复宿主、启动性能、官方插件和 SDK 契约门禁的可信基线;阶段 1a 已补齐 TaskRegistry owner 零债务门禁和诚实的关停超时语义;阶段 1b1 已收口整理 worker、pending 回放、失败通知、进程内 AI 重试、插件监控与事件投递的生命周期所有权。 ## 当前复核结论(2026-08-23) @@ -15,7 +15,7 @@ ### 长期整改阶段 0:治理门禁恢复(2026-08-23) -- 宿主依赖基线已审查 TaskRegistry、有界后台 owner 与插件变更准入接入后的语义差异:当前为 `805` 个模块、`6525` 条内部导入边,12 组重点禁止边继续全部为 `0`,唯一非平凡 SCC 仍是隔离的 TMDB 移植包。 +- 宿主依赖基线已审查 TaskRegistry、有界后台 owner 与插件变更准入接入后的语义差异:当前为 `805` 个模块、`6500` 条内部导入边,12 组重点禁止边继续全部为 `0`,唯一非平凡 SCC 仍是隔离的 TMDB 移植包。 - 启动性能探针会在隔离生命周期中真实创建并释放 TaskRegistry;normal/safe 组件数分别为 `23`/`11`,CI 只读检查使用稳定的宿主模块集合和生命周期组件顺序,不再把 Python/平台模块数量当作硬合同。 - 官方插件快照覆盖 `plugins.v3`、`plugins.v2` 以及 V3 实际会从 `package.json` 回退加载的 31 个默认实现;`app/plugins/**` 仍只是宿主运行副本,不进入扫描。 - SDK 快照以各模块显式 `__all__` 为公开合同,能够记录赋值别名;`typing`、`__future__` 等实现期导入不再被误冻结,既有数据库备份门面已补精确导出清单。 @@ -72,7 +72,7 @@ - 继续采用单进程控制面是正确选择,不建议现在拆成微服务;插件、调度器、工作流、事件和数据库共享进程内状态,拆分会放大部署、事务和兼容成本。 - `foundation/domain/runtime/adapters/application/chain/api/startup` 的职责方向基本成立;宿主架构基线、复杂度 ratchet、异步阻塞 ratchet 当前均通过。 -- 依赖图当前为 `805` 个 Python 模块、`6525` 条内部导入边;唯一非平凡 SCC 位于隔离的 TMDB 第三方移植包内部,不应为了指标归零重写。 +- 依赖图当前为 `805` 个 Python 模块、`6500` 条内部导入边;唯一非平凡 SCC 位于隔离的 TMDB 第三方移植包内部,不应为了指标归零重写。 - 当前主要风险已经从“目录和依赖失控”转移到运行时协议、后台副作用的可靠性和遗留兼容面。换言之,下一阶段重点应是**语义收口和可验证性**,而不是继续搬文件或机械拆大文件。 综合评价:架构方向可持续,生产可用性较高;可演进性仍处于中等水平。现阶段没有静态审计发现必须立即推倒重来的 P0 架构问题,但存在需要按 P1/P2 计划治理的真实债务。 @@ -81,7 +81,7 @@ 1. **后台任务的统一所有权已覆盖 API 入口,但仍有更深层任务机制待分级。** `app/runtime/tasks.py` 已建立 lifespan 级 TaskRegistry,启动收尾、插件 Release 刷新、Webhook E0 广播、CookieCloud E1 手工调度、消息入口、Seerr 订阅、整理历史 AI 重做、OpenAI/Anthropic 协议流和 WebAgent 断线后执行/快照保存均不再维护端点模块级任务集合或 Starlette 回调,shutdown 会停止接收、取消并有限等待,且生命周期清单明确登记其顺序。主仓 `app/` 已无裸 FastAPI `BackgroundTasks`;当前仍有约 `50` 个更底层 `create_task`/等价任务创建点,与线程池和 APScheduler 并存,后续需逐项确认 owner、取消、等待、重试、幂等和是否 durable,关键业务副作用优先接入已有 Outbox/恢复表。 2. **动态模块契约仍以 legacy 聚合语义为主。** 当前登记 `212` 个模块方法,其中 `194` 个仍使用 `legacy` aggregation,只有 `14` 个 `first_non_empty`、`4` 个 `ordered_list_merge`。`app/runtime/extensions/module/contracts.py:422-455` 已能登记 family、输入/结果标签和基础签名诊断,但 `193` 个方法没有 required parameters,调度器 `app/runtime/extensions/module/dispatcher.py:109-260` 仍主要依赖运行时反射、返回值形状和短路规则。未知第三方方法保留 legacy fallback 是兼容要求,不应删除;宿主高频能力则应逐族补齐可执行的输入校验、结果校验、超时和错误语义。 -3. **查询侧数据库兼容 ABI 已完成正式装饰器清零。** 写事务装饰器和正式 `db_query/async_db_query` 均为 `0`。站点、消息、用户、订阅、下载/整理历史、工作流、MediaServer、SiteUserData、AgentChat、AgentTaskRun、TransferPending、SystemConfig、PassKey 和 SubscribeHistory 的宿主查询已迁到显式 Session 路径;对应旧插件 Model 调用由独立 `legacy_*` 外壳保留,可同时接受显式 Session 与无 Session 的位置/关键字参数。后续重点转为减少 ORM 对象跨层流转,并保持正式装饰器零回退。 +3. **Model/Base 的数据库装饰器和隐式会话 ABI 已全部清零。** 查询、写事务和 `legacy_*` 装饰器均为 `0`;所有 Model `db` 参数要求显式 Session,Base CRUD 仅在调用方事务内查询或 stage。可无会话构造的入口统一留在 Oper,经组合根事务执行器运行;插件 SDK 不再导出宿主 Model。后续重点转为减少 ORM 对象跨层流转,并保持 Model 隐式事务零回退。 4. **组合根和全局状态仍形成复杂的隐式运行时图。** Singleton 实例、模块级 provider、`configure_*` 注册函数和兼容 Facade 同时存在;它们解决了旧 ABI 和启动顺序问题,但增加测试污染、重复装配、实例身份和初始化顺序风险。`app/startup/lifecycle/__init__.py` 已有声明式生命周期,`app/startup/initializers/modules.py` 也有分阶段关闭,但尚未做到所有进程级资源都只通过 typed HostRuntime 访问。后续应以“新代码禁止新增 Service Locator/Singleton 依赖、旧入口有命中观测”为 ratchet。 ### P2:中长期可演进性债务 @@ -108,7 +108,7 @@ 1. 未知第三方插件自定义模块方法继续走 `legacy` fallback,不能因宿主契约收口而拒绝加载旧插件。 2. `PluginManager`、`PluginHelper`、`MoviePilotServerHelper` 等 Facade 继续保留旧公开/私有调用面,并通过 `compat.facade.hit` 统计迁移命中。 3. `app/runtime/compat` 的精确旧导入映射、`app.sdk._legacy` 薄门面和插件 V1/V2/V3 三代索引继续存在,直到命中数据和发行策略支持删除。 -4. 既有查询 Model 方法保留只读兼容入口;宿主 Oper 必须走显式 Session,`legacy_*` 只服务旧插件 ABI,不得成为新 Model 方法的默认模式。 +4. 插件访问宿主持久化必须经过 Oper 或稳定 SDK;不再保留直接调用宿主 Model 的事务兼容。 ### 建议的后续治理顺序 @@ -184,7 +184,7 @@ MoviePilot V3 当前不是“目录混乱、必须推倒重来”的状态。第 | 专用 EventData model | 53 | Event Contract Registry 已为全部事件登记 typed payload/fallback 原因 | | 直接读取 `settings` 的文件 | 105 | 仍按模块族迁移,动态协议和安全端口暂保留 | | `SystemConfigOper()` | 1 个 | 仅组合根创建 `SystemConfigService` 时保留 | -| Model 上的正式 DB 查询装饰器 | 0 | 查询/写装饰器均保持为 0;旧插件只读 ABI 由 `legacy_*` 外壳承接 | +| Model/Base 上的 DB 装饰器 | 0 | 正式与 legacy 查询/写装饰器全部为 0;`db` 参数必须显式传入 | | 路由端点 | 335 | 11 个已装饰端点超过 80 行,最大 400 行 | | Chain 方法超过 150 行 | 18 | 最大 `TransferChain.do_transfer()` 885 行 | | Application 方法超过 150 行 | 8 | 最大 296 行 | @@ -230,8 +230,8 @@ MoviePilot V3 当前不是“目录混乱、必须推倒重来”的状态。第 | 对标来源 | 可复用实践 | MoviePilot 当前差距 | 采用方式 | | --- | --- | --- | --- | | [FastAPI:Bigger Applications](https://fastapi.tiangolo.com/tutorial/bigger-applications/) | Router、依赖和主应用分离;路由按领域聚合 | Router 已分文件,但 `app/api/deps.py` 集中 33 个依赖工厂,部分端点仍编排完整用例 | 保留现有 Router;按垂直切片拆依赖和 presentation mapper,不重做目录树 | -| [FastAPI 官方 Full Stack Template](https://github.com/fastapi/full-stack-fastapi-template/tree/master/backend/app) | 请求依赖提供 Session,测试和迁移入口明确 | MoviePilot 已有请求 Session 和 UoW,但大量 Model 方法仍自行取得 Session/commit | 将 Session 生命周期留在请求/作业边界,Repository 只登记变更 | -| [SQLAlchemy Session Basics](https://docs.sqlalchemy.org/en/20/orm/session_basics.html) | Session/事务生命周期应与具体数据操作分离;Session per thread、AsyncSession per task | `@db_update`/`@async_db_update` 隐式创建和提交,跨多个 Repository 的原子性不清晰 | 新写用例强制请求/任务级 UoW;Model 逐步变为映射和约束载体 | +| [FastAPI 官方 Full Stack Template](https://github.com/fastapi/full-stack-fastapi-template/tree/master/backend/app) | 请求依赖提供 Session,测试和迁移入口明确 | MoviePilot 已有请求 Session 和 UoW,Model 隐式事务已清零;仍需继续减少 ORM 对象跨层流转 | 将 Session 生命周期留在请求/作业边界,Repository 只登记变更 | +| [SQLAlchemy Session Basics](https://docs.sqlalchemy.org/en/20/orm/session_basics.html) | Session/事务生命周期应与具体数据操作分离;Session per thread、AsyncSession per task | Model/Base 已要求显式 Session;无会话 Oper 仍依赖组合根事务执行器 | 新写用例强制请求/任务级 UoW;持续禁止 Model 重新拥有事务 | | [Starlette Lifespan](https://www.starlette.io/lifespan/) | Lifespan 完成前不接流量;用 typed state 共享进程资源;用 task group 管理异步任务 | 已有声明式生命周期,但仍依赖多个模块全局注册表和裸 `create_task`/线程 | 建立类型化 `HostRuntime/AppState`,旧 provider 继续作兼容门面 | | [Uvicorn Deployment](https://www.uvicorn.org/deployment/) 与 [Lifespan](https://www.uvicorn.org/concepts/lifespan/) | reload/workers 使用 import string/factory;每个 worker 独立执行 lifespan | 当前 app 实例与 reload/workers 配置并存,多 worker 会重复控制面 | V3 先明确只支持单 worker;开发 reload 改为 factory/import string;未来再拆 control role | | [Home Assistant:Integration Quality Scale](https://developers.home-assistant.io/docs/core/integration-quality-scale/) | 插件/集成按可测试性、错误处理、异步安全、类型和文档分级;豁免必须说明 | Module 能力差异大,只有统一发现和方法名快照,没有每个集成的质量状态 | 为宿主 Module 建立轻量质量清单和逐项 ratchet,不阻塞历史模块运行 | @@ -532,22 +532,16 @@ flowchart TB - `app/application/subscription/write.py` 定义用例 Port,`app/db/adapters/subscription.py` 为每次规范新增创建独占同步/异步 Session,`app/startup/composition/subscription.py` 只负责注入; `CreateSubscriptionCommand` / `AsyncCreateSubscriptionCommand` 持有 UoW,Oper 只执行 查重、`add` 与 `flush`。 -- `SubscribeOper.stage_add()` 的查重 SQL 已收口到 Oper,不再调用 Model 自动会话装饰器; - 无会话构造 `SubscribeOper()` 的旧 SDK 路径保留原自动短会话和返回值,未扩散为规范入口。 +- `SubscribeOper.stage_add()` 的查重 SQL已收口到 Oper;无会话构造 `SubscribeOper()` 时由 + Oper 的 `_execute_*` 委托组合根事务执行器,Model 不再创建会话。 - Chain 把原有“成功消息 → `SubscribeAdded` 事件 → Server 统计”作为显式 post-commit 回调交给 Command;commit/flush 失败回滚,事件或上报失败只传播原异常,不回滚已提交记录。 - 同步/异步 `SubscribeChain.add` 方法长度从各 203 行降至 183/186 行;新增 9 个事务边界测试, 覆盖成功顺序、commit/flush 失败、重复请求、Oper 不提交、事件失败、上报失败与真实落库。 -- Model 查询装饰器此前为 123 个:本切片绕开了继承自 `Base.create/async_create` 的自动提交, - 并继续保留既有 Model/旧 SDK 查询兼容;本次 AgentTask 切片将查询装饰器减少到 121 个。 - 2026-08-23 已完成 AgentTask 查询切片:`AgentTaskOper.get/list` 直接在调用方 Session 中执行查询, - `AgentTask.get_for_user/list_for_user` 保留原签名和返回语义供旧调用方使用,但不再持有查询装饰器; - 无 Session 的旧 Oper 入口继续由组合根兼容事务执行器承接。随后 PassKey 的宿主同步查询迁移到 - `PassKeyOper`,其按用户/凭证的启用状态过滤由显式 Session 测试覆盖;异步 Model 查询保留旧 ABI。 - 查询装饰器低水位由 123 降至 119,归属过滤、启用状态过滤和创建时间/主键稳定排序由 canonical - Oper 测试覆盖。`PassKey.get_by_user_id/get_by_credential_id` 与 - `AgentTask.get_for_user/list_for_user` 同时保留旧插件省略 Session 的同步调用方式;该路径显式委托 - 一次性兼容查询会话,不重新增加 Model 查询装饰器,也不影响宿主显式 Session 的事务所有权。 +- Model 查询装饰器曾有 123 个,分切片迁移后已连同 Base 的 12 个 legacy 查询/写装饰器全部删除。 + `AgentTask`、PassKey 等 Model 方法保留查询语义,但签名统一要求显式 Session;无 Session 使用方式 + 只在对应 Oper 上存在,由组合根事务执行器承接。归属过滤、启用状态过滤和稳定排序继续由显式 + Session 的 Model 测试与无会话 Oper 测试共同覆盖。 #### ARCH-222:按风险迁移其余写用例 @@ -971,9 +965,9 @@ ADR 必须逐个映射当前 Event、BackgroundTasks、Scheduler job、Agent tas 的旧 Oper ABI 委托 Startup 注入的短事务执行器。当前 Model 正式查询装饰器仅剩 30 个(同步 16、异步 14), `db_update` 与 `async_db_update` 均为 0,Oper 自建 Session/直接提交仍为 0。 - 数据清理按批次显式提交 UoW,单表失败先回滚会话再继续汇总后续表;不再依赖删除 Model 的隐式提交。 -- 收尾批次进一步移除宿主 Oper 对 `Base.create/update/delete/truncate` 八个兼容包装器的调用:显式 - Session 只 stage,由 Application UoW 提交;无 Session 的旧 Oper 入口才委托 Startup 的短事务执行器。 - Base 包装器继续保留给插件/旧模型 ABI,新增 AST 门禁禁止宿主 Oper 回退到隐式提交。 +- 收尾批次进一步移除宿主 Oper 对 `Base.create/update/delete/truncate` 隐式提交语义的依赖:显式 + Session 只 stage,由 Application UoW 提交;无 Session 的 Oper 入口委托 Startup 的短事务执行器。 + Base 方法最终改成纯显式 Session 原语,AST 门禁禁止 Model/Base 再引入装饰器或可选 Session。 **禁止**:本阶段不引入 Celery、Kafka、RabbitMQ 等新基础设施。 @@ -1237,49 +1231,21 @@ Settings 读取作为基础设施边界,架构基线已明确记录该例外 2026-08-23 收口兼容回归:`RuntimeSettingsCompat` 补齐 `update_setting`、`update_settings` 和 `model_dump` 旧 Settings ABI,并由应用组合根注入服务对象,低层 runtime 不再反向导入 `app.application`; `SkillHelper` 的技能市场写入继续经过兼容代理,旧插件/测试的模块级替换语义保持。`UserConfigOper` 的 -无 Session 查询改为一次性兼容查询会话,显式 Session 仍由调用方持有。配置债务稳定为 8 个文件,Model -查询装饰器在消息、用户和订阅查询切片后曾降至 75 个且写装饰器为 0;四分片全量测试 `5492 passed, 3 skipped`,mypy、复杂度、异步阻塞、 +无 Session 查询由组合根事务执行器创建一次性会话,显式 Session 仍由调用方持有。配置债务稳定为 +8 个文件;Model 查询装饰器在消息、用户和订阅查询切片后曾降至 75 个,随后已全部清零。该阶段 +四分片全量测试 `5492 passed, 3 skipped`,mypy、复杂度、异步阻塞、 host/plugin 架构基线均通过。 -2026-08-23 完成下载/整理历史查询切片:`TransferHistoryOper`、`DownloadHistoryOper` 的正式入口统一 -通过 `_execute_sync_query` / `_execute_async_query` 复用调用方 Session,正式查询装饰器由 75 降至 -38 个且写装饰器保持 0。旧插件仍可直接调用 Model 方法;`legacy_db_query` / `legacy_async_db_query` -按签名插入一次性会话,兼容无 Session 的位置参数和关键字参数,同时显式 Session 不创建额外会话。 -历史查询、删除工具、类型门禁和插件架构专项共 `101 passed`,host/plugin 架构基线通过。 +2026-08-23 分阶段完成下载/整理历史、Workflow、MediaServer、SiteUserData、AgentChat、AgentTaskRun、 +TransferPending、SystemConfig、PassKey 与 SubscribeHistory 查询切片:宿主 Oper 统一通过 +`_execute_sync_query` / `_execute_async_query` 复用调用方 Session,正式 Model 查询装饰器由 75 逐步降至 +0。各阶段显式 Session 查询、过滤语义、架构基线和全量测试均有回归记录。 -2026-08-23 完成 Workflow 查询切片:`WorkflowOper` 的同步/异步查询入口统一通过 -`_execute_sync_query` / `_execute_async_query` 复用调用方 Session,正式查询装饰器由 38 降至 -30 个且写装饰器保持 0。旧插件仍可直接调用 Workflow Model 方法,显式 Session 与无 Session 的 -关键字调用均有回归覆盖;Workflow、架构基线专项共 `76 passed`,host/plugin 架构基线通过。 - -2026-08-23 完成 MediaServer 与 SiteUserData 查询切片:`MediaServerOper`、`SiteOper` 的同步/异步 -查询入口统一通过 `_execute_sync_query` / `_execute_async_query` 复用调用方 Session,正式查询装饰器 -由 30 降至 18 个(同步 9、异步 9),写装饰器保持 0。旧插件仍可直接调用对应 Model 方法,显式 -Session 不创建额外会话,无 Session 的位置参数和关键字参数继续由 `legacy_*` 外壳兼容;专项测试 -`158 passed`,四分片全量测试 `5539 passed, 3 skipped`,host/plugin 架构基线和 Pylint 均通过。 - -2026-08-23 完成 AgentChat 与 AgentTaskRun 查询切片:`AgentChatOper`、`AgentTaskOper` 的查询入口 -统一通过 `_execute_sync_query` / `_execute_async_query` 复用调用方 Session,正式查询装饰器由 18 -降至 12 个(同步 5、异步 7),写装饰器保持 0。旧插件仍可直接调用对应 Model 方法,显式 Session -不创建额外会话,无 Session 的关键字调用继续由 `legacy_*` 外壳兼容;专项测试 `44 passed`,四分片 -全量测试 `5543 passed, 4 skipped`。 - -2026-08-23 完成 TransferPending 与 SystemConfig 查询切片:待整理回放和系统配置的宿主查询统一 -复用调用方 Session,正式查询装饰器由 12 降至 9 个(同步 3、异步 6),写装饰器保持 0。旧插件 -仍可直接调用对应 Model 方法,显式 Session 不创建额外会话,无 Session 的关键字调用继续由 -`legacy_*` 外壳兼容;专项与架构测试 `146 passed`,四分片全量测试 `5547 passed, 3 skipped`, -host/plugin 架构基线和 Pylint 均通过。 - -2026-08-23 完成 PassKey 查询切片:三个异步查询与按 ID 同步查询改由 `legacy_*` 外壳承接, -正式查询装饰器由 9 降至 5 个(同步 2、异步 3),写装饰器保持 0。显式 Session/AsyncSession -不创建额外会话,旧插件无 Session 的位置与关键字调用仍保持兼容;专项与架构测试 `99 passed`, -四分片全量测试 `5549 passed, 3 skipped`,host/plugin 架构基线和 Pylint 均通过。 - -2026-08-23 完成 SubscribeHistory 查询切片:同步/异步分页、owner 筛选和存在性查询均由 -`legacy_*` 外壳保留旧插件 ABI,`SubscribeHistoryOper` 与 `SubscribeOper.exist_history` 统一复用 -显式 Session/AsyncSession,且按 ID 查询不再调用 Base 查询包装器。正式查询装饰器由 5 降至 0, -同步/异步写装饰器继续保持 0;专项与架构测试 `176 passed`(另有 11 个子测试),四分片全量 -测试 `5551 passed, 3 skipped`,host/plugin 架构基线和 Pylint 均通过。 +2026-08-23 在正式装饰器清零后继续删除过渡性的 `legacy_db_query`、`legacy_async_db_query`、 +`legacy_db_update`、`legacy_async_db_update`:Base 与全部 Model 只接受显式 Session,不再替无会话调用 +创建或提交事务。原先把 `self._db=None` 直传 Model 的 User、PluginData、Subscribe、Site、配置和下载失败 +Oper 已迁到 `_execute_*`,插件 SDK 也移除了 User、Subscribe、TransferHistory Model 导出。架构测试新增 +三项硬约束:Model/Base 不得导入 DB 装饰器、`db` 参数不得可选、插件 SDK 不得导入 `app.db.models`。 #### ARCH-272:异步阻塞检测 @@ -1458,7 +1424,7 @@ rollback: | 基线写入行为 | 默认命令可能覆盖 fixture | 所有默认/check 命令保证工作树不变;write 必须显式 scope | | 全功能 worker | 配置允许 >1,控制面会复制 | 启动期明确拒绝 >1;文档与配置一致 | | 健康接口 | 认证 `/system/ping` 为主 | 分离公开 live 与受限/安全 ready;失败原因可诊断 | -| Model 事务装饰器 | 正式查询/写装饰器均为 0 | 持续保持为 0;兼容外壳不得被宿主新增调用 | +| Model/Base 事务装饰器 | 正式与 legacy 查询/写装饰器均为 0 | 持续保持为 0;`db` 参数保持显式必传 | | 新写用例事务 | 宿主写 Oper 已脱离 Base 隐式提交 | 100% 由入口/Application 边界拥有 Session/UoW | | 高频 Module 契约 | 212 个宿主能力显式登记 | 新观察到的宿主方法必须同步登记完整契约 | | Event payload | 53 类型全部登记 typed payload 与可靠性 | 新事件必须同步登记,不回退裸 dict | diff --git a/docs/rules/10-data-and-persistent.md b/docs/rules/10-data-and-persistent.md index 2ec63db9c..095333ff4 100644 --- a/docs/rules/10-data-and-persistent.md +++ b/docs/rules/10-data-and-persistent.md @@ -91,21 +91,21 @@ adapters; it does not retain reusable repository implementations. - `tests/fixtures/architecture/transaction-debt-baseline.json` records formal decorators in concrete files under `app/db/models/`. Their count is zero and - must remain zero. Compatibility-only `legacy_*` shells are tracked separately - and must never be treated as the target design. -- `legacy_db_query` / `legacy_async_db_query` preserve an existing plugin-facing - Model method whose no-Session call shape cannot be removed yet. If a Model - method has no external ABI obligation, move the query into its Oper and remove - the Model method instead of adding `legacy_*`. -- `Base.create/get/update/delete/list/truncate` and their async forms are inherited - plugin ABI, so `app/db/base.py` deliberately uses legacy query/write wrappers. - New host code must not call these convenience methods; Oper staging methods and - explicit UoW are the canonical path. Removal requires plugin-usage evidence and - a separately announced compatibility break, not a mechanical rename. -- Host Oper code must pass an explicit Session through `_execute_sync_query` / - `_execute_async_query`; new Model methods must not add any legacy decorator. -- New Model methods must not use `db_query`, `db_update`, `async_db_query`, or - `async_db_update`, create a Session, or call `commit()` / `rollback()`. + must remain zero. Model/Base code may not import `app.db.decorators`; legacy + Model transaction shells have been removed and must not be recreated. +- Every Model method with a `db` parameter requires an explicit `Session` or + `AsyncSession`. The parameter may not default to `None`, accept displaced + business arguments, create a Session, or call `commit()` / `rollback()`. +- `Base.create/get/update/delete/list/truncate` and their async forms are plain + explicit-session primitives. They only query or stage changes in the caller's + transaction; they never own transaction lifecycle. +- Host Oper code routes optional-session entry points through + `_execute_sync_query` / `_execute_async_query` / `_execute_*_write`. Plugins + access host persistence through Oper or a curated SDK contract, never by + importing `app.db.models`. +- The public `db_query`, `db_update`, `async_db_query`, and `async_db_update` + exports remain available only for plugin-owned database functions. They are + forbidden on host Model/Base methods. - Oper receives a caller-owned Session and may query, add, update, delete, or flush. A composable Oper method must not create its own Session and must not commit or roll back. diff --git a/tests/conftest.py b/tests/conftest.py index 034f76d56..b24409123 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -5,8 +5,11 @@ """ import asyncio import sys +from collections.abc import Awaitable, Callable +from typing import TypeVar import pytest +from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Session # 必须早于首个牵入 app.runtime.config 的 import(app.db / app.chain.* 都会牵入):引擎本身已惰性, @@ -21,6 +24,9 @@ prepare_backend() from app.testing.network_guard import block_real_network # noqa: E402,F401 +TResult = TypeVar("TResult") + + class _TestDatabaseExecutor: """让绕过完整 lifespan 的测试仍通过线程执行同步数据库写入。""" @@ -89,9 +95,10 @@ def configure_plugin_system_services(): configure_token_runtime_config(lambda: build_token_runtime_config(settings)) database_executor = _TestDatabaseExecutor() system_config = SystemConfigOper() - system_config.load_snapshot() user_config = UserConfigOper() - user_config.load_snapshot() + with SessionFactory() as session: + system_config.load_snapshot(session) + user_config.load_snapshot(session) configure_system_config( SystemConfigService( repository=system_config, @@ -159,14 +166,12 @@ def configure_plugin_system_services(): from app.db.adapters.workflow import TransactionalWorkflowExecutionService from app.db.adapters.transaction import TransactionalWriteRunner - def compatibility_sync_session() -> Session: - """动态读取可被存量隔离数据库用例替换的 ScopedSession。""" - from app.db import decorators - - return decorators.ScopedSession() + def create_sync_session() -> Session: + """为无显式会话的 Oper 测试入口创建独占同步 Session。""" + return SessionFactory() transaction_runner = TransactionalWriteRunner( - sync_session=compatibility_sync_session, + sync_session=create_sync_session, async_session=async_session_scope, ) configure_transaction_runners( @@ -347,6 +352,20 @@ class DbHarness: self.session.commit() return rows[0] if len(rows) == 1 else list(rows) + def run_async_session( + self, + operation: Callable[[AsyncSession], Awaitable[TResult]], + ) -> TResult: + """在临时数据库的显式 AsyncSession 中执行被测操作。""" + from app.db.session import async_session_scope + + async def execute() -> TResult: + """打开异步会话并把事务所有权留在测试载具。""" + async with async_session_scope() as session: + return await operation(session) + + return asyncio.run(execute()) + def cleanup(self) -> None: """按水位删除本用例新增的全部行。""" from sqlalchemy import delete diff --git a/tests/fixtures/architecture/dependency-baseline.json b/tests/fixtures/architecture/dependency-baseline.json index c341aadd5..7d8a0abb8 100644 --- a/tests/fixtures/architecture/dependency-baseline.json +++ b/tests/fixtures/architecture/dependency-baseline.json @@ -13,8 +13,8 @@ "runtime_to_db": [], "workflow_to_db": [] }, - "edge_count": 6525, - "edge_sha256": "72b9416b5309bb51768b95b3c9da4b167242c3bc3f27d4e561a3f097267ca6a0", + "edge_count": 6500, + "edge_sha256": "f16cc04898ae0fc7602f2d154321035040c6e64155c59b0b0d72a7d707f59ec9", "edges": [ "app -> app.runtime", "app -> app.runtime.compat", @@ -3549,7 +3549,6 @@ "app.db.adapters.workflow -> app.db.oper.workflow", "app.db.adapters.workflow -> app.db.uow", "app.db.base -> app.db", - "app.db.base -> app.db.decorators", "app.db.base -> app.db.uow", "app.db.base -> app.runtime", "app.db.base -> app.runtime.config", @@ -3584,13 +3583,10 @@ "app.db.models._identity -> app.schemas.media", "app.db.models.agentchat -> app.db", "app.db.models.agentchat -> app.db.base", - "app.db.models.agentchat -> app.db.decorators", "app.db.models.agenttask -> app.db", "app.db.models.agenttask -> app.db.base", - "app.db.models.agenttask -> app.db.decorators", "app.db.models.agenttaskrun -> app.db", "app.db.models.agenttaskrun -> app.db.base", - "app.db.models.agenttaskrun -> app.db.decorators", "app.db.models.agenttaskrun -> app.db.models", "app.db.models.agenttaskrun -> app.db.models.agenttask", "app.db.models.downloadfailure -> app.db", @@ -3599,76 +3595,60 @@ "app.db.models.downloadfailure -> app.db.models._constraints", "app.db.models.downloadhistory -> app.db", "app.db.models.downloadhistory -> app.db.base", - "app.db.models.downloadhistory -> app.db.decorators", "app.db.models.downloadhistory -> app.db.models", "app.db.models.downloadhistory -> app.db.models._constraints", "app.db.models.downloadhistory -> app.schemas", "app.db.models.downloadhistory -> app.schemas.types", "app.db.models.mediaserver -> app.db", "app.db.models.mediaserver -> app.db.base", - "app.db.models.mediaserver -> app.db.decorators", "app.db.models.mediaserver -> app.db.models", "app.db.models.mediaserver -> app.db.models._constraints", "app.db.models.mediaserver -> app.schemas", "app.db.models.mediaserver -> app.schemas.types", "app.db.models.message -> app.db", "app.db.models.message -> app.db.base", - "app.db.models.message -> app.db.decorators", "app.db.models.outbox -> app.db", "app.db.models.outbox -> app.db.base", "app.db.models.passkey -> app.db", "app.db.models.passkey -> app.db.base", - "app.db.models.passkey -> app.db.decorators", "app.db.models.plugindata -> app.db", "app.db.models.plugindata -> app.db.base", - "app.db.models.plugindata -> app.db.decorators", "app.db.models.site -> app.db", "app.db.models.site -> app.db.base", - "app.db.models.site -> app.db.decorators", "app.db.models.siteicon -> app.db", "app.db.models.siteicon -> app.db.base", - "app.db.models.siteicon -> app.db.decorators", "app.db.models.sitestatistic -> app.db", "app.db.models.sitestatistic -> app.db.base", - "app.db.models.sitestatistic -> app.db.decorators", "app.db.models.siteuserdata -> app.db", "app.db.models.siteuserdata -> app.db.base", - "app.db.models.siteuserdata -> app.db.decorators", "app.db.models.subscribe -> app.db", "app.db.models.subscribe -> app.db.base", - "app.db.models.subscribe -> app.db.decorators", "app.db.models.subscribe -> app.db.models", "app.db.models.subscribe -> app.db.models._constraints", "app.db.models.subscribe -> app.schemas", "app.db.models.subscribe -> app.schemas.types", "app.db.models.subscribehistory -> app.db", "app.db.models.subscribehistory -> app.db.base", - "app.db.models.subscribehistory -> app.db.decorators", "app.db.models.subscribehistory -> app.db.models", "app.db.models.subscribehistory -> app.db.models._constraints", "app.db.models.subscribehistory -> app.schemas", "app.db.models.subscribehistory -> app.schemas.types", "app.db.models.systemconfig -> app.db", "app.db.models.systemconfig -> app.db.base", - "app.db.models.systemconfig -> app.db.decorators", "app.db.models.transferhistory -> app.db", "app.db.models.transferhistory -> app.db.base", - "app.db.models.transferhistory -> app.db.decorators", "app.db.models.transferhistory -> app.db.models", "app.db.models.transferhistory -> app.db.models._constraints", "app.db.models.transferhistory -> app.schemas", "app.db.models.transferhistory -> app.schemas.types", "app.db.models.transferpending -> app.db", "app.db.models.transferpending -> app.db.base", - "app.db.models.transferpending -> app.db.decorators", "app.db.models.user -> app.db", "app.db.models.user -> app.db.base", - "app.db.models.user -> app.db.decorators", "app.db.models.userconfig -> app.db", "app.db.models.userconfig -> app.db.base", "app.db.models.workflow -> app.db", "app.db.models.workflow -> app.db.base", - "app.db.models.workflow -> app.db.decorators", "app.db.oper.agentchat -> app.db", "app.db.oper.agentchat -> app.db.base", "app.db.oper.agentchat -> app.db.models", @@ -5901,8 +5881,6 @@ "app.sdk._legacy.history -> app.application", "app.sdk._legacy.history -> app.application.history", "app.sdk._legacy.history -> app.db", - "app.sdk._legacy.history -> app.db.models", - "app.sdk._legacy.history -> app.db.models.transferhistory", "app.sdk._legacy.history -> app.db.oper", "app.sdk._legacy.history -> app.db.oper.transferhistory", "app.sdk._legacy.history -> app.domain", @@ -5916,8 +5894,6 @@ "app.sdk._legacy.subscribe -> app.application.subscription", "app.sdk._legacy.subscribe -> app.application.subscription.write", "app.sdk._legacy.subscribe -> app.db", - "app.sdk._legacy.subscribe -> app.db.models", - "app.sdk._legacy.subscribe -> app.db.models.subscribe", "app.sdk._legacy.subscribe -> app.db.oper", "app.sdk._legacy.subscribe -> app.db.oper.subscribe", "app.sdk._legacy.subscribe -> app.domain", @@ -5927,8 +5903,6 @@ "app.sdk._legacy.user -> app.api", "app.sdk._legacy.user -> app.api.deps", "app.sdk._legacy.user -> app.db", - "app.sdk._legacy.user -> app.db.models", - "app.sdk._legacy.user -> app.db.models.user", "app.sdk._legacy.user -> app.db.oper", "app.sdk._legacy.user -> app.db.oper.user", "app.sdk.browser -> app.adapters", @@ -6362,6 +6336,7 @@ "app.testing.bootstrap -> app.db.oper", "app.testing.bootstrap -> app.db.oper.systemconfig", "app.testing.bootstrap -> app.db.oper.userconfig", + "app.testing.bootstrap -> app.db.session", "app.testing.bootstrap -> app.startup", "app.testing.bootstrap -> app.startup.initializers", "app.testing.bootstrap -> app.startup.initializers.cache", diff --git a/tests/fixtures/architecture/startup-performance-baseline.json b/tests/fixtures/architecture/startup-performance-baseline.json index 643acc39b..fab94f39d 100644 --- a/tests/fixtures/architecture/startup-performance-baseline.json +++ b/tests/fixtures/architecture/startup-performance-baseline.json @@ -1,41 +1,41 @@ { "schema_version": 2, - "generated_at": "2026-08-23T13:47:08.442528+00:00", + "generated_at": "2026-08-23T15:32:22.484415+00:00", "platform": "macOS-26.5.2-arm64-arm-64bit-Mach-O", "python": "3.14.3", "repeat": 3, "targets": { "app.startup.lifecycle": { - "loaded_app_module_count": 367, - "max_ms": 1591.175, - "median_ms": 1049.37, - "min_ms": 1022.351, + "loaded_app_module_count": 366, + "max_ms": 985.623, + "median_ms": 947.803, + "min_ms": 932.937, "samples_ms": [ - 1591.175, - 1049.37, - 1022.351 + 947.803, + 985.623, + 932.937 ] }, "app.factory": { - "loaded_app_module_count": 379, - "max_ms": 1198.166, - "median_ms": 1083.805, - "min_ms": 1009.309, + "loaded_app_module_count": 378, + "max_ms": 957.03, + "median_ms": 950.108, + "min_ms": 947.713, "samples_ms": [ - 1009.309, - 1198.166, - 1083.805 + 950.108, + 957.03, + 947.713 ] }, "app.main": { - "loaded_app_module_count": 381, - "max_ms": 1229.477, - "median_ms": 1179.859, - "min_ms": 1129.055, + "loaded_app_module_count": 380, + "max_ms": 1098.929, + "median_ms": 1095.482, + "min_ms": 1076.479, "samples_ms": [ - 1129.055, - 1179.859, - 1229.477 + 1076.479, + 1095.482, + 1098.929 ] } }, @@ -47,87 +47,87 @@ { "mode": "normal", "enabled_component_count": 23, - "startup_ms": 0.64, - "full_lifespan_ms": 0.805, + "startup_ms": 0.622, + "full_lifespan_ms": 0.781, "stage_ms": { - "后台任务登记器": 0.079, - "数据库准备": 0.039, - "HTTP 基础能力": 0.029, - "领域依赖装配": 0.027, - "数据库引擎预热": 0.024, - "数据库连接预算": 0.023, - "路由": 0.023, - "模块服务": 0.021, - "插件备份恢复": 0.024, - "插件": 0.021, - "定时器": 0.026, - "监控器": 0.022, - "待处理整理回放": 0.025, - "命令服务": 0.025, - "工作流": 0.021, - "插件同步与启动收尾": 0.037 - }, - "threads_before": 2, - "threads_started": 2, - "threads_after": 2, - "tasks_before": 1, - "tasks_started": 2, - "tasks_after": 1, - "database_connections_started": 0 - }, - { - "mode": "normal", - "enabled_component_count": 23, - "startup_ms": 0.652, - "full_lifespan_ms": 0.808, - "stage_ms": { - "后台任务登记器": 0.079, - "数据库准备": 0.039, - "HTTP 基础能力": 0.033, - "领域依赖装配": 0.029, - "数据库引擎预热": 0.026, - "数据库连接预算": 0.021, - "路由": 0.022, - "模块服务": 0.023, - "插件备份恢复": 0.021, - "插件": 0.022, - "定时器": 0.026, - "监控器": 0.023, - "待处理整理回放": 0.024, - "命令服务": 0.023, - "工作流": 0.025, - "插件同步与启动收尾": 0.035 - }, - "threads_before": 2, - "threads_started": 2, - "threads_after": 2, - "tasks_before": 1, - "tasks_started": 2, - "tasks_after": 1, - "database_connections_started": 0 - }, - { - "mode": "normal", - "enabled_component_count": 23, - "startup_ms": 0.739, - "full_lifespan_ms": 0.944, - "stage_ms": { - "后台任务登记器": 0.087, + "后台任务登记器": 0.076, "数据库准备": 0.037, "HTTP 基础能力": 0.029, - "领域依赖装配": 0.028, - "数据库引擎预热": 0.024, - "数据库连接预算": 0.021, - "路由": 0.02, + "领域依赖装配": 0.031, + "数据库引擎预热": 0.026, + "数据库连接预算": 0.025, + "路由": 0.022, "模块服务": 0.02, - "插件备份恢复": 0.024, - "插件": 0.019, + "插件备份恢复": 0.023, + "插件": 0.023, + "定时器": 0.024, + "监控器": 0.025, + "待处理整理回放": 0.02, + "命令服务": 0.022, + "工作流": 0.02, + "插件同步与启动收尾": 0.033 + }, + "threads_before": 2, + "threads_started": 2, + "threads_after": 2, + "tasks_before": 1, + "tasks_started": 2, + "tasks_after": 1, + "database_connections_started": 0 + }, + { + "mode": "normal", + "enabled_component_count": 23, + "startup_ms": 0.629, + "full_lifespan_ms": 0.782, + "stage_ms": { + "后台任务登记器": 0.075, + "数据库准备": 0.039, + "HTTP 基础能力": 0.029, + "领域依赖装配": 0.031, + "数据库引擎预热": 0.024, + "数据库连接预算": 0.022, + "路由": 0.023, + "模块服务": 0.021, + "插件备份恢复": 0.023, + "插件": 0.023, "定时器": 0.025, - "监控器": 0.054, - "待处理整理回放": 0.041, - "命令服务": 0.034, - "工作流": 0.025, - "插件同步与启动收尾": 0.037 + "监控器": 0.02, + "待处理整理回放": 0.024, + "命令服务": 0.021, + "工作流": 0.022, + "插件同步与启动收尾": 0.033 + }, + "threads_before": 2, + "threads_started": 2, + "threads_after": 2, + "tasks_before": 1, + "tasks_started": 2, + "tasks_after": 1, + "database_connections_started": 0 + }, + { + "mode": "normal", + "enabled_component_count": 23, + "startup_ms": 0.631, + "full_lifespan_ms": 0.785, + "stage_ms": { + "后台任务登记器": 0.08, + "数据库准备": 0.04, + "HTTP 基础能力": 0.031, + "领域依赖装配": 0.028, + "数据库引擎预热": 0.025, + "数据库连接预算": 0.023, + "路由": 0.022, + "模块服务": 0.022, + "插件备份恢复": 0.021, + "插件": 0.025, + "定时器": 0.022, + "监控器": 0.024, + "待处理整理回放": 0.02, + "命令服务": 0.024, + "工作流": 0.02, + "插件同步与启动收尾": 0.034 }, "threads_before": 2, "threads_started": 2, @@ -138,8 +138,8 @@ "database_connections_started": 0 } ], - "median_startup_ms": 0.652, - "median_full_lifespan_ms": 0.808, + "median_startup_ms": 0.629, + "median_full_lifespan_ms": 0.782, "enabled_component_count": 23, "enabled_components": [ "后台任务登记器", @@ -172,66 +172,66 @@ { "mode": "safe", "enabled_component_count": 11, - "startup_ms": 0.475, - "full_lifespan_ms": 0.619, - "stage_ms": { - "后台任务登记器": 0.079, - "数据库准备": 0.038, - "HTTP 基础能力": 0.03, - "领域依赖装配": 0.028, - "数据库引擎预热": 0.024, - "数据库连接预算": 0.023, - "路由": 0.023, - "模块服务": 0.022, - "插件同步与启动收尾": 0.036 - }, - "threads_before": 2, - "threads_started": 2, - "threads_after": 2, - "tasks_before": 1, - "tasks_started": 2, - "tasks_after": 1, - "database_connections_started": 0 - }, - { - "mode": "safe", - "enabled_component_count": 11, - "startup_ms": 0.493, - "full_lifespan_ms": 0.676, - "stage_ms": { - "后台任务登记器": 0.089, - "数据库准备": 0.037, - "HTTP 基础能力": 0.035, - "领域依赖装配": 0.028, - "数据库引擎预热": 0.025, - "数据库连接预算": 0.021, - "路由": 0.025, - "模块服务": 0.023, - "插件同步与启动收尾": 0.062 - }, - "threads_before": 2, - "threads_started": 2, - "threads_after": 2, - "tasks_before": 1, - "tasks_started": 2, - "tasks_after": 1, - "database_connections_started": 0 - }, - { - "mode": "safe", - "enabled_component_count": 11, - "startup_ms": 0.464, - "full_lifespan_ms": 0.638, + "startup_ms": 0.597, + "full_lifespan_ms": 0.772, "stage_ms": { "后台任务登记器": 0.075, - "数据库准备": 0.036, - "HTTP 基础能力": 0.028, - "领域依赖装配": 0.028, - "数据库引擎预热": 0.023, - "数据库连接预算": 0.022, + "数据库准备": 0.044, + "HTTP 基础能力": 0.033, + "领域依赖装配": 0.059, + "数据库引擎预热": 0.061, + "数据库连接预算": 0.032, + "路由": 0.027, + "模块服务": 0.046, + "插件同步与启动收尾": 0.044 + }, + "threads_before": 2, + "threads_started": 2, + "threads_after": 2, + "tasks_before": 1, + "tasks_started": 2, + "tasks_after": 1, + "database_connections_started": 0 + }, + { + "mode": "safe", + "enabled_component_count": 11, + "startup_ms": 0.461, + "full_lifespan_ms": 0.642, + "stage_ms": { + "后台任务登记器": 0.076, + "数据库准备": 0.037, + "HTTP 基础能力": 0.029, + "领域依赖装配": 0.029, + "数据库引擎预热": 0.026, + "数据库连接预算": 0.023, "路由": 0.024, - "模块服务": 0.021, - "插件同步与启动收尾": 0.057 + "模块服务": 0.023, + "插件同步与启动收尾": 0.06 + }, + "threads_before": 2, + "threads_started": 2, + "threads_after": 2, + "tasks_before": 1, + "tasks_started": 2, + "tasks_after": 1, + "database_connections_started": 0 + }, + { + "mode": "safe", + "enabled_component_count": 11, + "startup_ms": 0.483, + "full_lifespan_ms": 0.659, + "stage_ms": { + "后台任务登记器": 0.077, + "数据库准备": 0.039, + "HTTP 基础能力": 0.034, + "领域依赖装配": 0.035, + "数据库引擎预热": 0.025, + "数据库连接预算": 0.026, + "路由": 0.025, + "模块服务": 0.024, + "插件同步与启动收尾": 0.059 }, "threads_before": 2, "threads_started": 2, @@ -242,8 +242,8 @@ "database_connections_started": 0 } ], - "median_startup_ms": 0.475, - "median_full_lifespan_ms": 0.638, + "median_startup_ms": 0.483, + "median_full_lifespan_ms": 0.659, "enabled_component_count": 11, "enabled_components": [ "后台任务登记器", diff --git a/tests/test_agent_task_runs.py b/tests/test_agent_task_runs.py index 5d80d91bc..3c0aa3e66 100644 --- a/tests/test_agent_task_runs.py +++ b/tests/test_agent_task_runs.py @@ -10,11 +10,11 @@ from sqlalchemy.exc import IntegrityError from app.agent.orchestrator import AgentManager from app.agent.tools.impl.query_agent_tasks import QueryAgentTasksTool from app.db.engine import get_engine +from app.db import base as db_base from app.db.oper.agenttask import AgentTaskOper from app.db.models.agenttask import AgentTask from app.db.models.agenttaskrun import AgentTaskRun from app.db.session import SessionFactory -from app.db import decorators Engine = get_engine() @@ -141,9 +141,11 @@ def test_agenttaskrun_oper_reuses_explicit_query_session(db, monkeypatch): run = AgentTaskOper().begin_run(task.id) assert run monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (_ for _ in ()).throw(AssertionError("不应创建额外同步会话")), + db_base, + "run_sync_transaction", + lambda _operation: (_ for _ in ()).throw( + AssertionError("不应创建额外同步事务") + ), ) oper = AgentTaskOper(db.session) @@ -151,19 +153,6 @@ def test_agenttaskrun_oper_reuses_explicit_query_session(db, monkeypatch): assert oper.list_runs(task.id) -def test_agenttaskrun_model_legacy_query_keeps_keyword_abi(monkeypatch): - """旧插件以关键字直调 AgentTaskRun 时仍自动补入短会话。""" - opened = [] - monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (opened.append(True) or SessionFactory()), - ) - - assert AgentTaskRun.get_by_run_id(run_id="missing-legacy") is None - assert opened == [True] - - def test_begin_run_rolls_back_task_claim_when_run_insert_fails() -> None: """运行记录插入失败时,任务的 running 投影必须随事务回滚。""" first_task = _add_task("run-rollback-first") diff --git a/tests/test_architecture_dependencies.py b/tests/test_architecture_dependencies.py index 02e4d2e7d..1ca3c4273 100644 --- a/tests/test_architecture_dependencies.py +++ b/tests/test_architecture_dependencies.py @@ -432,46 +432,76 @@ def test_database_internals_do_not_import_db_facades(): assert violations == [] -def test_base_crud_is_explicitly_legacy_only(): - """Base 便利 CRUD 只能保留兼容壳,不得伪装成新的正式事务入口。""" - path = APP_ROOT / "db" / "base.py" - tree = ast.parse(path.read_text(encoding="utf-8-sig"), filename=str(path)) - base_class = next( - node - for node in tree.body - if isinstance(node, ast.ClassDef) and node.name == "Base" - ) - formal_decorators = { +def test_models_and_base_require_explicit_database_sessions(): + """Model/Base 不得装饰事务,且所有 db 参数必须由调用方显式传入。""" + decorator_names = { "db_query", "db_update", "async_db_query", "async_db_update", + "legacy_db_query", + "legacy_db_update", + "legacy_async_db_query", + "legacy_async_db_update", } violations: list[str] = [] - for node in base_class.body: - if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): - continue - decorators = { - decorator.id - for decorator in node.decorator_list - if isinstance(decorator, ast.Name) - } - if decorators & formal_decorators: - violations.append(node.name) + paths = [APP_ROOT / "db" / "base.py"] + paths.extend((APP_ROOT / "db" / "models").rglob("*.py")) + for path in paths: + tree = ast.parse(path.read_text(encoding="utf-8-sig"), filename=str(path)) + relative = str(path.relative_to(PROJECT_ROOT)) + nodes = list(ast.walk(tree)) + for node in nodes: + if isinstance(node, ast.ImportFrom) and node.module == "app.db.decorators": + violations.append(f"{relative}:{node.lineno}:decorator-import") + if path.name == "base.py": + base_class = next( + node + for node in tree.body + if isinstance(node, ast.ClassDef) and node.name == "Base" + ) + nodes = list(ast.walk(base_class)) + for node in nodes: + if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): + continue + for decorator in node.decorator_list: + name = ( + decorator.id + if isinstance(decorator, ast.Name) + else decorator.attr + if isinstance(decorator, ast.Attribute) + else None + ) + if name in decorator_names: + violations.append(f"{relative}:{node.lineno}:@{name}") + arguments = [*node.args.posonlyargs, *node.args.args] + defaults = [None] * (len(arguments) - len(node.args.defaults)) + list( + node.args.defaults + ) + for argument, default in zip(arguments, defaults): + if argument.arg != "db": + continue + annotation = ast.unparse(argument.annotation) if argument.annotation else "" + if default is not None or "None" in annotation: + violations.append( + f"{relative}:{node.lineno}:{node.name}:optional-db" + ) assert violations == [] -def test_models_use_one_legacy_query_compatibility_shell(): - """旧 Model 查询统一使用 legacy 装饰器,不得再手写隐式会话 runner。""" - retired_names = {"run_legacy_sync_query", "run_legacy_async_query"} +def test_plugin_sdk_does_not_import_or_export_host_models(): + """插件 SDK 只能暴露 Oper,不得把宿主 ORM Model 作为插件接口。""" violations: list[str] = [] - for path in (APP_ROOT / "db" / "models").glob("*.py"): + for path in (APP_ROOT / "sdk").rglob("*.py"): tree = ast.parse(path.read_text(encoding="utf-8-sig"), filename=str(path)) - if any( - isinstance(node, ast.Name) and node.id in retired_names - for node in ast.walk(tree) - ): - violations.append(str(path.relative_to(PROJECT_ROOT))) + for node in ast.walk(tree): + if isinstance(node, ast.ImportFrom) and node.module and ( + node.module == "app.db.models" + or node.module.startswith("app.db.models.") + ): + violations.append( + f"{path.relative_to(PROJECT_ROOT)}:{node.lineno}:{node.module}" + ) assert violations == [] diff --git a/tests/test_db_base_crud.py b/tests/test_db_base_crud.py index 0024197db..4164e4a47 100644 --- a/tests/test_db_base_crud.py +++ b/tests/test_db_base_crud.py @@ -11,6 +11,7 @@ import pytest from app.db.models.systemconfig import SystemConfig from app.db.models.userconfig import UserConfig +from app.db.uow import run_async_transaction @pytest.fixture(autouse=True) @@ -25,10 +26,13 @@ def test_create_persists_and_get_reads_back(db): """ row = SystemConfig(key="base-create", value={"n": 1}) row.create(db.session) + db.session.commit() assert row.id is not None assert SystemConfig.get(db.session, row.id).key == "base-create" - assert asyncio.run(SystemConfig.async_get(rid=row.id)).key == "base-create" + assert db.run_async_session( + lambda session: SystemConfig.async_get(session, rid=row.id) + ).key == "base-create" def test_get_returns_none_for_missing_id(db): @@ -36,7 +40,9 @@ def test_get_returns_none_for_missing_id(db): 主键不存在时返回 None,而不是抛异常或返回任意一行。 """ assert SystemConfig.get(db.session, -1) is None - assert asyncio.run(SystemConfig.async_get(rid=-1)) is None + assert db.run_async_session( + lambda session: SystemConfig.async_get(session, rid=-1) + ) is None def test_async_create_flushes_and_assigns_primary_key(db): @@ -45,7 +51,11 @@ def test_async_create_flushes_and_assigns_primary_key(db): 异步路径的调用方常常紧接着用 id 建立关联,拿到 None 会让关联静默丢失。 """ - created = asyncio.run(SystemConfig(key="base-async-create", value={"n": 2}).async_create()) + created = asyncio.run(run_async_transaction( + lambda session: SystemConfig( + key="base-async-create", value={"n": 2} + ).async_create(session) + )) assert created.id is not None assert SystemConfig.get(db.session, created.id).value == {"n": 2} @@ -57,11 +67,20 @@ def test_update_writes_payload_fields(db): """ row = SystemConfig(key="base-update", value={"n": 1}) row.create(db.session) + db.session.flush() row.update(db.session, {"value": {"n": 9}}) assert SystemConfig.get(db.session, row.id).value == {"n": 9} + db.session.commit() - asyncio.run(row.async_update(payload={"value": {"n": 10}})) + async def update_in_owned_transaction(session) -> None: + """在同一异步事务中读取并更新目标行。""" + async_row = await SystemConfig.async_get(session, row.id) + assert async_row is not None + await async_row.async_update(session, payload={"value": {"n": 10}}) + + asyncio.run(run_async_transaction(update_in_owned_transaction)) + db.session.expire_all() assert SystemConfig.get(db.session, row.id).value == {"n": 10} @@ -85,7 +104,9 @@ def test_async_delete_removes_only_the_given_row(db): dropped = db.add(SystemConfig(key="base-async-del", value={"n": 1})) kept = db.add(SystemConfig(key="base-async-keep", value={"n": 2})) - asyncio.run(SystemConfig.async_delete(rid=dropped.id)) + asyncio.run(run_async_transaction( + lambda session: SystemConfig.async_delete(session, rid=dropped.id) + )) assert SystemConfig.get(db.session, dropped.id) is None assert SystemConfig.get(db.session, kept.id) is not None @@ -95,7 +116,9 @@ def test_async_delete_tolerates_missing_row(db): """ 删除不存在的行不抛异常,保持调用方的幂等语义。 """ - asyncio.run(SystemConfig.async_delete(rid=-1)) + asyncio.run(run_async_transaction( + lambda session: SystemConfig.async_delete(session, rid=-1) + )) def test_list_returns_every_row_of_that_model_only(db): @@ -117,7 +140,7 @@ def test_async_list_matches_sync_list(db): db.add(UserConfig(username="base-list", key="k", value="v")) sync_ids = sorted(item.id for item in UserConfig.list(db.session)) - async_ids = sorted(item.id for item in asyncio.run(UserConfig.async_list())) + async_ids = sorted(item.id for item in db.run_async_session(UserConfig.async_list)) assert sync_ids == async_ids @@ -132,7 +155,7 @@ def test_truncate_empties_the_table(db): assert UserConfig.list(db.session) == [] db.add(UserConfig(username="base-truncate-async", key="k", value="v")) - asyncio.run(UserConfig.async_truncate()) + asyncio.run(run_async_transaction(UserConfig.async_truncate)) assert UserConfig.list(db.session) == [] diff --git a/tests/test_db_config_user_queries.py b/tests/test_db_config_user_queries.py index 26547f0ed..d0acff5a9 100644 --- a/tests/test_db_config_user_queries.py +++ b/tests/test_db_config_user_queries.py @@ -9,14 +9,14 @@ import asyncio import pytest -from app.db import decorators +from app.db import base as db_base from app.db.models.passkey import PassKey from app.db.models.systemconfig import SystemConfig from app.db.models.user import User from app.db.models.userconfig import UserConfig from app.db.oper.passkey import PassKeyOper from app.db.oper.user import UserOper -from app.db.session import SessionFactory, async_session_scope +from app.db.session import async_session_scope @pytest.fixture(autouse=True) @@ -39,7 +39,9 @@ def test_systemconfig_get_by_key_matches_async_twin(db): found = SystemConfig.get_by_key(db.session, "mp-test-a") assert found.value == {"n": 1} - async_found = asyncio.run(SystemConfig.async_get_by_key(key="mp-test-a")) + async_found = db.run_async_session( + lambda session: SystemConfig.async_get_by_key(session, "mp-test-a") + ) assert async_found.value == found.value @@ -54,9 +56,11 @@ def test_systemconfig_queries_reuse_explicit_sessions(db, monkeypatch): """SystemConfig 显式同步与异步会话不得触发兼容会话。""" db.add(SystemConfig(key="mp-explicit-config", value=True)) monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (_ for _ in ()).throw(AssertionError("不应创建额外同步会话")), + db_base, + "run_sync_transaction", + lambda _operation: (_ for _ in ()).throw( + AssertionError("不应创建额外同步事务") + ), ) assert SystemConfig.get_by_key(db.session, "mp-explicit-config") is not None @@ -64,9 +68,11 @@ def test_systemconfig_queries_reuse_explicit_sessions(db, monkeypatch): """验证异步配置查询复用显式 AsyncSession。""" async with async_session_scope() as session: monkeypatch.setattr( - decorators, - "async_session_scope", - lambda: (_ for _ in ()).throw(AssertionError("不应创建额外异步会话")), + db_base, + "run_async_transaction", + lambda _operation: (_ for _ in ()).throw( + AssertionError("不应创建额外异步事务") + ), ) assert await SystemConfig.async_get_by_key( session, @@ -76,20 +82,6 @@ def test_systemconfig_queries_reuse_explicit_sessions(db, monkeypatch): asyncio.run(check()) -def test_systemconfig_model_legacy_query_keeps_keyword_abi(db, monkeypatch): - """旧插件以关键字直调 SystemConfig 时仍自动补入短会话。""" - db.add(SystemConfig(key="mp-legacy-config", value=True)) - opened = [] - monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (opened.append(True) or SessionFactory()), - ) - - assert SystemConfig.get_by_key(key="mp-legacy-config") is not None - assert opened == [True] - - def test_systemconfig_delete_by_key_removes_only_that_key(db): """ 按键删除只能删掉那一个键,误删会静默丢失其他配置。 @@ -163,8 +155,12 @@ def test_user_lookup_by_name_and_id_matches_async_twin(db): by_id = User.get_by_id(db.session, created.id) assert by_name.id == by_id.id == created.id - assert asyncio.run(User.async_get_by_name(name="mp-test-user")).id == created.id - assert asyncio.run(User.async_get_by_id(user_id=created.id)).id == created.id + assert db.run_async_session( + lambda session: User.async_get_by_name(session, "mp-test-user") + ).id == created.id + assert db.run_async_session( + lambda session: User.async_get_by_id(session, created.id) + ).id == created.id def test_user_lookup_returns_none_when_absent(db): @@ -175,14 +171,6 @@ def test_user_lookup_returns_none_when_absent(db): assert User.get_by_id(db.session, -1) is None -def test_user_sync_queries_preserve_legacy_no_session_abi(db): - """旧插件省略 Session 时仍可按用户名和用户 ID 查询。""" - created = db.add(User(name="mp-legacy-query-user", hashed_password="secret")) - - assert User.get_by_name("mp-legacy-query-user").id == created.id - assert User.get_by_id(created.id).name == "mp-legacy-query-user" - - def test_user_delete_by_name_and_by_id_remove_only_the_target(db): """ 按名、按 ID 删除都只能删掉目标用户。 @@ -258,7 +246,9 @@ def test_passkey_listing_excludes_inactive_credentials(db): listed = PassKey.get_by_user_id(db.session, 9001) assert {p.credential_id for p in listed} == {"cred-active-1", "cred-active-2"} - assert {p.credential_id for p in asyncio.run(PassKey.async_get_by_user_id(user_id=9001))} == \ + assert {p.credential_id for p in db.run_async_session( + lambda session: PassKey.async_get_by_user_id(session, 9001) + )} == \ {"cred-active-1", "cred-active-2"} @@ -289,26 +279,20 @@ def test_passkey_lookup_by_credential_id_skips_inactive(db): assert PassKey.get_by_credential_id(db.session, "cred-live").user_id == 9003 assert PassKey.get_by_credential_id(db.session, "cred-dead") is None - assert asyncio.run(PassKey.async_get_by_credential_id(credential_id="cred-dead")) is None - - -def test_passkey_model_sync_queries_keep_no_session_plugin_abi(db): - """旧插件不传 Session 时仍由统一 legacy 装饰器获得短会话查询。""" - db.add(_passkey(9004, "cred-legacy")) - - assert [item.credential_id for item in PassKey.get_by_user_id(user_id=9004)] == [ - "cred-legacy" - ] - assert PassKey.get_by_credential_id("cred-legacy").user_id == 9004 + assert db.run_async_session( + lambda session: PassKey.async_get_by_credential_id(session, "cred-dead") + ) is None def test_passkey_remaining_queries_reuse_explicit_sessions(db, monkeypatch): """PassKey 其余同步/异步查询必须复用调用方会话。""" key = db.add(_passkey(9008, "cred-explicit")) monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (_ for _ in ()).throw(AssertionError("不应创建额外同步会话")), + db_base, + "run_sync_transaction", + lambda _operation: (_ for _ in ()).throw( + AssertionError("不应创建额外同步事务") + ), ) assert PassKey.get_by_id(db.session, key.id).credential_id == "cred-explicit" @@ -316,9 +300,11 @@ def test_passkey_remaining_queries_reuse_explicit_sessions(db, monkeypatch): """验证三个异步查询都复用显式 AsyncSession。""" async with async_session_scope() as session: monkeypatch.setattr( - decorators, - "async_session_scope", - lambda: (_ for _ in ()).throw(AssertionError("不应创建额外异步会话")), + db_base, + "run_async_transaction", + lambda _operation: (_ for _ in ()).throw( + AssertionError("不应创建额外异步事务") + ), ) assert [item.credential_id for item in await PassKey.async_get_by_user_id( session, @@ -333,35 +319,6 @@ def test_passkey_remaining_queries_reuse_explicit_sessions(db, monkeypatch): asyncio.run(check()) -def test_passkey_remaining_queries_keep_legacy_keyword_abi(db, monkeypatch): - """旧插件关键字直调 PassKey 其余查询时仍自动补入短会话。""" - key = db.add(_passkey(9009, "cred-keyword")) - opened_sync = [] - monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (opened_sync.append(True) or SessionFactory()), - ) - assert PassKey.get_by_id(passkey_id=key.id) is not None - assert opened_sync == [True] - - opened_async = [] - original_scope = async_session_scope - - def tracked_scope(): - """记录旧异步 ABI 创建的兼容会话作用域。""" - opened_async.append(True) - return original_scope() - - monkeypatch.setattr(decorators, "async_session_scope", tracked_scope) - assert asyncio.run(PassKey.async_get_by_user_id(user_id=9009)) - assert asyncio.run(PassKey.async_get_by_credential_id( - credential_id="cred-keyword", - )) is not None - assert asyncio.run(PassKey.async_get_by_id(passkey_id=key.id)) is not None - assert opened_async == [True, True, True] - - def test_passkey_get_by_id_ignores_active_flag(db): """ 按主键取记录是管理用途,不应过滤停用状态——否则管理端看不到自己刚停用的凭据。 @@ -369,7 +326,9 @@ def test_passkey_get_by_id_ignores_active_flag(db): dead = db.add(_passkey(9004, "cred-admin", is_active=False)) assert PassKey.get_by_id(db.session, dead.id).credential_id == "cred-admin" - assert asyncio.run(PassKey.async_get_by_id(passkey_id=dead.id)).credential_id == "cred-admin" + assert db.run_async_session( + lambda session: PassKey.async_get_by_id(session, dead.id) + ).credential_id == "cred-admin" def test_passkey_delete_requires_matching_owner(db): diff --git a/tests/test_db_downloadhistory_queries.py b/tests/test_db_downloadhistory_queries.py index 5c665c53c..3fd5e4466 100644 --- a/tests/test_db_downloadhistory_queries.py +++ b/tests/test_db_downloadhistory_queries.py @@ -108,8 +108,11 @@ def test_list_by_page_is_newest_first_and_paged(db): assert [h.title for h in page1] == ["p-3", "p-2"] assert [h.title for h in DownloadHistory.list_by_page(db.session, page=2, count=2)] == \ ["p-1", "p-0"] - assert [h.title for h in asyncio.run( - DownloadHistory.async_list_by_page(page=1, count=2))] == ["p-3", "p-2"] + assert [h.title for h in db.run_async_session( + lambda session: DownloadHistory.async_list_by_page( + session, page=1, count=2 + ) + )] == ["p-3", "p-2"] def test_get_by_path_finds_the_download_directory(db): @@ -312,10 +315,17 @@ def test_count_and_title_search_match_async_twins(db): """ db.add(_history("Unique Title Here", date="2026-08-13 10:00:00")) - assert asyncio.run(DownloadHistory.async_count()) >= 1 - assert asyncio.run(DownloadHistory.async_count_by_title(title="unique title")) == 1 - assert [h.title for h in asyncio.run(DownloadHistory.async_list_by_title( - title="UNIQUE TITLE"))] == ["Unique Title Here"] + assert db.run_async_session(DownloadHistory.async_count) >= 1 + assert db.run_async_session( + lambda session: DownloadHistory.async_count_by_title( + session, title="unique title" + ) + ) == 1 + assert [h.title for h in db.run_async_session( + lambda session: DownloadHistory.async_list_by_title( + session, title="UNIQUE TITLE" + ) + )] == ["Unique Title Here"] # --------------------------------------------------------------------------- # diff --git a/tests/test_db_media_identity_normalizer.py b/tests/test_db_media_identity_normalizer.py index 83bdef1e3..59fc47833 100644 --- a/tests/test_db_media_identity_normalizer.py +++ b/tests/test_db_media_identity_normalizer.py @@ -145,6 +145,7 @@ def test_normalization_also_applies_on_update(db): row = _write(db, _history(media_source=MediaSource.TMDB, media_id="550")) row.update(db.session, {"media_source": "douban", "media_id": " 1291546 "}) + db.session.commit() db.session.expire_all() updated = TransferHistory.get(db.session, row.id) diff --git a/tests/test_db_mediaserver_queries.py b/tests/test_db_mediaserver_queries.py index 22731f20e..de7da61bb 100644 --- a/tests/test_db_mediaserver_queries.py +++ b/tests/test_db_mediaserver_queries.py @@ -9,10 +9,10 @@ import asyncio import pytest -from app.db import decorators +from app.db import base as db_base from app.db.models.mediaserver import MediaServerItem from app.db.oper.mediaserver import MediaServerOper -from app.db.session import SessionFactory, async_session_scope +from app.db.session import async_session_scope from app.schemas.types import MediaSource @@ -39,7 +39,9 @@ def test_get_by_itemid_matches_async_twin(db): db.add(_item("emby", "it-1"), _item("plex", "it-2")) assert MediaServerItem.get_by_itemid(db.session, "it-1").server == "emby" - assert asyncio.run(MediaServerItem.async_get_by_itemid(item_id="it-1")).server == "emby" + assert db.run_async_session( + lambda session: MediaServerItem.async_get_by_itemid(session, "it-1") + ).server == "emby" assert MediaServerItem.get_by_itemid(db.session, "it-missing") is None @@ -47,9 +49,11 @@ def test_mediaserver_oper_reuses_explicit_query_sessions(db, monkeypatch): """媒体服务器 Oper 绑定调用方会话后不得再创建兼容查询会话。""" db.add(_item("emby", "explicit-ms", media_id="explicit-1001")) monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (_ for _ in ()).throw(AssertionError("不应创建额外同步会话")), + db_base, + "run_sync_transaction", + lambda _operation: (_ for _ in ()).throw( + AssertionError("不应创建额外同步事务") + ), ) assert MediaServerOper(db.session).exists( @@ -62,9 +66,11 @@ def test_mediaserver_oper_reuses_explicit_query_sessions(db, monkeypatch): """验证异步存在性查询复用显式 AsyncSession。""" async with async_session_scope() as session: monkeypatch.setattr( - decorators, - "async_session_scope", - lambda: (_ for _ in ()).throw(AssertionError("不应创建额外异步会话")), + db_base, + "run_async_transaction", + lambda _operation: (_ for _ in ()).throw( + AssertionError("不应创建额外异步事务") + ), ) assert await MediaServerOper(session).async_exists( media_source=MediaSource.TMDB, @@ -75,20 +81,6 @@ def test_mediaserver_oper_reuses_explicit_query_sessions(db, monkeypatch): asyncio.run(check()) -def test_mediaserver_model_legacy_query_keeps_keyword_abi(db, monkeypatch): - """旧插件以关键字直调媒体服务器 Model 时仍自动补入短会话。""" - db.add(_item("emby", "legacy-ms")) - opened = [] - monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (opened.append(True) or SessionFactory()), - ) - - assert MediaServerItem.get_by_itemid(item_id="legacy-ms") is not None - assert opened == [True] - - def test_get_by_server_itemid_scopes_by_server(db): """ 条目 ID 只在单个服务器内唯一,查找必须同时限定服务器。 @@ -118,8 +110,14 @@ def test_exist_by_media_identity_requires_source_id_and_type(db): assert MediaServerItem.exist_by_media_identity( db.session, MediaSource.TMDB, "556", "电影") is None - assert asyncio.run(MediaServerItem.async_exist_by_media_identity( - media_source=MediaSource.TMDB, media_id="555", mtype="电影")) is not None + assert db.run_async_session( + lambda session: MediaServerItem.async_exist_by_media_identity( + session, + media_source=MediaSource.TMDB, + media_id="555", + mtype="电影", + ) + ) is not None @pytest.mark.parametrize("mtype,year,expected", [ @@ -152,8 +150,11 @@ def test_exists_by_title_matches_async_twin(db): for mtype, year in ((None, None), ("电影", None), (None, "2026"), ("电影", "2026")): sync_found = MediaServerItem.exists_by_title(db.session, "并行标题", mtype, year) - async_found = asyncio.run(MediaServerItem.async_exists_by_title( - title="并行标题", mtype=mtype, year=year)) + async_found = db.run_async_session( + lambda session: MediaServerItem.async_exists_by_title( + session, title="并行标题", mtype=mtype, year=year + ) + ) assert (sync_found is None) == (async_found is None) diff --git a/tests/test_db_plugin_message_agent_queries.py b/tests/test_db_plugin_message_agent_queries.py index 3074615bf..2e817e2e6 100644 --- a/tests/test_db_plugin_message_agent_queries.py +++ b/tests/test_db_plugin_message_agent_queries.py @@ -8,14 +8,13 @@ import asyncio import pytest -from app.db import decorators +from app.db import base as db_base from app.db.models.agentchat import AgentChat from app.db.models.agenttask import AgentTask from app.db.models.downloadfailure import DownloadFailure from app.db.models.message import Message from app.db.models.plugindata import PluginData from app.db.oper.agenttask import AgentTaskOper -from app.db.session import SessionFactory @pytest.fixture(autouse=True) @@ -39,7 +38,9 @@ def test_plugindata_is_scoped_by_plugin_id(db): rows = PluginData.get_plugin_data(db.session, "PluginA") assert {r.key for r in rows} == {"k1", "k2"} - assert {r.key for r in asyncio.run(PluginData.async_get_plugin_data(plugin_id="PluginA"))} \ + assert {r.key for r in db.run_async_session( + lambda session: PluginData.async_get_plugin_data(session, "PluginA") + )} \ == {"k1", "k2"} @@ -53,8 +54,11 @@ def test_plugindata_get_by_key_needs_both_plugin_and_key(db): assert PluginData.get_plugin_data_by_key(db.session, "PluginA", "shared").value == {"v": 1} assert PluginData.get_plugin_data_by_key(db.session, "PluginB", "shared").value == {"v": 2} assert PluginData.get_plugin_data_by_key(db.session, "PluginC", "shared") is None - assert asyncio.run(PluginData.async_get_plugin_data_by_key( - plugin_id="PluginA", key="shared")).value == {"v": 1} + assert db.run_async_session( + lambda session: PluginData.async_get_plugin_data_by_key( + session, plugin_id="PluginA", key="shared" + ) + ).value == {"v": 1} def test_plugindata_delete_by_key_removes_only_that_entry(db): @@ -120,7 +124,9 @@ def test_message_list_by_page_matches_async_twin(db): db.add(_message(f"2026-08-13 11:00:0{index}", f"par-{index}")) sync_titles = [m.title for m in Message.list_by_page(db.session, page=1, count=3)] - async_titles = [m.title for m in asyncio.run(Message.async_list_by_page(page=1, count=3))] + async_titles = [m.title for m in db.run_async_session( + lambda session: Message.async_list_by_page(session, page=1, count=3) + )] assert sync_titles == async_titles @@ -181,7 +187,11 @@ def test_message_async_list_sent_excludes_the_clear_boundary(db): def _titles(**clears) -> set: """取本用例写入的消息标题集合,隔离其他用例可能残留的消息。""" - rows = asyncio.run(Message.async_list_sent_by_page(page=1, count=100, **clears)) + rows = db.run_async_session( + lambda session: Message.async_list_sent_by_page( + session, page=1, count=100, **clears + ) + ) return {m.title for m in rows if m.title.startswith("bd-")} # 全量清空水位:边界上的两条都属于被清空的那一批 @@ -226,7 +236,9 @@ def test_agentchat_get_by_session_takes_the_newest_row(db): newest = db.add(_chat("s-dup")) assert AgentChat.get_by_session(db.session, "s-dup").id == newest.id - assert asyncio.run(AgentChat.async_get_by_session(session_id="s-dup")).id == newest.id + assert db.run_async_session( + lambda session: AgentChat.async_get_by_session(session, "s-dup") + ).id == newest.id def test_agentchat_get_by_session_enforces_user_scope(db): @@ -237,16 +249,22 @@ def test_agentchat_get_by_session_enforces_user_scope(db): assert AgentChat.get_by_session(db.session, "s-owned", user_id="alice") is not None assert AgentChat.get_by_session(db.session, "s-owned", user_id="bob") is None - assert asyncio.run(AgentChat.async_get_by_session(session_id="s-owned", user_id="bob")) is None + assert db.run_async_session( + lambda session: AgentChat.async_get_by_session( + session, session_id="s-owned", user_id="bob" + ) + ) is None def test_agentchat_oper_reuses_explicit_query_sessions(db, monkeypatch): """AgentChatOper 的同步与异步查询必须复用调用方会话。""" db.add(_chat("s-explicit", user_id="explicit")) monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (_ for _ in ()).throw(AssertionError("不应创建额外同步会话")), + db_base, + "run_sync_transaction", + lambda _operation: (_ for _ in ()).throw( + AssertionError("不应创建额外同步事务") + ), ) from app.db.oper.agentchat import AgentChatOper @@ -259,29 +277,17 @@ def test_agentchat_oper_reuses_explicit_query_sessions(db, monkeypatch): async with async_session_scope() as session: monkeypatch.setattr( - decorators, - "async_session_scope", - lambda: (_ for _ in ()).throw(AssertionError("不应创建额外异步会话")), + db_base, + "run_async_transaction", + lambda _operation: (_ for _ in ()).throw( + AssertionError("不应创建额外异步事务") + ), ) assert await AgentChatOper(session).async_get("s-explicit", "explicit") asyncio.run(check()) -def test_agentchat_model_legacy_query_keeps_keyword_abi(db, monkeypatch): - """旧插件以关键字直调 AgentChat 时仍自动补入短会话。""" - db.add(_chat("s-legacy")) - opened = [] - monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (opened.append(True) or SessionFactory()), - ) - - assert AgentChat.get_by_session(session_id="s-legacy") is not None - assert opened == [True] - - def test_agentchat_list_by_page_matches_either_user_or_username(db): """ 同时给出用户 ID 与用户名时按「或」匹配。 @@ -324,8 +330,11 @@ def test_agentchat_list_by_page_is_newest_first_and_paged(db): assert [c.session_id for c in page1] == ["s-p3", "s-p2"] assert [c.session_id for c in page2] == ["s-p1", "s-p0"] - assert [c.session_id for c in asyncio.run( - AgentChat.async_list_by_page(page=1, count=2, user_id="uid-page"))] == ["s-p3", "s-p2"] + assert [c.session_id for c in db.run_async_session( + lambda session: AgentChat.async_list_by_page( + session, page=1, count=2, user_id="uid-page" + ) + )] == ["s-p3", "s-p2"] # --------------------------------------------------------------------------- # @@ -352,18 +361,6 @@ def test_agenttask_get_for_user_enforces_ownership(db): assert AgentTask.get_for_user(db.session, task_id, user_id="bob") is None -def test_agenttask_model_queries_keep_no_session_plugin_abi(db): - """旧插件省略 Session 时仍由统一 legacy 装饰器按原参数查询。""" - task_id = AgentTask.add_task(db.session, **_task("legacy", user_id="legacy-user")) - db.session.commit() - - assert AgentTask.get_for_user( - task_id=task_id, - user_id="legacy-user", - ).id == task_id - assert [task.id for task in AgentTask.list_for_user(user_id="legacy-user")] == [task_id] - - def test_agenttask_oper_reads_with_explicit_session(db, monkeypatch): """AgentTaskOper 的宿主查询使用调用方 Session,不再经过旧事务兼容执行器。""" task_id = AgentTask.add_task(db.session, **_task("canonical", user_id="alice")) diff --git a/tests/test_db_public_api.py b/tests/test_db_public_api.py index 406a8b827..409aeba94 100644 --- a/tests/test_db_public_api.py +++ b/tests/test_db_public_api.py @@ -22,8 +22,8 @@ def test_session_factories_are_not_part_of_the_public_contract(): """ 三个会话工厂不得出现在 ``__all__`` 里。 - 插件要访问数据库应走 ``DbOper`` 子类或 ``db_query`` / ``async_db_query`` 装饰器, - 由装饰器收口会话的提交、回滚与释放。 + 插件访问宿主数据应走 ``DbOper``;插件自有表可使用 ``db_query`` / ``async_db_query`` + 装饰器,由装饰器收口插件自有会话的提交、回滚与释放。 """ leaked = [name for name in INTERNAL_FACTORY_NAMES if name in db_package.__all__] assert not leaked, f"会话工厂被重新放进了对外契约:{leaked}" diff --git a/tests/test_db_site_queries.py b/tests/test_db_site_queries.py index c3ab7fbda..80ed55e11 100644 --- a/tests/test_db_site_queries.py +++ b/tests/test_db_site_queries.py @@ -9,13 +9,13 @@ import asyncio import pytest -from app.db import decorators +from app.db import base as db_base from app.db.models.site import Site from app.db.models.siteicon import SiteIcon from app.db.models.sitestatistic import SiteStatistic from app.db.models.siteuserdata import SiteUserData from app.db.oper.site import SiteOper -from app.db.session import SessionFactory, async_session_scope +from app.db.session import async_session_scope @pytest.fixture(autouse=True) @@ -40,8 +40,12 @@ def test_site_get_by_domain_matches_async_twin(db): db.add(_site("站点A", "a.test"), _site("站点B", "b.test")) assert Site.get_by_domain(db.session, "a.test").name == "站点A" - assert asyncio.run(Site.async_get_by_domain(domain="a.test")).name == "站点A" - assert asyncio.run(Site.async_get_by_name(name="站点B")).domain == "b.test" + assert db.run_async_session( + lambda session: Site.async_get_by_domain(session, "a.test") + ).name == "站点A" + assert db.run_async_session( + lambda session: Site.async_get_by_name(session, "站点B") + ).domain == "b.test" def test_site_get_by_domain_returns_none_when_absent(db): @@ -61,7 +65,9 @@ def test_site_get_actives_excludes_disabled_sites(db): _site("停用", "off.test", is_active=False)) assert {s.domain for s in Site.get_actives(db.session)} == {"on1.test", "on2.test"} - assert {s.domain for s in asyncio.run(Site.async_get_actives())} == {"on1.test", "on2.test"} + assert {s.domain for s in db.run_async_session(Site.async_get_actives)} == { + "on1.test", "on2.test" + } def test_site_list_order_by_pri_is_ascending(db): @@ -73,7 +79,7 @@ def test_site_list_order_by_pri_is_ascending(db): assert [s.domain for s in Site.list_order_by_pri(db.session)] == \ ["p1.test", "p2.test", "p3.test"] - assert [s.domain for s in asyncio.run(Site.async_list_order_by_pri())] == \ + assert [s.domain for s in db.run_async_session(Site.async_list_order_by_pri)] == \ ["p1.test", "p2.test", "p3.test"] @@ -123,7 +129,9 @@ def test_siteicon_get_by_domain_matches_async_twin(db): SiteIcon(name="站点B", domain="icon-b.test", url="https://icon-b.test/f.ico")) assert SiteIcon.get_by_domain(db.session, "icon-a.test").name == "站点A" - assert asyncio.run(SiteIcon.async_get_by_domain(domain="icon-a.test")).name == "站点A" + assert db.run_async_session( + lambda session: SiteIcon.async_get_by_domain(session, "icon-a.test") + ).name == "站点A" assert SiteIcon.get_by_domain(db.session, "icon-missing.test") is None @@ -135,7 +143,9 @@ def test_sitestatistic_get_by_domain_matches_async_twin(db): SiteStatistic(domain="stat-b.test", success=1, fail=0, seconds=1, lst_state=0)) assert SiteStatistic.get_by_domain(db.session, "stat-a.test").success == 3 - assert asyncio.run(SiteStatistic.async_get_by_domain(domain="stat-a.test")).success == 3 + assert db.run_async_session( + lambda session: SiteStatistic.async_get_by_domain(session, "stat-a.test") + ).success == 3 assert SiteStatistic.get_by_domain(db.session, "stat-missing.test") is None @@ -186,7 +196,11 @@ def test_userdata_get_by_domain_matches_async_twin(db): for kwargs in ({}, {"workdate": "2026-08-12"}, {"workdate": "2026-08-12", "worktime": "20:00:00"}): sync_rows = SiteUserData.get_by_domain(db.session, "ud2.test", **kwargs) - async_rows = asyncio.run(SiteUserData.async_get_by_domain(domain="ud2.test", **kwargs)) + async_rows = db.run_async_session( + lambda session: SiteUserData.async_get_by_domain( + session, domain="ud2.test", **kwargs + ) + ) assert len(sync_rows) == len(async_rows) @@ -194,9 +208,11 @@ def test_site_oper_reuses_explicit_userdata_query_sessions(db, monkeypatch): """站点用户数据 Oper 必须复用调用方同步与异步会话。""" db.add(_userdata("explicit-site.test", "2026-08-12", "10:00:00")) monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (_ for _ in ()).throw(AssertionError("不应创建额外同步会话")), + db_base, + "run_sync_transaction", + lambda _operation: (_ for _ in ()).throw( + AssertionError("不应创建额外同步事务") + ), ) assert SiteOper(db.session).get_userdata_by_domain("explicit-site.test") @@ -205,9 +221,11 @@ def test_site_oper_reuses_explicit_userdata_query_sessions(db, monkeypatch): """验证异步站点用户数据查询复用显式 AsyncSession。""" async with async_session_scope() as session: monkeypatch.setattr( - decorators, - "async_session_scope", - lambda: (_ for _ in ()).throw(AssertionError("不应创建额外异步会话")), + db_base, + "run_async_transaction", + lambda _operation: (_ for _ in ()).throw( + AssertionError("不应创建额外异步事务") + ), ) assert await SiteOper(session).async_get_userdata_by_domain( "explicit-site.test" @@ -216,20 +234,6 @@ def test_site_oper_reuses_explicit_userdata_query_sessions(db, monkeypatch): asyncio.run(check()) -def test_site_userdata_model_legacy_query_keeps_keyword_abi(db, monkeypatch): - """旧插件以关键字直调 SiteUserData 时仍自动补入短会话。""" - db.add(_userdata("legacy-site.test", "2026-08-12", "10:00:00")) - opened = [] - monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (opened.append(True) or SessionFactory()), - ) - - assert SiteUserData.get_by_domain(domain="legacy-site.test") - assert opened == [True] - - def test_userdata_get_by_date_returns_all_domains_of_that_day(db): """ 按日期查询应跨站点返回当天全部快照。 @@ -283,7 +287,10 @@ def test_userdata_get_latest_matches_async_twin(db): _userdata("par.test", "2026-08-12", "10:00:00")) sync_rows = [(r.domain, r.updated_day) for r in SiteUserData.get_latest(db.session)] - async_rows = [(r.domain, r.updated_day) for r in asyncio.run(SiteUserData.async_get_latest())] + async_rows = [ + (r.domain, r.updated_day) + for r in db.run_async_session(SiteUserData.async_get_latest) + ] assert sorted(sync_rows) == sorted(async_rows) diff --git a/tests/test_db_subscribe_queries.py b/tests/test_db_subscribe_queries.py index aef591471..a6d7fa291 100644 --- a/tests/test_db_subscribe_queries.py +++ b/tests/test_db_subscribe_queries.py @@ -10,11 +10,11 @@ import time as _time import pytest -from app.db import decorators +from app.db import base as db_base from app.db.models import subscribe as subscribe_module from app.db.models.subscribe import Subscribe from app.db.models.subscribehistory import SubscribeHistory -from app.db.session import SessionFactory, async_session_scope +from app.db.session import async_session_scope from app.schemas.types import MediaSource, MediaType TMDB = str(MediaSource.TMDB) @@ -67,8 +67,11 @@ def test_exists_matches_async_twin(db): db.add(_sub("并行", season=1)) sync_found = Subscribe.exists(db.session, MediaSource.TMDB, "9001", season=1) - async_found = asyncio.run(Subscribe.async_exists( - media_source=MediaSource.TMDB, media_id="9001", season=1)) + async_found = db.run_async_session( + lambda session: Subscribe.async_exists( + session, media_source=MediaSource.TMDB, media_id="9001", season=1 + ) + ) assert sync_found.id == async_found.id @@ -77,9 +80,11 @@ def test_history_queries_reuse_explicit_sessions(db, monkeypatch): """订阅历史同步/异步查询必须复用调用方会话。""" row = db.add(_history("显式历史", media_id="8501")) monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (_ for _ in ()).throw(AssertionError("不应创建额外同步会话")), + db_base, + "run_sync_transaction", + lambda _operation: (_ for _ in ()).throw( + AssertionError("不应创建额外同步事务") + ), ) assert SubscribeHistory.list_by_type( db.session, MediaType.TV.value, page=1, count=10 @@ -92,9 +97,11 @@ def test_history_queries_reuse_explicit_sessions(db, monkeypatch): """验证异步订阅历史查询复用显式 AsyncSession。""" async with async_session_scope() as session: monkeypatch.setattr( - decorators, - "async_session_scope", - lambda: (_ for _ in ()).throw(AssertionError("不应创建额外异步会话")), + db_base, + "run_async_transaction", + lambda _operation: (_ for _ in ()).throw( + AssertionError("不应创建额外异步事务") + ), ) assert await SubscribeHistory.async_list_by_type( session, MediaType.TV.value, page=1, count=10 @@ -109,44 +116,6 @@ def test_history_queries_reuse_explicit_sessions(db, monkeypatch): asyncio.run(check()) -def test_history_queries_keep_legacy_keyword_abi(db, monkeypatch): - """旧插件关键字直调订阅历史查询时仍自动补入兼容会话。""" - row = db.add(_history("关键字历史", media_id="8601")) - opened_sync = [] - monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (opened_sync.append(True) or SessionFactory()), - ) - assert SubscribeHistory.list_by_type( - mtype=MediaType.TV.value, page=1, count=10 - ) - assert SubscribeHistory.exists( - media_source=MediaSource.TMDB, media_id="8601", season=1 - ).id == row.id - assert opened_sync == [True, True] - - opened_async = [] - original_scope = async_session_scope - - def tracked_scope(): - """记录旧异步 ABI 创建的兼容会话作用域。""" - opened_async.append(True) - return original_scope() - - monkeypatch.setattr(decorators, "async_session_scope", tracked_scope) - assert asyncio.run(SubscribeHistory.async_list_by_type( - mtype=MediaType.TV.value, page=1, count=10 - )) - assert asyncio.run(SubscribeHistory.async_list_by_type_and_username( - mtype=MediaType.TV.value, username="alice", page=1, count=10 - )) - assert asyncio.run(SubscribeHistory.async_exists( - media_source=MediaSource.TMDB, media_id="8601", season=1 - )) is not None - assert opened_async == [True, True, True] - - @pytest.mark.parametrize("media_id", [None, "", " "]) def test_exists_rejects_blank_media_id(db, media_id): """ @@ -232,7 +201,9 @@ def test_get_by_state_splits_comma_separated_states(db): assert states == {"N", "R"} assert len(Subscribe.get_by_state(db.session, "")) >= 3 - assert {s.state for s in asyncio.run(Subscribe.async_get_by_state(state="N,R"))} == {"N", "R"} + assert {s.state for s in db.run_async_session( + lambda session: Subscribe.async_get_by_state(session, state="N,R") + )} == {"N", "R"} def test_get_by_title_optionally_narrows_by_season(db): @@ -281,8 +252,11 @@ def test_list_by_username_matches_async_twin(db): ("N", MediaType.TV.value)): sync_names = sorted(s.name for s in Subscribe.list_by_username(db.session, "alice", state, mtype)) - async_names = sorted(s.name for s in asyncio.run( - Subscribe.async_list_by_username(username="alice", state=state, mtype=mtype))) + async_names = sorted(s.name for s in db.run_async_session( + lambda session: Subscribe.async_list_by_username( + session, username="alice", state=state, mtype=mtype + ) + )) assert sync_names == async_names @@ -319,8 +293,11 @@ def test_list_by_type_includes_the_window_start_boundary(db, frozen_now): date=one_second_earlier)) names = {s.name for s in Subscribe.list_by_type(db.session, MediaType.TV.value, days=7)} - async_names = {s.name for s in asyncio.run( - Subscribe.async_list_by_type(mtype=MediaType.TV.value, days=7))} + async_names = {s.name for s in db.run_async_session( + lambda session: Subscribe.async_list_by_type( + session, mtype=MediaType.TV.value, days=7 + ) + )} assert "窗口起点上" in names and "窗口起点前一秒" not in names assert "窗口起点上" in async_names and "窗口起点前一秒" not in async_names @@ -363,8 +340,11 @@ def test_history_list_by_type_matches_async_twin(db): sync_names = [h.name for h in SubscribeHistory.list_by_type( db.session, MediaType.TV.value, page=1, count=10)] - async_names = [h.name for h in asyncio.run(SubscribeHistory.async_list_by_type( - mtype=MediaType.TV.value, page=1, count=10))] + async_names = [h.name for h in db.run_async_session( + lambda session: SubscribeHistory.async_list_by_type( + session, mtype=MediaType.TV.value, page=1, count=10 + ) + )] assert sync_names == async_names @@ -392,7 +372,13 @@ def test_history_exists_matches_async_twin(db): db.add(_history("并行历史", season=1, media_id="8401")) sync_found = SubscribeHistory.exists(db.session, MediaSource.TMDB, "8401", season=1) - async_found = asyncio.run(SubscribeHistory.async_exists( - media_source=MediaSource.TMDB, media_id="8401", season=1)) + async_found = db.run_async_session( + lambda session: SubscribeHistory.async_exists( + session, + media_source=MediaSource.TMDB, + media_id="8401", + season=1, + ) + ) assert sync_found.id == async_found.id diff --git a/tests/test_db_transferhistory_queries.py b/tests/test_db_transferhistory_queries.py index d9c5a3436..5bd9bdb94 100644 --- a/tests/test_db_transferhistory_queries.py +++ b/tests/test_db_transferhistory_queries.py @@ -401,8 +401,11 @@ def test_list_by_title_matches_async_twin(db): sync_titles = [h.title for h in TransferHistory.list_by_title( db.session, "ParallelSearch", count=-1)] - async_titles = [h.title for h in asyncio.run(TransferHistory.async_list_by_title( - title="ParallelSearch", count=-1))] + async_titles = [h.title for h in db.run_async_session( + lambda session: TransferHistory.async_list_by_title( + session, title="ParallelSearch", count=-1 + ) + )] assert sync_titles == async_titles @@ -415,13 +418,21 @@ def test_count_and_count_by_title_match_async_twins(db): db.add(_hist("CountMe", status=True, src="/downloads/c1.mkv", dest="/media/c1.mkv"), _hist("CountMe", status=False, src="/downloads/c2.mkv", dest="/media/c2.mkv")) - assert TransferHistory.count(db.session) == asyncio.run(TransferHistory.async_count()) + assert TransferHistory.count(db.session) == db.run_async_session( + TransferHistory.async_count + ) assert TransferHistory.count(db.session, status=True) == \ - asyncio.run(TransferHistory.async_count(status=True)) + db.run_async_session( + lambda session: TransferHistory.async_count(session, status=True) + ) assert TransferHistory.count_by_title(db.session, "CountMe") == 2 assert TransferHistory.count_by_title(db.session, "CountMe", status=False) == 1 assert TransferHistory.count_by_title(db.session, "CountMe") == \ - asyncio.run(TransferHistory.async_count_by_title(title="CountMe")) + db.run_async_session( + lambda session: TransferHistory.async_count_by_title( + session, title="CountMe" + ) + ) def test_statistic_groups_by_day_within_the_window(db): @@ -452,7 +463,9 @@ def test_statistic_includes_the_window_start_boundary(db, frozen_now): db.add(_hist("窗口起点上", src="/data/bstat.mkv", date=window_start)) rows = dict(TransferHistory.statistic(db.session, days=7)) - async_rows = dict(asyncio.run(TransferHistory.async_statistic(days=7))) + async_rows = dict(db.run_async_session( + lambda session: TransferHistory.async_statistic(session, days=7) + )) assert rows.get(boundary_day, 0) == 1 assert async_rows.get(boundary_day, 0) == 1 diff --git a/tests/test_db_transferpending_queries.py b/tests/test_db_transferpending_queries.py index 2c34d13a9..d7f3afe09 100644 --- a/tests/test_db_transferpending_queries.py +++ b/tests/test_db_transferpending_queries.py @@ -7,10 +7,9 @@ """ import pytest -from app.db import decorators +from app.db import base as db_base from app.db.models.transferpending import TransferPending from app.db.oper.transferpending import TransferPendingOper -from app.db.session import SessionFactory @pytest.fixture(autouse=True) @@ -160,32 +159,16 @@ def test_oper_reuses_explicit_query_session(db, monkeypatch): created_at="2026-08-13 10:00:00", )) monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (_ for _ in ()).throw(AssertionError("不应创建额外同步会话")), + db_base, + "run_sync_transaction", + lambda _operation: (_ for _ in ()).throw( + AssertionError("不应创建额外同步事务") + ), ) assert ("local", "/mnt/explicit.mkv") in TransferPendingOper(db.session).list_all() -def test_model_legacy_query_keeps_keyword_abi(db, monkeypatch): - """旧插件以关键字直调 TransferPending 时仍自动补入短会话。""" - db.add(TransferPending( - storage="local", - src_path="/mnt/legacy.mkv", - created_at="2026-08-13 10:00:00", - )) - opened = [] - monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (opened.append(True) or SessionFactory()), - ) - - assert TransferPending.list_all(limit=1) - assert opened == [True] - - def test_oper_drops_rows_with_missing_fields(db): """ 回放时必须跳过字段残缺的历史遗留行,不能把空存储送进整理链。 diff --git a/tests/test_db_workflow_queries.py b/tests/test_db_workflow_queries.py index f204e26bb..4b1164be2 100644 --- a/tests/test_db_workflow_queries.py +++ b/tests/test_db_workflow_queries.py @@ -9,10 +9,10 @@ import asyncio import pytest -from app.db import decorators +from app.db import base as db_base from app.db.models.workflow import Workflow from app.db.oper.workflow import WorkflowOper -from app.db.session import SessionFactory, async_session_scope +from app.db.session import async_session_scope @pytest.fixture(autouse=True) @@ -52,11 +52,13 @@ def test_list_and_get_by_name_match_async_twins(db): created = db.add(_flow("wf-name")) assert Workflow.get_by_name(db.session, "wf-name").id == created.id - assert asyncio.run(Workflow.async_get_by_name(name="wf-name")).id == created.id + assert db.run_async_session( + lambda session: Workflow.async_get_by_name(session, "wf-name") + ).id == created.id assert Workflow.get_by_name(db.session, "wf-missing") is None sync_ids = sorted(w.id for w in Workflow.list(db.session)) - async_ids = sorted(w.id for w in asyncio.run(Workflow.async_list())) + async_ids = sorted(w.id for w in db.run_async_session(Workflow.async_list)) assert sync_ids == async_ids @@ -64,9 +66,11 @@ def test_workflow_oper_reuses_explicit_query_sessions(db, monkeypatch): """WorkflowOper 绑定显式会话后不得再创建兼容查询会话。""" created = db.add(_flow("wf-explicit-session")) monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (_ for _ in ()).throw(AssertionError("不应创建额外同步会话")), + db_base, + "run_sync_transaction", + lambda _operation: (_ for _ in ()).throw( + AssertionError("不应创建额外同步事务") + ), ) assert WorkflowOper(db.session).get_by_name(created.name).id == created.id @@ -75,29 +79,17 @@ def test_workflow_oper_reuses_explicit_query_sessions(db, monkeypatch): """验证异步 Oper 同样复用调用方会话。""" async with async_session_scope() as session: monkeypatch.setattr( - decorators, - "async_session_scope", - lambda: (_ for _ in ()).throw(AssertionError("不应创建额外异步会话")), + db_base, + "run_async_transaction", + lambda _operation: (_ for _ in ()).throw( + AssertionError("不应创建额外异步事务") + ), ) assert (await WorkflowOper(session).async_get_by_name(created.name)).id == created.id asyncio.run(check()) -def test_workflow_model_legacy_queries_keep_no_session_abi(db, monkeypatch): - """旧插件直接调用 Workflow Model 时仍应按签名自动补入短会话。""" - created = db.add(_flow("wf-legacy-query")) - opened = [] - monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (opened.append(True) or SessionFactory()), - ) - - assert Workflow.get_by_name(name=created.name).id == created.id - assert opened == [True] - - def test_enabled_workflows_exclude_paused(db): """ 启用列表排除暂停状态。 @@ -111,8 +103,9 @@ def test_enabled_workflows_exclude_paused(db): assert {"wf-waiting", "wf-running"} <= names assert "wf-paused" not in names - assert "wf-paused" not in {w.name for w in - asyncio.run(Workflow.async_get_enabled_workflows())} + assert "wf-paused" not in { + w.name for w in db.run_async_session(Workflow.async_get_enabled_workflows) + } def test_timer_triggered_includes_legacy_null_trigger_type(db): @@ -156,9 +149,13 @@ def test_trigger_lists_match_async_twins(db): db.add(_flow("wf-t", trigger_type="timer"), _flow("wf-e", trigger_type="event")) assert sorted(w.id for w in Workflow.get_timer_triggered_workflows(db.session)) == \ - sorted(w.id for w in asyncio.run(Workflow.async_get_timer_triggered_workflows())) + sorted(w.id for w in db.run_async_session( + Workflow.async_get_timer_triggered_workflows + )) assert sorted(w.id for w in Workflow.get_event_triggered_workflows(db.session)) == \ - sorted(w.id for w in asyncio.run(Workflow.async_get_event_triggered_workflows())) + sorted(w.id for w in db.run_async_session( + Workflow.async_get_event_triggered_workflows + )) # --------------------------------------------------------------------------- # diff --git a/tests/test_mediaserver_sync_incremental.py b/tests/test_mediaserver_sync_incremental.py index 6054fd956..711ce442b 100644 --- a/tests/test_mediaserver_sync_incremental.py +++ b/tests/test_mediaserver_sync_incremental.py @@ -100,12 +100,18 @@ def test_sync_persists_music_without_querying_tv_episodes(database): ) chain.episodes = lambda *_args, **_kwargs: pytest.fail("音乐条目不应查询电视剧分集") - with patch("app.db.decorators.ScopedSession", database), patch.object( - MEDIA_SERVER_CHAIN_MODULE.ServiceConfigHelper, - "get_mediaserver_configs", - return_value=[SimpleNamespace(name="navidrome", enabled=True, sync_libraries=["all"])], - ): - chain.sync() + with database() as session: + with patch.object( + MEDIA_SERVER_CHAIN_MODULE, + "MediaServerOper", + lambda: MediaServerOper(session), + ), patch.object( + MEDIA_SERVER_CHAIN_MODULE.ServiceConfigHelper, + "get_mediaserver_configs", + return_value=[SimpleNamespace(name="navidrome", enabled=True, sync_libraries=["all"])], + ): + chain.sync() + session.commit() with database() as db: item = db.query(MediaServerItem).one() @@ -187,12 +193,18 @@ def test_sync_updates_rows_and_removes_stale_entries(database): ) chain.episodes = lambda *_args, **_kwargs: [] - with patch("app.db.decorators.ScopedSession", database), patch.object( - MEDIA_SERVER_CHAIN_MODULE.ServiceConfigHelper, - "get_mediaserver_configs", - return_value=[SimpleNamespace(name="plex", enabled=True, sync_libraries=["movies"])], - ): - chain.sync() + with database() as session: + with patch.object( + MEDIA_SERVER_CHAIN_MODULE, + "MediaServerOper", + lambda: MediaServerOper(session), + ), patch.object( + MEDIA_SERVER_CHAIN_MODULE.ServiceConfigHelper, + "get_mediaserver_configs", + return_value=[SimpleNamespace(name="plex", enabled=True, sync_libraries=["movies"])], + ): + chain.sync() + session.commit() with database() as db: items = ( @@ -265,17 +277,23 @@ def test_sync_queries_counts_before_items_and_reports_media_progress(database): chain.items = items chain.episodes = lambda *_args, **_kwargs: [] - with patch("app.db.decorators.ScopedSession", database), patch.object( - MEDIA_SERVER_CHAIN_MODULE.ServiceConfigHelper, - "get_mediaserver_configs", - return_value=[ - SimpleNamespace(name="plex-a", enabled=True, sync_libraries=["all"]), - SimpleNamespace(name="plex-b", enabled=True, sync_libraries=["all"]), - ], - ): - chain.sync( - progress_callback=lambda **kwargs: progress_snapshots.append(kwargs) - ) + with database() as session: + with patch.object( + MEDIA_SERVER_CHAIN_MODULE, + "MediaServerOper", + lambda: MediaServerOper(session), + ), patch.object( + MEDIA_SERVER_CHAIN_MODULE.ServiceConfigHelper, + "get_mediaserver_configs", + return_value=[ + SimpleNamespace(name="plex-a", enabled=True, sync_libraries=["all"]), + SimpleNamespace(name="plex-b", enabled=True, sync_libraries=["all"]), + ], + ): + chain.sync( + progress_callback=lambda **kwargs: progress_snapshots.append(kwargs) + ) + session.commit() assert events == [ "count:plex-a", diff --git a/tests/test_subscribe_create_command.py b/tests/test_subscribe_create_command.py index 286866316..3b970875f 100644 --- a/tests/test_subscribe_create_command.py +++ b/tests/test_subscribe_create_command.py @@ -190,16 +190,13 @@ def test_default_sync_writer_persists_once_and_reuses_duplicate(db) -> None: ] -def test_stage_add_executes_identity_sql_in_oper(db, monkeypatch) -> None: - """规范新增路径直接由 Oper 查询,不能退回 Model 自动会话装饰器。""" +def test_stage_add_reuses_explicit_session_without_commit(db, monkeypatch) -> None: + """Oper 将调用方 Session 传给 Model 查询原语,暂存期间不自行提交。""" db.watermark(Subscribe) commit = Mock(wraps=db.session.commit) + exists = Mock(wraps=Subscribe.exists) monkeypatch.setattr(db.session, "commit", commit) - monkeypatch.setattr( - Subscribe, - "exists", - Mock(side_effect=AssertionError("model query must not run")), - ) + monkeypatch.setattr(Subscribe, "exists", exists) oper = SubscribeOper(db.session) identity = { "media_source": str(MediaSource.TMDB), @@ -221,6 +218,7 @@ def test_stage_add_executes_identity_sql_in_oper(db, monkeypatch) -> None: assert staged.created is True assert staged.subscribe_id > 0 + assert exists.call_args.args[0] is db.session commit.assert_not_called() db.session.rollback() diff --git a/tests/test_subscribe_oper.py b/tests/test_subscribe_oper.py index a6fa14566..0f81f4ee9 100644 --- a/tests/test_subscribe_oper.py +++ b/tests/test_subscribe_oper.py @@ -418,11 +418,10 @@ def test_exists_defaults_to_main_season_episode_group(): assert history_model.exists.call_args.kwargs["episode_group"] == "eg-1" -def test_subscribe_exists_distinguishes_same_season_episode_groups(): +def test_subscribe_exists_distinguishes_same_season_episode_groups(db): """同一媒体同一季的主季、自定义剧集组应分别命中各自订阅。""" - oper = SubscribeOper() + db.watermark(Subscribe) media_id = str(-(900_000_000 + os.getpid())) - created_ids = [] rows = [ Subscribe(name="主季订阅", type=MediaType.TV.value, state="N", media_source=MediaSource.TMDB.value, media_id=media_id, @@ -431,43 +430,38 @@ def test_subscribe_exists_distinguishes_same_season_episode_groups(): media_source=MediaSource.TMDB.value, media_id=media_id, season=1, episode_group="eg-1"), ] - try: - for row in rows: - row.create(oper._db) + for row in rows: + row.create(db.session) + db.session.commit() - main_season = Subscribe.exists( - oper._db, media_source=MediaSource.TMDB, - media_id=media_id, season=1, episode_group=None, - ) - created_ids.append(main_season.id) - main_name = main_season.name - episode_group = Subscribe.exists( - oper._db, media_source=MediaSource.TMDB, - media_id=media_id, season=1, episode_group="eg-1", - ) - created_ids.append(episode_group.id) - episode_group_name = episode_group.name + main_season = Subscribe.exists( + db.session, media_source=MediaSource.TMDB, + media_id=media_id, season=1, episode_group=None, + ) + main_name = main_season.name + episode_group = Subscribe.exists( + db.session, media_source=MediaSource.TMDB, + media_id=media_id, season=1, episode_group="eg-1", + ) + episode_group_name = episode_group.name - assert main_name == "主季订阅" - assert episode_group_name == "剧集组订阅" + assert main_name == "主季订阅" + assert episode_group_name == "剧集组订阅" - Subscribe.delete(oper._db, rid=created_ids.pop(0)) - assert Subscribe.exists( - oper._db, - media_source=MediaSource.TMDB, - media_id=media_id, - season=1, - ) is None - finally: - for subscribe_id in created_ids: - Subscribe.delete(oper._db, rid=subscribe_id) + Subscribe.delete(db.session, rid=main_season.id) + db.session.commit() + assert Subscribe.exists( + db.session, + media_source=MediaSource.TMDB, + media_id=media_id, + season=1, + ) is None -def test_subscribe_exists_distinguishes_music_entities_with_same_source_id(): +def test_subscribe_exists_distinguishes_music_entities_with_same_source_id(db): """统一来源 ID 相同时,单曲与专辑仍是两条独立订阅身份。""" - oper = SubscribeOper() + db.watermark(Subscribe) media_id = f"music-shared-{os.getpid()}" - created_ids = [] rows = [ Subscribe( name="同名单曲", @@ -487,29 +481,24 @@ def test_subscribe_exists_distinguishes_music_entities_with_same_source_id(): total_tracks=10, ), ] - try: - for row in rows: - row.create(oper._db) + for row in rows: + row.create(db.session) + db.session.commit() - recording = Subscribe.exists( - oper._db, - media_source="musicbrainz", - media_id=media_id, - music_type="recording", - ) - created_ids.append(recording.id) - album = Subscribe.exists( - oper._db, - media_source="musicbrainz", - media_id=media_id, - music_type="album", - ) - created_ids.append(album.id) - assert recording.name == "同名单曲" - assert album.name == "同名专辑" - finally: - for subscribe_id in created_ids: - Subscribe.delete(oper._db, rid=subscribe_id) + recording = Subscribe.exists( + db.session, + media_source="musicbrainz", + media_id=media_id, + music_type="recording", + ) + album = Subscribe.exists( + db.session, + media_source="musicbrainz", + media_id=media_id, + music_type="album", + ) + assert recording.name == "同名单曲" + assert album.name == "同名专辑" def test_subscribe_chain_exists_forwards_episode_group(): diff --git a/tests/test_systemconfig_oper.py b/tests/test_systemconfig_oper.py index 5e6b77d61..076ccb2d3 100644 --- a/tests/test_systemconfig_oper.py +++ b/tests/test_systemconfig_oper.py @@ -7,6 +7,7 @@ import pytest from app.db.models.systemconfig import SystemConfig from app.db.oper.systemconfig import SystemConfigOper +from app.db.session import SessionFactory from app.schemas.types import SystemConfigKey from app.foundation.singleton import Singleton @@ -24,6 +25,12 @@ def _fresh_oper() -> SystemConfigOper: return oper +def _stored_config(key: str) -> SystemConfig | None: + """使用显式短会话回读系统配置持久化结果。""" + with SessionFactory() as session: + return SystemConfig.get_by_key(session, key) + + def test_constructor_does_not_query_database(monkeypatch): """构造配置对象时不打开数据库会话。""" Singleton._instances.pop((SystemConfigOper, (), frozenset()), None) @@ -167,7 +174,7 @@ def test_set_creates_record_for_falsy_value(): assert oper.set(key, False) is True assert oper.get(key) is False - assert SystemConfig.get_by_key(oper._db, key) is not None + assert _stored_config(key) is not None def test_set_persists_falsy_value_on_existing_record(): @@ -178,7 +185,7 @@ def test_set_persists_falsy_value_on_existing_record(): oper.set(key, True) assert oper.set(key, False) is True assert oper.get(key) is False - assert SystemConfig.get_by_key(oper._db, key).value is False + assert _stored_config(key).value is False @pytest.mark.asyncio @@ -200,7 +207,7 @@ async def test_async_set_persists_falsy_value_on_existing_record(): ) assert await service.async_set(key, False) is True assert oper.get(key) is False - assert SystemConfig.get_by_key(oper._db, key).value is False + assert _stored_config(key).value is False @pytest.mark.asyncio @@ -221,7 +228,7 @@ async def test_async_set_creates_record_for_falsy_value(): ) assert await service.async_set(key, 0) is True assert oper.get(key) == 0 - assert SystemConfig.get_by_key(oper._db, key).value == 0 + assert _stored_config(key).value == 0 def test_delete_removes_record_explicitly(): @@ -232,7 +239,7 @@ def test_delete_removes_record_explicitly(): oper.set(key, False) assert oper.delete(key) is True assert oper.get(key) is None - assert SystemConfig.get_by_key(oper._db, key) is None + assert _stored_config(key) is None def test_mounted_local_disk_delete_empty_dirs_off_is_persisted(): diff --git a/tests/test_transfer_download_history_oper_sessions.py b/tests/test_transfer_download_history_oper_sessions.py new file mode 100644 index 000000000..67e451340 --- /dev/null +++ b/tests/test_transfer_download_history_oper_sessions.py @@ -0,0 +1,51 @@ +"""Transfer/Download History Oper 的显式会话复用验证。""" + +import asyncio + +from app.db import base as db_base +from app.db.models.transferhistory import TransferHistory +from app.db.oper.downloadhistory import DownloadHistoryOper +from app.db.oper.transferhistory import TransferHistoryOper +from app.db.models.downloadhistory import DownloadHistory +from app.db.session import async_session_scope + + +def test_oper_reuses_explicit_sync_session(db, monkeypatch): + """显式同步会话绑定到 Oper 后,查询不能再创建兼容会话。""" + row = db.add(TransferHistory(src="/compat/transfer.mkv", src_storage="local")) + monkeypatch.setattr( + db_base, + "run_sync_transaction", + lambda _operation: (_ for _ in ()).throw( + AssertionError("不应创建额外同步事务") + ), + ) + + assert TransferHistoryOper(db.session).get_by_src("/compat/transfer.mkv").id == row.id + assert DownloadHistoryOper(db.session).get_by_hash("missing") is None + + +def test_oper_reuses_explicit_async_session(db, monkeypatch): + """显式异步会话绑定到 Oper 后,异步查询不能再创建兼容作用域。""" + db.add( + DownloadHistory( + path="/compat/async-download", + type="电视剧", + title="异步兼容", + download_hash="async-compat", + ) + ) + + async def check() -> None: + async with async_session_scope() as session: + monkeypatch.setattr( + db_base, + "run_async_transaction", + lambda _operation: (_ for _ in ()).throw( + AssertionError("不应创建额外异步事务") + ), + ) + result = await DownloadHistoryOper(session).async_list_by_page(count=10) + assert any(item.download_hash == "async-compat" for item in result) + + asyncio.run(check()) diff --git a/tests/test_transfer_download_history_query_compat.py b/tests/test_transfer_download_history_query_compat.py deleted file mode 100644 index a85eda5f2..000000000 --- a/tests/test_transfer_download_history_query_compat.py +++ /dev/null @@ -1,111 +0,0 @@ -"""Transfer/Download History 查询兼容层的会话与旧插件 ABI 验证。""" - -import asyncio - -from app.db import decorators -from app.db.models.downloadhistory import DownloadHistory -from app.db.models.transferhistory import TransferHistory -from app.db.oper.downloadhistory import DownloadHistoryOper -from app.db.oper.transferhistory import TransferHistoryOper -from app.db.session import SessionFactory, async_session_scope - - -def test_oper_reuses_explicit_sync_session(db, monkeypatch): - """显式同步会话绑定到 Oper 后,查询不能再创建兼容会话。""" - row = db.add(TransferHistory(src="/compat/transfer.mkv", src_storage="local")) - monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (_ for _ in ()).throw(AssertionError("不应创建额外同步会话")), - ) - - assert TransferHistoryOper(db.session).get_by_src("/compat/transfer.mkv").id == row.id - assert DownloadHistoryOper(db.session).get_by_hash("missing") is None - - -def test_model_legacy_sync_calls_preserve_business_arguments(db, monkeypatch): - """旧插件省略 db 时,第一个位置参数仍须作为业务参数传入。""" - row = db.add(TransferHistory(src="/compat/legacy.mkv", src_storage="local")) - created = [] - monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (created.append(True) or SessionFactory()), - ) - - assert TransferHistory.get_by_src("/compat/legacy.mkv").id == row.id - assert created == [True] - - -def test_download_model_legacy_sync_call_preserves_keyword_arguments(db, monkeypatch): - """旧插件使用关键字查询时,兼容层仍须自动补入 db。""" - row = db.add( - DownloadHistory( - path="/compat/download", - type="电视剧", - download_hash="compat-hash", - title="兼容", - ) - ) - created = [] - monkeypatch.setattr( - decorators, - "ScopedSession", - lambda: (created.append(True) or SessionFactory()), - ) - - assert DownloadHistory.get_by_hash(download_hash="compat-hash").id == row.id - assert created == [True] - - -def test_oper_reuses_explicit_async_session(db, monkeypatch): - """显式异步会话绑定到 Oper 后,异步查询不能再创建兼容作用域。""" - db.add( - DownloadHistory( - path="/compat/async-download", - type="电视剧", - title="异步兼容", - download_hash="async-compat", - ) - ) - - async def check() -> None: - async with async_session_scope() as session: - monkeypatch.setattr( - decorators, - "async_session_scope", - lambda: (_ for _ in ()).throw(AssertionError("不应创建额外异步会话")), - ) - result = await DownloadHistoryOper(session).async_list_by_page(count=10) - assert any(item.download_hash == "async-compat" for item in result) - - asyncio.run(check()) - - -def test_model_legacy_async_calls_support_explicit_and_implicit_sessions(db, monkeypatch): - """异步 Model 查询同时保留显式会话调用与旧插件无会话调用。""" - db.add( - DownloadHistory( - path="/compat/async-legacy", - type="电视剧", - title="异步旧 ABI", - download_hash="async-legacy", - ) - ) - original_scope = decorators.async_session_scope - created = [] - - def tracked_scope(): - """记录兼容层是否创建了异步会话作用域。""" - created.append(True) - return original_scope() - - async def check() -> None: - async with original_scope() as session: - assert await DownloadHistory.async_count(session) >= 1 - monkeypatch.setattr(decorators, "async_session_scope", tracked_scope) - result = await DownloadHistory.async_list_by_title(title="异步旧 ABI") - assert result[0].download_hash == "async-legacy" - - asyncio.run(check()) - assert created == [True] diff --git a/tests/test_web_agent_stream.py b/tests/test_web_agent_stream.py index 687349532..101ef7a07 100644 --- a/tests/test_web_agent_stream.py +++ b/tests/test_web_agent_stream.py @@ -1206,7 +1206,7 @@ def test_web_agent_stream_emits_secret_result_only_as_protected_event(): worker = agent_manager._session_workers.pop(session_id, None) if worker: worker.cancel() - AgentChat.delete(rid=existing_chat.id) + AgentChatOper().delete_by_id(existing_chat.id) def test_web_agent_cancel_keeps_existing_display_history(): @@ -1265,7 +1265,7 @@ def test_web_agent_cancel_keeps_existing_display_history(): assert preserved_chat.message_count == 2 assert preserved_chat.preview == "保留的回答" finally: - AgentChat.delete(rid=existing_chat.id) + AgentChatOper().delete_by_id(existing_chat.id) def test_web_agent_stream_rejects_confirmation_without_protected_capability(): @@ -1411,7 +1411,7 @@ def test_web_agent_stream_drops_secret_result_after_disconnect(): assert preserved_chat.message_count == 2 assert preserved_chat.preview == "断线前的回答" finally: - AgentChat.delete(rid=existing_chat.id) + AgentChatOper().delete_by_id(existing_chat.id) def test_web_agent_stream_emits_heartbeat_during_idle_tool_wait():