AgentEvalTool/backend/agenteval/storage/repository.py
sinohqb 5db0ede4f4
Some checks failed
CI / test (push) Failing after 50s
refactor(judgement): converge case-pass decision into one deep module
「用例是否通过」此前散落 8 处且互相矛盾:engine 权威判定焊死在持久化里
不可单测;report 聚合/compare/markdown 各自从规则结果反推,规则还不一致
(markdown 用 all([]) 把故障用例误渲染成 )。

- 新增纯函数 evaluation/judgement.combine_case_outcome(RuleOutcome/
  CaseOutcome),判定组合脱离通道与 DB 可单测(判定矩阵 14 例)
- engine 调用它一次,逐用例权威结果写入 summary.case_outcomes(JSON,
  零迁移);report/compare/markdown 只读权威值,老 run fallback 反推
- 故障用例判 False(ADR-0002):修正 markdown 的  bug 与 compare 的
  None;顺带修 engine 连通用例无回复也算通过的 bug
- pass_rate 口径改为用例级(CONTEXT.md 词条),规则级保留在
  passed_rules/total_rules;CLI 对比标签同步更正
- 修 RunRepository.update 漏拷 scenario_version/triggered_by 的字段漂移
2026-07-29 19:45:02 +08:00

364 lines
12 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,
scenario_version=run.scenario_version,
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,
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 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.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 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)