"""Repository layer for database access.""" from dataclasses import dataclass from datetime import datetime from enum import Enum from typing import Generic, Optional, TypeVar from sqlalchemy import insert as sql_insert from sqlalchemy import literal from sqlalchemy import update as sql_update from sqlalchemy.exc import IntegrityError from sqlmodel import Session, select from agenteval.exploration.models import ExplorationMessage, ExplorationSession, ExplorationSessionStatus from agenteval.models import ( Campaign, CampaignStatus, CampaignSummary, Case, EvalResult, EvalRun, EvalTarget, RunStatus, RunTrigger, Scenario, ) from agenteval.services.model_configs import ModelConfigService from agenteval.storage.db import ( CampaignAnalysisDB, CampaignDB, CampaignPeriodComparisonDB, EvalResultDB, EvalRunDB, EvalTargetDB, ExplorationMessageDB, ExplorationSessionDB, ScenarioDB, TurnDB, get_session, new_uuid, utc_now, ) from agenteval.storage.model_config_repository import ScenarioModelBindingRepository M = TypeVar("M") # domain model 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 BaseRepository(Generic[M, DB]): """Shared CRUD skeleton for id-keyed entity repositories. Subclasses declare the table (``_table``) and the ``list_all`` ordering column name (``_order_by``, newest-first), and implement the ``_to_db`` / ``_from_db`` converter pair. The converters are instance methods so a subclass whose ``_from_db`` needs cross-table reads (e.g. Scenario's model bindings) can reach ``self.session``. Entities with bespoke create/update (binding validation, versioning) override just those methods. """ _table: type _order_by: str def __init__(self, session: Optional[Session] = None): self.session = session or get_session() def _to_db(self, obj: M) -> DB: raise NotImplementedError def _from_db(self, db: DB) -> M: raise NotImplementedError def list_all(self) -> list[M]: column = getattr(self._table, self._order_by) statement = select(self._table).order_by(column.desc()) return [self._from_db(r) for r in self.session.exec(statement).all()] def get(self, entity_id: str) -> Optional[M]: db = self.session.get(self._table, entity_id) return self._from_db(db) if db else None def create(self, obj: M) -> M: db = self._to_db(obj) self.session.add(db) self.session.commit() self.session.refresh(db) return self._from_db(db) def delete(self, entity_id: str) -> bool: db = self.session.get(self._table, entity_id) if not db: return False self.session.delete(db) self.session.commit() return True def _result_to_db(result: EvalResult) -> EvalResultDB: return EvalResultDB( id=result.id, run_id=result.run_id, case_id=result.case_id, turn_id=result.turn_id, rule_type=result.rule_type, passed=result.passed, score=result.score, reason=result.reason, ) def _result_from_db(db: EvalResultDB) -> EvalResult: return EvalResult( id=db.id, run_id=db.run_id, case_id=db.case_id, turn_id=db.turn_id, rule_type=db.rule_type, passed=db.passed, score=db.score, reason=db.reason, ) class TargetRepository(BaseRepository[EvalTarget, EvalTargetDB]): """Repository for evaluation targets.""" _table = EvalTargetDB _order_by = "created_at" def _to_db(self, target: EvalTarget) -> EvalTargetDB: db = EvalTargetDB( id=target.id, name=target.name, description=target.description, platform=target.platform.value, channel_type=target.channel_type.value, status=target.status.value, created_at=target.created_at, updated_at=target.updated_at or utc_now(), ) db.set_config(target.channel_config) return db def _from_db(self, db: EvalTargetDB) -> EvalTarget: return EvalTarget( id=db.id, name=db.name, description=db.description, platform=db.platform, channel_type=db.channel_type, channel_config=db.get_config(), status=db.status, created_at=db.created_at, updated_at=db.updated_at, ) def update(self, target: EvalTarget) -> Optional[EvalTarget]: existing = self.session.get(EvalTargetDB, target.id) if not existing: return None existing.name = target.name existing.description = target.description existing.platform = target.platform.value existing.channel_type = target.channel_type.value existing.status = target.status.value existing.set_config(target.channel_config) existing.updated_at = utc_now() self.session.add(existing) self.session.commit() self.session.refresh(existing) return self._from_db(existing) class ScenarioRepository(BaseRepository[Scenario, ScenarioDB]): """Repository for evaluation scenarios.""" _table = ScenarioDB _order_by = "created_at" def _to_db(self, scenario: Scenario) -> ScenarioDB: db = ScenarioDB( id=scenario.id, name=scenario.name, description=scenario.description, created_at=scenario.created_at, updated_at=scenario.updated_at or utc_now(), ) db.set_tags(scenario.tags) # mode="json" 与 update() 的考纲比较保持同一序列化形态,避免假升版 db.set_cases([case.model_dump(mode="json") for case in scenario.cases]) db.set_llm_config(scenario.llm_config) return db def _from_db(self, db: ScenarioDB) -> Scenario: bindings = ScenarioModelBindingRepository(self.session).get_for_scenario(db.id or "") return Scenario( id=db.id, name=db.name, description=db.description, tags=db.get_tags(), cases=[Case(**case) for case in db.get_cases()], model_bindings=bindings, llm_config=db.get_llm_config(), version=db.version or 1, created_at=db.created_at, updated_at=db.updated_at, ) def create(self, scenario: Scenario) -> Scenario: db = self._to_db(scenario) bindings = {purpose.value: config_id for purpose, config_id in scenario.model_bindings.items()} try: ModelConfigService(self.session).validate_bindings(bindings) self.session.add(db) self.session.flush() ScenarioModelBindingRepository(self.session).replace_for_scenario(db.id or "", bindings) self.session.commit() self.session.refresh(db) except Exception: self.session.rollback() raise return self._from_db(db) def update(self, scenario: Scenario) -> Optional[Scenario]: existing = self.session.get(ScenarioDB, scenario.id) if not existing: return None bindings = {purpose.value: config_id for purpose, config_id in scenario.model_bindings.items()} try: ModelConfigService(self.session).validate_bindings(bindings) # 考纲字段(cases / model_bindings / llm_config)变更才升版(ADR-0001); # 版本由系统维护,忽略 scenario.version 的外部传入值。 new_cases = [case.model_dump(mode="json") for case in scenario.cases] old_bindings = ScenarioModelBindingRepository(self.session).get_for_scenario(existing.id or "") syllabus_changed = ( existing.get_cases() != new_cases or existing.get_llm_config() != scenario.llm_config or old_bindings != bindings ) if syllabus_changed: existing.version = (existing.version or 1) + 1 existing.name = scenario.name existing.description = scenario.description existing.set_tags(scenario.tags) existing.set_cases(new_cases) existing.set_llm_config(scenario.llm_config) existing.updated_at = utc_now() self.session.add(existing) ScenarioModelBindingRepository(self.session).replace_for_scenario(existing.id or "", bindings) self.session.commit() self.session.refresh(existing) except Exception: self.session.rollback() raise return self._from_db(existing) def delete(self, scenario_id: str) -> bool: db = self.session.get(ScenarioDB, scenario_id) if not db: return False try: ScenarioModelBindingRepository(self.session).delete_for_scenario(scenario_id) self.session.delete(db) self.session.commit() except Exception: self.session.rollback() raise return True def name_map(self) -> dict[str, str]: """scenario_id → 名称映射:报告 / 时间线 / 列表等读路径共用的场景名取法。""" return {sid: name for sid, name in self.session.exec(select(ScenarioDB.id, ScenarioDB.name)).all()} class CampaignRunClaimStatus(str, Enum): """Outcome of a durable Campaign child-Run claim.""" CLAIMED = "claimed" EXISTING = "existing" NOT_FOUND = "not_found" CONFLICT = "conflict" @dataclass(frozen=True) class CampaignRunClaimResult: """Typed result for the idempotent child-Run claim seam.""" status: CampaignRunClaimStatus run: Optional[EvalRun] = None @property def accepted(self) -> bool: return self.status in (CampaignRunClaimStatus.CLAIMED, CampaignRunClaimStatus.EXISTING) @property def created(self) -> bool: return self.status is CampaignRunClaimStatus.CLAIMED class RunRepository(BaseRepository[EvalRun, EvalRunDB]): """Repository for evaluation runs.""" _table = EvalRunDB _order_by = "started_at" def claim_campaign_run( self, *, campaign_id: str, scenario_id: str, scenario_version: int, plan_index: int, occurrence_index: int, ) -> "CampaignRunClaimResult": """Claim one durable child Run occurrence while its Campaign is running. The conditional ``INSERT ... SELECT`` makes the Campaign status check part of the write. The composite unique index turns retries or a competing scheduler into an ``EXISTING`` result instead of a second pending Run. No external messages are sent by this operation. """ if plan_index < 0 or occurrence_index < 0: raise ValueError("plan_index and occurrence_index must be non-negative") run_id = new_uuid() now = utc_now() statement = sql_insert(EvalRunDB).from_select( [ "id", "target_id", "scenario_id", "scenario_version", "campaign_id", "campaign_plan_index", "campaign_occurrence_index", "status", "triggered_by", "started_at", ], select( literal(run_id), CampaignDB.target_id, literal(scenario_id), literal(scenario_version), CampaignDB.id, literal(plan_index), literal(occurrence_index), literal(RunStatus.PENDING.value), literal(RunTrigger.CAMPAIGN.value), literal(now), ).where( CampaignDB.id == campaign_id, CampaignDB.status == CampaignStatus.RUNNING.value, CampaignDB.target_id.is_not(None), ), ) try: result = self.session.exec(statement) if result.rowcount == 1: self.session.commit() db = self.session.get(EvalRunDB, run_id) return CampaignRunClaimResult(CampaignRunClaimStatus.CLAIMED, self._from_db(db)) self.session.rollback() except IntegrityError: self.session.rollback() except Exception: self.session.rollback() raise existing = self._get_campaign_identity(campaign_id, plan_index, occurrence_index) campaign = self.session.get(CampaignDB, campaign_id) if campaign is None: return CampaignRunClaimResult(CampaignRunClaimStatus.NOT_FOUND) if campaign.status != CampaignStatus.RUNNING.value: return CampaignRunClaimResult(CampaignRunClaimStatus.CONFLICT) if existing is not None: return CampaignRunClaimResult(CampaignRunClaimStatus.EXISTING, self._from_db(existing)) return CampaignRunClaimResult(CampaignRunClaimStatus.CONFLICT) def _get_campaign_identity( self, campaign_id: str, plan_index: int, occurrence_index: int ) -> Optional[EvalRunDB]: statement = select(EvalRunDB).where( EvalRunDB.campaign_id == campaign_id, EvalRunDB.campaign_plan_index == plan_index, EvalRunDB.campaign_occurrence_index == occurrence_index, ) return self.session.exec(statement).first() def _to_db(self, run: EvalRun) -> EvalRunDB: db = EvalRunDB( id=run.id, target_id=run.target_id, scenario_id=run.scenario_id, scenario_version=run.scenario_version, campaign_id=run.campaign_id, campaign_plan_index=run.campaign_plan_index, campaign_occurrence_index=run.campaign_occurrence_index, status=run.status.value, triggered_by=run.triggered_by.value, started_at=run.started_at, completed_at=run.completed_at, ) if run.summary is not None: db.set_summary(run.summary.model_dump(mode="json")) return db def _from_db(self, db: EvalRunDB) -> EvalRun: return EvalRun( id=db.id, target_id=db.target_id, scenario_id=db.scenario_id, scenario_version=db.scenario_version or 1, campaign_id=db.campaign_id, campaign_plan_index=db.campaign_plan_index, campaign_occurrence_index=db.campaign_occurrence_index, status=db.status, triggered_by=db.triggered_by or "manual", started_at=db.started_at, completed_at=db.completed_at, summary=db.get_summary(), ) def list_by_campaign(self, campaign_id: str) -> list[EvalRun]: statement = select(EvalRunDB).where(EvalRunDB.campaign_id == campaign_id).order_by(EvalRunDB.started_at) return [self._from_db(r) for r in self.session.exec(statement).all()] def list_by_campaigns(self, campaign_ids: list[str]) -> dict[str, list[EvalRun]]: """Load child Runs for multiple Campaigns in one query.""" grouped = {campaign_id: [] for campaign_id in campaign_ids} if not campaign_ids: return grouped statement = ( select(EvalRunDB) .where(EvalRunDB.campaign_id.in_(campaign_ids)) # type: ignore[union-attr] .order_by(EvalRunDB.started_at) ) for row in self.session.exec(statement).all(): if row.campaign_id in grouped: grouped[row.campaign_id].append(self._from_db(row)) return grouped def mark_orphans_failed(self) -> int: """Fail process-orphaned Runs while preserving recoverable child claims. Every RUNNING Run may already have sent messages and must not replay. A PENDING Run is preserved only when it has a complete Campaign plan identity and its parent Campaign is still running; the Campaign loop can safely resume those claims before any external send occurred. """ statement = select(EvalRunDB).where(EvalRunDB.status.in_(["running", "pending"])) # type: ignore[attr-defined] candidates = self.session.exec(statement).all() failed: list[EvalRunDB] = [] for db in candidates: if db.status == RunStatus.PENDING.value and self._is_recoverable_campaign_pending(db): continue self._set_interrupted(db) failed.append(db) if failed: self.session.commit() return len(failed) def _is_recoverable_campaign_pending(self, db: EvalRunDB) -> bool: if ( db.campaign_id is None or db.campaign_plan_index is None or db.campaign_occurrence_index is None ): return False campaign = self.session.get(CampaignDB, db.campaign_id) return campaign is not None and campaign.status == CampaignStatus.RUNNING.value def _set_interrupted(self, db: EvalRunDB) -> None: db.status = RunStatus.FAILED.value db.completed_at = db.completed_at or utc_now() summary = db.get_summary() or {} summary["error"] = {"code": "interrupted", "message": "服务重启导致评测中断"} db.set_summary(summary) self.session.add(db) def list_pending_campaign_children(self, campaign_id: str) -> list[EvalRun]: """Return durable pending claims in deterministic plan order.""" statement = ( select(EvalRunDB) .where( EvalRunDB.campaign_id == campaign_id, EvalRunDB.campaign_plan_index.is_not(None), EvalRunDB.campaign_occurrence_index.is_not(None), EvalRunDB.status == RunStatus.PENDING.value, ) .order_by(EvalRunDB.campaign_plan_index, EvalRunDB.campaign_occurrence_index) ) return [self._from_db(db) for db in self.session.exec(statement).all()] def mark_campaign_running_interrupted(self, campaign_id: str) -> list[str]: """Fail running child Runs from a previous process without replaying.""" statement = select(EvalRunDB).where( EvalRunDB.campaign_id == campaign_id, EvalRunDB.campaign_plan_index.is_not(None), EvalRunDB.campaign_occurrence_index.is_not(None), EvalRunDB.status == RunStatus.RUNNING.value, ) rows = list(self.session.exec(statement).all()) for db in rows: self._set_interrupted(db) if rows: self.session.commit() return [db.id or "" for db in rows] def mark_pending_campaign_child_interrupted(self, run_id: str) -> Optional[EvalRun]: """Fail one still-pending child claim that cannot be safely recovered.""" db = self.session.get(EvalRunDB, run_id) if db is None or db.status != RunStatus.PENDING.value: return self._from_db(db) if db is not None else None self._set_interrupted(db) self.session.commit() self.session.refresh(db) return self._from_db(db) def update(self, run: EvalRun) -> Optional[EvalRun]: existing = self.session.get(EvalRunDB, run.id) if not existing: return None existing.target_id = run.target_id existing.scenario_id = run.scenario_id existing.scenario_version = run.scenario_version existing.campaign_id = run.campaign_id existing.campaign_plan_index = run.campaign_plan_index existing.campaign_occurrence_index = run.campaign_occurrence_index existing.status = run.status.value existing.triggered_by = run.triggered_by.value existing.completed_at = run.completed_at if run.summary is not None: existing.set_summary(run.summary.model_dump(mode="json")) self.session.add(existing) self.session.commit() self.session.refresh(existing) return self._from_db(existing) def get_turns(self, run_id: str) -> list[TurnDB]: statement = select(TurnDB).where(TurnDB.run_id == run_id).order_by(TurnDB.sent_at) return list(self.session.exec(statement).all()) def get_results(self, run_id: str) -> list[EvalResult]: statement = select(EvalResultDB).where(EvalResultDB.run_id == run_id) return [_result_from_db(r) for r in self.session.exec(statement).all()] class CampaignWriteStatus(str, Enum): """Outcome of a conditional Campaign lifecycle write.""" APPLIED = "applied" NOT_FOUND = "not_found" CONFLICT = "conflict" @dataclass(frozen=True) class CampaignWriteResult: """Typed result returned by Campaign compare-and-set operations.""" status: CampaignWriteStatus campaign: Optional[Campaign] = None @property def applied(self) -> bool: return self.status is CampaignWriteStatus.APPLIED class CampaignRepository(BaseRepository[Campaign, CampaignDB]): """Repository for evaluation campaigns (评估活动).""" _table = CampaignDB _order_by = "created_at" def _to_db(self, campaign: Campaign) -> CampaignDB: db = CampaignDB( id=campaign.id, name=campaign.name, target_id=campaign.target_id, window_seconds=campaign.window_seconds, time_scale=campaign.time_scale, status=campaign.status.value, started_at=campaign.started_at, completed_at=campaign.completed_at, created_at=campaign.created_at, analysis_model_config_id=campaign.analysis_model_config_id, last_patrolled_at=campaign.last_patrolled_at, ) db.set_plan([entry.model_dump(mode="json") for entry in campaign.plan]) if campaign.summary: db.set_summary(campaign.summary.model_dump(mode="json")) if campaign.exploration_seeds is not None: db.set_exploration_seeds(campaign.exploration_seeds.model_dump(mode="json")) if campaign.exploration_budget is not None: db.set_exploration_budget(campaign.exploration_budget.model_dump(mode="json")) return db def _from_db(self, db: CampaignDB) -> Campaign: return Campaign( id=db.id, name=db.name, target_id=db.target_id, window_seconds=db.window_seconds, time_scale=db.time_scale, plan=db.get_plan(), status=db.status, started_at=db.started_at, completed_at=db.completed_at, created_at=db.created_at, summary=db.get_summary(), analysis_model_config_id=db.analysis_model_config_id, exploration_seeds=db.get_exploration_seeds(), exploration_budget=db.get_exploration_budget(), last_patrolled_at=db.last_patrolled_at, ) def update(self, campaign: Campaign) -> Optional[Campaign]: existing = self.session.get(CampaignDB, campaign.id) if not existing: return None existing.name = campaign.name existing.target_id = campaign.target_id existing.window_seconds = campaign.window_seconds existing.time_scale = campaign.time_scale existing.set_plan([entry.model_dump(mode="json") for entry in campaign.plan]) existing.status = campaign.status.value existing.started_at = campaign.started_at existing.completed_at = campaign.completed_at existing.analysis_model_config_id = campaign.analysis_model_config_id existing.last_patrolled_at = campaign.last_patrolled_at if campaign.summary is not None: existing.set_summary(campaign.summary.model_dump(mode="json")) if campaign.exploration_seeds is not None: existing.set_exploration_seeds(campaign.exploration_seeds.model_dump(mode="json")) if campaign.exploration_budget is not None: existing.set_exploration_budget(campaign.exploration_budget.model_dump(mode="json")) self.session.add(existing) self.session.commit() self.session.refresh(existing) return self._from_db(existing) def touch_patrol_watermark(self, campaign_id: str, at: datetime) -> None: """窄口径原子更新:只写巡检水位,不覆写并发的 status / summary 变更。""" db = self.session.get(CampaignDB, campaign_id) if not db: return db.last_patrolled_at = at self.session.add(db) self.session.commit() self.session.refresh(db) def mark_cancelled(self, campaign_id: str, at: datetime) -> Optional[Campaign]: """窄口径终态迁移:只写 status + completed_at,不抹掉水位与调度进度。""" db = self.session.get(CampaignDB, campaign_id) if not db: return None db.status = CampaignStatus.CANCELLED.value db.completed_at = at self.session.add(db) self.session.commit() self.session.refresh(db) return self._from_db(db) def _compare_and_set_status( self, campaign_id: str, *, expected_statuses: tuple[CampaignStatus, ...], target_status: CampaignStatus, at: datetime, settle_exploration: bool = False, ) -> CampaignWriteResult: """Apply one conditional Campaign transition and optional settlement. Final Campaign status and expiration of running exploration sessions share one transaction. A settlement failure therefore rolls the Campaign back to its prior runnable state. """ values: dict[str, object] = {"status": target_status.value} if target_status is CampaignStatus.RUNNING: values["started_at"] = at if target_status in (CampaignStatus.CANCELLED, CampaignStatus.COMPLETED, CampaignStatus.FAILED): values["completed_at"] = at statement = ( sql_update(CampaignDB) .where( CampaignDB.id == campaign_id, CampaignDB.status.in_([status.value for status in expected_statuses]), ) .values(**values) ) try: result = self.session.exec(statement) if result.rowcount != 1: self.session.rollback() db = self.session.get(CampaignDB, campaign_id) status = CampaignWriteStatus.NOT_FOUND if db is None else CampaignWriteStatus.CONFLICT return CampaignWriteResult(status=status, campaign=self._from_db(db) if db else None) if settle_exploration: self.session.exec( sql_update(ExplorationSessionDB) .where( ExplorationSessionDB.campaign_id == campaign_id, ExplorationSessionDB.status == ExplorationSessionStatus.RUNNING.value, ) .values(status=ExplorationSessionStatus.EXPIRED.value, closed_at=at) ) self.session.commit() except Exception: self.session.rollback() raise self.session.expire_all() db = self.session.get(CampaignDB, campaign_id) return CampaignWriteResult( status=CampaignWriteStatus.APPLIED, campaign=self._from_db(db) if db else None, ) def cancel_if_active(self, campaign_id: str, at: datetime) -> CampaignWriteResult: """Atomically cancel and expire running exploration sessions.""" return self._compare_and_set_status( campaign_id, expected_statuses=(CampaignStatus.PLANNED, CampaignStatus.RUNNING), target_status=CampaignStatus.CANCELLED, at=at, settle_exploration=True, ) def complete_if_running(self, campaign_id: str, at: datetime) -> CampaignWriteResult: """Atomically complete and expire running exploration sessions.""" return self._compare_and_set_status( campaign_id, expected_statuses=(CampaignStatus.RUNNING,), target_status=CampaignStatus.COMPLETED, at=at, settle_exploration=True, ) def start_if_planned(self, campaign_id: str, at: datetime) -> CampaignWriteResult: """Atomically stamp a planned Campaign as running.""" return self._compare_and_set_status( campaign_id, expected_statuses=(CampaignStatus.PLANNED,), target_status=CampaignStatus.RUNNING, at=at, ) def save_scheduler_state(self, campaign_id: str, summary: CampaignSummary) -> None: """窄口径调度进度持久化:只写 summary 列,不覆写并发的状态 / 水位变更。""" db = self.session.get(CampaignDB, campaign_id) if not db: return db.set_summary(summary.model_dump(mode="json")) self.session.add(db) self.session.commit() self.session.refresh(db) 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("服务重启导致周期对比生成中断") class ResultRepository: """Repository for evaluation results.""" def __init__(self, session: Optional[Session] = None): self.session = session or get_session() def save_turn(self, turn) -> TurnDB: db = TurnDB( id=turn.id, run_id=turn.run_id, case_id=turn.case_id, round_index=turn.round_index, question_msg_id=turn.question_msg_id, sent_at=turn.sent_at, received_at=turn.received_at, latency_ms=turn.latency_ms, ) db.set_sent_message(turn.sent_message) db.set_reply(turn.reply) self.session.add(db) self.session.commit() self.session.refresh(db) return db def update_turn_exchange( self, turn_id: str, *, question_msg_id: Optional[str], reply: Optional[dict], received_at: Optional[datetime], latency_ms: Optional[int], ) -> Optional[TurnDB]: """Attach exchange facts to an already-persisted Turn. This is deliberately narrower than updating a whole Turn: the sent message, case identity, ordering and send timestamp remain untouched. A caller can therefore commit the sent fact before polling and safely complete the same ledger row after a reply, timeout or poll failure. """ db = self.session.get(TurnDB, turn_id) if db is None: return None db.question_msg_id = question_msg_id db.set_reply(reply) db.received_at = received_at db.latency_ms = latency_ms self.session.add(db) self.session.commit() self.session.refresh(db) return db def save_result(self, result: EvalResult) -> EvalResult: db = _result_to_db(result) self.session.add(db) self.session.commit() self.session.refresh(db) return _result_from_db(db) class ExplorationSessionRepository(BaseRepository[ExplorationSession, ExplorationSessionDB]): """Repository for virtual-user exploration sessions (探索会话).""" _table = ExplorationSessionDB _order_by = "created_at" def _copy_mutable(self, db: ExplorationSessionDB, session_obj: ExplorationSession) -> None: db.goal = session_obj.goal db.status = session_obj.status.value db.triggered_by = session_obj.triggered_by.value db.turn_count = session_obj.turn_count db.error = session_obj.error db.closed_at = session_obj.closed_at db.set_persona(session_obj.persona) if session_obj.seed_ref is not None: db.set_seed_ref(session_obj.seed_ref) if session_obj.experience is not None: db.set_experience(session_obj.experience) if session_obj.judge_review is not None: db.set_judge_review(session_obj.judge_review) def _to_db(self, session_obj: ExplorationSession) -> ExplorationSessionDB: db = ExplorationSessionDB( id=session_obj.id, campaign_id=session_obj.campaign_id, target_id=session_obj.target_id, created_at=session_obj.created_at, ) self._copy_mutable(db, session_obj) return db def _from_db(self, db: ExplorationSessionDB) -> ExplorationSession: return ExplorationSession( id=db.id, campaign_id=db.campaign_id, target_id=db.target_id, persona=db.get_persona(), goal=db.goal, seed_ref=db.get_seed_ref(), status=db.status, triggered_by=db.triggered_by, experience=db.get_experience(), judge_review=db.get_judge_review(), turn_count=db.turn_count, error=db.error, created_at=db.created_at, closed_at=db.closed_at, ) def update(self, session_obj: ExplorationSession) -> Optional[ExplorationSession]: existing = self.session.get(ExplorationSessionDB, session_obj.id) if not existing: return None self._copy_mutable(existing, session_obj) self.session.add(existing) self.session.commit() self.session.refresh(existing) return self._from_db(existing) def list_by_campaign(self, campaign_id: str) -> list[ExplorationSession]: statement = ( select(ExplorationSessionDB) .where(ExplorationSessionDB.campaign_id == campaign_id) .order_by(ExplorationSessionDB.created_at) ) return [self._from_db(r) for r in self.session.exec(statement).all()] def expire_running_sessions(self, campaign_id: str) -> int: """Expire every still-running exploration session of a finalized campaign. Returns the number of sessions expired. Completed/failed sessions keep their evidence untouched. """ expired = 0 for session_obj in self.list_by_campaign(campaign_id): if session_obj.status != ExplorationSessionStatus.RUNNING: continue session_obj.status = ExplorationSessionStatus.EXPIRED session_obj.closed_at = utc_now() self.update(session_obj) expired += 1 return expired class ExplorationMessageRepository: """Append-only repository for exploration session chat rows.""" def __init__(self, session: Optional[Session] = None): self.session = session or get_session() def save_message(self, message: ExplorationMessage) -> ExplorationMessage: db = ExplorationMessageDB( id=message.id, session_id=message.session_id, round_index=message.round_index, role=message.role, content=message.content, latency_ms=message.latency_ms, created_at=message.created_at, ) self.session.add(db) self.session.commit() self.session.refresh(db) message.id = db.id return message def list_by_session(self, session_id: str) -> list[ExplorationMessage]: statement = ( select(ExplorationMessageDB) .where(ExplorationMessageDB.session_id == session_id) .order_by(ExplorationMessageDB.created_at) ) return [ ExplorationMessage( id=r.id, session_id=r.session_id, round_index=r.round_index, role=r.role, content=r.content, latency_ms=r.latency_ms, created_at=r.created_at, ) for r in self.session.exec(statement).all() ]