refactor(repository): extract AsyncJobRepository base class for analysis/comparison

提取 AsyncJobRepository 泛型基类,消除 CampaignAnalysisRepository 和
CampaignPeriodComparisonRepository 的重复代码。基类提供 get_by_campaign
和 mark_orphans_failed 通用逻辑,子类只需指定 _table 类型和实现 upsert。

- 新增 AsyncJobRepository[DB] 泛型基类
- mark_orphans_failed 接受 error_message 参数,子类传入特定错误信息
- 删除约 60 行重复代码(两个 __init__、两个 get_by_campaign、两个 mark_orphans_failed 实现)
- 所有测试通过,行为不变
This commit is contained in:
sinohqb 2026-08-04 13:25:42 +08:00
parent cbfdf86b36
commit 0aa3ef81c5

View File

@ -445,16 +445,46 @@ class CampaignRepository(BaseRepository[Campaign, CampaignDB]):
self.session.refresh(db)
class CampaignAnalysisRepository:
"""Repository for campaign analysis rows (one per campaign, upserted)."""
class AsyncJobRepository(Generic[DB]):
"""Base class for async job repositories (analysis, comparison).
Provides common pattern: get_by_campaign, mark_orphans_failed.
Subclasses implement upsert with their specific fields.
"""
_table: type[DB]
def __init__(self, session: Optional[Session] = None):
self.session = session or get_session()
def get_by_campaign(self, campaign_id: str) -> Optional[CampaignAnalysisDB]:
statement = select(CampaignAnalysisDB).where(CampaignAnalysisDB.campaign_id == campaign_id)
def get_by_campaign(self, campaign_id: str) -> Optional[DB]:
statement = select(self._table).where(self._table.campaign_id == campaign_id) # type: ignore[attr-defined]
return self.session.exec(statement).first()
def mark_orphans_failed(self, error_message: str) -> int:
"""服务启动时清理:把滞留的 generating 行标记为 failed。
异步任务是进程内 asyncio 任务服务重启后不会恢复不清理则这些
行永远停留在 generating僵尸状态
"""
rows = self.session.exec(
select(self._table).where(self._table.status == "generating") # type: ignore[attr-defined]
).all()
for row in rows:
row.status = "failed" # type: ignore[attr-defined]
row.error = error_message # type: ignore[attr-defined]
row.updated_at = utc_now() # type: ignore[attr-defined]
self.session.add(row)
if rows:
self.session.commit()
return len(rows)
class CampaignAnalysisRepository(AsyncJobRepository[CampaignAnalysisDB]):
"""Repository for campaign analysis rows (one per campaign, upserted)."""
_table = CampaignAnalysisDB
def upsert(
self,
campaign_id: str,
@ -483,31 +513,14 @@ class CampaignAnalysisRepository:
return row
def mark_orphans_failed(self) -> int:
"""服务启动时清理:把滞留的 generating 分析行标记为 failed。
分析任务是进程内 asyncio 任务服务重启后不会恢复不清理则这些
行永远停留在 generating僵尸状态
"""
rows = self.session.exec(select(CampaignAnalysisDB).where(CampaignAnalysisDB.status == "generating")).all()
for row in rows:
row.status = "failed"
row.error = "服务重启导致分析生成中断"
row.updated_at = utc_now()
self.session.add(row)
if rows:
self.session.commit()
return len(rows)
"""服务启动时清理:把滞留的 generating 分析行标记为 failed。"""
return super().mark_orphans_failed("服务重启导致分析生成中断")
class CampaignPeriodComparisonRepository:
class CampaignPeriodComparisonRepository(AsyncJobRepository[CampaignPeriodComparisonDB]):
"""Repository for period-comparison rows (one per campaign, upserted)."""
def __init__(self, session: Optional[Session] = None):
self.session = session or get_session()
def get_by_campaign(self, campaign_id: str) -> Optional[CampaignPeriodComparisonDB]:
statement = select(CampaignPeriodComparisonDB).where(CampaignPeriodComparisonDB.campaign_id == campaign_id)
return self.session.exec(statement).first()
_table = CampaignPeriodComparisonDB
def upsert(
self,
@ -543,22 +556,8 @@ class CampaignPeriodComparisonRepository:
return row
def mark_orphans_failed(self) -> int:
"""服务启动时清理:把滞留的 generating 周期对比行标记为 failed。
对比任务是进程内 asyncio 任务服务重启后不会恢复不清理则这些
行永远停留在 generating僵尸状态基线配对保留可直接重新触发
"""
rows = self.session.exec(
select(CampaignPeriodComparisonDB).where(CampaignPeriodComparisonDB.status == "generating")
).all()
for row in rows:
row.status = "failed"
row.error = "服务重启导致周期对比生成中断"
row.updated_at = utc_now()
self.session.add(row)
if rows:
self.session.commit()
return len(rows)
"""服务启动时清理:把滞留的 generating 周期对比行标记为 failed。"""
return super().mark_orphans_failed("服务重启导致周期对比生成中断")
class ResultRepository: