439 lines
16 KiB
Python
439 lines
16 KiB
Python
"""Agent 评测脚本:执行单轮和多轮评测集,输出 JSON + Markdown 报告。
|
||
|
||
用法:
|
||
python scripts/eval/evaluate_agent.py --dataset docs/开发文档/51-Agent单轮评测集.md
|
||
python scripts/eval/evaluate_agent.py --dataset docs/开发文档/52-Agent多轮场景评测集.md
|
||
python scripts/eval/evaluate_agent.py --all
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import json
|
||
import re
|
||
import sys
|
||
from dataclasses import dataclass
|
||
from datetime import datetime
|
||
from pathlib import Path
|
||
from typing import Any
|
||
|
||
|
||
@dataclass
|
||
class EvalResult:
|
||
"""单条评测结果"""
|
||
case_id: str
|
||
agent_type: str
|
||
user_input: str
|
||
expected_tool: str
|
||
expected_params: dict[str, Any]
|
||
actual_tool: str | None
|
||
actual_params: dict[str, Any]
|
||
actual_response: str
|
||
tool_match: bool
|
||
param_match: bool
|
||
passed: bool
|
||
error: str | None = None
|
||
|
||
|
||
class AgentEvaluator:
|
||
"""Agent 评测器"""
|
||
|
||
def __init__(self, base_url: str = "http://127.0.0.1:8000"):
|
||
self.base_url = base_url
|
||
self.results: list[EvalResult] = []
|
||
|
||
def _login(self, username: str, password: str) -> str:
|
||
"""登录获取 token"""
|
||
import httpx
|
||
resp = httpx.post(
|
||
f"{self.base_url}/api/v1/auth/login",
|
||
json={"username": username, "password": password},
|
||
)
|
||
resp.raise_for_status()
|
||
return resp.json()["data"]["access_token"]
|
||
|
||
def _chat(self, token: str, message: str, agent_type: str) -> dict:
|
||
"""调用 chat API"""
|
||
import httpx
|
||
resp = httpx.post(
|
||
f"{self.base_url}/api/v1/chat",
|
||
headers={
|
||
"Authorization": f"Bearer {token}",
|
||
"X-Agent-Type": agent_type,
|
||
"Content-Type": "application/json",
|
||
},
|
||
json={"message": message},
|
||
timeout=30.0,
|
||
)
|
||
resp.raise_for_status()
|
||
return resp.json()["data"]
|
||
|
||
def parse_single_turn_cases(self, md_path: Path) -> list[dict]:
|
||
"""解析单轮评测集 Markdown"""
|
||
content = md_path.read_text(encoding="utf-8")
|
||
cases = []
|
||
|
||
# 匹配表格行:| ID | 用户输入 | 预期工具 | 预期参数 | 预期回复要点 |
|
||
# 或:| ID | 用户输入 | 预期工具 | 预期回复要点 |
|
||
table_pattern = re.compile(
|
||
r'\|\s*([A-Z]+-\d+)\s*\|\s*(.+?)\s*\|\s*(\w+)\s*\|\s*(.+?)\s*\|'
|
||
)
|
||
|
||
agent_type_map = {
|
||
"C": "customer",
|
||
"A": "advisor",
|
||
"AN": "analyst",
|
||
"X": "customer", # 通用场景默认用 customer
|
||
}
|
||
|
||
for match in table_pattern.finditer(content):
|
||
case_id = match.group(1)
|
||
user_input = match.group(2).strip()
|
||
expected_tool = match.group(3).strip()
|
||
rest = match.group(4).strip()
|
||
|
||
# 尝试解析参数
|
||
expected_params = {}
|
||
if rest.startswith("{"):
|
||
try:
|
||
expected_params = json.loads(rest)
|
||
rest = ""
|
||
except json.JSONDecodeError:
|
||
pass
|
||
|
||
agent_prefix = case_id.split("-")[0]
|
||
agent_type = agent_type_map.get(agent_prefix, "customer")
|
||
|
||
cases.append({
|
||
"case_id": case_id,
|
||
"agent_type": agent_type,
|
||
"user_input": user_input,
|
||
"expected_tool": expected_tool if expected_tool != "无" else None,
|
||
"expected_params": expected_params,
|
||
"notes": rest,
|
||
})
|
||
|
||
return cases
|
||
|
||
def parse_multi_turn_cases(self, md_path: Path) -> list[dict]:
|
||
"""解析多轮场景评测集 Markdown"""
|
||
content = md_path.read_text(encoding="utf-8")
|
||
scenarios = []
|
||
|
||
# 按场景分割
|
||
scenario_pattern = re.compile(
|
||
r'### 场景 ([A-Z]+-M\d+):(.+?)\n.*?```text\n(.*?)```',
|
||
re.DOTALL
|
||
)
|
||
|
||
agent_type_map = {
|
||
"C": "customer",
|
||
"A": "advisor",
|
||
"AN": "analyst",
|
||
"X": "customer",
|
||
}
|
||
|
||
for match in scenario_pattern.finditer(content):
|
||
scenario_id = match.group(1)
|
||
scenario_name = match.group(2).strip()
|
||
scenario_content = match.group(3)
|
||
|
||
# 解析轮次
|
||
turns = []
|
||
turn_pattern = re.compile(
|
||
r'轮次 (\d+):\n\s*用户:(.+?)\n\s*预期:(.+?)(?:\n\s*验证点:(.+?))?(?=\n\n轮次|\Z)',
|
||
re.DOTALL
|
||
)
|
||
|
||
for turn_match in turn_pattern.finditer(scenario_content):
|
||
turn_num = int(turn_match.group(1))
|
||
user_input = turn_match.group(2).strip()
|
||
expected = turn_match.group(3).strip()
|
||
checkpoint = turn_match.group(4).strip() if turn_match.group(4) else ""
|
||
|
||
# 特殊处理:[模拟会话中断]
|
||
if user_input.startswith("["):
|
||
user_input = user_input.split("]")[-1].strip()
|
||
|
||
turns.append({
|
||
"turn": turn_num,
|
||
"user_input": user_input,
|
||
"expected": expected,
|
||
"checkpoint": checkpoint,
|
||
})
|
||
|
||
agent_prefix = scenario_id.split("-")[0]
|
||
agent_type = agent_type_map.get(agent_prefix, "customer")
|
||
|
||
scenarios.append({
|
||
"scenario_id": scenario_id,
|
||
"scenario_name": scenario_name,
|
||
"agent_type": agent_type,
|
||
"turns": turns,
|
||
})
|
||
|
||
return scenarios
|
||
|
||
def evaluate_single_turn(self, cases: list[dict], token: str) -> list[EvalResult]:
|
||
"""执行单轮评测"""
|
||
results = []
|
||
|
||
for case in cases:
|
||
try:
|
||
response = self._chat(token, case["user_input"], case["agent_type"])
|
||
actual_tool = response.get("tool_calls", [{}])[0].get("tool_name") if response.get("tool_calls") else None
|
||
actual_params = response.get("tool_calls", [{}])[0].get("tool_params", {}) if response.get("tool_calls") else {}
|
||
|
||
tool_match = actual_tool == case["expected_tool"]
|
||
param_match = True
|
||
if case["expected_params"]:
|
||
param_match = all(
|
||
actual_params.get(k) == v
|
||
for k, v in case["expected_params"].items()
|
||
)
|
||
|
||
results.append(EvalResult(
|
||
case_id=case["case_id"],
|
||
agent_type=case["agent_type"],
|
||
user_input=case["user_input"],
|
||
expected_tool=case["expected_tool"] or "",
|
||
expected_params=case["expected_params"],
|
||
actual_tool=actual_tool,
|
||
actual_params=actual_params,
|
||
actual_response=response.get("response", ""),
|
||
tool_match=tool_match,
|
||
param_match=param_match,
|
||
passed=tool_match and param_match,
|
||
))
|
||
except Exception as e:
|
||
results.append(EvalResult(
|
||
case_id=case["case_id"],
|
||
agent_type=case["agent_type"],
|
||
user_input=case["user_input"],
|
||
expected_tool=case["expected_tool"] or "",
|
||
expected_params=case["expected_params"],
|
||
actual_tool=None,
|
||
actual_params={},
|
||
actual_response="",
|
||
tool_match=False,
|
||
param_match=False,
|
||
passed=False,
|
||
error=str(e),
|
||
))
|
||
|
||
return results
|
||
|
||
def evaluate_multi_turn(self, scenarios: list[dict], token: str) -> list[EvalResult]:
|
||
"""执行多轮场景评测"""
|
||
results = []
|
||
|
||
for scenario in scenarios:
|
||
scenario_results = []
|
||
|
||
for turn in scenario["turns"]:
|
||
try:
|
||
response = self._chat(token, turn["user_input"], scenario["agent_type"])
|
||
|
||
# 简单验证:检查响应非空且无错误
|
||
has_response = bool(response.get("response"))
|
||
no_error = "error" not in response.get("response", "").lower()
|
||
|
||
scenario_results.append(EvalResult(
|
||
case_id=f"{scenario['scenario_id']}-T{turn['turn']}",
|
||
agent_type=scenario["agent_type"],
|
||
user_input=turn["user_input"],
|
||
expected_tool=turn["expected"][:50], # 截取前50字符
|
||
expected_params={},
|
||
actual_tool=None,
|
||
actual_params={},
|
||
actual_response=response.get("response", ""),
|
||
tool_match=has_response,
|
||
param_match=no_error,
|
||
passed=has_response and no_error,
|
||
))
|
||
except Exception as e:
|
||
scenario_results.append(EvalResult(
|
||
case_id=f"{scenario['scenario_id']}-T{turn['turn']}",
|
||
agent_type=scenario["agent_type"],
|
||
user_input=turn["user_input"],
|
||
expected_tool=turn["expected"][:50],
|
||
expected_params={},
|
||
actual_tool=None,
|
||
actual_params={},
|
||
actual_response="",
|
||
tool_match=False,
|
||
param_match=False,
|
||
passed=False,
|
||
error=str(e),
|
||
))
|
||
|
||
results.extend(scenario_results)
|
||
|
||
return results
|
||
|
||
def generate_report(self, results: list[EvalResult], output_path: Path) -> None:
|
||
"""生成评测报告"""
|
||
total = len(results)
|
||
passed = sum(1 for r in results if r.passed)
|
||
tool_correct = sum(1 for r in results if r.tool_match)
|
||
param_correct = sum(1 for r in results if r.param_match)
|
||
|
||
# 按 Agent 统计
|
||
by_agent = {}
|
||
for r in results:
|
||
if r.agent_type not in by_agent:
|
||
by_agent[r.agent_type] = {"total": 0, "passed": 0}
|
||
by_agent[r.agent_type]["total"] += 1
|
||
if r.passed:
|
||
by_agent[r.agent_type]["passed"] += 1
|
||
|
||
report = {
|
||
"timestamp": datetime.now().isoformat(),
|
||
"summary": {
|
||
"total": total,
|
||
"passed": passed,
|
||
"failed": total - passed,
|
||
"pass_rate": round(passed / total, 4) if total > 0 else 0,
|
||
"tool_accuracy": round(tool_correct / total, 4) if total > 0 else 0,
|
||
"param_accuracy": round(param_correct / total, 4) if total > 0 else 0,
|
||
},
|
||
"by_agent": by_agent,
|
||
"failed_cases": [
|
||
{
|
||
"case_id": r.case_id,
|
||
"user_input": r.user_input,
|
||
"expected_tool": r.expected_tool,
|
||
"actual_tool": r.actual_tool,
|
||
"error": r.error,
|
||
}
|
||
for r in results if not r.passed
|
||
],
|
||
}
|
||
|
||
# 保存 JSON
|
||
json_path = output_path.with_suffix(".json")
|
||
json_path.write_text(json.dumps(report, ensure_ascii=False, indent=2), encoding="utf-8")
|
||
|
||
# 生成 Markdown
|
||
md_lines = [
|
||
"# Agent 评测报告",
|
||
"",
|
||
f"**评测时间**:{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}",
|
||
"",
|
||
"## 总体统计",
|
||
"",
|
||
f"- **总用例数**:{total}",
|
||
f"- **通过数**:{passed}",
|
||
f"- **失败数**:{total - passed}",
|
||
f"- **通过率**:{report['summary']['pass_rate']:.2%}",
|
||
f"- **工具准确率**:{report['summary']['tool_accuracy']:.2%}",
|
||
f"- **参数准确率**:{report['summary']['param_accuracy']:.2%}",
|
||
"",
|
||
"## 按 Agent 统计",
|
||
"",
|
||
]
|
||
|
||
for agent, stats in by_agent.items():
|
||
rate = stats["passed"] / stats["total"] if stats["total"] > 0 else 0
|
||
md_lines.append(f"- **{agent}**:{stats['passed']}/{stats['total']} ({rate:.2%})")
|
||
|
||
md_lines.extend([
|
||
"",
|
||
"## 失败用例",
|
||
"",
|
||
])
|
||
|
||
for r in results:
|
||
if not r.passed:
|
||
md_lines.extend([
|
||
f"### {r.case_id}",
|
||
"",
|
||
f"- **输入**:{r.user_input}",
|
||
f"- **预期工具**:{r.expected_tool}",
|
||
f"- **实际工具**:{r.actual_tool or 'None'}",
|
||
])
|
||
if r.error:
|
||
md_lines.append(f"- **错误**:{r.error}")
|
||
md_lines.append("")
|
||
|
||
output_path.write_text("\n".join(md_lines), encoding="utf-8")
|
||
print(f"报告已生成:{output_path}")
|
||
print(f"JSON 报告:{json_path}")
|
||
|
||
|
||
def main():
|
||
parser = argparse.ArgumentParser(description="Agent 评测脚本")
|
||
parser.add_argument("--dataset", type=Path, help="评测集 Markdown 文件路径")
|
||
parser.add_argument("--all", action="store_true", help="运行所有评测集")
|
||
parser.add_argument("--base-url", default="http://127.0.0.1:8000", help="API 基础 URL")
|
||
parser.add_argument("--output", type=Path, help="报告输出路径")
|
||
args = parser.parse_args()
|
||
|
||
evaluator = AgentEvaluator(args.base_url)
|
||
|
||
# 登录获取 token
|
||
print("正在登录...")
|
||
try:
|
||
token = evaluator._login("advisor_test", "advisor_test")
|
||
except Exception as e:
|
||
print(f"登录失败:{e}")
|
||
print("请确保服务已启动:uvicorn app.main:app --reload")
|
||
sys.exit(1)
|
||
|
||
all_results = []
|
||
|
||
if args.all or args.dataset:
|
||
datasets = []
|
||
|
||
if args.all:
|
||
datasets = [
|
||
Path("docs/开发文档/51-Agent单轮评测集.md"),
|
||
Path("docs/开发文档/52-Agent多轮场景评测集.md"),
|
||
]
|
||
else:
|
||
datasets = [args.dataset]
|
||
|
||
for dataset_path in datasets:
|
||
if not dataset_path.exists():
|
||
print(f"评测集不存在:{dataset_path}")
|
||
continue
|
||
|
||
print(f"\n正在评测:{dataset_path}")
|
||
|
||
if "单轮" in dataset_path.name:
|
||
cases = evaluator.parse_single_turn_cases(dataset_path)
|
||
print(f"解析到 {len(cases)} 条单轮用例")
|
||
results = evaluator.evaluate_single_turn(cases, token)
|
||
else:
|
||
scenarios = evaluator.parse_multi_turn_cases(dataset_path)
|
||
print(f"解析到 {len(scenarios)} 个多轮场景")
|
||
results = evaluator.evaluate_multi_turn(scenarios, token)
|
||
|
||
all_results.extend(results)
|
||
print(f"完成:{sum(1 for r in results if r.passed)}/{len(results)} 通过")
|
||
|
||
if all_results:
|
||
output_path = args.output or Path("docs/评测报告/agent_eval_report.md")
|
||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||
evaluator.generate_report(all_results, output_path)
|
||
|
||
# 输出摘要
|
||
total = len(all_results)
|
||
passed = sum(1 for r in all_results if r.passed)
|
||
print(f"\n总计:{passed}/{total} 通过 ({passed/total:.2%})")
|
||
|
||
# 判断是否达标
|
||
if passed / total >= 0.9:
|
||
print("✓ 达到通过标准(≥90%)")
|
||
sys.exit(0)
|
||
else:
|
||
print("✗ 未达到通过标准(<90%)")
|
||
sys.exit(1)
|
||
else:
|
||
print("未执行任何评测")
|
||
sys.exit(0)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|