"""Repositories for durable async job rows (活动分析 / 周期对比). DB 行是任务的持久权威:queued/generating/completed/failed 状态机与 claim 竞争都住在这里;进程内任务句柄另由 TaskRegistry 管理。 """ from dataclasses import dataclass from enum import Enum from typing import Generic, Optional, TypeVar from sqlalchemy import update as sql_update from sqlmodel import Session, select from agenteval.storage.db import CampaignAnalysisDB, CampaignPeriodComparisonDB, get_session, utc_now DB = TypeVar("DB") # persisted table row class AsyncJobClaimStatus(str, Enum): CLAIMED = "claimed" NOT_FOUND = "not_found" ALREADY_CLAIMED = "already_claimed" NOT_QUEUED = "not_queued" @dataclass(frozen=True) class AsyncJobClaimResult: status: AsyncJobClaimStatus @property def claimed(self) -> bool: return self.status is AsyncJobClaimStatus.CLAIMED 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[DB]: statement = select(self._table).where(self._table.campaign_id == campaign_id) # type: ignore[attr-defined] return self.session.exec(statement).first() def list_queued(self) -> list[DB]: statement = select(self._table).where(self._table.status == "queued") # type: ignore[attr-defined] return list(self.session.exec(statement).all()) def claim_queued(self, campaign_id: str) -> AsyncJobClaimResult: """Atomically move one queued job to generating. The status predicate is the durable idempotency authority. Competing workers may observe the same queued row, but only one can claim it. """ statement = ( sql_update(self._table) .where( self._table.campaign_id == campaign_id, # type: ignore[attr-defined] self._table.status == "queued", # type: ignore[attr-defined] ) .values(status="generating", error=None, updated_at=utc_now()) ) result = self.session.exec(statement) self.session.commit() self.session.expire_all() if result.rowcount == 1: return AsyncJobClaimResult(AsyncJobClaimStatus.CLAIMED) row = self.get_by_campaign(campaign_id) if row is None: return AsyncJobClaimResult(AsyncJobClaimStatus.NOT_FOUND) if row.status == "generating": # type: ignore[attr-defined] return AsyncJobClaimResult(AsyncJobClaimStatus.ALREADY_CLAIMED) return AsyncJobClaimResult(AsyncJobClaimStatus.NOT_QUEUED) def prepare_queued_recovery(self, max_attempts: int, exhausted_error: str) -> list[DB]: """Increment queued recovery attempts and fail exhausted jobs.""" recoverable: list[DB] = [] rows = self.list_queued() for row in rows: if row.recovery_attempts >= max_attempts: # type: ignore[attr-defined] row.status = "failed" # type: ignore[attr-defined] row.error = exhausted_error # type: ignore[attr-defined] else: row.recovery_attempts += 1 # type: ignore[attr-defined] recoverable.append(row) row.updated_at = utc_now() # type: ignore[attr-defined] self.session.add(row) if rows: self.session.commit() self.session.expire_all() return recoverable 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 enqueue(self, campaign_id: str, *, triggered_by: str = "manual") -> CampaignAnalysisDB: """Persist an analysis job before launching its process-local task.""" row = self.get_by_campaign(campaign_id) if row is None: row = CampaignAnalysisDB(campaign_id=campaign_id, status="queued", triggered_by=triggered_by) elif row.status in {"queued", "generating"}: return row else: row.status = "queued" row.result = None row.model_config_id = None row.error = None row.triggered_by = triggered_by row.recovery_attempts = 0 row.updated_at = utc_now() self.session.add(row) self.session.commit() self.session.refresh(row) return row def upsert( self, campaign_id: str, *, status: str, result: Optional[dict] = None, model_config_id: Optional[str] = None, error: Optional[str] = None, triggered_by: str = "manual", ) -> CampaignAnalysisDB: row = self.get_by_campaign(campaign_id) if row is None: row = CampaignAnalysisDB(campaign_id=campaign_id) row.status = status if result is not None: row.set_result(result) else: row.result = None row.model_config_id = model_config_id row.error = error row.triggered_by = triggered_by row.updated_at = utc_now() self.session.add(row) self.session.commit() self.session.refresh(row) return row def mark_orphans_failed(self) -> int: """服务启动时清理:把滞留的 generating 分析行标记为 failed。""" return super().mark_orphans_failed("服务重启导致分析生成中断") class CampaignPeriodComparisonRepository(AsyncJobRepository[CampaignPeriodComparisonDB]): """Repository for period-comparison rows (one per campaign, upserted).""" _table = CampaignPeriodComparisonDB def enqueue( self, campaign_id: str, *, baseline_campaign_id: str, triggered_by: str = "manual", ) -> CampaignPeriodComparisonDB: """Persist a comparison job before launching its process-local task.""" row = self.get_by_campaign(campaign_id) if row is None: row = CampaignPeriodComparisonDB( campaign_id=campaign_id, baseline_campaign_id=baseline_campaign_id, status="queued", triggered_by=triggered_by, ) elif row.status in {"queued", "generating"}: return row else: row.baseline_campaign_id = baseline_campaign_id row.status = "queued" row.result = None row.model_config_id = None row.error = None row.triggered_by = triggered_by row.recovery_attempts = 0 row.updated_at = utc_now() self.session.add(row) self.session.commit() self.session.refresh(row) return row def upsert( self, campaign_id: str, *, status: str, baseline_campaign_id: Optional[str] = None, result: Optional[dict] = None, model_config_id: Optional[str] = None, error: Optional[str] = None, triggered_by: str = "manual", ) -> CampaignPeriodComparisonDB: row = self.get_by_campaign(campaign_id) if row is None: row = CampaignPeriodComparisonDB( campaign_id=campaign_id, baseline_campaign_id=baseline_campaign_id or "", ) row.status = status if baseline_campaign_id is not None: row.baseline_campaign_id = baseline_campaign_id if result is not None: row.set_result(result) else: row.result = None row.model_config_id = model_config_id row.error = error row.triggered_by = triggered_by row.updated_at = utc_now() self.session.add(row) self.session.commit() self.session.refresh(row) return row def mark_orphans_failed(self) -> int: """服务启动时清理:把滞留的 generating 周期对比行标记为 failed。""" return super().mark_orphans_failed("服务重启导致周期对比生成中断")