Files
iAOP/core/llm-gateway/evaluate_routing.py

173 lines
7.9 KiB
Python
Raw Permalink 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.
# -*- coding: utf-8 -*-
"""iAOP-Core · LLM 网关 —— 敏感度路由准确率评测报告脚本(Issue #49)。
对应 PRD 5.4 / EPIC #6「敏感度路由准确率 ≥ 96.5%」:把评测集 JSON +
路由规则模板配置跑一遍 `SensitivityRouter.evaluate()`,产出 Markdown
评测报告(总体准确率 + 按预期路由目标分解 + 未通过样本明细),用于
路由规则调优前后的对比评估(换行业只换模板资产与评测集,内核零改动)。
用法示例:
# 使用内置演示评测集,输出报告到文件(utf-8)
python evaluate_routing.py --demo --output reports/routing.md
# 使用自定义评测集 JSON(见下方 DEMO_SAMPLES 字段说明)
python evaluate_routing.py --samples eval_set.json --output report.md
# 不指定 --output:报告打印到 stdout
python evaluate_routing.py --demo
评测集 JSON 格式(顶层为数组):
[
{
"query": "炉温当前是多少", // 用户 query
"expected": "local", // 期望路由目标 local / cloud / block
"dlp_blocked": false // 可选,模拟上游 DLP 出站拦截
}
]
测试:`python -m unittest discover -s tests -v`(在 core/llm-gateway 目录下执行)。
"""
from __future__ import annotations
import argparse
import json
import os
import sys
import types
from typing import List
# ---------------------------------------------------------------------------
# 包挂载(目录名含连字符,无法直接以包名 import;与 tests/_bootstrap.py 同款)
# ---------------------------------------------------------------------------
_LLM_GW_DIR = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, _LLM_GW_DIR)
if "llm_gateway" not in sys.modules:
_pkg = types.ModuleType("llm_gateway")
_pkg.__path__ = [_LLM_GW_DIR]
sys.modules["llm_gateway"] = _pkg
from llm_gateway.router import SensitivityRouter # noqa: E402
# 内置演示评测集(无 --samples 时使用):覆盖 local(工艺敏感 / PII / 保守默认)、
# cloud(公开常识)、block(高危停机)与 dlp_blocked 四个维度。
# 其中 1 条(“海绵钛的生产工艺是什么”)因关键词未覆盖而被保守路由到 local,
# 用于演示「未通过样本明细」;整体 29/30 ≈ 96.7%,仍达 ≥ 96.5% 验收线。
DEMO_SAMPLES: List[dict] = [
# ---- local:工艺敏感参数(模板规则 rt_proc_*,数据不出厂) ----
{"query": "炉温当前是多少", "expected": "local"},
{"query": "炉温偏高如何处理", "expected": "local"},
{"query": "氯气流量超限报警", "expected": "local"},
{"query": "加料比如何调整", "expected": "local"},
{"query": "钛纯度检测结果如何", "expected": "local"},
{"query": "查询氯气流量历史曲线", "expected": "local"},
{"query": "炉温报警原因分析", "expected": "local"},
{"query": "加料比偏差过大怎么办", "expected": "local"},
{"query": "钛纯度不达标原因分析", "expected": "local"},
{"query": "氯气流量调节阀开度", "expected": "local"},
{"query": "氯气流量与炉温的关联趋势", "expected": "local"},
{"query": "钛纯度标准参照国标如何执行", "expected": "local"},
# ---- local:内置保底 PII(身份证 / 手机号) ----
{"query": "员工身份证 110101199003071234 入职登记", "expected": "local"},
{"query": "联系人手机号 13800138000 请查收", "expected": "local"},
# ---- local:无规则命中 → 保守默认本地(未知 = 敏感,数据不出厂) ----
{"query": "今天天气怎么样", "expected": "local"},
{"query": "给我讲讲三国演义", "expected": "local"},
{"query": "车间排班表安排", "expected": "local"},
# ---- cloud:公开常识(模板规则 rt_common_knowledge,脱敏/通用) ----
{"query": "海绵钛是什么", "expected": "cloud"},
{"query": "海绵钛是什么材料", "expected": "cloud"},
{"query": "海绵钛是什么?", "expected": "cloud"},
{"query": "海绵钛是什么物质", "expected": "cloud"},
{"query": "海绵钛是什么用途", "expected": "cloud"},
# ---- block:高危安全指令(内置 rt_emergency_cmd / 模板 rt_safety_emergency) ----
{"query": "请执行停机操作", "expected": "block"},
{"query": "现场出现紧急停机指令", "expected": "block"},
{"query": "紧急停机怎么操作", "expected": "block"},
{"query": "立即停机", "expected": "block"},
{"query": "发现设备异常请停机", "expected": "block"},
# ---- block:上游 DLP 已拦截 → fail-closed 强制 block ----
{"query": "炉温当前是多少", "expected": "block", "dlp_blocked": True},
{"query": "海绵钛是什么", "expected": "block", "dlp_blocked": True},
# ---- 演示未通过样本:关键词未覆盖 → 误路由 local(期望 cloud) ----
{"query": "海绵钛的生产工艺是什么", "expected": "cloud"},
]
def load_samples(path: str) -> List[dict]:
"""从 JSON 文件加载评测集(顶层为样本数组)。"""
with open(path, "r", encoding="utf-8") as fh:
data = json.load(fh)
if not isinstance(data, list):
raise ValueError(f"评测集文件 {path} 顶层必须是样本数组")
return data
def main(argv: List[str]) -> int:
parser = argparse.ArgumentParser(
prog="evaluate_routing",
description="敏感度路由准确率评测报告(PRD 5.4 / EPIC #6,目标 ≥ 96.5%)",
)
parser.add_argument(
"--router-config", default="config/router.template.yaml",
help="敏感度路由规则模板资产路径(默认 config/router.template.yaml)",
)
parser.add_argument(
"--samples", default=None,
help="评测集 JSON 文件路径(与 --demo 二选一)",
)
parser.add_argument(
"--demo", action="store_true",
help="使用内置演示评测集(未指定 --samples 时默认开启)",
)
parser.add_argument(
"--threshold", type=float, default=0.965,
help="验收准确率阈值(默认 0.965,即 96.5%)",
)
parser.add_argument(
"--output", default=None,
help="报告输出文件路径(utf-8);缺省打印到 stdout",
)
args = parser.parse_args(argv)
# 评测集:--samples 优先,否则内置演示集
if args.samples:
samples = load_samples(args.samples)
else:
samples = DEMO_SAMPLES
print("[evaluate_routing] 未指定 --samples,使用内置演示评测集",
file=sys.stderr)
# 路由引擎(模板资产 + 内置保底),跑评测并渲染报告
router = SensitivityRouter.from_template_config(args.router_config)
report = router.evaluate(samples)
text = router.render_evaluation_report(report, threshold=args.threshold)
if args.output:
os.makedirs(os.path.dirname(os.path.abspath(args.output)) or ".", exist_ok=True)
with open(args.output, "w", encoding="utf-8") as fh:
fh.write(text)
print(f"[evaluate_routing] 评测报告已写入:{args.output}", file=sys.stderr)
else:
# Windows 控制台可能为 GBK:显式用 utf-8 输出避免编码错误
try:
sys.stdout.reconfigure(encoding="utf-8", errors="replace")
except (AttributeError, ValueError): # pragma: no cover - 旧版解释器
pass
print(text)
# 简报(stderr,便于定时任务抓取结论)
acc = float(report["accuracy"]) if report["total"] else 0.0
passed = acc >= args.threshold
print(
"[evaluate_routing] 样本 {total}|路由准确率 {acc:.1%}|{verdict}"
"(验收目标 ≥ {target:.1%})".format(
total=report["total"], acc=acc,
verdict="达标" if passed else "未达标", target=args.threshold),
file=sys.stderr,
)
return 0
if __name__ == "__main__":
sys.exit(main(sys.argv[1:]))