合并两个不可分割的深化: Phase 2 — 智能作业结算统一(ADR-0012) - intelligence_jobs.execute(job_kind, campaign_id, ...) 作为结算的 唯一实现:建行 → 认领 → 校验 → generating → 落账,一处编排、 一处截断(500 字符)。两个 executor 退化为 ensure_queued / validate / work_fn 三个小 adapter。 - analysis.validate_analysis_request() 共享校验入口(活动终态 → 模型),路由捕获映射 400、executor 捕获落 failed 行,与 validate_comparison_request 先例同构。 - campaign_runner._auto_start_analysis 的跳过守卫收敛至 auto_intelligence_eligible 单一判断点。 - comparison.py 删除零调用的 build_comparison_payload; load_comparison_view 投影归位至 campaign_read_model。 - 新增 characterization 测试(认领竞争、重复触发、截断、恢复上限)。 Phase 3 — storage/repository.py 拆分 - AsyncJobRepository 及两个子类迁至 storage/async_job_repository.py(Phase 2 的 intelligence_jobs 与 comparison 必须 import 自该路径,故与 Phase 2 同 commit)。 - ExplorationSession / ExplorationMessage 迁至 storage/exploration_repository.py;repository.py 由 1180 行降至 约 814 行,grep 确认无残留符号。 - exploration 子模块与路由 import 全部更新;测试 import 跟随。 刻意不做:CAS 共享原语、app.py 五 registry 关停顺序归一 (ADR-0006 精神,等真实需求出现再议)。
248 lines
9.0 KiB
Python
248 lines
9.0 KiB
Python
"""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("服务重启导致周期对比生成中断")
|