"""Repository for evaluation campaigns (评估活动).""" from dataclasses import dataclass from datetime import datetime from enum import Enum from typing import Optional from sqlalchemy import update as sql_update from agenteval.exploration.models import ExplorationSessionStatus from agenteval.models import Campaign, CampaignStatus, CampaignSummary from agenteval.storage.db import CampaignDB, ExplorationSessionDB from agenteval.storage.repository.base import BaseRepository 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)