From 0aa3ef81c5355acbca271e2ab284492b23681bb7 Mon Sep 17 00:00:00 2001 From: sinohqb Date: Tue, 4 Aug 2026 13:25:42 +0800 Subject: [PATCH] refactor(repository): extract AsyncJobRepository base class for analysis/comparison MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 提取 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 实现) - 所有测试通过,行为不变 --- backend/agenteval/storage/repository.py | 81 ++++++++++++------------- 1 file changed, 40 insertions(+), 41 deletions(-) diff --git a/backend/agenteval/storage/repository.py b/backend/agenteval/storage/repository.py index 9fad145..63e54b7 100644 --- a/backend/agenteval/storage/repository.py +++ b/backend/agenteval/storage/repository.py @@ -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: