AgentEvalTool/backend/agenteval/evaluation/engine.py
sinohqb 867d4e3ff1 fix(engine): dynamic 生成失败原因持久化到 run.summary.case_errors
## 背景
用户反馈动态问诊评测「执行不下去」。诊断发现:dynamic 用例的 LLM 消息
生成 API 调用失败(0.2s 瞬间 failed,凭证/参数问题),引擎正确地标记
case 失败——但失败的具体原因(如 401 详情)只 emit 到 WebSocket,从不
写入 run.summary。导致 run 记录只有 total_rules:0 failed,DB/报告查不到
任何原因,用户和排查者都无从下手。

## 修复
- EvalEngine 新增 self._case_errors 收集致命的 case 级错误
- _generate_messages 的 8 个失败点统一走 _fail() helper:既 emit 到
  WebSocket,也记录到 _case_errors(含 case_id + stage + 具体 error)
- run() 汇总时把 _case_errors 写入 summary["case_errors"]
- 新增测试:dynamic 生成失败时 summary.case_errors 必须含原因(补上
  之前 KNOWN-2 记录的 _generate_messages 测试盲区)

## 注
这不是导致失败的 bug(失败源于外部 API 凭证/参数),而是让失败「可诊断」
的可用性修复。用户需自查 llm_config 的 api_key 是否有效/model 是否被
该端点接受。

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-17 16:29:51 +08:00

562 lines
20 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
import httpx
from agenteval.channels.base import EvalChannel
from agenteval.channels.factory import ChannelFactory
from agenteval.evaluation.rules import get_rule
from agenteval.models import (
Case,
CaseType,
EvalResult,
EvalRun,
EvalTarget,
RuleLogic,
RunStatus,
Scenario,
Turn,
)
from agenteval.storage.db import get_session, utc_now
from agenteval.storage.repository import ResultRepository, RunRepository
from agenteval.utils.llm import extract_content_from_llm_response, 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,
):
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._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]] = []
# ── 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 "",
status=RunStatus.RUNNING,
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
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).
Combination logic (case.rule_logic):
ALL — all rules must pass (default)
ANY — at least one rule must pass
WEIGHTED — weighted average score >= case.rule_pass_threshold
"""
from agenteval.models import EvalRuleConfig
rules_config: list[EvalRuleConfig] = list(case.eval_rules)
# If no explicit rules, derive implicit rules from expectations.
if not rules_config:
if case.expectations.response_time_max_ms:
rules_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:
rules_config.append(
EvalRuleConfig(
type="keyword_match",
params={
"keywords": case.expectations.keywords_include,
"exclude_keywords": case.expectations.keywords_exclude,
},
)
)
if not rules_config:
# No rules defined and no expectations → case passes with no checks
return True, 0, 0
passed_count = 0
total_count = 0
weighted_score = 0.0
total_weight = 0.0
for rule_config in rules_config:
rule = get_rule(rule_config.type, rule_config.params)
result = await rule.evaluate(case, dialog)
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=result.reason,
)
self.result_repo.save_result(eval_result)
total_count += 1
if result.passed:
passed_count += 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": result.reason,
"weight": weight,
},
)
# Determine case pass/fail based on rule_logic
logic = case.rule_logic
if logic == RuleLogic.ALL:
case_passed = passed_count == total_count
elif logic == RuleLogic.ANY:
case_passed = passed_count > 0
elif logic == RuleLogic.WEIGHTED:
avg = weighted_score / total_weight if total_weight > 0 else 0.0
case_passed = avg >= case.rule_pass_threshold
else:
case_passed = passed_count == total_count
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 []
llm_config = self.scenario.llm_config
if not llm_config:
return await _fail("动态用例需要配置 llm_config")
api_url = llm_config.get("api_url")
api_key = llm_config.get("api_key")
model = llm_config.get("model", "doubao-seed-2.0-lite")
if not api_url:
return await _fail("llm_config 缺少 api_url")
turns = case.turns or 3
prompt = case.prompt or "请生成一些测试问题"
system_prompt = (
f"你需要扮演一个真实的用户/患者,根据以下要求生成 {turns} 条独立的测试问题。\n\n"
f"要求:{prompt}\n\n"
"输出格式要求:只输出一个 JSON 数组,包含 " + str(turns) + " 个字符串,每个字符串是一条消息。"
"不要输出任何解释、markdown 或其他内容。"
)
headers = {"Content-Type": "application/json"}
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
payload = {
"model": model,
"messages": [
{"role": "system", "content": system_prompt},
{"role": "user", "content": f"请生成 {turns} 条测试消息"},
],
"temperature": 0.7,
}
try:
async with httpx.AsyncClient(timeout=self.timeout_config.llm_generate) as client:
resp = await client.post(api_url, headers=headers, json=payload)
resp.raise_for_status()
content = extract_content_from_llm_response(resp.json())
if not content:
return await _fail("LLM 返回内容为空或无法解析")
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}")
# ── 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