Introduce the 评估活动 (Campaign) aggregate above Run: a single-target, service-cycle window driving a static plan. Adds Campaign/CampaignPlanEntry models, CampaignDB table, nullable eval_runs.campaign_id, CampaignRepository, Alembic migration, and POST/GET /api/campaigns with validation. Ticket 01 of v0.6; no scheduling or child-run spawning yet (ADR-0003 v1).
424 lines
14 KiB
Python
424 lines
14 KiB
Python
"""Repository layer for database access."""
|
||
|
||
from typing import Optional
|
||
|
||
from sqlmodel import Session, select
|
||
|
||
from agenteval.models import Campaign, Case, EvalResult, EvalRun, EvalTarget, Scenario
|
||
from agenteval.services.model_configs import ModelConfigService
|
||
from agenteval.storage.db import (
|
||
CampaignDB,
|
||
EvalResultDB,
|
||
EvalRunDB,
|
||
EvalTargetDB,
|
||
ScenarioDB,
|
||
TurnDB,
|
||
get_session,
|
||
utc_now,
|
||
)
|
||
from agenteval.storage.model_config_repository import ScenarioModelBindingRepository
|
||
|
||
|
||
def _target_to_db(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 _target_from_db(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 _scenario_to_db(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 _scenario_from_db(db: ScenarioDB, session: Session) -> Scenario:
|
||
bindings = ScenarioModelBindingRepository(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 _run_to_db(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,
|
||
status=run.status.value,
|
||
triggered_by=run.triggered_by.value,
|
||
started_at=run.started_at,
|
||
completed_at=run.completed_at,
|
||
)
|
||
if run.summary:
|
||
db.set_summary(run.summary)
|
||
return db
|
||
|
||
|
||
def _run_from_db(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,
|
||
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 _campaign_to_db(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,
|
||
)
|
||
db.set_plan([entry.model_dump(mode="json") for entry in campaign.plan])
|
||
if campaign.summary:
|
||
db.set_summary(campaign.summary)
|
||
return db
|
||
|
||
|
||
def _campaign_from_db(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(),
|
||
)
|
||
|
||
|
||
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:
|
||
"""Repository for evaluation targets."""
|
||
|
||
def __init__(self, session: Optional[Session] = None):
|
||
self.session = session or get_session()
|
||
|
||
def list_all(self) -> list[EvalTarget]:
|
||
statement = select(EvalTargetDB).order_by(EvalTargetDB.created_at.desc())
|
||
return [_target_from_db(r) for r in self.session.exec(statement).all()]
|
||
|
||
def get(self, target_id: str) -> Optional[EvalTarget]:
|
||
db = self.session.get(EvalTargetDB, target_id)
|
||
return _target_from_db(db) if db else None
|
||
|
||
def create(self, target: EvalTarget) -> EvalTarget:
|
||
db = _target_to_db(target)
|
||
self.session.add(db)
|
||
self.session.commit()
|
||
self.session.refresh(db)
|
||
return _target_from_db(db)
|
||
|
||
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 _target_from_db(existing)
|
||
|
||
def delete(self, target_id: str) -> bool:
|
||
db = self.session.get(EvalTargetDB, target_id)
|
||
if not db:
|
||
return False
|
||
self.session.delete(db)
|
||
self.session.commit()
|
||
return True
|
||
|
||
|
||
class ScenarioRepository:
|
||
"""Repository for evaluation scenarios."""
|
||
|
||
def __init__(self, session: Optional[Session] = None):
|
||
self.session = session or get_session()
|
||
|
||
def list_all(self) -> list[Scenario]:
|
||
statement = select(ScenarioDB).order_by(ScenarioDB.created_at.desc())
|
||
return [_scenario_from_db(r, self.session) for r in self.session.exec(statement).all()]
|
||
|
||
def get(self, scenario_id: str) -> Optional[Scenario]:
|
||
db = self.session.get(ScenarioDB, scenario_id)
|
||
return _scenario_from_db(db, self.session) if db else None
|
||
|
||
def create(self, scenario: Scenario) -> Scenario:
|
||
db = _scenario_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 _scenario_from_db(db, self.session)
|
||
|
||
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 _scenario_from_db(existing, self.session)
|
||
|
||
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
|
||
|
||
|
||
class RunRepository:
|
||
"""Repository for evaluation runs."""
|
||
|
||
def __init__(self, session: Optional[Session] = None):
|
||
self.session = session or get_session()
|
||
|
||
def list_all(self) -> list[EvalRun]:
|
||
statement = select(EvalRunDB).order_by(EvalRunDB.started_at.desc())
|
||
return [_run_from_db(r) for r in self.session.exec(statement).all()]
|
||
|
||
def get(self, run_id: str) -> Optional[EvalRun]:
|
||
db = self.session.get(EvalRunDB, run_id)
|
||
return _run_from_db(db) if db else None
|
||
|
||
def create(self, run: EvalRun) -> EvalRun:
|
||
db = _run_to_db(run)
|
||
self.session.add(db)
|
||
self.session.commit()
|
||
self.session.refresh(db)
|
||
return _run_from_db(db)
|
||
|
||
def mark_orphans_failed(self) -> int:
|
||
"""服务启动时清理:把遗留的 running/pending 运行标记为 failed。
|
||
|
||
评测任务是进程内 asyncio 任务,服务重启后不会恢复;不清理则这些
|
||
运行永远停留在 running(僵尸运行)。
|
||
"""
|
||
statement = select(EvalRunDB).where(EvalRunDB.status.in_(["running", "pending"])) # type: ignore[attr-defined]
|
||
orphans = self.session.exec(statement).all()
|
||
for db in orphans:
|
||
db.status = "failed"
|
||
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)
|
||
if orphans:
|
||
self.session.commit()
|
||
return len(orphans)
|
||
|
||
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.status = run.status.value
|
||
existing.triggered_by = run.triggered_by.value
|
||
existing.completed_at = run.completed_at
|
||
if run.summary:
|
||
existing.set_summary(run.summary)
|
||
self.session.add(existing)
|
||
self.session.commit()
|
||
self.session.refresh(existing)
|
||
return _run_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()]
|
||
|
||
def delete(self, run_id: str) -> bool:
|
||
"""Delete a run. ORM-level cascade removes associated turns/results."""
|
||
db = self.session.get(EvalRunDB, run_id)
|
||
if not db:
|
||
return False
|
||
self.session.delete(db)
|
||
self.session.commit()
|
||
return True
|
||
|
||
|
||
class CampaignRepository:
|
||
"""Repository for evaluation campaigns (评估活动)."""
|
||
|
||
def __init__(self, session: Optional[Session] = None):
|
||
self.session = session or get_session()
|
||
|
||
def list_all(self) -> list[Campaign]:
|
||
statement = select(CampaignDB).order_by(CampaignDB.created_at.desc())
|
||
return [_campaign_from_db(r) for r in self.session.exec(statement).all()]
|
||
|
||
def get(self, campaign_id: str) -> Optional[Campaign]:
|
||
db = self.session.get(CampaignDB, campaign_id)
|
||
return _campaign_from_db(db) if db else None
|
||
|
||
def create(self, campaign: Campaign) -> Campaign:
|
||
db = _campaign_to_db(campaign)
|
||
self.session.add(db)
|
||
self.session.commit()
|
||
self.session.refresh(db)
|
||
return _campaign_from_db(db)
|
||
|
||
|
||
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 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)
|