AgentEvalTool/backend/agenteval/storage/async_job_repository.py
sinohqb 7eae6de52d refactor(evaluation/storage): 结算统一与 repository 拆分(Phase 2 + 3)
合并两个不可分割的深化:

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 精神,等真实需求出现再议)。
2026-08-24 05:50:27 +08:00

248 lines
9.0 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""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("服务重启导致周期对比生成中断")