## 新规则(共 6 种,增加 3 种) ### semantic_similarity - 调用 OpenAI 兼容 embedding API(asyncio.gather 并发两路请求) - 余弦相似度与 reference 比对,min_score 可配置(默认 0.7) - API 异常时明确返回失败原因,不隐藏错误 ### json_schema - 验证回复是否为合法 JSON(支持 markdown 代码块剥离) - required_keys / forbidden_keys / key_types 三维校验 - dot-path 支持嵌套字段("data.id") - strict_json=false 模式非阻断校验 ### safety - 双层检测:关键词黑名单(零延迟)+ 可选 moderation API - API 不可用时自动降级黑名单,不中止评测 - 支持自定义 flagged_categories ## 规则组合逻辑(rule_logic + rule_pass_threshold) - models.py: EvalRuleConfig 增加 weight 字段;Case 增加 rule_logic / rule_pass_threshold - models.py: 新增 RuleLogic 枚举(all / any / weighted) - engine._save_rule_results: 按 rule_logic 决定 case 通过/失败 - ALL:全部通过才通过(原有行为,向下兼容) - ANY:至少一条通过即通过 - WEIGHTED:加权平均分 >= rule_pass_threshold ## 测试(43 → 76,新增 33) - test_s2_rules_and_logic.py:3 个新规则的 pass/fail/边界/API 降级 + 5 个组合逻辑集成测试 Co-Authored-By: Claude <noreply@anthropic.com>
571 lines
20 KiB
Python
571 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))
|
|
|
|
# ── 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,
|
|
}
|
|
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."""
|
|
llm_config = self.scenario.llm_config
|
|
if not llm_config:
|
|
await self._emit(
|
|
progress_callback,
|
|
"error",
|
|
{
|
|
"error": "动态用例需要配置 llm_config",
|
|
},
|
|
)
|
|
return []
|
|
|
|
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:
|
|
await self._emit(progress_callback, "error", {"error": "llm_config 缺少 api_url"})
|
|
return []
|
|
|
|
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:
|
|
await self._emit(
|
|
progress_callback,
|
|
"error",
|
|
{
|
|
"error": "LLM 返回内容为空或无法解析",
|
|
},
|
|
)
|
|
return []
|
|
|
|
try:
|
|
parsed = parse_json_from_llm_text(content)
|
|
except (ValueError, Exception) as parse_exc:
|
|
await self._emit(
|
|
progress_callback,
|
|
"error",
|
|
{
|
|
"error": f"LLM 返回无法解析为数组: {parse_exc}",
|
|
},
|
|
)
|
|
return []
|
|
|
|
if not isinstance(parsed, list):
|
|
await self._emit(progress_callback, "error", {"error": "LLM 返回的不是数组"})
|
|
return []
|
|
|
|
messages = [str(m) for m in parsed if isinstance(m, str) and m.strip()]
|
|
if not messages:
|
|
await self._emit(progress_callback, "error", {"error": "LLM 返回的消息为空"})
|
|
return []
|
|
|
|
await self._emit(
|
|
progress_callback,
|
|
"messages_generated",
|
|
{
|
|
"case_id": case.id,
|
|
"messages": messages,
|
|
},
|
|
)
|
|
return messages
|
|
|
|
except Exception as exc:
|
|
await self._emit(progress_callback, "error", {"error": f"LLM 生成消息失败: {exc}"})
|
|
return []
|
|
|
|
# ── 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
|