fix(site): use explicit transactions for connectivity paths

This commit is contained in:
jxxghp
2026-08-23 10:33:14 +08:00
parent 38ddb59879
commit a8b3dd3988
7 changed files with 357 additions and 70 deletions
+13 -3
View File
@@ -130,6 +130,7 @@ from app.startup.subscription import (
)
from app.startup.chain_events import TransactionalChainDurableEventWriter
from app.startup.download_failure import TransactionalDownloadFailureRepository
from app.startup.site import TransactionalSiteRepository
from app.startup.workflow import TransactionalWorkflowExecutionService
from app.startup.transaction import TransactionalWriteRunner
from app.startup.context import (
@@ -685,7 +686,10 @@ async def init_modules() -> HostRuntime:
workflow_execution = TransactionalWorkflowExecutionService(SessionFactory)
configure_workflow_legacy_writer(workflow_execution)
configure_chain_data_ports(
site=lambda: SiteOper(),
site=lambda: TransactionalSiteRepository(
sync_session=SessionFactory,
async_session=async_session_scope,
),
subscribe=lambda: SubscribeOper(),
workflow=lambda: WorkflowOper(),
download_history=lambda: DownloadHistoryOper(),
@@ -720,13 +724,19 @@ async def init_modules() -> HostRuntime:
configure_passkey_service(PasskeyService(repository=PassKeyOper()))
configure_transfer_history_provider(lambda: TransferHistoryOper())
configure_site_query_service(SiteQueryService(repository=SiteOper()))
configure_site_health_service(SiteHealthService(repository=SiteOper()))
configure_site_health_service(SiteHealthService(repository=TransactionalSiteRepository(
sync_session=SessionFactory,
async_session=async_session_scope,
)))
configure_workflow_query(WorkflowQueryService(repository=WorkflowOper()))
configure_agent_data_ports(
agent_chat=lambda: AgentChatOper(),
agent_task=lambda: AgentTaskOper(),
user=lambda: UserOper(),
site=lambda: SiteOper(),
site=lambda: TransactionalSiteRepository(
sync_session=SessionFactory,
async_session=async_session_scope,
),
subscribe=lambda: SubscribeOper(),
subscribe_history=lambda: SubscribeHistoryOper(),
transfer_history=lambda: TransferHistoryOper(),
+193
View File
@@ -0,0 +1,193 @@
"""站点 Chain 端口的显式会话与事务适配器。"""
from __future__ import annotations
from collections.abc import Awaitable, Callable
from contextlib import AbstractAsyncContextManager
from typing import Any, TypeVar
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import Session
from app.db.oper.site import SiteOper
from app.db.uow import SqlAlchemyAsyncUnitOfWork, SqlAlchemyUnitOfWork
T = TypeVar("T")
class TransactionalSiteRepository:
"""为同步 Chain 站点端口和异步健康统计提供短生命周期会话。"""
def __init__(
self,
*,
sync_session: Callable[[], Session],
async_session: Callable[[], AbstractAsyncContextManager[AsyncSession]],
) -> None:
"""保存同步会话工厂和异步会话上下文工厂。"""
self._sync_session = sync_session
self._async_session = async_session
def _read(self, operation: Callable[[SiteOper], T]) -> T:
"""在独立同步会话中执行只读站点操作。"""
with self._sync_session() as session:
return operation(SiteOper(db=session))
def _write(self, operation: Callable[[SiteOper], T]) -> T:
"""在独立同步 UoW 中执行站点写操作。"""
with self._sync_session() as session:
session.expire_on_commit = False
unit_of_work = SqlAlchemyUnitOfWork(session)
try:
result = operation(SiteOper(db=session))
unit_of_work.commit()
return result
except Exception:
unit_of_work.rollback()
raise
async def _async_write(
self,
operation: Callable[[SiteOper], Awaitable[T]],
) -> T:
"""在独立异步 UoW 中执行站点写操作。"""
async with self._async_session() as session:
session.sync_session.expire_on_commit = False
unit_of_work = SqlAlchemyAsyncUnitOfWork(session)
try:
result = await operation(SiteOper(db=session))
await unit_of_work.commit()
return result
except Exception:
await unit_of_work.rollback()
raise
async def _async_read(self, operation: Callable[[SiteOper], Awaitable[T]]) -> T:
"""在独立异步会话中执行只读站点操作。"""
async with self._async_session() as session:
return await operation(SiteOper(db=session))
def add(self, **kwargs: Any) -> tuple[bool, str]:
"""新增站点并提交事务。"""
return self._write(lambda repository: repository.add(**kwargs))
def get(self, site_id: int) -> Any:
"""按 ID 查询站点。"""
return self._read(lambda repository: repository.get(site_id))
def get_by_domain(self, domain: str) -> Any:
"""按域名查询站点。"""
return self._read(lambda repository: repository.get_by_domain(domain))
def get_domains_by_ids(self, ids: list[int]) -> list[str | None]:
"""查询一组站点 ID 对应的域名。"""
return self._read(lambda repository: repository.get_domains_by_ids(ids))
def list(self) -> list[Any]:
"""查询全部站点。"""
return self._read(lambda repository: repository.list())
async def async_get(self, site_id: int) -> Any:
"""异步按 ID 查询站点。"""
return await self._async_read(lambda repository: repository.async_get(site_id))
async def async_get_by_domain(self, domain: str) -> Any:
"""异步按域名查询站点。"""
return await self._async_read(
lambda repository: repository.async_get_by_domain(domain)
)
async def async_get_by_name(self, name: str) -> Any:
"""异步按名称查询站点。"""
return await self._async_read(
lambda repository: repository.async_get_by_name(name)
)
async def async_list(self) -> list[Any]:
"""异步查询全部站点。"""
return await self._async_read(lambda repository: repository.async_list())
async def async_list_order_by_pri(self) -> list[Any]:
"""异步按优先级查询站点。"""
return await self._async_read(
lambda repository: repository.async_list_order_by_pri()
)
async def async_update(self, site_id: int, payload: dict[str, Any]) -> Any:
"""异步更新站点并提交事务。"""
return await self._async_write(
lambda repository: repository.async_update(site_id, payload)
)
async def async_get_userdata_by_domain(
self,
domain: str,
workdate: str | None = None,
) -> list[Any]:
"""异步查询站点用户数据。"""
return await self._async_read(
lambda repository: repository.async_get_userdata_by_domain(domain, workdate)
)
def update(self, site_id: int, payload: dict[str, Any]) -> Any:
"""更新站点并提交事务。"""
return self._write(lambda repository: repository.update(site_id, payload))
def update_cookie(self, domain: str, cookies: str) -> tuple[bool, str]:
"""更新站点 Cookie 并提交事务。"""
return self._write(
lambda repository: repository.update_cookie(domain, cookies)
)
def update_rss(self, domain: str, rss: str) -> tuple[bool, str]:
"""更新站点 RSS 地址并提交事务。"""
return self._write(lambda repository: repository.update_rss(domain, rss))
def update_userdata(
self,
domain: str,
name: str,
payload: dict[str, Any],
) -> tuple[bool, str]:
"""更新站点用户数据并提交事务。"""
return self._write(
lambda repository: repository.update_userdata(domain, name, payload)
)
def update_icon(
self,
name: str,
domain: str,
icon_url: str,
icon_base64: str,
) -> bool:
"""更新站点图标并提交事务。"""
return self._write(
lambda repository: repository.update_icon(
name,
domain,
icon_url,
icon_base64,
)
)
def success(self, domain: str, seconds: int | None = None) -> Any:
"""记录站点访问成功并提交事务。"""
return self._write(lambda repository: repository.success(domain, seconds))
def fail(self, domain: str) -> Any:
"""记录站点访问失败并提交事务。"""
return self._write(lambda repository: repository.fail(domain))
async def async_success(self, domain: str, seconds: int | None = None) -> Any:
"""异步记录站点访问成功并提交事务。"""
return await self._async_write(
lambda repository: repository.async_success(domain, seconds)
)
async def async_fail(self, domain: str) -> Any:
"""异步记录站点访问失败并提交事务。"""
return await self._async_write(
lambda repository: repository.async_fail(domain)
)