refactor: govern background tasks and query ownership

This commit is contained in:
jxxghp
2026-08-23 13:24:04 +08:00
parent 43c173a0e7
commit f1e542bef0
37 changed files with 1570 additions and 510 deletions
+127 -47
View File
@@ -297,13 +297,24 @@ class SubscribeOper(DbOper):
"""
获取订阅
"""
return Subscribe.get(self._db, rid=sid)
return self._execute_sync_query(
lambda session: session.execute(
select(Subscribe).where(Subscribe.id == sid)
).scalars().first()
)
async def async_get(self, sid: int) -> Optional[Subscribe]:
"""
获取订阅
"""
return await Subscribe.async_get(self._db, rid=sid)
if self._db is not None and not isinstance(self._db, (Session, AsyncSession)):
# 保留旧测试替身与插件注入对象对 Model ABI 的兼容入口。
return await Subscribe.async_get(self._db, rid=sid)
async def query(session: AsyncSession) -> Optional[Subscribe]:
"""在调用方异步会话中执行订阅主键查询。"""
result = await session.execute(select(Subscribe).where(Subscribe.id == sid))
return result.scalars().first()
return await self._execute_async_query(query)
async def async_list_by_media_identity(
self,
@@ -312,12 +323,18 @@ class SubscribeOper(DbOper):
music_type: Optional[str] = None,
) -> List[Subscribe]:
"""异步按规范媒体身份读取订阅。"""
return await Subscribe.async_list_by_media_identity(
self._db,
media_source=media_source,
media_id=media_id,
music_type=music_type,
)
async def query(session: AsyncSession) -> List[Subscribe]:
"""在调用方异步会话中执行媒体身份列表查询。"""
condition = Subscribe._identity_condition( # pylint: disable=protected-access
media_source, media_id, music_type
)
if condition is None:
return []
result = await session.execute(select(Subscribe).where(condition))
return list(result.scalars().all())
if isinstance(self._db, AsyncSession):
return await query(self._db)
return await self._execute_async_query(query)
def list_by_media_identity(
self,
@@ -326,12 +343,15 @@ class SubscribeOper(DbOper):
music_type: Optional[str] = None,
) -> List[Subscribe]:
"""同步按规范媒体身份读取订阅。"""
return Subscribe.list_by_media_identity(
self._db,
media_source=media_source,
media_id=media_id,
music_type=music_type,
)
def query(session: Session) -> List[Subscribe]:
"""在调用方同步会话中执行媒体身份列表查询。"""
condition = Subscribe._identity_condition( # pylint: disable=protected-access
media_source, media_id, music_type
)
if condition is None:
return []
return list(session.execute(select(Subscribe).where(condition)).scalars().all())
return self._execute_sync_query(query)
async def get_candidate(
self,
@@ -360,11 +380,8 @@ class SubscribeOper(DbOper):
music_type: Optional[str],
) -> List[SubscribeDeletionCandidate]:
"""按媒体身份读取去重后的订阅删除快照。"""
subscribes = await Subscribe.async_list_by_media_identity(
self._db,
media_source=media_source,
media_id=media_id,
music_type=music_type,
subscribes = await self.async_list_by_media_identity(
media_source, media_id, music_type
)
candidates = []
seen_ids = set()
@@ -395,11 +412,7 @@ class SubscribeOper(DbOper):
async def list_search_ids(self, username: str, state: str) -> List[int]:
"""返回用户指定状态的订阅编号,不向应用用例暴露 ORM 列表。"""
subscribes = await Subscribe.async_list_by_username(
self._db,
username,
state=state,
)
subscribes = await self.async_list_by_username(username, state=state)
return [subscribe.id for subscribe in subscribes if subscribe.id]
def get_by(
@@ -410,9 +423,18 @@ class SubscribeOper(DbOper):
"""
根据条件查询订阅
"""
return Subscribe.get_by(
self._db, type, media_source, media_id, season, music_type,
)
def query(session: Session) -> Optional[Subscribe]:
"""在调用方同步会话中执行类型媒体查询。"""
condition = Subscribe._identity_condition( # pylint: disable=protected-access
media_source, media_id, music_type
)
if condition is None:
return None
statement = select(Subscribe).where(condition, Subscribe.type == type)
if season is not None:
statement = statement.where(Subscribe.season == season)
return session.execute(statement).scalars().first()
return self._execute_sync_query(query)
async def async_get_by(
self, type: str, media_source: MediaSource, media_id: str,
@@ -422,25 +444,55 @@ class SubscribeOper(DbOper):
"""
根据条件查询订阅
"""
return await Subscribe.async_get_by(
self._db, type, media_source, media_id, season, music_type,
)
async def query(session: AsyncSession) -> Optional[Subscribe]:
"""在调用方异步会话中执行类型媒体查询。"""
condition = Subscribe._identity_condition( # pylint: disable=protected-access
media_source, media_id, music_type
)
if condition is None:
return None
statement = select(Subscribe).where(condition, Subscribe.type == type)
if season is not None:
statement = statement.where(Subscribe.season == season)
result = await session.execute(statement)
return result.scalars().first()
return await self._execute_async_query(query)
def list(self, state: Optional[str] = None) -> List[Subscribe]:
"""
获取订阅列表
"""
if state:
return Subscribe.get_by_state(self._db, state)
return Subscribe.list(self._db)
return self._execute_sync_query(
lambda session: list(session.execute(
select(Subscribe).where(Subscribe.state.in_(state.split(',')))
).scalars().all())
)
return self._execute_sync_query(
lambda session: list(session.execute(select(Subscribe)).scalars().all())
)
async def async_list(self, state: Optional[str] = None) -> List[Subscribe]:
"""
异步获取订阅列表
"""
if self._db is not None and not isinstance(self._db, (Session, AsyncSession)):
if state:
return await Subscribe.async_get_by_state(self._db, state)
return await Subscribe.async_list(self._db)
if state:
return await Subscribe.async_get_by_state(self._db, state)
return await Subscribe.async_list(self._db)
async def query(session: AsyncSession) -> List[Subscribe]:
"""在调用方异步会话中执行状态列表查询。"""
result = await session.execute(
select(Subscribe).where(Subscribe.state.in_(state.split(',')))
)
return list(result.scalars().all())
return await self._execute_async_query(query)
async def query_all(session: AsyncSession) -> List[Subscribe]:
"""在调用方异步会话中执行全量订阅查询。"""
result = await session.execute(select(Subscribe))
return list(result.scalars().all())
return await self._execute_async_query(query_all)
async def async_list_by_username(
self,
@@ -449,12 +501,20 @@ class SubscribeOper(DbOper):
mtype: Optional[str] = None,
) -> List[Subscribe]:
"""异步按用户获取订阅。"""
return await Subscribe.async_list_by_username(
self._db,
username=username,
state=state,
mtype=mtype,
)
if self._db is not None and not isinstance(self._db, (Session, AsyncSession)):
return await Subscribe.async_list_by_username(
self._db, username=username, state=state, mtype=mtype
)
async def query(session: AsyncSession) -> List[Subscribe]:
"""在调用方异步会话中执行用户筛选查询。"""
statement = select(Subscribe).where(Subscribe.username == username)
if state:
statement = statement.where(Subscribe.state == state)
if mtype:
statement = statement.where(Subscribe.type == mtype)
result = await session.execute(statement)
return list(result.scalars().all())
return await self._execute_async_query(query)
async def async_list_by_title(
self,
@@ -462,11 +522,14 @@ class SubscribeOper(DbOper):
season: Optional[int] = None,
) -> List[Subscribe]:
"""异步按标题获取订阅,供旧查询测试和迁移调用兼容。"""
return await Subscribe.async_list_by_title(
self._db,
title=title,
season=season,
)
async def query(session: AsyncSession) -> List[Subscribe]:
"""在调用方异步会话中执行标题列表查询。"""
statement = select(Subscribe).where(Subscribe.name == title)
if season is not None:
statement = statement.where(Subscribe.season == season)
result = await session.execute(statement)
return list(result.scalars().all())
return await self._execute_async_query(query)
def delete(self, sid: int):
"""
@@ -535,13 +598,30 @@ class SubscribeOper(DbOper):
"""
获取指定用户的订阅
"""
return Subscribe.list_by_username(self._db, username=username, state=state, mtype=mtype)
def query(session: Session) -> List[Subscribe]:
"""在调用方同步会话中执行用户筛选查询。"""
statement = select(Subscribe).where(Subscribe.username == username)
if state:
statement = statement.where(Subscribe.state == state)
if mtype:
statement = statement.where(Subscribe.type == mtype)
return list(session.execute(statement).scalars().all())
return self._execute_sync_query(query)
def list_by_type(self, mtype: str, days: int = 7) -> List[Subscribe]:
"""
获取指定类型的订阅
"""
return Subscribe.list_by_type(self._db, mtype=mtype, days=days)
def query(session: Session) -> List[Subscribe]:
"""在调用方同步会话中执行时间窗订阅查询。"""
cutoff = time.strftime(
"%Y-%m-%d %H:%M:%S",
time.localtime(time.time() - 86400 * int(days)),
)
return list(session.execute(select(Subscribe).where(
Subscribe.type == mtype, Subscribe.date >= cutoff
)).scalars().all())
return self._execute_sync_query(query)
def add_history(self, **kwargs):
"""