diff --git a/app/api/context.py b/app/api/context.py index 8b9686eca..3cd2f7bd0 100644 --- a/app/api/context.py +++ b/app/api/context.py @@ -1,6 +1,6 @@ """从 FastAPI AppState 读取类型化宿主能力。""" -from collections.abc import AsyncGenerator +from collections.abc import AsyncGenerator, Generator from typing import cast from fastapi import Depends, Request @@ -39,6 +39,21 @@ def get_api_runtime_config( return runtime.configuration.api() +def get_sync_session( + runtime: HostRuntime = Depends(get_host_runtime), +) -> Generator[object, None, None]: + """从 HostRuntime 生成请求独占的同步数据库会话。""" + yield from runtime.persistence.sync_session() + + +async def get_async_session( + runtime: HostRuntime = Depends(get_host_runtime), +) -> AsyncGenerator[object, None]: + """从 HostRuntime 生成请求独占的异步数据库会话。""" + async for session in runtime.persistence.async_session(): + yield session + + def resolve_api_runtime_config(value: object) -> ApiRuntimeConfig: """兼容直接调用 endpoint 的旧入口,并统一返回真实配置快照。""" if isinstance(value, ApiRuntimeConfig): diff --git a/app/api/dependencies/agent.py b/app/api/dependencies/agent.py index d5e8dd2f0..fcac8e71f 100644 --- a/app/api/dependencies/agent.py +++ b/app/api/dependencies/agent.py @@ -3,15 +3,19 @@ from fastapi import Depends from sqlalchemy.ext.asyncio import AsyncSession -from app.api.context import get_agent_chat_repository, get_agent_chat_transaction -from app.api.data import get_async_db -from app.api.dependencies.data import repository +from app.api.context import ( + get_agent_chat_repository, + get_agent_chat_transaction, + get_async_session, + get_host_runtime, +) from app.application.messaging.chat import ( AgentChatService, AsyncAgentChatRepository, AsyncUnitOfWork, ) from app.application.messaging.message import MessageQueryService +from app.startup.context import HostRuntime def get_agent_chat_service( @@ -23,7 +27,8 @@ def get_agent_chat_service( def get_message_query_service( - db: AsyncSession = Depends(get_async_db), + db: AsyncSession = Depends(get_async_session), + runtime: HostRuntime = Depends(get_host_runtime), ) -> MessageQueryService: """组装消息历史异步查询服务。""" - return MessageQueryService(repository=repository("message", db)) + return MessageQueryService(repository=runtime.messaging.repository(db)) diff --git a/app/api/dependencies/auth.py b/app/api/dependencies/auth.py index 5dfc19ea0..ab83a4660 100644 --- a/app/api/dependencies/auth.py +++ b/app/api/dependencies/auth.py @@ -1,58 +1,89 @@ """用户身份、授权与认证服务依赖。""" -from typing import Any +from typing import Any, cast from fastapi import Depends, HTTPException from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Session from app.adapters.web.security.access import verify_token -from app.api.data import get_async_db, get_db -from app.api.dependencies.data import repository, standalone_repository -from app.application.security.auth import AuthService -from app.application.security.passkeys import PasskeyService -from app.application.security.user import UserService +from app.api.context import get_async_session, get_host_runtime, get_sync_session +from app.application.security.auth import ( + AuthConfigRepository, + AuthPasskeyRepository, + AuthService, + AuthUserRepository, +) +from app.application.security.passkeys import PasskeyRepository, PasskeyService +from app.application.security.user import ( + AsyncUnitOfWork, + UserRepository, + UserService, +) from app.schemas.token import TokenPayload as _SchemaTokenPayload +from app.startup.context import HostRuntime def get_user_service( - db: AsyncSession = Depends(get_async_db), + db: AsyncSession = Depends(get_async_session), + runtime: HostRuntime = Depends(get_host_runtime), ) -> UserService: """组装用户管理应用服务。""" - return UserService(repository=repository("user", db)) - - -def get_auth_service() -> AuthService: - """组装同步认证应用服务。""" - return AuthService( - users=standalone_repository("user"), - config=standalone_repository("system_config"), - passkeys=standalone_repository("passkey"), + return UserService( + repository=cast( + UserRepository, runtime.authentication.user_repository(db) + ), + unit_of_work=cast( + AsyncUnitOfWork, runtime.persistence.async_transaction(db) + ), ) -def get_passkey_service() -> PasskeyService: +def get_auth_service( + runtime: HostRuntime = Depends(get_host_runtime), +) -> AuthService: + """组装同步认证应用服务。""" + return AuthService( + users=cast(AuthUserRepository, runtime.authentication.standalone_user()), + config=cast(AuthConfigRepository, runtime.authentication.system_config()), + passkeys=cast(AuthPasskeyRepository, runtime.authentication.passkey()), + ) + + +def get_passkey_service( + runtime: HostRuntime = Depends(get_host_runtime), +) -> PasskeyService: """组装 PassKey 应用服务。""" - return PasskeyService(repository=standalone_repository("passkey")) + return PasskeyService(repository=cast( + PasskeyRepository, runtime.authentication.passkey() + )) def get_current_user( - db: Session = Depends(get_db), + db: Session = Depends(get_sync_session), token_data: _SchemaTokenPayload = Depends(verify_token), + runtime: HostRuntime = Depends(get_host_runtime), ) -> Any: """读取令牌对应用户,不存在时返回 403。""" - user = repository("user", db).get_by_id(token_data.sub) + user_repository = cast( + AuthUserRepository, runtime.authentication.user_repository(db) + ) + user = user_repository.get_by_id(token_data.sub) if not user: raise HTTPException(status_code=403, detail="用户不存在") return user async def get_current_user_async( - db: AsyncSession = Depends(get_async_db), + db: AsyncSession = Depends(get_async_session), token_data: _SchemaTokenPayload = Depends(verify_token), + runtime: HostRuntime = Depends(get_host_runtime), ) -> Any: """异步读取令牌对应用户,不存在时返回 403。""" - user = await repository("user", db).async_get_by_id(token_data.sub) + user_repository = cast( + UserRepository, runtime.authentication.user_repository(db) + ) + user = await user_repository.async_get_by_id(token_data.sub) if not user: raise HTTPException(status_code=403, detail="用户不存在") return user diff --git a/app/api/dependencies/history.py b/app/api/dependencies/history.py index 36bc43d88..3367ac8af 100644 --- a/app/api/dependencies/history.py +++ b/app/api/dependencies/history.py @@ -4,8 +4,7 @@ from fastapi import Depends from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Session -from app.api.data import get_async_db, get_db -from app.api.dependencies.data import repository, transaction +from app.api.context import get_async_session, get_host_runtime, get_sync_session from app.application.dashboard import DashboardQueryService from app.application.history import ( DownloadHistoryMutationCommand, @@ -19,63 +18,74 @@ from app.chain.storage import StorageChain from app.runtime.events import eventmanager from app.schemas.types import EventType from app.schemas.workflow import FileItem as _SchemaFileItem +from app.startup.context import HostRuntime def get_mediaserver_query_service( - db: AsyncSession = Depends(get_async_db), + db: AsyncSession = Depends(get_async_session), + runtime: HostRuntime = Depends(get_host_runtime), ) -> MediaServerQueryService: """组装媒体服务器本地条目异步查询服务。""" - return MediaServerQueryService(repository=repository("media_server", db)) + return MediaServerQueryService( + repository=runtime.history.media_server_repository(db) + ) def get_dashboard_query_service( - db: Session = Depends(get_db), + db: Session = Depends(get_sync_session), + runtime: HostRuntime = Depends(get_host_runtime), ) -> DashboardQueryService: """组装 Dashboard 媒体与整理历史统计查询服务。""" from app.chain.dashboard import DashboardChain return DashboardQueryService( - repository=repository("transfer_history", db), + repository=runtime.history.transfer_repository(db), media_statistics=DashboardChain().media_statistic, ) def get_download_history_mutation_command( - db: Session = Depends(get_db), + db: Session = Depends(get_sync_session), + runtime: HostRuntime = Depends(get_host_runtime), ) -> DownloadHistoryMutationCommand: """组装下载历史删除用例及其请求级事务。""" return DownloadHistoryMutationCommand( - repository=repository("download_history", db), - unit_of_work=transaction("sync", db), + repository=runtime.history.download_repository(db), + unit_of_work=runtime.persistence.sync_transaction(db), ) def get_history_query_service( - db: AsyncSession = Depends(get_async_db), + db: AsyncSession = Depends(get_async_session), + runtime: HostRuntime = Depends(get_host_runtime), ) -> HistoryQueryService: """组装历史列表和详情异步查询服务。""" return HistoryQueryService( - download_repository=repository("download_history", db), - transfer_repository=repository("transfer_history", db), + download_repository=runtime.history.download_repository(db), + transfer_repository=runtime.history.transfer_repository(db), ) def get_transfer_history_lookup_service( - db: Session = Depends(get_db), + db: Session = Depends(get_sync_session), + runtime: HostRuntime = Depends(get_host_runtime), ) -> TransferHistoryLookupService: """组装手动整理使用的同步历史投影服务。""" - return TransferHistoryLookupService(repository("transfer_history", db)) + return TransferHistoryLookupService( + runtime.history.transfer_repository(db) + ) def get_transfer_history_mutation_command( - db: Session = Depends(get_db), + db: Session = Depends(get_sync_session), + runtime: HostRuntime = Depends(get_host_runtime), ) -> TransferHistoryMutationCommand: """组装整理历史删除、文件处理和事件发布用例。""" storage_chain = StorageChain() return TransferHistoryMutationCommand( - repository=repository("transfer_history", db), - download_repository=repository("download_history", db), - unit_of_work=transaction("sync", db), + repository=runtime.history.transfer_repository(db), + download_repository=runtime.history.download_repository(db), + unit_of_work=runtime.persistence.sync_transaction(db), file_item_factory=lambda payload: _SchemaFileItem(**payload), delete_media_file=storage_chain.delete_media_file, publish_download_file_deleted=lambda payload: eventmanager.send_event( diff --git a/app/api/dependencies/site.py b/app/api/dependencies/site.py index 7b0eab9ea..01e383e0e 100644 --- a/app/api/dependencies/site.py +++ b/app/api/dependencies/site.py @@ -1,11 +1,12 @@ """站点领域的请求级 command/query 依赖。""" +from typing import Any + from fastapi import Depends from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Session -from app.api.data import get_async_db, get_db -from app.api.dependencies.data import repository, transaction +from app.api.context import get_async_session, get_host_runtime, get_sync_session from app.application.site.mutation import SiteMutationCommand from app.application.site.query import SiteQueryService from app.application.site.sites import SitesHelper # pylint: disable=import-error,no-name-in-module @@ -13,20 +14,22 @@ from app.domain import site as site_rules from app.foundation import url as url_tools from app.runtime.events import eventmanager from app.schemas.types import EventType +from app.startup.context import HostRuntime -async def _publish_site_updated(payload: dict) -> None: +async def _publish_site_updated(payload: dict[str, Any]) -> None: """发布已提交的站点更新事件。""" await eventmanager.async_send_event(EventType.SiteUpdated, payload) -async def _publish_site_deleted(payload: dict) -> None: +async def _publish_site_deleted(payload: dict[str, Any]) -> None: """发布已提交的站点删除事件。""" await eventmanager.async_send_event(EventType.SiteDeleted, payload) def get_site_mutation_command( - db: AsyncSession = Depends(get_async_db), + db: AsyncSession = Depends(get_async_session), + runtime: HostRuntime = Depends(get_host_runtime), ) -> SiteMutationCommand: """组装请求级站点写用例及其事务和外部目录依赖。""" sites_helper = SitesHelper() @@ -37,8 +40,8 @@ def get_site_mutation_command( return f"{scheme}://{netloc}/" return SiteMutationCommand( - repository=repository("site", db), - unit_of_work=transaction("async", db), + repository=runtime.site.repository(db), + unit_of_work=runtime.persistence.async_transaction(db), auth_level_provider=lambda: sites_helper.auth_level, indexer_loader=sites_helper.async_get_indexer, domain_extractor=site_rules.extract_domain, @@ -49,14 +52,16 @@ def get_site_mutation_command( def get_site_query_service( - db: AsyncSession = Depends(get_async_db), + db: AsyncSession = Depends(get_async_session), + runtime: HostRuntime = Depends(get_host_runtime), ) -> SiteQueryService: """组装站点异步查询服务。""" - return SiteQueryService(repository=repository("site", db)) + return SiteQueryService(repository=runtime.site.repository(db)) def get_site_sync_query_service( - db: Session = Depends(get_db), + db: Session = Depends(get_sync_session), + runtime: HostRuntime = Depends(get_host_runtime), ) -> SiteQueryService: """组装站点同步查询服务,用于同步 Chain 路由。""" - return SiteQueryService(repository=repository("site", db)) + return SiteQueryService(repository=runtime.site.repository(db)) diff --git a/app/api/dependencies/subscription.py b/app/api/dependencies/subscription.py index 906925ca6..2897e35c1 100644 --- a/app/api/dependencies/subscription.py +++ b/app/api/dependencies/subscription.py @@ -8,13 +8,14 @@ from sqlalchemy.orm import Session from app.adapters.external.server import MoviePilotServerHelper from app.api.context import ( + get_async_session, + get_host_runtime, get_subscription_history_repository, get_subscription_outbox, get_subscription_repository, get_subscription_transaction, + get_sync_session, ) -from app.api.data import get_async_db, get_db -from app.api.dependencies.data import repository from app.application.outbox import AsyncOutboxTransaction from app.application.scheduling import start_scheduler_job from app.application.servarr import ServarrSubscriptionService @@ -36,6 +37,7 @@ from app.application.subscription.search import SearchSubscriptionsCommand from app.runtime.events import eventmanager from app.runtime.log import logger from app.schemas.types import EventType +from app.startup.context import HostRuntime async def _publish_subscribe_deleted( @@ -93,7 +95,8 @@ def get_delete_subscriptions_by_identity_command( def get_search_subscriptions_command( background_tasks: BackgroundTasks, - db: AsyncSession = Depends(get_async_db), + db: AsyncSession = Depends(get_async_session), + runtime: HostRuntime = Depends(get_host_runtime), ) -> SearchSubscriptionsCommand: """组装手工订阅搜索用例,并把调度延迟到响应后的后台任务。""" def schedule_search(subscribe_id: int | None, state: str | None) -> None: @@ -107,19 +110,20 @@ def get_search_subscriptions_command( ) return SearchSubscriptionsCommand( - repository=repository("subscribe", db), + repository=runtime.subscription.repository(db), schedule_search=schedule_search, ) def get_subscription_query_service( - db: AsyncSession = Depends(get_async_db), + db: AsyncSession = Depends(get_async_session), + runtime: HostRuntime = Depends(get_host_runtime), ) -> SubscriptionQueryService: """组装订阅和订阅历史异步查询服务。""" return SubscriptionQueryService( - repository=repository("subscribe", db), - async_repository=repository("subscribe", db), - history_repository=repository("subscribe_history", db), + repository=runtime.subscription.repository(db), + async_repository=runtime.subscription.repository(db), + history_repository=runtime.subscription.history_repository(db), ) @@ -142,18 +146,25 @@ def get_subscription_mutation_service( def get_subscription_sync_mutation_service( - db: Session = Depends(get_db), + db: Session = Depends(get_sync_session), + runtime: HostRuntime = Depends(get_host_runtime), ) -> SubscriptionMutationService: """组装同步订阅查询服务,供文件信息接口使用。""" - return SubscriptionMutationService(repository=repository("subscribe", db)) + return SubscriptionMutationService( + repository=cast( + SubscriptionMutationRepository, + runtime.subscription.repository(db), + ) + ) def get_servarr_subscription_service( - async_db: AsyncSession = Depends(get_async_db), - db: Session = Depends(get_db), + async_db: AsyncSession = Depends(get_async_session), + db: Session = Depends(get_sync_session), + runtime: HostRuntime = Depends(get_host_runtime), ) -> ServarrSubscriptionService: """组装 Servarr 兼容路由的请求级订阅数据用例。""" return ServarrSubscriptionService( - async_repository=repository("subscribe", async_db), - sync_repository=repository("subscribe", db), + async_repository=runtime.subscription.repository(async_db), + sync_repository=runtime.subscription.repository(db), ) diff --git a/app/api/dependencies/workflow.py b/app/api/dependencies/workflow.py index a7b03921e..ca156a55f 100644 --- a/app/api/dependencies/workflow.py +++ b/app/api/dependencies/workflow.py @@ -1,12 +1,13 @@ """工作流领域的请求级 command/query 依赖。""" +from typing import Any, cast + from fastapi import Depends from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Session from app.adapters.external.server import MoviePilotServerHelper -from app.api.data import get_async_db, get_db -from app.api.dependencies.data import repository, standalone_repository, transaction +from app.api.context import get_async_session, get_host_runtime, get_sync_session from app.application.scheduling import Scheduler from app.application.workflow import ( WorkflowDefinitionCommand, @@ -15,46 +16,50 @@ from app.application.workflow import ( ) from app.runtime.config import global_vars from app.workflow import WorkFlowManager +from app.startup.context import HostRuntime def get_workflow_mutation_command( - db: Session = Depends(get_db), + db: Session = Depends(get_sync_session), + runtime: HostRuntime = Depends(get_host_runtime), ) -> WorkflowMutationCommand: """组装请求级工作流写用例和提交后的调度副作用。""" scheduler = Scheduler() workflow_manager = WorkFlowManager() return WorkflowMutationCommand( - repository=repository("workflow", db), - unit_of_work=transaction("sync", db), + repository=runtime.workflow.repository(db), + unit_of_work=runtime.persistence.sync_transaction(db), add_timer=scheduler.update_workflow_job, remove_timer=scheduler.remove_workflow_job, load_event=workflow_manager.load_workflow_events, remove_event=workflow_manager.remove_workflow_event, refresh_event=workflow_manager.update_workflow_event, stop_running=global_vars.stop_workflow, - delete_cache=lambda workflow_id: standalone_repository( - "system_config" + delete_cache=lambda workflow_id: cast( + Any, runtime.workflow.system_config() ).delete(f"WorkflowCache-{workflow_id}"), ) def get_workflow_definition_command( - db: AsyncSession = Depends(get_async_db), + db: AsyncSession = Depends(get_async_session), + runtime: HostRuntime = Depends(get_host_runtime), ) -> WorkflowDefinitionCommand: """组装工作流创建、复用和重置的异步写用例。""" return WorkflowDefinitionCommand( - repository=repository("workflow", db), - unit_of_work=transaction("async", db), + repository=runtime.workflow.repository(db), + unit_of_work=runtime.persistence.async_transaction(db), stop_running=global_vars.stop_workflow, - delete_cache=lambda workflow_id: standalone_repository( - "system_config" + delete_cache=lambda workflow_id: cast( + Any, runtime.workflow.system_config() ).delete(f"WorkflowCache-{workflow_id}"), report_fork=MoviePilotServerHelper.async_workflow_fork_by_id, ) def get_workflow_query_service( - db: AsyncSession = Depends(get_async_db), + db: AsyncSession = Depends(get_async_session), + runtime: HostRuntime = Depends(get_host_runtime), ) -> WorkflowQueryService: """组装工作流只读查询用例,避免端点直接持有数据库操作器。""" - return WorkflowQueryService(repository=repository("workflow", db)) + return WorkflowQueryService(repository=runtime.workflow.repository(db)) diff --git a/app/application/maintenance.py b/app/application/maintenance.py index 94d31a432..61264a666 100644 --- a/app/application/maintenance.py +++ b/app/application/maintenance.py @@ -45,6 +45,10 @@ class CleanupRepository(Protocol): """返回一次维护运行共用的数据库会话上下文。""" ... + def unit_of_work(self, db: Any) -> "CleanupUnitOfWork": + """返回绑定到当前维护会话的事务边界。""" + ... + def delete_messages(self, db: Any, cutoff: str, limit: int) -> int: """删除早于截止时间的消息。""" ... @@ -70,6 +74,18 @@ class CleanupRepository(Protocol): ... +class CleanupUnitOfWork(Protocol): + """数据维护每一批删除所需的最小事务能力。""" + + def commit(self) -> None: + """提交当前批次。""" + ... + + def rollback(self) -> None: + """回滚失败批次并恢复会话可用状态。""" + ... + + class DataCleanupService: """按配置执行分批数据清理并生成兼容报告。""" @@ -175,16 +191,19 @@ class DataCleanupService: value=plan_index / total_plans * 100, text=f"正在清理数据表 {plan.name} ...", ) + unit_of_work = self._repository.unit_of_work(db) table_report = self._cleanup_in_batches( db=db, table_name=plan.name, delete_batch=plan.delete_batch, + unit_of_work=unit_of_work, ) table_report["cutoff"] = plan.cutoff table_report["retention_days"] = plan.retention_days report["tables"][plan.name] = table_report report["total_deleted"] += table_report["deleted"] except Exception as err: + self._repository.unit_of_work(db).rollback() errors.append(f"{plan.name}: {str(err)}") logger.error(f"数据表 {plan.name} 清理失败:{str(err)}") report["tables"][plan.name] = { @@ -279,12 +298,13 @@ class DataCleanupService: ), ] - @staticmethod def _cleanup_in_batches( + self, *, db: Any, table_name: str, delete_batch: Callable[[Any], int], + unit_of_work: CleanupUnitOfWork, ) -> Dict[str, int]: """循环执行单表分批删除,直到持久化端口返回零。""" total_deleted = 0 @@ -293,6 +313,7 @@ class DataCleanupService: deleted = delete_batch(db) or 0 if deleted <= 0: break + unit_of_work.commit() batches += 1 total_deleted += deleted logger.info( diff --git a/app/application/outbox.py b/app/application/outbox.py index 16db0c8a7..ee5f3cfe7 100644 --- a/app/application/outbox.py +++ b/app/application/outbox.py @@ -146,14 +146,16 @@ class OutboxDispatcher: lease_seconds: int = 60, clock: Callable[[], datetime] | None = None, close: Callable[[], None] | None = None, + failure_observer: Callable[[bool], None] | None = None, ) -> None: - """注入持久端口、topic handler 和有界重试策略。""" + """注入持久端口、topic handler、有界重试策略与失败观测端口。""" self._repository = repository self._handlers = handlers self._max_attempts = max_attempts self._lease_seconds = lease_seconds self._clock = clock or (lambda: datetime.now(timezone.utc)) self._close = close or (lambda: None) + self._failure_observer = failure_observer or (lambda _dead: None) def dispatch_one(self) -> bool: """处理一条到期消息;无消息返回 False,handler 失败留待重试。""" @@ -176,6 +178,7 @@ class OutboxDispatcher: last_error=str(error)[:4000], dead=dead, ) + self._failure_observer(dead) return True self._repository.complete(message.message_id, now) return True diff --git a/app/application/security/user.py b/app/application/security/user.py index 52cb6bd5a..71a2d0fa5 100644 --- a/app/application/security/user.py +++ b/app/application/security/user.py @@ -4,7 +4,7 @@ 避免 API 层同时承担 HTTP 编排和 ORM 适配职责。 """ -from collections.abc import Callable +from collections.abc import Awaitable, Callable from typing import Any, Protocol @@ -33,12 +33,27 @@ class UserRepository(Protocol): """更新用户 OTP 状态。""" +class AsyncUnitOfWork(Protocol): + """用户写用例所需的异步事务边界。""" + + async def commit(self) -> None: + """提交用户写入。""" + + async def rollback(self) -> None: + """回滚失败的用户写入。""" + + class UserService: """用户管理应用服务。""" - def __init__(self, repository: UserRepository) -> None: - """创建用户服务。""" + def __init__( + self, + repository: UserRepository, + unit_of_work: AsyncUnitOfWork | None = None, + ) -> None: + """创建用户服务;旧独立仓储可暂不提供请求级 UoW。""" self._repository = repository + self._unit_of_work = unit_of_work async def list(self) -> list[Any]: """返回用户列表。""" @@ -54,19 +69,35 @@ class UserService: async def create(self, payload: dict[str, Any]) -> Any | None: """创建用户。""" - return await self._repository.async_create(payload) + return await self._write(lambda: self._repository.async_create(payload)) async def update(self, user_id: int, payload: dict[str, Any]) -> Any | None: """更新用户。""" - return await self._repository.async_update(user_id, payload) + return await self._write( + lambda: self._repository.async_update(user_id, payload) + ) async def delete(self, user_id: int) -> None: """删除用户。""" - await self._repository.async_delete(user_id) + await self._write(lambda: self._repository.async_delete(user_id)) async def update_otp(self, name: str, otp: bool, secret: str) -> None: """更新用户 OTP 状态。""" - await self._repository.async_update_otp_by_name(name, otp, secret) + await self._write( + lambda: self._repository.async_update_otp_by_name(name, otp, secret) + ) + + async def _write(self, operation: Callable[[], Awaitable[Any]]) -> Any: + """执行用户写入,并在正式请求路径统一提交或回滚。""" + try: + result = await operation() + if self._unit_of_work is not None: + await self._unit_of_work.commit() + return result + except Exception: + if self._unit_of_work is not None: + await self._unit_of_work.rollback() + raise _configured_user_id_lookup: Callable[[int], Any | None] | None = None diff --git a/app/db/base.py b/app/db/base.py index 1dd79ec3d..2952f64ea 100644 --- a/app/db/base.py +++ b/app/db/base.py @@ -4,7 +4,8 @@ ORM 基类与数据访问基类。 Base 提供声明式基类与通用的行为(字典转换、增删改查便利方法); DbOper 是各业务 Oper 的基类,持有一个可注入的会话。 """ -from typing import Any, List, Optional, Self, Union, cast +from collections.abc import Awaitable, Callable +from typing import Any, List, Optional, Self, TypeVar, Union, cast from sqlalchemy import (CursorResult, Executable, Identity, Integer, Sequence, and_, delete, inspect, select) @@ -13,6 +14,10 @@ from sqlalchemy.orm import DeclarativeBase, Mapped, Session, declared_attr, mapp from app.runtime.config import settings from app.db.decorators import async_db_query, async_db_update, db_query, db_update +from app.db.uow import run_async_transaction, run_sync_transaction + + +T = TypeVar("T") def execute_dml(db: Session, statement: Executable, @@ -147,4 +152,24 @@ 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: + return run_sync_transaction(operation) + if not isinstance(self._db, Session): + raise TypeError("同步写操作不能使用 AsyncSession") + return operation(self._db) + + async def _execute_async_write( + self, + operation: Callable[[AsyncSession], Awaitable[T]], + ) -> T: + """在当前异步会话暂存,或委托组合根创建兼容事务。""" + if self._db is None: + return await run_async_transaction(operation) + if not isinstance(self._db, AsyncSession): + raise TypeError("异步写操作不能使用同步 Session") + return await operation(self._db) diff --git a/app/db/maintenance.py b/app/db/maintenance.py index f6c23c5bd..248caac4a 100644 --- a/app/db/maintenance.py +++ b/app/db/maintenance.py @@ -7,6 +7,7 @@ from app.db.models.downloadhistory import DownloadFiles, DownloadHistory from app.db.models.message import Message from app.db.models.siteuserdata import SiteUserData from app.db.models.transferhistory import TransferHistory +from app.db.uow import SqlAlchemyUnitOfWork class DatabaseCleanupRepository: @@ -20,6 +21,11 @@ class DatabaseCleanupRepository: """创建一次维护运行共用的数据库会话。""" return self._session_factory() + @staticmethod + def unit_of_work(db: Any) -> SqlAlchemyUnitOfWork: + """把当前维护 Session 适配成显式批次事务边界。""" + return SqlAlchemyUnitOfWork(db) + @staticmethod def delete_messages(db: Any, cutoff: str, limit: int) -> int: """删除早于截止时间的消息。""" diff --git a/app/db/models/agenttask.py b/app/db/models/agenttask.py index 1aebab44e..0bc1d8f82 100644 --- a/app/db/models/agenttask.py +++ b/app/db/models/agenttask.py @@ -4,7 +4,7 @@ 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 db_query, db_update +from app.db.decorators import db_query class AgentTask(Base): @@ -49,7 +49,6 @@ class AgentTask(Base): ) @classmethod - @db_update def add_task(cls, db: Session, **kwargs: object) -> int: """ 新增 Agent 定时任务并返回任务 ID。 @@ -96,7 +95,6 @@ class AgentTask(Base): ).scalars().all()) @classmethod - @db_update def update_task( cls, db: Session, diff --git a/app/db/models/agenttaskrun.py b/app/db/models/agenttaskrun.py index 5fe8c6812..8c37607db 100644 --- a/app/db/models/agenttaskrun.py +++ b/app/db/models/agenttaskrun.py @@ -4,7 +4,7 @@ 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 db_query, db_update +from app.db.decorators import db_query from app.db.models.agenttask import AgentTask @@ -41,7 +41,6 @@ class AgentTaskRun(Base): ) @classmethod - @db_update def begin_run( cls, db: Session, @@ -107,7 +106,6 @@ class AgentTaskRun(Base): return run_id @classmethod - @db_update def finish_run( cls, db: Session, @@ -174,7 +172,6 @@ class AgentTaskRun(Base): return True @classmethod - @db_update def interrupt_task( cls, db: Session, @@ -226,7 +223,6 @@ class AgentTaskRun(Base): )) @classmethod - @db_update def delete_task_and_runs( cls, db: Session, diff --git a/app/db/models/downloadfailure.py b/app/db/models/downloadfailure.py index 3ab0d3a63..d4b8d48c9 100644 --- a/app/db/models/downloadfailure.py +++ b/app/db/models/downloadfailure.py @@ -4,7 +4,6 @@ from sqlalchemy import Float, Index, Integer, 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 db_update from app.db.models._constraints import media_identity_constraint @@ -115,7 +114,6 @@ class DownloadFailure(Base): return failure @classmethod - @db_update def delete_expired( cls, db: Session, diff --git a/app/db/models/downloadhistory.py b/app/db/models/downloadhistory.py index 4f0cab880..da875424e 100644 --- a/app/db/models/downloadhistory.py +++ b/app/db/models/downloadhistory.py @@ -6,7 +6,7 @@ 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 async_db_query, db_query, db_update +from app.db.decorators import async_db_query, db_query from app.db.models._constraints import media_identity_constraint from app.schemas.types import MediaSource @@ -295,7 +295,6 @@ class DownloadHistory(Base): ).scalars().all()) @classmethod - @db_update def delete_before( cls, db: Session, @@ -367,14 +366,12 @@ class DownloadFiles(Base): return list(db.execute(select(cls).where(cls.savepath == savepath)).scalars().all()) @classmethod - @db_update def delete_by_fullpath(cls, db: Session, fullpath: str): db.execute( update(cls).where(cls.fullpath == fullpath, cls.state == 1).values(state=0) ) @classmethod - @db_update def delete_orphans( cls, db: Session, diff --git a/app/db/models/mediaserver.py b/app/db/models/mediaserver.py index cfbd93c56..6e85b47a5 100644 --- a/app/db/models/mediaserver.py +++ b/app/db/models/mediaserver.py @@ -7,7 +7,7 @@ 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 async_db_query, db_query, db_update +from app.db.decorators import async_db_query, db_query from app.db.models._constraints import media_identity_constraint from app.schemas.types import MediaSource @@ -65,16 +65,16 @@ class MediaServerItem(Base): ).scalars().first() @classmethod - @db_update def empty(cls, db: Session, server: Optional[str] = None): + """在调用方事务中暂存媒体服务器条目清空操作。""" statement = delete(cls) if server is not None: statement = statement.where(cls.server == server) db.execute(statement, execution_options={"synchronize_session": False}) @classmethod - @db_update def delete_stale(cls, db: Session, server: str, sync_time: str): + """在调用方事务中删除本轮同步未更新的条目。""" return execute_dml( db, delete(cls).where( @@ -85,8 +85,8 @@ class MediaServerItem(Base): ) @classmethod - @db_update def delete_excluded_servers(cls, db: Session, servers: List[str]): + """在调用方事务中删除不属于启用服务器的条目。""" statement = delete(cls) if servers: statement = statement.where( diff --git a/app/db/models/message.py b/app/db/models/message.py index 882efdb79..21eb97852 100644 --- a/app/db/models/message.py +++ b/app/db/models/message.py @@ -5,7 +5,7 @@ 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 async_db_query, db_query, db_update +from app.db.decorators import async_db_query, db_query class Message(Base): @@ -40,7 +40,6 @@ class Message(Base): Index('ix_message_reg_time_id', 'reg_time', 'id'), ) - @db_update def create_and_to_dict(self, db: Session) -> dict: """ 创建消息记录并返回写入后的字段字典。 @@ -134,7 +133,6 @@ class Message(Base): return list(result.scalars().all()) @classmethod - @db_update def delete_before( cls, db: Session, diff --git a/app/db/models/passkey.py b/app/db/models/passkey.py index 430107de8..23f8812e8 100644 --- a/app/db/models/passkey.py +++ b/app/db/models/passkey.py @@ -1,11 +1,11 @@ from typing import Optional -from sqlalchemy import Integer, String, Boolean, DateTime, Text, select, ForeignKey +from sqlalchemy import Integer, String, Boolean, DateTime, Text, select, ForeignKey, update from sqlalchemy.ext.asyncio import AsyncSession 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 db_query, db_update, async_db_query, async_db_update +from app.db.decorators import db_query, async_db_query class PassKey(Base): @@ -85,19 +85,17 @@ class PassKey(Base): return result.scalars().first() @classmethod - @db_update def delete_by_id(cls, db: Session, passkey_id: int, user_id: int): """删除指定用户的PassKey""" passkey = db.execute( select(cls).where(cls.id == passkey_id, cls.user_id == user_id) ).scalars().first() if passkey: - passkey.delete(db, passkey.id) + db.delete(passkey) return True return False @classmethod - @async_db_update async def async_delete_by_id(cls, db: AsyncSession, passkey_id: int, user_id: int): """异步删除指定用户的PassKey""" result = await db.execute( @@ -108,24 +106,22 @@ class PassKey(Base): ) passkey = result.scalars().first() if passkey: - await passkey.async_delete(db, passkey.id) + await db.delete(passkey) return True return False - @db_update def update_last_used(self, db: Session, sign_count: int): """更新最后使用时间和签名计数""" - self.update(db, { - 'last_used_at': datetime.now(), - 'sign_count': sign_count - }) + db.execute(update(type(self)).where(type(self).id == self.id).values( + last_used_at=datetime.now(), + sign_count=sign_count, + )) return True - @async_db_update async def async_update_last_used(self, db: AsyncSession, sign_count: int): """异步更新最后使用时间和签名计数""" - await self.async_update(db, { - 'last_used_at': datetime.now(), - 'sign_count': sign_count - }) + await db.execute(update(type(self)).where(type(self).id == self.id).values( + last_used_at=datetime.now(), + sign_count=sign_count, + )) return True diff --git a/app/db/models/plugindata.py b/app/db/models/plugindata.py index ca666abde..d626c91d2 100644 --- a/app/db/models/plugindata.py +++ b/app/db/models/plugindata.py @@ -4,7 +4,7 @@ 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 db_query, db_update, async_db_query +from app.db.decorators import db_query, async_db_query class PluginData(Base): @@ -49,13 +49,13 @@ class PluginData(Base): return result.scalar_one_or_none() @classmethod - @db_update def del_plugin_data_by_key(cls, db: Session, plugin_id: str, key: str): + """在调用方事务中暂存单个插件键删除。""" db.execute(delete(cls).where(cls.plugin_id == plugin_id, cls.key == key)) @classmethod - @db_update def del_plugin_data(cls, db: Session, plugin_id: str): + """在调用方事务中暂存插件全部数据删除。""" db.execute(delete(cls).where(cls.plugin_id == plugin_id)) @classmethod diff --git a/app/db/models/site.py b/app/db/models/site.py index bf2848bf5..7ed4e6cbb 100644 --- a/app/db/models/site.py +++ b/app/db/models/site.py @@ -6,7 +6,7 @@ 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 db_query, db_update, async_db_query, async_db_update +from app.db.decorators import db_query, async_db_query class Site(Base): @@ -102,11 +102,11 @@ class Site(Base): return list(db.execute(select(cls.domain).where(cls.id.in_(ids))).scalars().all()) @classmethod - @db_update def reset(cls, db: Session): + """在调用方持有的同步事务中暂存清空操作。""" db.execute(delete(cls)) @classmethod - @async_db_update async def async_reset(cls, db: AsyncSession): + """在调用方持有的异步事务中暂存清空操作。""" await db.execute(delete(cls)) diff --git a/app/db/models/sitestatistic.py b/app/db/models/sitestatistic.py index 2406ad9c8..69b645ac8 100644 --- a/app/db/models/sitestatistic.py +++ b/app/db/models/sitestatistic.py @@ -6,7 +6,7 @@ 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 db_query, db_update, async_db_query +from app.db.decorators import db_query, async_db_query class SiteStatistic(Base): @@ -41,6 +41,6 @@ class SiteStatistic(Base): return result.scalar_one_or_none() @classmethod - @db_update def reset(cls, db: Session): + """在调用方持有的事务中暂存统计表清空操作。""" db.execute(delete(cls)) diff --git a/app/db/models/siteuserdata.py b/app/db/models/siteuserdata.py index 455e39519..0823b594c 100644 --- a/app/db/models/siteuserdata.py +++ b/app/db/models/siteuserdata.py @@ -6,7 +6,7 @@ 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 async_db_query, db_query, db_update +from app.db.decorators import async_db_query, db_query class SiteUserData(Base): @@ -138,7 +138,6 @@ class SiteUserData(Base): return list(result.scalars().all()) @classmethod - @db_update def delete_before( cls, db: Session, diff --git a/app/db/models/systemconfig.py b/app/db/models/systemconfig.py index 24d6872ee..405ae468f 100644 --- a/app/db/models/systemconfig.py +++ b/app/db/models/systemconfig.py @@ -4,7 +4,7 @@ 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 db_query, db_update, async_db_query +from app.db.decorators import db_query, async_db_query class SystemConfig(Base): @@ -28,9 +28,9 @@ class SystemConfig(Base): result = await db.execute(select(cls).where(cls.key == key)) return result.scalar_one_or_none() - @db_update def delete_by_key(self, db: Session, key: str): + """在调用方持有的事务中暂存指定配置删除。""" systemconfig = self.get_by_key(db, key) if systemconfig: - systemconfig.delete(db, systemconfig.id) + db.delete(systemconfig) return True diff --git a/app/db/models/transferhistory.py b/app/db/models/transferhistory.py index 060ec657a..271b0b41a 100644 --- a/app/db/models/transferhistory.py +++ b/app/db/models/transferhistory.py @@ -8,7 +8,7 @@ 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 async_db_query, db_query, db_update +from app.db.decorators import async_db_query, db_query from app.db.models._constraints import media_identity_constraint from app.schemas.types import MUSIC_ENTITY_ALBUM, MUSIC_ENTITY_RECORDING, MediaSource, MediaType @@ -555,14 +555,13 @@ class TransferHistory(Base): )).scalars().first() @classmethod - @db_update def update_download_hash(cls, db: Session, historyid: Optional[int] = None, download_hash: Optional[str] = None): + """在调用方事务中暂存下载任务哈希更新。""" db.execute( update(cls).where(cls.id == historyid).values(download_hash=download_hash) ) @classmethod - @db_update def replace_by_src(cls, db: Session, **kwargs) -> "TransferHistory": """ 用同源存储的新记录原子替换旧整理历史。 @@ -600,7 +599,6 @@ class TransferHistory(Base): ).scalars().all()) @classmethod - @db_update def delete_before( cls, db: Session, diff --git a/app/db/models/transferpending.py b/app/db/models/transferpending.py index 1dbdb1914..993360f01 100644 --- a/app/db/models/transferpending.py +++ b/app/db/models/transferpending.py @@ -4,7 +4,7 @@ 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 db_query, db_update +from app.db.decorators import db_query class TransferPending(Base): @@ -35,7 +35,6 @@ class TransferPending(Base): ) @classmethod - @db_update def register(cls, db: Session, storage: str, src_path: str, now_time: str) -> Optional["TransferPending"]: """ @@ -58,7 +57,6 @@ class TransferPending(Base): return pending @classmethod - @db_update def discard(cls, db: Session, storage: str, src_path: str) -> int: """ 注销一个待整理文件登记,整理到达终态(成功或失败)时调用。 @@ -93,7 +91,6 @@ class TransferPending(Base): ).scalars().all()) @classmethod - @db_update def clear(cls, db: Session) -> int: """ 清空全部待整理登记。 diff --git a/app/db/models/user.py b/app/db/models/user.py index 8f9fc0ae5..22fd6fca9 100644 --- a/app/db/models/user.py +++ b/app/db/models/user.py @@ -4,7 +4,7 @@ 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 db_query, db_update, async_db_query, async_db_update +from app.db.decorators import db_query, async_db_query class User(Base): @@ -60,56 +60,46 @@ class User(Base): ) return result.scalars().first() - @db_update def delete_by_name(self, db: Session, name: str): user = self.get_by_name(db, name) if user: - user.delete(db, user.id) + db.delete(user) return True - @async_db_update async def async_delete_by_name(self, db: AsyncSession, name: str): user = await self.async_get_by_name(db, name) if user: - await user.async_delete(db, user.id) + await db.delete(user) return True - @db_update def delete_by_id(self, db: Session, user_id: int): user = self.get_by_id(db, user_id) if user: - user.delete(db, user.id) + db.delete(user) return True @classmethod - @async_db_update async def async_delete_by_id(cls, db: AsyncSession, user_id: int): """异步按用户 ID 删除用户,供 UserOper 通过类方法调用。""" user = await cls.async_get_by_id(db, user_id) if user: - await user.async_delete(db, user.id) + await db.delete(user) return True - @db_update def update_otp_by_name(self, db: Session, name: str, otp: bool, secret: str): user = self.get_by_name(db, name) if user: - user.update(db, { - 'is_otp': otp, - 'otp_secret': secret - }) + user.is_otp = otp + user.otp_secret = secret return True return False @classmethod - @async_db_update async def async_update_otp_by_name(cls, db: AsyncSession, name: str, otp: bool, secret: str): """异步按用户名更新 OTP 状态,供 UserOper 通过类方法调用。""" user = await cls.async_get_by_name(db, name) if user: - await user.async_update(db, { - 'is_otp': otp, - 'otp_secret': secret - }) + user.is_otp = otp + user.otp_secret = secret return True return False diff --git a/app/db/models/userconfig.py b/app/db/models/userconfig.py index 59cc09936..01c5129de 100644 --- a/app/db/models/userconfig.py +++ b/app/db/models/userconfig.py @@ -3,7 +3,7 @@ from sqlalchemy import String, UniqueConstraint, JSON, select from sqlalchemy.orm import Mapped, Session, mapped_column from app.db.base import get_id_column, Base -from app.db.decorators import db_query, db_update +from app.db.decorators import db_query class UserConfig(Base): @@ -30,9 +30,9 @@ class UserConfig(Base): select(cls).where(cls.username == username, cls.key == key) ).scalars().first() - @db_update def delete_by_key(self, db: Session, username: str, key: str): + """在调用方持有的事务中暂存指定用户配置删除。""" userconfig = self.get_by_key(db=db, username=username, key=key) if userconfig: - userconfig.delete(db=db, rid=userconfig.id) + db.delete(userconfig) return True diff --git a/app/db/models/workflow.py b/app/db/models/workflow.py index 90561893e..00473f509 100644 --- a/app/db/models/workflow.py +++ b/app/db/models/workflow.py @@ -7,7 +7,7 @@ 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 db_query, async_db_query, async_db_update +from app.db.decorators import db_query, async_db_query class Workflow(Base): @@ -135,8 +135,8 @@ class Workflow(Base): return True @classmethod - @async_db_update async def async_update_state(cls, db: AsyncSession, wid: int, state: str): + """在调用方持有的异步事务中暂存工作流状态。""" await db.execute(update(cls).where(cls.id == wid).values(state=state)) return True @@ -146,8 +146,8 @@ class Workflow(Base): return True @classmethod - @async_db_update async def async_start(cls, db: AsyncSession, wid: int): + """在调用方持有的异步事务中暂存运行中状态。""" await db.execute(update(cls).where(cls.id == wid).values(state='R')) return True @@ -163,8 +163,8 @@ class Workflow(Base): return True @classmethod - @async_db_update async def async_fail(cls, db: AsyncSession, wid: int, result: str): + """在调用方持有的异步事务中暂存失败结果。""" await db.execute(update(cls).where( and_(cls.id == wid, cls.state != "P") ).values( @@ -187,8 +187,8 @@ class Workflow(Base): return True @classmethod - @async_db_update async def async_success(cls, db: AsyncSession, wid: int, result: Optional[str] = None): + """在调用方持有的异步事务中暂存成功结果。""" await db.execute(update(cls).where( and_(cls.id == wid, cls.state != "P") ).values( @@ -212,8 +212,8 @@ class Workflow(Base): return True @classmethod - @async_db_update async def async_reset(cls, db: AsyncSession, wid: int, reset_count: Optional[bool] = False): + """在调用方持有的异步事务中暂存执行状态重置。""" await db.execute(update(cls).where(cls.id == wid).values( state='W', result=None, @@ -243,9 +243,9 @@ class Workflow(Base): return True @classmethod - @async_db_update async def async_update_current_action(cls, db: AsyncSession, wid: int, action_id: str, context: dict, execution_state: Optional[dict] = None): + """在调用方持有的异步事务中暂存动作进度。""" # 先获取当前current_action result = await db.execute(select(cls.current_action).where(cls.id == wid)) current_action = result.scalar() diff --git a/app/db/oper/agenttask.py b/app/db/oper/agenttask.py index a33fc4efb..4ce59eea7 100644 --- a/app/db/oper/agenttask.py +++ b/app/db/oper/agenttask.py @@ -24,14 +24,16 @@ class AgentTaskOper(DbOper): 新增 Agent 定时任务。 """ now = self._now() - task_id = AgentTask.add_task( - self._db, - **kwargs, - enabled=True, - last_status="waiting", - run_count=0, - created_at=now, - updated_at=now, + task_id = self._execute_sync_write( + lambda session: AgentTask.add_task( + session, + **kwargs, + enabled=True, + last_status="waiting", + run_count=0, + created_at=now, + updated_at=now, + ) ) return self.get(task_id) @@ -81,38 +83,50 @@ class AgentTaskOper(DbOper): if not normalized_payload: return False normalized_payload["updated_at"] = self._now() - return AgentTask.update_task( - self._db, - task_id=task_id, - payload=normalized_payload, - user_id=user_id, + return self._execute_sync_write( + lambda session: AgentTask.update_task( + session, + task_id=task_id, + payload=normalized_payload, + user_id=user_id, + ) ) def delete(self, task_id: int, user_id: Optional[str] = None) -> bool: """ 删除非运行中的 Agent 定时任务及其运行历史。 """ - return AgentTaskRun.delete_task_and_runs( - self._db, - task_id=task_id, - user_id=user_id, + return self._execute_sync_write( + lambda session: AgentTaskRun.delete_task_and_runs( + session, + task_id=task_id, + user_id=user_id, + ) ) def begin_run( self, task_id: int, trigger_source: str = "scheduled", + *, + run_id: Optional[str] = None, + started_at: Optional[str] = None, ) -> Optional[AgentTaskRun]: """ 原子创建一次运行并返回其任务快照。 + + 可选运行 ID 和开始时间用于恢复/幂等验证;正常调度入口由本方法生成。 """ - run_id = uuid4().hex - created_run_id = AgentTaskRun.begin_run( - self._db, - task_id=task_id, - run_id=run_id, - trigger_source=trigger_source, - started_at=self._now(), + resolved_run_id = run_id or uuid4().hex + resolved_started_at = started_at or self._now() + created_run_id = self._execute_sync_write( + lambda session: AgentTaskRun.begin_run( + session, + task_id=task_id, + run_id=resolved_run_id, + trigger_source=trigger_source, + started_at=resolved_started_at, + ) ) return self.get_run(created_run_id) if created_run_id else None @@ -124,11 +138,15 @@ class AgentTaskOper(DbOper): """ 将遗留的运行中任务标记为中断且结果未知。 """ - return AgentTaskRun.interrupt_task( - self._db, - task_id=task_id, - result=(result or "")[:20000], - finished_at=self._now(), + finished_at = self._now() + normalized_result = (result or "")[:20000] + return self._execute_sync_write( + lambda session: AgentTaskRun.interrupt_task( + session, + task_id=task_id, + result=normalized_result, + finished_at=finished_at, + ) ) def get_run(self, run_id: str) -> Optional[AgentTaskRun]: @@ -157,13 +175,17 @@ class AgentTaskOper(DbOper): disable_date_task: bool = False, ) -> bool: """收口精确运行并更新仍匹配的任务投影。""" - return AgentTaskRun.finish_run( - self._db, - run_id=run_id, - success=success, - result=(result or "")[:20000], - finished_at=self._now(), - disable_date_task=disable_date_task, + finished_at = self._now() + normalized_result = (result or "")[:20000] + return self._execute_sync_write( + lambda session: AgentTaskRun.finish_run( + session, + run_id=run_id, + success=success, + result=normalized_result, + finished_at=finished_at, + disable_date_task=disable_date_task, + ) ) def finish( diff --git a/app/db/oper/downloadfailure.py b/app/db/oper/downloadfailure.py index e80b5d317..8798f18ee 100644 --- a/app/db/oper/downloadfailure.py +++ b/app/db/oper/downloadfailure.py @@ -54,8 +54,10 @@ class DownloadFailureOper(DbOper): """ 删除已过期较久的失败记录。 """ - return DownloadFailure.delete_expired( - self._db, - before_time=before_time, - limit=limit, + return self._execute_sync_write( + lambda session: DownloadFailure.delete_expired( + session, + before_time=before_time, + limit=limit, + ) ) diff --git a/app/db/oper/downloadhistory.py b/app/db/oper/downloadhistory.py index 59e307321..3213438c7 100644 --- a/app/db/oper/downloadhistory.py +++ b/app/db/oper/downloadhistory.py @@ -127,7 +127,9 @@ class DownloadHistoryOper(DbOper): 按fullpath删除下载文件记录 :param fullpath: 数据key """ - DownloadFiles.delete_by_fullpath(self._db, fullpath) + self._execute_sync_write( + lambda session: DownloadFiles.delete_by_fullpath(session, fullpath) + ) def stage_delete_file_by_fullpath(self, fullpath: str) -> None: """暂存指定完整路径的下载文件记录删除。""" diff --git a/app/db/oper/mediaserver.py b/app/db/oper/mediaserver.py index e6a08802e..64913ee29 100644 --- a/app/db/oper/mediaserver.py +++ b/app/db/oper/mediaserver.py @@ -61,19 +61,32 @@ class MediaServerOper(DbOper): """ 清空媒体服务器数据 """ - MediaServerItem.empty(self._db, server) + self._execute_sync_write( + lambda session: MediaServerItem.empty(session, server) + ) def delete_stale(self, server: str, sync_time: str) -> int: """ 删除本轮同步未更新的旧数据 """ - return MediaServerItem.delete_stale(self._db, server, sync_time) + return self._execute_sync_write( + lambda session: MediaServerItem.delete_stale( + session, + server, + sync_time, + ) + ) def delete_excluded_servers(self, servers: list[str]) -> int: """ 删除未启用或已移除媒体服务器的数据 """ - return MediaServerItem.delete_excluded_servers(self._db, servers) + return self._execute_sync_write( + lambda session: MediaServerItem.delete_excluded_servers( + session, + servers, + ) + ) def exists(self, **kwargs) -> Optional[MediaServerItem]: """ diff --git a/app/db/oper/message.py b/app/db/oper/message.py index f81a6d833..b298f11f4 100644 --- a/app/db/oper/message.py +++ b/app/db/oper/message.py @@ -62,7 +62,8 @@ class MessageOper(DbOper): if k not in Message.__table__.columns.keys(): # noqa kwargs.pop(k) - return Message(**kwargs).create_and_to_dict(self._db) + message = Message(**kwargs) + return self._execute_sync_write(message.create_and_to_dict) async def async_add(self, channel: Optional[NotificationChannel] = None, diff --git a/app/db/oper/passkey.py b/app/db/oper/passkey.py index 40aca00e9..a4f20ba19 100644 --- a/app/db/oper/passkey.py +++ b/app/db/oper/passkey.py @@ -24,13 +24,31 @@ class PassKeyOper(DbOper): def create(self, payload: dict[str, Any]) -> PassKey: """创建 PassKey 凭证。""" passkey = PassKey(**payload) - passkey.create(self._db) + self._execute_sync_write(lambda session: self._stage_create(session, passkey)) return passkey + @staticmethod + def _stage_create(session: Any, passkey: PassKey) -> None: + """在调用方事务中暂存凭证并分配主键。""" + session.add(passkey) + session.flush() + def update_last_used(self, passkey: PassKey, sign_count: int) -> bool: """更新凭证最后使用时间和签名计数。""" - return bool(passkey.update_last_used(self._db, sign_count)) + return bool(self._execute_sync_write( + lambda session: passkey.update_last_used(session, sign_count) + )) def delete_by_id(self, passkey_id: int, user_id: int) -> bool: """删除指定用户的凭证。""" - return bool(PassKey.delete_by_id(self._db, passkey_id, user_id)) + return bool(self._execute_sync_write( + lambda session: PassKey.delete_by_id(session, passkey_id, user_id) + )) + + async def async_delete_by_id(self, passkey_id: int, user_id: int) -> bool: + """在独立异步事务中删除指定用户的凭证。""" + return bool(await self._execute_async_write( + lambda session: PassKey.async_delete_by_id( + session, passkey_id, user_id + ) + )) diff --git a/app/db/oper/plugindata.py b/app/db/oper/plugindata.py index 09edfc07f..9dd76425b 100644 --- a/app/db/oper/plugindata.py +++ b/app/db/oper/plugindata.py @@ -80,10 +80,14 @@ class PluginDataOper(DbOper): :param plugin_id: 插件id :param key: 数据key """ - if key: - PluginData.del_plugin_data_by_key(self._db, plugin_id, key) - else: - PluginData.del_plugin_data(self._db, plugin_id) + def stage(session: Session) -> None: + """把兼容删除入口映射到调用方或组合根持有的事务。""" + if key: + PluginData.del_plugin_data_by_key(session, plugin_id, key) + else: + PluginData.del_plugin_data(session, plugin_id) + + self._execute_sync_write(stage) def stage_delete(self, plugin_id: str) -> None: """暂存目标插件全部数据删除并 flush,不提交调用方事务。""" diff --git a/app/db/oper/site.py b/app/db/oper/site.py index 14a85e5fa..6c8a7f152 100644 --- a/app/db/oper/site.py +++ b/app/db/oper/site.py @@ -116,8 +116,8 @@ class SiteOper(DbOper): Site.delete(self._db, sid) def reset(self) -> None: - """清空站点表,保留站点模型细节在数据库适配层。""" - Site.reset(self._db) + """清空站点表;兼容入口的事务由组合根统一持有。""" + self._execute_sync_write(Site.reset) async def stage_reset(self) -> None: """暂存清空站点表,由应用事务统一提交。""" diff --git a/app/db/oper/transferhistory.py b/app/db/oper/transferhistory.py index fc91231ba..aeeffadc9 100644 --- a/app/db/oper/transferhistory.py +++ b/app/db/oper/transferhistory.py @@ -264,14 +264,18 @@ class TransferHistoryOper(DbOper): kwargs.update({ "date": time.strftime("%Y-%m-%d %H:%M:%S", time.localtime()) }) - TransferHistory.replace_by_src(self._db, **kwargs) + def stage(session: Session) -> Optional[TransferHistory]: + """在同一事务替换记录并返回兼容查询投影。""" + TransferHistory.replace_by_src(session, **kwargs) + return TransferHistory.get_by_src( + session, + kwargs.get("src"), + kwargs["src_storage"], + ) + # 保持 add_force 的既有返回契约:返回可被调用方安全读取字段的查询结果, # 而非事务提交后可能已脱离会话的新建实例。 - return TransferHistory.get_by_src( - self._db, - kwargs.get("src"), - kwargs["src_storage"], - ) + return self._execute_sync_write(stage) def stage_replace_by_src(self, **kwargs) -> TransferHistory: """在调用方事务内按源路径替换整理历史并返回已分配 ID 的新记录。""" @@ -295,7 +299,13 @@ class TransferHistoryOper(DbOper): """ 补充转移记录download_hash """ - TransferHistory.update_download_hash(self._db, historyid, download_hash) + self._execute_sync_write( + lambda session: TransferHistory.update_download_hash( + session, + historyid, + download_hash, + ) + ) def list_by_date(self, date: str) -> List[TransferHistory]: """ diff --git a/app/db/oper/transferpending.py b/app/db/oper/transferpending.py index 6324f5687..e66db8f29 100644 --- a/app/db/oper/transferpending.py +++ b/app/db/oper/transferpending.py @@ -20,11 +20,14 @@ class TransferPendingOper(DbOper): :param src_path: 源文件路径 :return: 登记记录 """ - return TransferPending.register( - self._db, - storage=storage, - src_path=src_path, - now_time=datetime.now().strftime("%Y-%m-%d %H:%M:%S"), + now_time = datetime.now().strftime("%Y-%m-%d %H:%M:%S") + return self._execute_sync_write( + lambda session: TransferPending.register( + session, + storage=storage, + src_path=src_path, + now_time=now_time, + ) ) def discard(self, storage: str, src_path: str) -> int: @@ -34,7 +37,13 @@ class TransferPendingOper(DbOper): :param src_path: 源文件路径 :return: 删除的记录数 """ - return TransferPending.discard(self._db, storage=storage, src_path=src_path) + return self._execute_sync_write( + lambda session: TransferPending.discard( + session, + storage=storage, + src_path=src_path, + ) + ) def list_all(self, limit: Optional[int] = 5000) -> List[Tuple[str, str]]: """ @@ -56,4 +65,4 @@ class TransferPendingOper(DbOper): 清空全部待整理登记。 :return: 删除的记录数 """ - return TransferPending.clear(self._db) + return self._execute_sync_write(TransferPending.clear) diff --git a/app/db/oper/user.py b/app/db/oper/user.py index eea7b5833..3e226f295 100644 --- a/app/db/oper/user.py +++ b/app/db/oper/user.py @@ -11,6 +11,8 @@ runtime 兼容映射指向 SDK 薄门面;canonical 数据访问模块仍只依 """ from typing import List, Optional +from sqlalchemy.ext.asyncio import AsyncSession + from app.db.base import DbOper from app.db.models.user import User @@ -49,27 +51,53 @@ class UserOper(DbOper): async def async_create(self, payload: dict) -> Optional[User]: """异步创建用户。""" - return await User(**payload).async_create(self._db) + user = User(**payload) + + async def stage(session: AsyncSession) -> User: + """在当前异步事务中暂存用户并分配主键。""" + session.add(user) + await session.flush() + return user + + return await self._execute_async_write(stage) async def async_update(self, user_id: int, payload: dict) -> Optional[User]: """异步更新用户。""" user = await self.async_get_by_id(user_id) if user: - await user.async_update(self._db, payload) + async def stage(session: AsyncSession) -> User: + """在当前事务中更新用户字段,必要时重新附加游离对象。""" + for key, value in payload.items(): + setattr(user, key, value) + return await session.merge(user) + + await self._execute_async_write(stage) return user - async def async_delete(self, user_id: int) -> None: + async def async_delete(self, user_id: int) -> bool: """异步删除用户。""" - await User.async_delete_by_id(self._db, user_id) + return bool(await self._execute_async_write( + lambda session: User.async_delete_by_id(session, user_id) + )) + + async def async_delete_by_name(self, name: str) -> bool: + """在独立异步事务中按用户名删除用户。""" + return bool(await self._execute_async_write( + lambda session: User().async_delete_by_name(session, name) + )) async def async_update_otp_by_name( self, name: str, otp: bool, secret: str, - ) -> None: + ) -> bool: """异步更新用户 OTP 状态。""" - await User.async_update_otp_by_name(self._db, name, otp, secret) + return bool(await self._execute_async_write( + lambda session: User.async_update_otp_by_name( + session, name, otp, secret + ) + )) async def async_get_by_name(self, name: str) -> Optional[User]: """ diff --git a/app/db/uow.py b/app/db/uow.py index 52fbb6884..0e6ab270d 100644 --- a/app/db/uow.py +++ b/app/db/uow.py @@ -1,9 +1,65 @@ -"""SQLAlchemy 请求级事务适配器。""" +"""SQLAlchemy 请求级事务适配器与旧 Oper 事务执行端口。""" + +from collections.abc import Awaitable, Callable +from typing import Protocol, TypeVar from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Session +T = TypeVar("T") + + +class SyncTransactionRunner(Protocol): + """为无显式 Session 的兼容写入口提供独占同步事务。""" + + def __call__(self, operation: Callable[[Session], T]) -> T: + """在一个独占会话中执行并提交操作。""" + ... + + +class AsyncTransactionRunner(Protocol): + """为无显式 Session 的兼容写入口提供独占异步事务。""" + + def __call__( + self, + operation: Callable[[AsyncSession], Awaitable[T]], + ) -> Awaitable[T]: + """在一个独占异步会话中执行并提交操作。""" + ... + + +_sync_transaction_runner: SyncTransactionRunner | None = None +_async_transaction_runner: AsyncTransactionRunner | None = None + + +def configure_transaction_runners( + *, + sync: SyncTransactionRunner, + async_: AsyncTransactionRunner, +) -> None: + """由组合根登记旧 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: + """委托组合根在独占同步事务中执行兼容写操作。""" + if _sync_transaction_runner is None: + raise RuntimeError("同步事务执行器尚未配置") + return _sync_transaction_runner(operation) + + +async def run_async_transaction( + operation: Callable[[AsyncSession], Awaitable[T]], +) -> T: + """委托组合根在独占异步事务中执行兼容写操作。""" + if _async_transaction_runner is None: + raise RuntimeError("异步事务执行器尚未配置") + return await _async_transaction_runner(operation) + + class SqlAlchemyUnitOfWork: """把同步 Session 的提交与回滚能力适配为应用层事务端口。""" diff --git a/app/startup/context.py b/app/startup/context.py index 178b2a1f7..c9e8311ad 100644 --- a/app/startup/context.py +++ b/app/startup/context.py @@ -81,22 +81,27 @@ class SyncSessionProvider(Protocol): ... -class CompatibilityApiData(Protocol): - """未迁移 API 领域继续使用的结构化兼容 Facade。""" +class RepositoryFactory(Protocol): + """由请求 Session 构造某一明确领域仓储的通用工厂。""" - sync_session: SyncSessionProvider - async_session: AsyncSessionProvider - - def repository(self, name: str, session: object) -> object: - """按旧能力名构造请求级仓储。""" + def __call__(self, session: object) -> object: + """绑定请求会话并返回领域仓储。""" ... - def standalone_repository(self, name: str) -> object: - """按旧能力名构造独立仓储。""" + +class StandaloneRepositoryFactory(Protocol): + """构造自持有兼容事务边界的领域仓储。""" + + def __call__(self) -> object: + """返回无需请求 Session 的领域仓储。""" ... - def transaction(self, name: str, session: object) -> object: - """按旧能力名构造事务端口。""" + +class SyncUnitOfWorkFactory(Protocol): + """由同步请求 Session 构造事务端口的工厂。""" + + def __call__(self, session: object) -> object: + """绑定请求会话并返回同步事务端口。""" ... @@ -109,6 +114,57 @@ class AgentChatRuntime: transaction: AsyncUnitOfWorkFactory +@dataclass(frozen=True, slots=True) +class PersistenceRuntime: + """全部 HTTP 业务领域共享的请求会话与事务工厂。""" + + sync_session: SyncSessionProvider + async_session: AsyncSessionProvider + sync_transaction: SyncUnitOfWorkFactory + async_transaction: AsyncUnitOfWorkFactory + + +@dataclass(frozen=True, slots=True) +class AuthenticationRuntime: + """认证、用户管理与 PassKey API 的显式数据工厂。""" + + user_repository: RepositoryFactory + standalone_user: StandaloneRepositoryFactory + system_config: StandaloneRepositoryFactory + passkey: StandaloneRepositoryFactory + + +@dataclass(frozen=True, slots=True) +class MessagingRuntime: + """消息历史 API 的显式仓储工厂。""" + + repository: RepositoryFactory + + +@dataclass(frozen=True, slots=True) +class HistoryRuntime: + """下载、整理、媒体服务器与 Dashboard 领域的数据工厂。""" + + download_repository: RepositoryFactory + transfer_repository: RepositoryFactory + media_server_repository: RepositoryFactory + + +@dataclass(frozen=True, slots=True) +class SiteRuntime: + """站点读写领域的显式仓储工厂。""" + + repository: RepositoryFactory + + +@dataclass(frozen=True, slots=True) +class WorkflowRuntime: + """工作流定义、状态与缓存操作所需的数据工厂。""" + + repository: RepositoryFactory + system_config: StandaloneRepositoryFactory + + @dataclass(frozen=True, slots=True) class SubscriptionRuntime: """订阅 API 可见的请求级写事务运行时。""" @@ -125,6 +181,11 @@ class HostRuntime: """宿主组合根构建且在一个 FastAPI lifespan 内共享的运行时对象。""" agent_chat: AgentChatRuntime + persistence: PersistenceRuntime + authentication: AuthenticationRuntime + messaging: MessagingRuntime + history: HistoryRuntime + site: SiteRuntime subscription: SubscriptionRuntime + workflow: WorkflowRuntime configuration: RuntimeConfiguration - compatibility_api_data: CompatibilityApiData diff --git a/app/startup/modules_initializer.py b/app/startup/modules_initializer.py index bc80f6593..74b3b7daf 100644 --- a/app/startup/modules_initializer.py +++ b/app/startup/modules_initializer.py @@ -23,6 +23,7 @@ from app.runtime.extensions.module_manager import ModuleManager from app.runtime.extensions.module.dispatcher import ModuleInvocationDispatcher from app.runtime.extensions.plugin_manager import PluginManager from app.runtime.events import EventHandlerBinding, EventManager +from app.runtime.observability import record_metric from app.runtime.state import SystemHelper from app.runtime.thread import ThreadHelper from app.adapters.network.doh import DohHelper @@ -80,7 +81,11 @@ from app.db.session import ( get_async_db, get_db, ) -from app.db.uow import SqlAlchemyAsyncUnitOfWork, SqlAlchemyUnitOfWork +from app.db.uow import ( + SqlAlchemyAsyncUnitOfWork, + SqlAlchemyUnitOfWork, + configure_transaction_runners, +) from app.db.oper.subscribe import SubscribeOper from app.db.oper.agentchat import AgentChatOper from app.db.oper.agenttask import AgentTaskOper @@ -114,7 +119,18 @@ from app.startup.subscription import ( from app.startup.chain_events import TransactionalChainDurableEventWriter from app.startup.download_failure import TransactionalDownloadFailureRepository from app.startup.workflow import TransactionalWorkflowExecutionService -from app.startup.context import AgentChatRuntime, HostRuntime, SubscriptionRuntime +from app.startup.transaction import TransactionalWriteRunner +from app.startup.context import ( + AgentChatRuntime, + AuthenticationRuntime, + HistoryRuntime, + HostRuntime, + MessagingRuntime, + PersistenceRuntime, + SiteRuntime, + SubscriptionRuntime, + WorkflowRuntime, +) from app.adapters.web.security.access import set_superuser_token_payload_provider from app.application.security.auth import build_superuser_token_payload from app.application.image import configure_wallpaper_providers @@ -320,6 +336,10 @@ def _build_outbox_dispatcher() -> OutboxDispatcher: ), }, close=session.close, + failure_observer=lambda dead: record_metric( + "scheduler.job.dead_letter" if dead else "scheduler.job.retry", + owner="outbox", + ), ) @@ -535,6 +555,15 @@ async def init_modules() -> HostRuntime: """ 启动模块并返回本次 lifespan 唯一的类型化 HostRuntime。 """ + # 兼容 Oper 的无 Session 写入口仍由组合根持有事务,避免模型恢复自动提交。 + transaction_runner = TransactionalWriteRunner( + sync_session=SessionFactory, + async_session=async_session_scope, + ) + configure_transaction_runners( + sync=transaction_runner.sync, + async_=transaction_runner.async_, + ) # 数据访问能力统一在启动组合根注入,Runtime 和 Adapter 不再直接依赖 Oper。 api_data = ApiDataPorts( sync_session=get_db, @@ -572,6 +601,25 @@ async def init_modules() -> HostRuntime: repository=AgentChatOper, transaction=SqlAlchemyAsyncUnitOfWork, ), + persistence=PersistenceRuntime( + sync_session=get_db, + async_session=get_async_db, + sync_transaction=SqlAlchemyUnitOfWork, + async_transaction=SqlAlchemyAsyncUnitOfWork, + ), + authentication=AuthenticationRuntime( + user_repository=UserOper, + standalone_user=UserOper, + system_config=SystemConfigOper, + passkey=PassKeyOper, + ), + messaging=MessagingRuntime(repository=MessageOper), + history=HistoryRuntime( + download_repository=DownloadHistoryOper, + transfer_repository=TransferHistoryOper, + media_server_repository=MediaServerOper, + ), + site=SiteRuntime(repository=SiteOper), subscription=SubscriptionRuntime( async_session=get_async_db, repository=SubscribeOper, @@ -579,11 +627,15 @@ async def init_modules() -> HostRuntime: transaction=SqlAlchemyAsyncUnitOfWork, outbox=SqlAlchemyAsyncOutboxStager, ), + workflow=WorkflowRuntime( + repository=WorkflowOper, + system_config=SystemConfigOper, + ), configuration=runtime_configuration, - compatibility_api_data=api_data, ) configure_runtime_configuration(host_runtime.configuration) - configure_api_data_runtime(host_runtime.compatibility_api_data) + # 旧 app.api.data 导入只保留 ABI 转发,正式 API 依赖全部读取 HostRuntime。 + configure_api_data_runtime(api_data) configure_runtime_data_providers() workflow_execution = TransactionalWorkflowExecutionService(SessionFactory) configure_workflow_legacy_writer(workflow_execution) diff --git a/app/startup/transaction.py b/app/startup/transaction.py new file mode 100644 index 000000000..a24156f97 --- /dev/null +++ b/app/startup/transaction.py @@ -0,0 +1,61 @@ +"""旧 Oper 写入口的 SQLAlchemy 事务执行适配器。""" + +from collections.abc import Awaitable, Callable +from contextlib import AbstractAsyncContextManager +from typing import TypeVar + +from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy.orm import Session + +from app.db.uow import SqlAlchemyAsyncUnitOfWork, SqlAlchemyUnitOfWork + + +T = TypeVar("T") + + +class TransactionalWriteRunner: + """为兼容写入口创建独占会话,并用 UoW 明确提交或回滚。""" + + def __init__( + self, + *, + sync_session: Callable[[], Session], + async_session: Callable[[], AbstractAsyncContextManager[AsyncSession]], + ) -> None: + """保存同步会话工厂和异步会话上下文工厂。""" + self._sync_session = sync_session + self._async_session = async_session + + def sync(self, operation: Callable[[Session], T]) -> T: + """在独占同步 Session 中执行操作并统一收口事务。""" + session = self._sync_session() + # 兼容 Oper 历史上会返回刚写入的 ORM 对象;提交后若过期,Session 关闭后连主键 + # 都无法读取。独占短会话没有后续一致性读取需求,因此保留已 flush 的字段快照。 + session.expire_on_commit = False + unit_of_work = SqlAlchemyUnitOfWork(session) + try: + result = operation(session) + unit_of_work.commit() + return result + except Exception: + unit_of_work.rollback() + raise + finally: + session.close() + + async def async_( + self, + operation: Callable[[AsyncSession], Awaitable[T]], + ) -> T: + """在独占 AsyncSession 中执行操作并统一收口事务。""" + async with self._async_session() as session: + # 与同步兼容入口保持相同的返回对象生命周期。 + session.sync_session.expire_on_commit = False + unit_of_work = SqlAlchemyAsyncUnitOfWork(session) + try: + result = await operation(session) + await unit_of_work.commit() + return result + except Exception: + await unit_of_work.rollback() + raise diff --git a/docs/architecture-overview.md b/docs/architecture-overview.md index 38343e105..d6d241b18 100644 --- a/docs/architecture-overview.md +++ b/docs/architecture-overview.md @@ -246,9 +246,9 @@ sequenceDiagram - **引擎预热 fail-fast**:同步/异步数据库引擎在单线程期完成首次创建, 避免调度器放出大量线程后再创建引擎导致连接锁竞争。 - **类型化请求装配**:`startup/context.py` 的 frozen slots `HostRuntime` 是 lifespan 内唯一宿主 - 上下文,`api/context.py` 从 `app.state` 收窄到具体领域能力。Agent 会话已迁移,不再通过 - 字符串仓储键定位;API、Scheduler、Chain 从 `HostRuntime.configuration` 获取 frozen 配置快照, - `ApiDataPorts` 暂作未迁移领域的同实例兼容 Facade。 + 上下文,`api/context.py` 从 `app.state` 收窄到具体领域能力。认证、消息、历史、媒体服务器、站点、 + 订阅、工作流和请求事务均使用命名 runtime 字段,不再通过字符串仓储键定位;API、Scheduler、Chain + 从 `HostRuntime.configuration` 获取 frozen 配置快照。`ApiDataPorts` 仅保留旧导入 ABI,不参与正式请求链路。 - **安全模式**:`MOVIEPILOT_SAFE_MODE` 会跳过插件、定时器、监控器、命令与工作流,用于故障自救。 - **进程拓扑**:全功能 V3 强制 `API_WORKERS=1`,避免每个 worker 重复启动插件和后台控制面;安全模式可临时使用多 worker 诊断,但不是正式扩容方案。 - **健康语义**:`/health/live` 只确认进程和事件循环可响应;`/health/ready` 仅在数据库 @@ -376,7 +376,8 @@ flowchart LR 成功后执行。订阅新增样板由 `startup/subscription.py` 创建独占 Session, `application/subscription/write.py` 决定事务与 post-commit 边界,`SubscribeOper.stage_add()` 只查重、`add` 和 `flush`。旧 SDK 显式构造的无会话 Oper 暂留兼容自动短会话,不得被新代码复用。 - `transaction-debt-baseline.json` 将存量 168 个 Model 事务装饰器冻结为只降不增低水位。 + `transaction-debt-baseline.json` 当前冻结 123 个只读查询装饰器;原有 45 个同步/异步写装饰器 + 已全部移除,`db_update` 与 `async_db_update` 必须持续保持为 0。 - 站点、历史、工作流、Agent 会话删除和插件数据重置已经形成同构事务切片;对应 Application Command/Service 持有 UoW,Oper 的 `stage_*` 方法只修改当前会话。插件数据重置从 `startup/plugins_initializer.py` 创建独占会话,插件直接使用 `PluginDataOper` 的旧 ABI 仅作兼容。 diff --git a/docs/refactor/backend-architecture-next-stage.md b/docs/refactor/backend-architecture-next-stage.md index 698b1b8e7..ee17870ed 100644 --- a/docs/refactor/backend-architecture-next-stage.md +++ b/docs/refactor/backend-architecture-next-stage.md @@ -501,6 +501,15 @@ app/api/dependencies/ # 按领域拆分依赖工厂 - fake Runtime 请求测试证明仓储与 UoW 共享同一请求会话,且无需加载真实 DB engine、 PluginManager 或其他运行时服务;旧 `configure_api_data_ports()` 调用形态继续可用。 +**收口记录(2026-08-22)**: + +- `HostRuntime` 已覆盖认证/用户/PassKey、消息、下载与整理历史、媒体服务器、站点、订阅、 + 工作流、请求 Session/UoW 和配置快照等全部正式 API 业务领域。每个能力均为命名字段, + 不再由 `repository("name")` 或 `transaction("name")` 在运行时猜测。 +- `app/api/dependencies/` 的正式领域模块已清除 `app.api.data` 与 + `app.api.dependencies.data` 依赖,并增加静态测试防止回退。旧 `ApiDataPorts` 只作为旧导入 + ABI 的全局转发保留,不再挂入 `HostRuntime`,也不参与正式 FastAPI 请求装配。 + #### ARCH-231:按领域拆分 API dependency 与 presentation **目标**:`app/api/deps.py` 从 512 行集中装配点变成兼容聚合入口,端点只负责 HTTP 解析、鉴权依赖和结果映射。 @@ -786,6 +795,11 @@ ADR 必须逐个映射当前 Event、BackgroundTasks、Scheduler job、Agent tas 失败和重置均由 Application command 显式 commit/rollback。`WorkflowOper()` 的旧方法名、参数和返回值 继续可用,无 Session 调用委托组合根服务,显式 Session 调用只暂存;同步 Model 自动提交装饰器移除 6 个, 事务低水位从 174 降到 168,Oper 仍不创建 Session、也不直接 commit/rollback。 +- 剩余 45 个同步/异步 Model 写装饰器已全部迁移:AgentTask、PassKey、User、消息、历史清理、 + 站点快照、媒体服务器、插件数据、TransferPending 等写入由调用方 Session 和 UoW 收口;无 Session + 的旧 Oper ABI 委托 Startup 注入的短事务执行器。当前 Model 装饰器仅剩 123 个查询装饰器, + `db_update` 与 `async_db_update` 均为 0,Oper 自建 Session/直接提交仍为 0。 +- 数据清理按批次显式提交 UoW,单表失败先回滚会话再继续汇总后续表;不再依赖删除 Model 的隐式提交。 **禁止**:本阶段不引入 Celery、Kafka、RabbitMQ 等新基础设施。 @@ -868,8 +882,10 @@ OTel 初始化只能位于 Startup/Adapter;Domain/Application 只依赖 no-op- request path 充当 label。 - 2026-08-22 扩展接线覆盖 SQLAlchemy checkout/checkin、异步回退配额 wait/timeout、Module 真实 `TimeoutError`、插件 start/initialize/stop/reload,以及 Agent 活跃任务、取消结果、供应商耗时和输入/ - 输出 token。自定义 Agent provider 统一归类为 `custom`,不会暴露配置名称;Scheduler retry/dead-letter - 属于本轮明确暂停的 Outbox worker 范围,目录合同保留但不在本轮接线。 + 输出 token。自定义 Agent provider 统一归类为 `custom`,不会暴露配置名称。 +- Outbox dispatcher 的有限重试和 dead-letter 已分别接入 `scheduler.job.retry` 与 + `scheduler.job.dead_letter`,只使用固定 `owner=outbox` 低基数标签;观测失败端口由 Startup 注入, + Application 不依赖具体 OTel SDK。 - 专项测试覆盖 exporter 缺失、非法标签、全目录高基数审计、成功/失败 outcome、动态 URL 路由模板; 既有 API、Event、Module、Scheduler 与健康探针回归保持通过。 diff --git a/docs/rules/10-data-and-persistent.md b/docs/rules/10-data-and-persistent.md index 2a5a98684..a4d8caae2 100644 --- a/docs/rules/10-data-and-persistent.md +++ b/docs/rules/10-data-and-persistent.md @@ -84,9 +84,9 @@ Oper classes accept and return persistence values. Turning a `MediaInfo` or ### Transaction ownership ratchet - `tests/fixtures/architecture/transaction-debt-baseline.json` records the - existing Model transaction decorators. The current 168 legacy decorators are + existing Model transaction decorators. The current 123 decorators are query-only migration debt: they may decrease but must never increase or move to a new - Model method. + Model method. Both `db_update` and `async_db_update` must remain at zero. - New Model methods must not use `db_query`, `db_update`, `async_db_query`, or `async_db_update`, create a Session, or call `commit()` / `rollback()`. - Oper receives a caller-owned Session and may query, add, update, delete, or diff --git a/tests/conftest.py b/tests/conftest.py index a90ef0655..d95871dbb 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -6,6 +6,7 @@ import sys import pytest +from sqlalchemy.orm import Session # 必须早于首个牵入 app.runtime.config 的 import(app.db / app.chain.* 都会牵入):引擎本身已惰性, # import app.db 不再连库,但 settings 在 import 期就把 CONFIG_DIR 读进字段并建好配置目录,之后 @@ -47,7 +48,11 @@ def configure_plugin_system_services(): get_async_db, get_db, ) - from app.db.uow import SqlAlchemyAsyncUnitOfWork, SqlAlchemyUnitOfWork + from app.db.uow import ( + SqlAlchemyAsyncUnitOfWork, + SqlAlchemyUnitOfWork, + configure_transaction_runners, + ) from app.db.oper.systemconfig import SystemConfigOper configure_token_codec(create_access_token, decode_access_token) @@ -113,6 +118,22 @@ def configure_plugin_system_services(): from app.db.oper.passkey import PassKeyOper from app.startup.subscription import TransactionalSubscribeWriter from app.startup.workflow import TransactionalWorkflowExecutionService + from app.startup.transaction import TransactionalWriteRunner + + def compatibility_sync_session() -> Session: + """动态读取可被存量隔离数据库用例替换的 ScopedSession。""" + from app.db import decorators + + return decorators.ScopedSession() + + transaction_runner = TransactionalWriteRunner( + sync_session=compatibility_sync_session, + async_session=async_session_scope, + ) + configure_transaction_runners( + sync=transaction_runner.sync, + async_=transaction_runner.async_, + ) configure_workflow_legacy_writer( TransactionalWorkflowExecutionService(SessionFactory) diff --git a/tests/fixtures/architecture/dependency-baseline.json b/tests/fixtures/architecture/dependency-baseline.json index b5d7733e4..b5f23e0e7 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": 6376, - "edge_sha256": "11152e1c89a0d5f07d5460e8a21a8399b9280b8e3027f66621d839817ac3db21", + "edge_count": 6379, + "edge_sha256": "1c619d72157004590838a85c497330b8ca462d3bc1eb490fd5e7bac97e6190b8", "edges": [ "app -> app.runtime", "app -> app.runtime.compat", @@ -1506,21 +1506,18 @@ "app.api.context -> app.startup.context", "app.api.dependencies.agent -> app.api", "app.api.dependencies.agent -> app.api.context", - "app.api.dependencies.agent -> app.api.data", - "app.api.dependencies.agent -> app.api.dependencies", - "app.api.dependencies.agent -> app.api.dependencies.data", "app.api.dependencies.agent -> app.application", "app.api.dependencies.agent -> app.application.messaging", "app.api.dependencies.agent -> app.application.messaging.chat", "app.api.dependencies.agent -> app.application.messaging.message", + "app.api.dependencies.agent -> app.startup", + "app.api.dependencies.agent -> app.startup.context", "app.api.dependencies.auth -> app.adapters", "app.api.dependencies.auth -> app.adapters.web", "app.api.dependencies.auth -> app.adapters.web.security", "app.api.dependencies.auth -> app.adapters.web.security.access", "app.api.dependencies.auth -> app.api", - "app.api.dependencies.auth -> app.api.data", - "app.api.dependencies.auth -> app.api.dependencies", - "app.api.dependencies.auth -> app.api.dependencies.data", + "app.api.dependencies.auth -> app.api.context", "app.api.dependencies.auth -> app.application", "app.api.dependencies.auth -> app.application.security", "app.api.dependencies.auth -> app.application.security.auth", @@ -1528,12 +1525,12 @@ "app.api.dependencies.auth -> app.application.security.user", "app.api.dependencies.auth -> app.schemas", "app.api.dependencies.auth -> app.schemas.token", + "app.api.dependencies.auth -> app.startup", + "app.api.dependencies.auth -> app.startup.context", "app.api.dependencies.data -> app.api", "app.api.dependencies.data -> app.api.data", "app.api.dependencies.history -> app.api", - "app.api.dependencies.history -> app.api.data", - "app.api.dependencies.history -> app.api.dependencies", - "app.api.dependencies.history -> app.api.dependencies.data", + "app.api.dependencies.history -> app.api.context", "app.api.dependencies.history -> app.application", "app.api.dependencies.history -> app.application.dashboard", "app.api.dependencies.history -> app.application.history", @@ -1546,6 +1543,8 @@ "app.api.dependencies.history -> app.schemas", "app.api.dependencies.history -> app.schemas.types", "app.api.dependencies.history -> app.schemas.workflow", + "app.api.dependencies.history -> app.startup", + "app.api.dependencies.history -> app.startup.context", "app.api.dependencies.plugin -> app.application", "app.api.dependencies.plugin -> app.application.commands", "app.api.dependencies.plugin -> app.application.plugin", @@ -1559,9 +1558,7 @@ "app.api.dependencies.plugin -> app.schemas.event", "app.api.dependencies.plugin -> app.schemas.types", "app.api.dependencies.site -> app.api", - "app.api.dependencies.site -> app.api.data", - "app.api.dependencies.site -> app.api.dependencies", - "app.api.dependencies.site -> app.api.dependencies.data", + "app.api.dependencies.site -> app.api.context", "app.api.dependencies.site -> app.application", "app.api.dependencies.site -> app.application.site", "app.api.dependencies.site -> app.application.site.mutation", @@ -1574,14 +1571,13 @@ "app.api.dependencies.site -> app.runtime.events", "app.api.dependencies.site -> app.schemas", "app.api.dependencies.site -> app.schemas.types", + "app.api.dependencies.site -> app.startup", + "app.api.dependencies.site -> app.startup.context", "app.api.dependencies.subscription -> app.adapters", "app.api.dependencies.subscription -> app.adapters.external", "app.api.dependencies.subscription -> app.adapters.external.server", "app.api.dependencies.subscription -> app.api", "app.api.dependencies.subscription -> app.api.context", - "app.api.dependencies.subscription -> app.api.data", - "app.api.dependencies.subscription -> app.api.dependencies", - "app.api.dependencies.subscription -> app.api.dependencies.data", "app.api.dependencies.subscription -> app.application", "app.api.dependencies.subscription -> app.application.outbox", "app.api.dependencies.subscription -> app.application.scheduling", @@ -1597,18 +1593,20 @@ "app.api.dependencies.subscription -> app.runtime.log", "app.api.dependencies.subscription -> app.schemas", "app.api.dependencies.subscription -> app.schemas.types", + "app.api.dependencies.subscription -> app.startup", + "app.api.dependencies.subscription -> app.startup.context", "app.api.dependencies.workflow -> app.adapters", "app.api.dependencies.workflow -> app.adapters.external", "app.api.dependencies.workflow -> app.adapters.external.server", "app.api.dependencies.workflow -> app.api", - "app.api.dependencies.workflow -> app.api.data", - "app.api.dependencies.workflow -> app.api.dependencies", - "app.api.dependencies.workflow -> app.api.dependencies.data", + "app.api.dependencies.workflow -> app.api.context", "app.api.dependencies.workflow -> app.application", "app.api.dependencies.workflow -> app.application.scheduling", "app.api.dependencies.workflow -> app.application.workflow", "app.api.dependencies.workflow -> app.runtime", "app.api.dependencies.workflow -> app.runtime.config", + "app.api.dependencies.workflow -> app.startup", + "app.api.dependencies.workflow -> app.startup.context", "app.api.dependencies.workflow -> app.workflow", "app.api.deps -> app.api", "app.api.deps -> app.api.dependencies", @@ -3422,6 +3420,7 @@ "app.command -> app.schemas.types", "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", "app.db.decorators -> app.db", @@ -3445,6 +3444,7 @@ "app.db.maintenance -> app.db.models.message", "app.db.maintenance -> app.db.models.siteuserdata", "app.db.maintenance -> app.db.models.transferhistory", + "app.db.maintenance -> app.db.uow", "app.db.models -> app.db", "app.db.models -> app.db.models._identity", "app.db.models._identity -> app.runtime", @@ -3464,7 +3464,6 @@ "app.db.models.agenttaskrun -> app.db.models.agenttask", "app.db.models.downloadfailure -> app.db", "app.db.models.downloadfailure -> app.db.base", - "app.db.models.downloadfailure -> app.db.decorators", "app.db.models.downloadfailure -> app.db.models", "app.db.models.downloadfailure -> app.db.models._constraints", "app.db.models.downloadhistory -> app.db", @@ -6114,6 +6113,7 @@ "app.startup.modules_initializer -> app.runtime.extensions.plugin_manager", "app.startup.modules_initializer -> app.runtime.extensions.service_config", "app.startup.modules_initializer -> app.runtime.log", + "app.startup.modules_initializer -> app.runtime.observability", "app.startup.modules_initializer -> app.runtime.state", "app.startup.modules_initializer -> app.runtime.thread", "app.startup.modules_initializer -> app.scheduler", @@ -6129,6 +6129,7 @@ "app.startup.modules_initializer -> app.startup.managed_resources_initializer", "app.startup.modules_initializer -> app.startup.outbox", "app.startup.modules_initializer -> app.startup.subscription", + "app.startup.modules_initializer -> app.startup.transaction", "app.startup.modules_initializer -> app.startup.workflow", "app.startup.monitor_initializer -> app.monitor", "app.startup.outbox -> app.application", @@ -6208,6 +6209,8 @@ "app.startup.subscription -> app.schemas.types", "app.startup.subscription -> app.startup", "app.startup.subscription -> app.startup.outbox", + "app.startup.transaction -> app.db", + "app.startup.transaction -> app.db.uow", "app.startup.transfer_initializer -> app.chain", "app.startup.transfer_initializer -> app.chain.transfer", "app.startup.workflow -> app.application", @@ -6393,7 +6396,7 @@ "app.workflow.actions.transfer_file -> app.workflow", "app.workflow.actions.transfer_file -> app.workflow.actions" ], - "module_count": 790, + "module_count": 791, "modules": [ "app", "app.adapters", @@ -7161,6 +7164,7 @@ "app.startup.routers_initializer", "app.startup.scheduler_initializer", "app.startup.subscription", + "app.startup.transaction", "app.startup.transfer_initializer", "app.startup.workflow", "app.startup.workflow_initializer", diff --git a/tests/fixtures/architecture/transaction-debt-baseline.json b/tests/fixtures/architecture/transaction-debt-baseline.json index b95014416..71c7ac16d 100644 --- a/tests/fixtures/architecture/transaction-debt-baseline.json +++ b/tests/fixtures/architecture/transaction-debt-baseline.json @@ -2,11 +2,11 @@ "model_decorators": { "by_kind": { "async_db_query": 49, - "async_db_update": 12, + "async_db_update": 0, "db_query": 74, - "db_update": 33 + "db_update": 0 }, - "count": 168, + "count": 123, "methods": [ { "decorator": "async_db_query", @@ -28,11 +28,6 @@ "file": "app/db/models/agentchat.py", "method": "AgentChat.list_by_page" }, - { - "decorator": "db_update", - "file": "app/db/models/agenttask.py", - "method": "AgentTask.add_task" - }, { "decorator": "db_query", "file": "app/db/models/agenttask.py", @@ -43,56 +38,16 @@ "file": "app/db/models/agenttask.py", "method": "AgentTask.list_for_user" }, - { - "decorator": "db_update", - "file": "app/db/models/agenttask.py", - "method": "AgentTask.update_task" - }, - { - "decorator": "db_update", - "file": "app/db/models/agenttaskrun.py", - "method": "AgentTaskRun.begin_run" - }, - { - "decorator": "db_update", - "file": "app/db/models/agenttaskrun.py", - "method": "AgentTaskRun.delete_task_and_runs" - }, - { - "decorator": "db_update", - "file": "app/db/models/agenttaskrun.py", - "method": "AgentTaskRun.finish_run" - }, { "decorator": "db_query", "file": "app/db/models/agenttaskrun.py", "method": "AgentTaskRun.get_by_run_id" }, - { - "decorator": "db_update", - "file": "app/db/models/agenttaskrun.py", - "method": "AgentTaskRun.interrupt_task" - }, { "decorator": "db_query", "file": "app/db/models/agenttaskrun.py", "method": "AgentTaskRun.list_for_task" }, - { - "decorator": "db_update", - "file": "app/db/models/downloadfailure.py", - "method": "DownloadFailure.delete_expired" - }, - { - "decorator": "db_update", - "file": "app/db/models/downloadhistory.py", - "method": "DownloadFiles.delete_by_fullpath" - }, - { - "decorator": "db_update", - "file": "app/db/models/downloadhistory.py", - "method": "DownloadFiles.delete_orphans" - }, { "decorator": "db_query", "file": "app/db/models/downloadhistory.py", @@ -128,11 +83,6 @@ "file": "app/db/models/downloadhistory.py", "method": "DownloadHistory.async_list_by_title" }, - { - "decorator": "db_update", - "file": "app/db/models/downloadhistory.py", - "method": "DownloadHistory.delete_before" - }, { "decorator": "db_query", "file": "app/db/models/downloadhistory.py", @@ -193,21 +143,6 @@ "file": "app/db/models/mediaserver.py", "method": "MediaServerItem.async_get_by_itemid" }, - { - "decorator": "db_update", - "file": "app/db/models/mediaserver.py", - "method": "MediaServerItem.delete_excluded_servers" - }, - { - "decorator": "db_update", - "file": "app/db/models/mediaserver.py", - "method": "MediaServerItem.delete_stale" - }, - { - "decorator": "db_update", - "file": "app/db/models/mediaserver.py", - "method": "MediaServerItem.empty" - }, { "decorator": "db_query", "file": "app/db/models/mediaserver.py", @@ -238,16 +173,6 @@ "file": "app/db/models/message.py", "method": "Message.async_list_sent_by_page" }, - { - "decorator": "db_update", - "file": "app/db/models/message.py", - "method": "Message.create_and_to_dict" - }, - { - "decorator": "db_update", - "file": "app/db/models/message.py", - "method": "Message.delete_before" - }, { "decorator": "db_query", "file": "app/db/models/message.py", @@ -258,11 +183,6 @@ "file": "app/db/models/message.py", "method": "Message.list_by_page" }, - { - "decorator": "async_db_update", - "file": "app/db/models/passkey.py", - "method": "PassKey.async_delete_by_id" - }, { "decorator": "async_db_query", "file": "app/db/models/passkey.py", @@ -278,16 +198,6 @@ "file": "app/db/models/passkey.py", "method": "PassKey.async_get_by_user_id" }, - { - "decorator": "async_db_update", - "file": "app/db/models/passkey.py", - "method": "PassKey.async_update_last_used" - }, - { - "decorator": "db_update", - "file": "app/db/models/passkey.py", - "method": "PassKey.delete_by_id" - }, { "decorator": "db_query", "file": "app/db/models/passkey.py", @@ -303,11 +213,6 @@ "file": "app/db/models/passkey.py", "method": "PassKey.get_by_user_id" }, - { - "decorator": "db_update", - "file": "app/db/models/passkey.py", - "method": "PassKey.update_last_used" - }, { "decorator": "async_db_query", "file": "app/db/models/plugindata.py", @@ -323,16 +228,6 @@ "file": "app/db/models/plugindata.py", "method": "PluginData.async_get_plugin_data_by_plugin_id" }, - { - "decorator": "db_update", - "file": "app/db/models/plugindata.py", - "method": "PluginData.del_plugin_data" - }, - { - "decorator": "db_update", - "file": "app/db/models/plugindata.py", - "method": "PluginData.del_plugin_data_by_key" - }, { "decorator": "db_query", "file": "app/db/models/plugindata.py", @@ -368,11 +263,6 @@ "file": "app/db/models/site.py", "method": "Site.async_list_order_by_pri" }, - { - "decorator": "async_db_update", - "file": "app/db/models/site.py", - "method": "Site.async_reset" - }, { "decorator": "db_query", "file": "app/db/models/site.py", @@ -393,11 +283,6 @@ "file": "app/db/models/site.py", "method": "Site.list_order_by_pri" }, - { - "decorator": "db_update", - "file": "app/db/models/site.py", - "method": "Site.reset" - }, { "decorator": "async_db_query", "file": "app/db/models/siteicon.py", @@ -418,11 +303,6 @@ "file": "app/db/models/sitestatistic.py", "method": "SiteStatistic.get_by_domain" }, - { - "decorator": "db_update", - "file": "app/db/models/sitestatistic.py", - "method": "SiteStatistic.reset" - }, { "decorator": "async_db_query", "file": "app/db/models/siteuserdata.py", @@ -433,11 +313,6 @@ "file": "app/db/models/siteuserdata.py", "method": "SiteUserData.async_get_latest" }, - { - "decorator": "db_update", - "file": "app/db/models/siteuserdata.py", - "method": "SiteUserData.delete_before" - }, { "decorator": "db_query", "file": "app/db/models/siteuserdata.py", @@ -568,11 +443,6 @@ "file": "app/db/models/systemconfig.py", "method": "SystemConfig.async_get_by_key" }, - { - "decorator": "db_update", - "file": "app/db/models/systemconfig.py", - "method": "SystemConfig.delete_by_key" - }, { "decorator": "db_query", "file": "app/db/models/systemconfig.py", @@ -613,11 +483,6 @@ "file": "app/db/models/transferhistory.py", "method": "TransferHistory.count_by_title" }, - { - "decorator": "db_update", - "file": "app/db/models/transferhistory.py", - "method": "TransferHistory.delete_before" - }, { "decorator": "db_query", "file": "app/db/models/transferhistory.py", @@ -683,51 +548,16 @@ "file": "app/db/models/transferhistory.py", "method": "TransferHistory.monthly_media_statistics" }, - { - "decorator": "db_update", - "file": "app/db/models/transferhistory.py", - "method": "TransferHistory.replace_by_src" - }, { "decorator": "db_query", "file": "app/db/models/transferhistory.py", "method": "TransferHistory.statistic" }, - { - "decorator": "db_update", - "file": "app/db/models/transferhistory.py", - "method": "TransferHistory.update_download_hash" - }, - { - "decorator": "db_update", - "file": "app/db/models/transferpending.py", - "method": "TransferPending.clear" - }, - { - "decorator": "db_update", - "file": "app/db/models/transferpending.py", - "method": "TransferPending.discard" - }, { "decorator": "db_query", "file": "app/db/models/transferpending.py", "method": "TransferPending.list_all" }, - { - "decorator": "db_update", - "file": "app/db/models/transferpending.py", - "method": "TransferPending.register" - }, - { - "decorator": "async_db_update", - "file": "app/db/models/user.py", - "method": "User.async_delete_by_id" - }, - { - "decorator": "async_db_update", - "file": "app/db/models/user.py", - "method": "User.async_delete_by_name" - }, { "decorator": "async_db_query", "file": "app/db/models/user.py", @@ -738,21 +568,6 @@ "file": "app/db/models/user.py", "method": "User.async_get_by_name" }, - { - "decorator": "async_db_update", - "file": "app/db/models/user.py", - "method": "User.async_update_otp_by_name" - }, - { - "decorator": "db_update", - "file": "app/db/models/user.py", - "method": "User.delete_by_id" - }, - { - "decorator": "db_update", - "file": "app/db/models/user.py", - "method": "User.delete_by_name" - }, { "decorator": "db_query", "file": "app/db/models/user.py", @@ -763,26 +578,11 @@ "file": "app/db/models/user.py", "method": "User.get_by_name" }, - { - "decorator": "db_update", - "file": "app/db/models/user.py", - "method": "User.update_otp_by_name" - }, - { - "decorator": "db_update", - "file": "app/db/models/userconfig.py", - "method": "UserConfig.delete_by_key" - }, { "decorator": "db_query", "file": "app/db/models/userconfig.py", "method": "UserConfig.get_by_key" }, - { - "decorator": "async_db_update", - "file": "app/db/models/workflow.py", - "method": "Workflow.async_fail" - }, { "decorator": "async_db_query", "file": "app/db/models/workflow.py", @@ -803,31 +603,6 @@ "file": "app/db/models/workflow.py", "method": "Workflow.async_get_timer_triggered_workflows" }, - { - "decorator": "async_db_update", - "file": "app/db/models/workflow.py", - "method": "Workflow.async_reset" - }, - { - "decorator": "async_db_update", - "file": "app/db/models/workflow.py", - "method": "Workflow.async_start" - }, - { - "decorator": "async_db_update", - "file": "app/db/models/workflow.py", - "method": "Workflow.async_success" - }, - { - "decorator": "async_db_update", - "file": "app/db/models/workflow.py", - "method": "Workflow.async_update_current_action" - }, - { - "decorator": "async_db_update", - "file": "app/db/models/workflow.py", - "method": "Workflow.async_update_state" - }, { "decorator": "db_query", "file": "app/db/models/workflow.py", diff --git a/tests/test_agent_task_runs.py b/tests/test_agent_task_runs.py index 5d328a00b..abb6190d5 100644 --- a/tests/test_agent_task_runs.py +++ b/tests/test_agent_task_runs.py @@ -7,12 +7,16 @@ import pytest from sqlalchemy import event from sqlalchemy.exc import IntegrityError -from app.agent import AgentManager +from app.agent.orchestrator import AgentManager from app.agent.tools.impl.query_agent_tasks import QueryAgentTasksTool -from app.db import Engine, SessionFactory +from app.db.engine import get_engine 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 + + +Engine = get_engine() def _add_task(prefix: str, *, trigger_type: str = "cron") -> AgentTask: @@ -135,17 +139,15 @@ def test_begin_run_rolls_back_task_claim_when_run_insert_fails() -> None: first_task = _add_task("run-rollback-first") second_task = _add_task("run-rollback-second") run_id = uuid4().hex - assert AgentTaskRun.begin_run( - None, + assert AgentTaskOper().begin_run( task_id=first_task.id, run_id=run_id, trigger_source="scheduled", started_at="2026-08-13 20:00:00", - ) == run_id + ).run_id == run_id with pytest.raises(IntegrityError): - AgentTaskRun.begin_run( - None, + AgentTaskOper().begin_run( task_id=second_task.id, run_id=run_id, trigger_source="manual", diff --git a/tests/test_architecture_contract_baseline.py b/tests/test_architecture_contract_baseline.py index 3959b577c..af484d1b9 100644 --- a/tests/test_architecture_contract_baseline.py +++ b/tests/test_architecture_contract_baseline.py @@ -120,13 +120,15 @@ def test_runtime_contract_baseline_excludes_diagnostic_line_numbers(): def test_transaction_debt_baseline_is_a_model_and_oper_ratchet() -> None: - """事务 fixture 必须冻结存量 Model 自动提交,并保持 Oper 自提交为零。""" + """事务 fixture 必须保持 Model 写装饰器归零,并冻结剩余查询债务。""" baseline_path = BASELINE_ROOT / "transaction-debt-baseline.json" baseline = json.loads(baseline_path.read_text(encoding="utf-8")) assert baseline["schema_version"] == 1 - assert baseline["model_decorators"]["count"] == 168 - assert sum(baseline["model_decorators"]["by_kind"].values()) == 168 + assert baseline["model_decorators"]["count"] == 123 + assert sum(baseline["model_decorators"]["by_kind"].values()) == 123 + assert baseline["model_decorators"]["by_kind"]["db_update"] == 0 + assert baseline["model_decorators"]["by_kind"]["async_db_update"] == 0 assert baseline["model_transaction_calls"] == {"count": 0, "calls": []} assert baseline["model_session_factories"] == {"count": 0, "calls": []} assert baseline["oper_transaction_calls"] == {"count": 0, "calls": []} diff --git a/tests/test_data_cleanup_service.py b/tests/test_data_cleanup_service.py index 9f5aefb9d..adfda6f9f 100644 --- a/tests/test_data_cleanup_service.py +++ b/tests/test_data_cleanup_service.py @@ -18,11 +18,25 @@ class FakeCleanupRepository: self.failing_table = failing_table self.calls: list[str] = [] self._message_results = iter((2, 1, 0)) + self.commits = 0 + self.rollbacks = 0 def session(self): """返回无需真实数据库的上下文。""" return nullcontext(object()) + def unit_of_work(self, db): + """返回记录提交和回滚次数的测试事务边界。""" + return self + + def commit(self) -> None: + """记录一个成功清理批次。""" + self.commits += 1 + + def rollback(self) -> None: + """记录一个失败清理批次。""" + self.rollbacks += 1 + def _delete(self, name: str) -> int: """记录删除调用并按配置模拟结果或异常。""" self.calls.append(name) @@ -87,6 +101,8 @@ def test_cleanup_service_owns_batching_report_and_progress() -> None: assert report["tables"]["message"]["deleted"] == 3 assert report["tables"]["message"]["batches"] == 2 assert report["total_deleted"] == 3 + assert repository.commits == 2 + assert repository.rollbacks == 0 assert repository.calls == [ "message", "message", @@ -113,6 +129,7 @@ def test_cleanup_service_finishes_other_tables_before_raising_partial_failure() service.execute(batch_size=2) assert repository.calls[-1] == "downloadfailure" + assert repository.rollbacks == 1 def test_scheduler_cleanup_is_a_compatibility_delegate() -> None: diff --git a/tests/test_db_config_user_queries.py b/tests/test_db_config_user_queries.py index fb6664f5e..eaa0fc08d 100644 --- a/tests/test_db_config_user_queries.py +++ b/tests/test_db_config_user_queries.py @@ -13,6 +13,7 @@ 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 @@ -171,12 +172,12 @@ def test_user_async_mutations_match_sync_behaviour(db): db.add(User(name="mp-test-async-otp", hashed_password="x", is_otp=False)) oper = UserOper() - assert asyncio.run(User.async_update_otp_by_name( + assert asyncio.run(oper.async_update_otp_by_name( name="mp-test-async-otp", otp=True, secret="S2")) is True - assert asyncio.run(User.async_update_otp_by_name( + assert asyncio.run(oper.async_update_otp_by_name( name="mp-test-nobody", otp=True, secret="S2")) is False - assert asyncio.run(User().async_delete_by_name(name="mp-test-async-otp")) is True + assert asyncio.run(oper.async_delete_by_name(name="mp-test-async-otp")) is True assert User.get_by_name(db.session, "mp-test-async-otp") is None asyncio.run(oper.async_delete(async_id_user.id)) @@ -253,8 +254,9 @@ def test_passkey_async_delete_enforces_the_same_ownership_rule(db): """ victim = db.add(_passkey(9006, "cred-async-victim")) - assert asyncio.run(PassKey.async_delete_by_id(passkey_id=victim.id, user_id=9999)) is False - assert asyncio.run(PassKey.async_delete_by_id(passkey_id=victim.id, user_id=9006)) is True + oper = PassKeyOper() + assert asyncio.run(oper.async_delete_by_id(passkey_id=victim.id, user_id=9999)) is False + assert asyncio.run(oper.async_delete_by_id(passkey_id=victim.id, user_id=9006)) is True assert PassKey.get_by_id(db.session, victim.id) is None diff --git a/tests/test_db_oper_layer_extra.py b/tests/test_db_oper_layer_extra.py index 0b133caeb..b47b884b6 100644 --- a/tests/test_db_oper_layer_extra.py +++ b/tests/test_db_oper_layer_extra.py @@ -258,6 +258,7 @@ def test_message_oper_listing_entry_points(db): oper = MessageOper(db=db.session) oper.add(title="分页消息", text="正文", source="op-msg-2", reg_time="2026-08-13 10:00:00") + db.session.commit() assert [m.title for m in oper.list_by_page(page=1, count=1)] == ["分页消息"] assert [m.title for m in asyncio.run(oper.async_list_by_page(page=1, count=1))] == \ diff --git a/tests/test_db_workflow_queries.py b/tests/test_db_workflow_queries.py index 0981e57a8..7bfb70f11 100644 --- a/tests/test_db_workflow_queries.py +++ b/tests/test_db_workflow_queries.py @@ -10,6 +10,7 @@ import asyncio import pytest from app.db.models.workflow import Workflow +from app.db.session import async_session_scope @pytest.fixture(autouse=True) @@ -26,6 +27,18 @@ def _flow(name: str, trigger_type: str = "timer", state: str = "W", actions=[], flows=[], context={}, execution_state={}) +async def _stage_async_action(workflow_id: int, action_id: str) -> None: + """用独占异步会话提交一次模型级暂存,模拟 Application UoW 边界。""" + async with async_session_scope() as session: + await Workflow.async_update_current_action( + session, + wid=workflow_id, + action_id=action_id, + context={}, + ) + await session.commit() + + # --------------------------------------------------------------------------- # # 列表查询 # --------------------------------------------------------------------------- # @@ -230,8 +243,7 @@ def test_update_current_action_matches_async_twin(db): Workflow.update_current_action(db.session, sync_flow.id, action, {}) # 同步 Model 方法只暂存 SQL;由测试持有的事务边界先提交,避免与异步会话争锁。 db.session.commit() - asyncio.run(Workflow.async_update_current_action( - wid=async_flow.id, action_id=action, context={})) + asyncio.run(_stage_async_action(async_flow.id, action)) assert Workflow.get_by_name(db.session, "wf-sync-action").current_action == \ Workflow.get_by_name(db.session, "wf-async-action").current_action diff --git a/tests/test_host_runtime_context.py b/tests/test_host_runtime_context.py index e4db0ecb6..aef67fcd1 100644 --- a/tests/test_host_runtime_context.py +++ b/tests/test_host_runtime_context.py @@ -1,6 +1,8 @@ """类型化 HostRuntime 与 FastAPI AppState 注入测试。""" +import ast from dataclasses import FrozenInstanceError +from pathlib import Path from types import SimpleNamespace import pytest @@ -11,13 +13,18 @@ from app.api.context import ( get_agent_chat_repository, get_agent_chat_transaction, ) -from app.api.data import ( - ApiDataPorts, - configure_api_data_runtime, - get_api_data_ports, -) from app.startup import lifecycle -from app.startup.context import AgentChatRuntime, HostRuntime, SubscriptionRuntime +from app.startup.context import ( + AgentChatRuntime, + AuthenticationRuntime, + HistoryRuntime, + HostRuntime, + MessagingRuntime, + PersistenceRuntime, + SiteRuntime, + SubscriptionRuntime, + WorkflowRuntime, +) from app.application.configuration import ( ApiRuntimeConfig, ChainRuntimeConfig, @@ -26,6 +33,9 @@ from app.application.configuration import ( ) +PROJECT_ROOT = Path(__file__).parents[1] + + class _Repository: """记录绑定会话的 Agent 会话仓储替身。""" @@ -48,6 +58,20 @@ class _UnitOfWork: """模拟回滚。""" +class _SyncUnitOfWork: + """记录绑定会话的同步事务替身。""" + + def __init__(self, session: object) -> None: + """保存与仓储相同的请求会话。""" + self.session = session + + def commit(self) -> None: + """模拟提交。""" + + def rollback(self) -> None: + """模拟回滚。""" + + class _Outbox: """记录绑定会话的异步 outbox 替身。""" @@ -73,19 +97,31 @@ def _runtime() -> HostRuntime: if False: yield object() - compatibility = ApiDataPorts( - sync_session=sync_session, - async_session=async_session, - repositories={}, - standalone={}, - unit_of_work={}, - ) return HostRuntime( agent_chat=AgentChatRuntime( async_session=async_session, repository=_Repository, transaction=_UnitOfWork, ), + persistence=PersistenceRuntime( + sync_session=sync_session, + async_session=async_session, + sync_transaction=_SyncUnitOfWork, + async_transaction=_UnitOfWork, + ), + authentication=AuthenticationRuntime( + user_repository=_Repository, + standalone_user=lambda: _Repository(object()), + system_config=lambda: _Repository(object()), + passkey=lambda: _Repository(object()), + ), + messaging=MessagingRuntime(repository=_Repository), + history=HistoryRuntime( + download_repository=_Repository, + transfer_repository=_Repository, + media_server_repository=_Repository, + ), + site=SiteRuntime(repository=_Repository), subscription=SubscriptionRuntime( async_session=async_session, repository=_Repository, @@ -93,6 +129,10 @@ def _runtime() -> HostRuntime: transaction=_UnitOfWork, outbox=_Outbox, ), + workflow=WorkflowRuntime( + repository=_Repository, + system_config=lambda: _Repository(object()), + ), configuration=RuntimeConfiguration( api=lambda: ApiRuntimeConfig(False, 60, False, True), scheduler=lambda: SchedulerRuntimeConfig( @@ -101,18 +141,20 @@ def _runtime() -> HostRuntime: ), chain=lambda: ChainRuntimeConfig(media_extensions=(".mkv",)), ), - compatibility_api_data=compatibility, ) -def test_host_runtime_is_frozen_slotted_and_reuses_compatibility_facade() -> None: - """运行时不可动态扩字段,旧 Facade 必须指向同一个端口实例。""" +def test_host_runtime_is_frozen_slotted_and_covers_all_api_domains() -> None: + """运行时不可动态扩字段,且全部正式 API 领域都有命名能力。""" runtime = _runtime() - configure_api_data_runtime(runtime.compatibility_api_data) - assert not hasattr(runtime, "__dict__") - assert get_api_data_ports() is runtime.compatibility_api_data + assert runtime.authentication.user_repository is _Repository + assert runtime.messaging.repository is _Repository + assert runtime.history.download_repository is _Repository + assert runtime.site.repository is _Repository + assert runtime.subscription.repository is _Repository + assert runtime.workflow.repository is _Repository with pytest.raises(FrozenInstanceError): runtime.agent_chat = runtime.agent_chat @@ -137,6 +179,30 @@ def test_fastapi_dependencies_use_fake_runtime_without_real_services() -> None: assert response.json() == {"same_session": True} +def test_official_api_dependencies_do_not_use_string_data_locator() -> None: + """正式业务依赖只能读取 HostRuntime 命名领域,禁止回退字符串注册表。""" + dependency_root = PROJECT_ROOT / "app" / "api" / "dependencies" + official_modules = { + "agent.py", + "auth.py", + "history.py", + "site.py", + "subscription.py", + "workflow.py", + } + for filename in official_modules: + tree = ast.parse( + (dependency_root / filename).read_text(encoding="utf-8") + ) + imported_modules = { + node.module + for node in ast.walk(tree) + if isinstance(node, ast.ImportFrom) and node.module + } + assert "app.api.data" not in imported_modules + assert "app.api.dependencies.data" not in imported_modules + + @pytest.mark.asyncio async def test_lifecycle_component_attaches_init_modules_result(monkeypatch) -> None: """模块组件把 init_modules 的构建结果发布到当前 AppState。""" diff --git a/tests/test_outbox.py b/tests/test_outbox.py index 2c8992bab..2f7fab758 100644 --- a/tests/test_outbox.py +++ b/tests/test_outbox.py @@ -62,11 +62,13 @@ def test_dispatcher_retries_then_dead_letters_with_stable_key() -> None: ClaimedOutboxMessage(1, "subscribe.added:42:v1", "subscribe.added", {}, 1, 2), ] handler = MagicMock(side_effect=RuntimeError("temporary")) + failure_observer = MagicMock() dispatcher = OutboxDispatcher( repository, {"subscribe.added": handler}, max_attempts=2, clock=lambda: now, + failure_observer=failure_observer, ) assert dispatcher.dispatch_one() is True @@ -77,6 +79,10 @@ def test_dispatcher_retries_then_dead_letters_with_stable_key() -> None: "subscribe.added:42:v1", "subscribe.added:42:v1", ] + assert [call.args[0] for call in failure_observer.call_args_list] == [ + False, + True, + ] def test_dispatcher_marks_success_and_closes_owned_resource() -> None: diff --git a/tests/test_user_service.py b/tests/test_user_service.py new file mode 100644 index 000000000..d18108642 --- /dev/null +++ b/tests/test_user_service.py @@ -0,0 +1,39 @@ +"""用户应用服务的请求级事务边界测试。""" + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from app.application.security.user import UserService + + +@pytest.mark.asyncio +async def test_user_service_commits_staged_mutation() -> None: + """正式用户写用例必须在仓储暂存成功后提交请求 UoW。""" + repository = MagicMock() + repository.async_create = AsyncMock(return_value={"id": 7}) + unit_of_work = MagicMock() + unit_of_work.commit = AsyncMock() + unit_of_work.rollback = AsyncMock() + service = UserService(repository, unit_of_work) + + assert await service.create({"name": "demo"}) == {"id": 7} + unit_of_work.commit.assert_awaited_once_with() + unit_of_work.rollback.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_user_service_rolls_back_failed_mutation() -> None: + """用户仓储写入失败时不得提交部分事务。""" + repository = MagicMock() + repository.async_delete = AsyncMock(side_effect=RuntimeError("write failed")) + unit_of_work = MagicMock() + unit_of_work.commit = AsyncMock() + unit_of_work.rollback = AsyncMock() + service = UserService(repository, unit_of_work) + + with pytest.raises(RuntimeError, match="write failed"): + await service.delete(7) + + unit_of_work.rollback.assert_awaited_once_with() + unit_of_work.commit.assert_not_awaited()