AgentEvalTool/backend/agenteval/storage/repository.py
sinohqb 43e05ee38d feat(scenario): system-maintained syllabus version (ticket 03)
场景新增整型 version(迁移回填 1,batch mode)。仅考纲字段
(cases / model_bindings / llm_config)变更时升版,元数据编辑不升版,
API 传入的 version 被忽略(ADR-0001)。前端场景列表展示版本标签。
2026-07-29 10:42:35 +08:00

341 lines
11 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.

"""Repository layer for database access."""
from typing import Optional
from sqlmodel import Session, select
from agenteval.models import Case, EvalResult, EvalRun, EvalTarget, Scenario
from agenteval.services.model_configs import ModelConfigService
from agenteval.storage.db import (
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,
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,
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 _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 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.status = run.status.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 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)