架构重构(候选 1-6): - storage/repository.py 按域拆分为包(target/scenario/run/campaign/result) - storage/db.py 按域拆分为包(eval/campaign/file/model_config/intelligent_eval) - intelligent_eval/lifecycle.py 按状态机阶段拆分为包 - services/runs.py 编排逻辑下沉 - Campaigns.tsx 拆分为 campaigns/ 子组件 测试补全(候选 7): 前端(+125 用例,107→232): - utils/ 纯函数:date/campaignTime/ruleLabels/fileTree/fileFormat/colors - stores/tabStore 状态管理 - 核心组件:FormDrawer/PageWrapper/ChatBubble/GeneratedMessages/SectionHeader/StatCard/TurnList - 业务组件:CaseBlock/CaseDetail/RuleOverview/WindowTimeline/RunList/TabBar/CampaignRunTimeline - 文件管理:FileCategoryTree/FileTable - hooks:sessionReducer/useFiles/useRunSession 后端(+38 用例,916→954): - targets API CRUD + 404 路径 - WebSocket 连接管理器 - proxy 头部重写(CSP/X-Frame-Options) - target 仓储 update 方法 - app 健康检查 + SPA 404 - scenarios 模板端点 + 404 - files API 边缘分支(404 场景 + 500 兜底) - files service update_category - 智能评估状态机迁移测试 门禁状态: - 前端:tsc 干净 + 232 passed - 后端:954 passed + ruff 全绿
231 lines
8.9 KiB
Python
231 lines
8.9 KiB
Python
"""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)
|