AgentEvalTool/backend/agenteval/evaluation/engine.py
sinohqb 0a47260237 feat(run): snapshot scenario version at run creation (ticket 04)
运行创建时快照场景考纲版本,三种触发来源(手动/AI 助手/CLI)一致;
迁移回填存量运行为其场景当前版本,孤儿运行回填 1。运行列表、
报告头与对比卡片展示 v{n} 版本标签。
2026-07-29 10:59:44 +08:00

627 lines
23 KiB
Python

"""Evaluation execution engine.
Async-first implementation: channels and LLM calls are awaited cooperatively,
so multiple cases can run concurrently and a run can be cancelled mid-flight
via an ``asyncio.Event`` cancel token.
"""
import asyncio
import uuid
from dataclasses import dataclass
from typing import Any, Callable, Optional
from agenteval.channels.base import EvalChannel
from agenteval.channels.factory import ChannelFactory
from agenteval.evaluation.rules import RuleResult, get_rule
from agenteval.model_gateway import ModelGateway
from agenteval.models import (
Case,
CaseType,
EvalResult,
EvalRun,
EvalTarget,
ModelCapability,
ModelPurpose,
RuleLogic,
RunStatus,
RunTrigger,
Scenario,
Turn,
)
from agenteval.services.model_configs import ModelConfigService, ModelRuntimeConfig
from agenteval.storage.db import get_session, utc_now
from agenteval.storage.repository import ResultRepository, RunRepository
from agenteval.utils.llm import extract_reply_text, parse_json_from_llm_text
# Progress callbacks may be sync or async; the engine awaits the result if
# it is a coroutine, otherwise treats it as a plain function.
ProgressCallback = Callable[[str, dict[str, Any]], Any]
class CancelledError(RuntimeError):
"""Raised inside the engine when the cancel token fires."""
@dataclass
class TimeoutConfig:
"""Per-operation timeouts (all in seconds)."""
poll_reply: float = 30.0
llm_generate: float = 60.0
def _build_send_message(content: str) -> dict[str, Any]:
return {
"msgType": "text",
"msgBody": {"content": content},
}
class EvalEngine:
"""Execute evaluation scenarios against targets.
The engine is async so the (slow, network-bound) channel and LLM calls
can be awaited cooperatively. Database writes remain synchronous for now
(SQLite + StaticPool); they are fast enough not to block the event loop
in practice.
"""
def __init__(
self,
target: EvalTarget,
scenario: Scenario,
session=None,
run_repo: Optional[RunRepository] = None,
result_repo: Optional[ResultRepository] = None,
cancel_token: Optional[asyncio.Event] = None,
timeout_config: Optional[TimeoutConfig] = None,
max_concurrent_cases: int = 1,
triggered_by: RunTrigger = RunTrigger.MANUAL,
):
self.target = target
self.scenario = scenario
self.channel: EvalChannel = ChannelFactory.create(target)
self.session = session or get_session()
self.run_repo = run_repo or RunRepository(self.session)
self.result_repo = result_repo or ResultRepository(self.session)
self.cancel_token = cancel_token or asyncio.Event()
self.timeout_config = timeout_config or TimeoutConfig()
self.triggered_by = triggered_by
self._case_semaphore = asyncio.Semaphore(max(1, max_concurrent_cases))
# Collects fatal case-level errors (e.g. dynamic message generation
# failures) so their cause is persisted into run.summary — not just
# emitted transiently over WebSocket.
self._case_errors: list[dict[str, str]] = []
self.model_gateway = ModelGateway(timeout=self.timeout_config.llm_generate)
self.model_service = ModelConfigService(self.session)
self._resolved_models: dict[ModelPurpose, ModelRuntimeConfig] = {}
# ── public entry point ────────────────────────────────────────────
async def run(
self,
progress_callback: Optional[ProgressCallback] = None,
existing_run: Optional[EvalRun] = None,
) -> EvalRun:
"""Run the evaluation and return the completed run record.
Cancellation is cooperative: set ``cancel_token`` and the engine will
mark the run as FAILED with ``cancelled_by_user`` at the next checkpoint.
"""
if existing_run:
run = existing_run
run.status = RunStatus.RUNNING
run.started_at = utc_now()
run = self.run_repo.update(run) or run
else:
run = EvalRun(
id=str(uuid.uuid4()),
target_id=self.target.id or "",
scenario_id=self.scenario.id or "",
scenario_version=self.scenario.version or 1,
status=RunStatus.RUNNING,
triggered_by=self.triggered_by,
started_at=utc_now(),
)
run = self.run_repo.create(run)
try:
total_cases = len(self.scenario.cases)
passed_cases = 0
failed_cases = 0
for idx, case in enumerate(self.scenario.cases, start=1):
self._check_cancel()
await self._emit(
progress_callback,
"case_start",
{
"index": idx,
"total": total_cases,
"case_id": case.id,
},
)
async with self._case_semaphore:
case_passed, rule_pass, rule_total = await self._run_case(
run,
case,
progress_callback,
)
if case_passed:
passed_cases += 1
else:
failed_cases += 1
await self._emit(
progress_callback,
"case_end",
{
"index": idx,
"total": total_cases,
"case_id": case.id,
"passed": case_passed,
"rule_pass_count": rule_pass,
"rule_total": rule_total,
},
)
results = self.run_repo.get_results(run.id)
total_rules = len(results)
passed_rules = sum(1 for r in results if r.passed)
summary = {
"total_cases": total_cases,
"passed_cases": passed_cases,
"failed_cases": failed_cases,
"total_rules": total_rules,
"passed_rules": passed_rules,
"pass_rate": round(passed_rules / total_rules, 4) if total_rules else 0.0,
}
# Surface fatal case-level errors (e.g. dynamic generation failures)
# so the report / DB record shows *why* a run produced no results.
if self._case_errors:
summary["case_errors"] = self._case_errors
if self._resolved_models:
summary["model_configs"] = {
purpose.value: config.snapshot() for purpose, config in self._resolved_models.items()
}
run.status = RunStatus.COMPLETED
run.completed_at = utc_now()
run.summary = summary
await self._emit(
progress_callback,
"run_completed",
{
"status": "completed",
"summary": summary,
},
)
except CancelledError:
run.status = RunStatus.FAILED
run.completed_at = utc_now()
run.summary = {
"error": {"code": "cancelled_by_user", "message": "评测已手动停止"},
}
await self._emit(
progress_callback,
"run_completed",
{
"status": "failed",
"reason": "cancelled",
"error": {"code": "cancelled_by_user", "message": "评测已手动停止"},
},
)
except Exception as exc:
run.status = RunStatus.FAILED
run.completed_at = utc_now()
run.summary = {"error": str(exc)}
await self._emit(progress_callback, "error", {"error": str(exc)})
await self._emit(
progress_callback,
"run_completed",
{
"status": "failed",
"error": str(exc),
},
)
raise
finally:
run = self.run_repo.update(run) or run
# Best-effort cleanup of the channel's HTTP client.
close = getattr(self.channel, "close", None)
if callable(close):
try:
result = close()
if asyncio.iscoroutine(result):
await result
except Exception:
pass
try:
self.session.close()
except Exception:
pass
return run
# ── case / turn execution ─────────────────────────────────────────
async def _run_case(
self,
run: EvalRun,
case: Case,
progress_callback: Optional[ProgressCallback],
) -> tuple[bool, int, int]:
"""Run a single case; returns (all_rules_passed, passed_rules, total_rules)."""
if case.type == CaseType.DYNAMIC:
generated = await self._generate_messages(case, progress_callback)
if not generated:
await self._emit(
progress_callback,
"error",
{
"error": "LLM 未能生成测试消息",
"case_id": case.id,
},
)
return False, 0, 0
case = case.model_copy(update={"messages": generated})
dialog: list[Turn] = []
for round_index, message in enumerate(case.messages, start=1):
self._check_cancel()
await self._emit(
progress_callback,
"turn_start",
{
"run_id": run.id,
"case_id": case.id,
"round": round_index,
"message": message,
},
)
sent_at = utc_now()
send_result = await self.channel.send(message)
if not send_result.ok:
turn = Turn(
id=str(uuid.uuid4()),
run_id=run.id,
case_id=case.id,
round_index=round_index,
sent_message=_build_send_message(message),
sent_at=sent_at,
)
self.result_repo.save_turn(turn)
await self._save_rule_results(run, case, turn, [], progress_callback)
await self._emit(
progress_callback,
"turn_error",
{
"case_id": case.id,
"round": round_index,
"error": send_result.error,
},
)
return False, 0, 0
try:
reply = await self.channel.poll_reply(
send_result.question_msg_id or "",
timeout=self.timeout_config.poll_reply,
)
except Exception as poll_exc:
received_at = utc_now()
turn = Turn(
id=str(uuid.uuid4()),
run_id=run.id,
case_id=case.id,
round_index=round_index,
sent_message=_build_send_message(message),
sent_at=sent_at,
question_msg_id=send_result.question_msg_id,
received_at=received_at,
)
self.result_repo.save_turn(turn)
await self._emit(
progress_callback,
"turn_error",
{
"case_id": case.id,
"round": round_index,
"error": f"poll_reply 异常: {poll_exc}",
},
)
return False, 0, 0
received_at = utc_now()
latency_ms = None
if sent_at and received_at:
latency_ms = int((received_at - sent_at).total_seconds() * 1000)
turn = Turn(
id=str(uuid.uuid4()),
run_id=run.id,
case_id=case.id,
round_index=round_index,
sent_message=_build_send_message(message),
sent_at=sent_at,
question_msg_id=send_result.question_msg_id,
reply=reply.raw_message if reply else None,
received_at=received_at,
latency_ms=latency_ms,
)
self.result_repo.save_turn(turn)
dialog.append(turn)
await self._emit(
progress_callback,
"turn_end",
{
"run_id": run.id,
"case_id": case.id,
"round": round_index,
"latency_ms": latency_ms,
"reply_text": extract_reply_text(reply.raw_message if reply else None),
},
)
if not dialog:
return False, 0, 0
return await self._save_rule_results(run, case, dialog[-1], dialog, progress_callback)
async def _save_rule_results(
self,
run: EvalRun,
case: Case,
turn: Turn,
dialog: list[Turn],
progress_callback: Optional[ProgressCallback],
) -> tuple[bool, int, int]:
"""Apply rules and save results; returns (case_passed, passed_count, total_count).
Judgement semantics (spec v0.5 / CONTEXT.md):
- Explicit rules are combined by case.rule_logic (ALL / ANY / WEIGHTED).
- Expectations always derive implicit checks, additive to explicit
rules. They are hard constraints: they never join the rule_logic
combination, and any implicit failure fails the case.
"""
from agenteval.models import EvalRuleConfig
rules_config: list[EvalRuleConfig] = list(case.eval_rules)
implicit_config: list[EvalRuleConfig] = []
if case.expectations.response_time_max_ms:
implicit_config.append(
EvalRuleConfig(
type="response_time",
params={
"max_ms": case.expectations.response_time_max_ms,
},
)
)
if case.expectations.keywords_include or case.expectations.keywords_exclude:
implicit_config.append(
EvalRuleConfig(
type="keyword_match",
params={
"keywords": case.expectations.keywords_include,
"exclude_keywords": case.expectations.keywords_exclude,
},
)
)
if not rules_config and not implicit_config:
# 连通用例:无任何判定标准,收到回复即通过
return True, 0, 0
passed_count = 0
total_count = 0
explicit_passed = 0
explicit_total = 0
weighted_score = 0.0
total_weight = 0.0
implicit_all_passed = True
all_rules = [(cfg, False) for cfg in rules_config] + [(cfg, True) for cfg in implicit_config]
for rule_config, is_implicit in all_rules:
purpose = {
"llm_score": ModelPurpose.JUDGE,
"semantic_similarity": ModelPurpose.EMBEDDING,
"safety": ModelPurpose.MODERATION,
}.get(rule_config.type)
try:
model_config = self._resolve_model(purpose) if purpose else None
rule = get_rule(
rule_config.type,
rule_config.params,
model_config=model_config,
gateway=self.model_gateway if model_config else None,
)
result = await rule.evaluate(case, dialog)
except Exception as exc:
result = RuleResult(passed=False, reason=f"模型配置解析失败: {exc}")
reason = f"[期望] {result.reason}" if is_implicit else result.reason
eval_result = EvalResult(
id=str(uuid.uuid4()),
run_id=run.id,
case_id=case.id,
turn_id=turn.id or "",
rule_type=rule_config.type,
passed=result.passed,
score=result.score,
reason=reason,
)
self.result_repo.save_result(eval_result)
total_count += 1
if result.passed:
passed_count += 1
if is_implicit:
if not result.passed:
implicit_all_passed = False
else:
explicit_total += 1
if result.passed:
explicit_passed += 1
# Weighted scoring: use rule score (default 1.0 if passed, 0.0 if failed)
score_val = result.score if result.score is not None else (1.0 if result.passed else 0.0)
weight = rule_config.weight
weighted_score += score_val * weight
total_weight += weight
await self._emit(
progress_callback,
"rule_result",
{
"run_id": run.id,
"case_id": case.id,
"rule_type": rule_config.type,
"passed": result.passed,
"score": result.score,
"reason": reason,
"weight": rule_config.weight,
},
)
# Combine explicit rules by rule_logic; no explicit rules → vacuously true.
logic = case.rule_logic
if not rules_config:
explicit_ok = True
elif logic == RuleLogic.ANY:
explicit_ok = explicit_passed > 0
elif logic == RuleLogic.WEIGHTED:
avg = weighted_score / total_weight if total_weight > 0 else 0.0
explicit_ok = avg >= case.rule_pass_threshold
else: # ALL and fallback
explicit_ok = explicit_passed == explicit_total
case_passed = explicit_ok and implicit_all_passed
return case_passed, passed_count, total_count
async def _generate_messages(
self,
case: Case,
progress_callback: Optional[ProgressCallback],
) -> list[str]:
"""Use LLM to generate test messages for dynamic cases."""
async def _fail(msg: str) -> list[str]:
# Persist the reason into run-level case_errors (surfaced in summary),
# not just a transient WebSocket emit that's lost after the run.
self._case_errors.append({"case_id": case.id, "stage": "generate_messages", "error": msg})
await self._emit(progress_callback, "error", {"error": msg, "case_id": case.id})
return []
turns = case.turns or 3
prompt = case.prompt or "请生成一些测试问题"
system_prompt = (
f"你需要扮演一个真实的用户/患者,根据以下要求生成 {turns} 条独立的测试问题。\n\n"
f"要求:{prompt}\n\n"
"输出格式要求:只输出一个 JSON 数组,包含 " + str(turns) + " 个字符串,每个字符串是一条消息。"
"不要输出任何解释、markdown 或其他内容。"
)
messages_payload = [
{"role": "system", "content": system_prompt},
{"role": "user", "content": f"请生成 {turns} 条测试消息"},
]
try:
model_config = self._resolve_model(ModelPurpose.GENERATOR)
if model_config:
content = await self.model_gateway.chat(model_config, messages_payload, temperature=0.7)
else:
content = await self._generate_messages_legacy(messages_payload)
try:
parsed = parse_json_from_llm_text(content)
except (ValueError, Exception) as parse_exc:
return await _fail(f"LLM 返回无法解析为数组: {parse_exc}")
if not isinstance(parsed, list):
return await _fail("LLM 返回的不是数组")
messages = [str(m) for m in parsed if isinstance(m, str) and m.strip()]
if not messages:
return await _fail("LLM 返回的消息为空")
await self._emit(
progress_callback,
"messages_generated",
{
"case_id": case.id,
"messages": messages,
},
)
return messages
except Exception as exc:
return await _fail(f"LLM 生成消息失败: {exc}")
async def _generate_messages_legacy(self, messages: list[dict[str, str]]) -> str:
"""Temporary fallback for scenarios not yet migrated to model bindings."""
import httpx
from agenteval.utils.llm import extract_content_from_llm_response
llm_config = self.scenario.llm_config
if not llm_config or not llm_config.get("api_url"):
raise ValueError("动态用例未绑定生成模型,且兼容 llm_config 缺少 api_url")
headers = {"Content-Type": "application/json"}
if llm_config.get("api_key"):
headers["Authorization"] = f"Bearer {llm_config['api_key']}"
payload = {
"model": llm_config.get("model", "doubao-seed-2.0-lite"),
"messages": messages,
"temperature": 0.7,
}
async with httpx.AsyncClient(timeout=self.timeout_config.llm_generate) as client:
response = await client.post(llm_config["api_url"], headers=headers, json=payload)
response.raise_for_status()
content = extract_content_from_llm_response(response.json())
if not content:
raise ValueError("LLM 返回内容为空或无法解析")
return content
def _resolve_model(self, purpose: ModelPurpose | None) -> ModelRuntimeConfig | None:
if purpose is None:
return None
if purpose in self._resolved_models:
return self._resolved_models[purpose]
config_id = self.scenario.model_bindings.get(purpose)
if not config_id:
return None
expected = {
ModelPurpose.GENERATOR: ModelCapability.CHAT,
ModelPurpose.JUDGE: ModelCapability.CHAT,
ModelPurpose.EMBEDDING: ModelCapability.EMBEDDING,
ModelPurpose.MODERATION: ModelCapability.MODERATION,
}[purpose]
runtime = self.model_service.resolve(config_id, expected)
self._resolved_models[purpose] = runtime
return runtime
# ── helpers ────────────────────────────────────────────────────────
def _check_cancel(self) -> None:
if self.cancel_token.is_set():
raise CancelledError("run cancelled")
async def _emit(
self,
callback: Optional[ProgressCallback],
event: str,
data: dict[str, Any],
) -> None:
if not callback:
return
try:
result = callback(event, data)
if asyncio.iscoroutine(result):
await result
except Exception:
pass