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: