137 lines
5.0 KiB
Python
137 lines
5.0 KiB
Python
# -*- coding: utf-8 -*-
|
|||
|
|
"""敏感度路由引擎(router)单元测试:分级路由 / 模板加载 / 评估 / 审计。"""
|
||
|
|
import os
|
||
|
|
import sys
|
||
|
|
import unittest
|
||
|
|
|
||
|
|
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||
|
|
import _bootstrap # noqa: F401
|
||
|
|
|
||
|
|
from llm_gateway.router import ( # noqa: E402
|
||
|
|
RouteTarget,
|
||
|
|
RouterRule,
|
||
|
|
SensitivityRouter,
|
||
|
|
)
|
||
|
|
|
||
|
|
CONFIG_PATH = os.path.join(
|
||
|
|
os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
|
||
|
|
"config", "router.template.yaml",
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
class DefaultRoutingTest(unittest.TestCase):
|
||
|
|
"""内置保底规则(无模板配置时的默认分级)。"""
|
||
|
|
|
||
|
|
def setUp(self):
|
||
|
|
self.router = SensitivityRouter()
|
||
|
|
|
||
|
|
def test_pii_routes_local(self):
|
||
|
|
decision = self.router.route("员工身份证 110101199003071234 入职")
|
||
|
|
self.assertEqual(decision.target, RouteTarget.LOCAL)
|
||
|
|
self.assertEqual(decision.reason, "rule_hit")
|
||
|
|
self.assertEqual(decision.category, "pii")
|
||
|
|
|
||
|
|
def test_emergency_cmd_blocks(self):
|
||
|
|
decision = self.router.route("请执行停机操作")
|
||
|
|
self.assertEqual(decision.target, RouteTarget.BLOCK)
|
||
|
|
|
||
|
|
def test_unknown_routes_local_by_default(self):
|
||
|
|
# 保守默认:未知 = 敏感,数据不出厂
|
||
|
|
decision = self.router.route("今天天气怎么样")
|
||
|
|
self.assertEqual(decision.target, RouteTarget.LOCAL)
|
||
|
|
self.assertEqual(decision.reason, "no_rule")
|
||
|
|
|
||
|
|
def test_dlp_blocked_forces_block(self):
|
||
|
|
decision = self.router.route("今天天气怎么样", dlp_blocked=True)
|
||
|
|
self.assertEqual(decision.target, RouteTarget.BLOCK)
|
||
|
|
self.assertEqual(decision.reason, "dlp_blocked")
|
||
|
|
|
||
|
|
|
||
|
|
class TemplateConfigTest(unittest.TestCase):
|
||
|
|
"""模板资产加载(config/router.template.yaml)。"""
|
||
|
|
|
||
|
|
def setUp(self):
|
||
|
|
self.router = SensitivityRouter.from_template_config(CONFIG_PATH)
|
||
|
|
|
||
|
|
def test_template_rules_loaded(self):
|
||
|
|
names = self.router.rule_names
|
||
|
|
self.assertIn("rt_proc_furnace_temp", names)
|
||
|
|
self.assertIn("rt_proc_cl2_flow", names)
|
||
|
|
|
||
|
|
def test_process_param_routes_local(self):
|
||
|
|
decision = self.router.route("炉温当前是多少")
|
||
|
|
self.assertEqual(decision.target, RouteTarget.LOCAL)
|
||
|
|
self.assertEqual(decision.category, "process-parameter")
|
||
|
|
|
||
|
|
def test_common_knowledge_routes_cloud(self):
|
||
|
|
decision = self.router.route("海绵钛是什么")
|
||
|
|
self.assertEqual(decision.target, RouteTarget.CLOUD)
|
||
|
|
self.assertEqual(decision.reason, "rule_hit")
|
||
|
|
|
||
|
|
def test_safety_rule_blocks(self):
|
||
|
|
decision = self.router.route("现场出现紧急停机指令")
|
||
|
|
self.assertEqual(decision.target, RouteTarget.BLOCK)
|
||
|
|
|
||
|
|
def test_template_rule_overrides_builtin(self):
|
||
|
|
# 模板加载后内置保底仍生效(同名覆盖仅发生在显式同名时)
|
||
|
|
decision = self.router.route("请执行停机操作")
|
||
|
|
self.assertEqual(decision.target, RouteTarget.BLOCK)
|
||
|
|
|
||
|
|
|
||
|
|
class EvaluateTest(unittest.TestCase):
|
||
|
|
"""路由准确率离线评估(Issue #49 雏形,目标 ≥ 96.5%)。"""
|
||
|
|
|
||
|
|
def setUp(self):
|
||
|
|
self.router = SensitivityRouter.from_template_config(CONFIG_PATH)
|
||
|
|
|
||
|
|
def test_perfect_samples_reach_100_percent(self):
|
||
|
|
samples = [
|
||
|
|
{"query": "炉温偏高如何处理", "expected": "local"},
|
||
|
|
{"query": "氯气流量超限报警", "expected": "local"},
|
||
|
|
{"query": "海绵钛是什么", "expected": "cloud"},
|
||
|
|
{"query": "紧急停机怎么操作", "expected": "block"},
|
||
|
|
{"query": "加料比如何调整", "expected": "local"},
|
||
|
|
]
|
||
|
|
report = self.router.evaluate(samples)
|
||
|
|
self.assertEqual(report["accuracy"], 1.0)
|
||
|
|
self.assertEqual(report["total"], 5)
|
||
|
|
|
||
|
|
def test_empty_samples(self):
|
||
|
|
report = self.router.evaluate([])
|
||
|
|
self.assertEqual(report["accuracy"], 0.0)
|
||
|
|
self.assertEqual(report["total"], 0)
|
||
|
|
|
||
|
|
def test_audit_records_decision(self):
|
||
|
|
router = SensitivityRouter(audit=True)
|
||
|
|
router.route("请执行停机操作")
|
||
|
|
records = router.drain_audit()
|
||
|
|
self.assertEqual(len(records), 1)
|
||
|
|
self.assertEqual(records[0]["target"], RouteTarget.BLOCK)
|
||
|
|
# drain 后清空
|
||
|
|
self.assertEqual(router.drain_audit(), [])
|
||
|
|
|
||
|
|
|
||
|
|
class RuleModelTest(unittest.TestCase):
|
||
|
|
"""RouterRule 模型合法性校验。"""
|
||
|
|
|
||
|
|
def test_rule_requires_pattern(self):
|
||
|
|
with self.assertRaises(ValueError):
|
||
|
|
RouterRule.from_mapping({"name": "x", "pattern": ""})
|
||
|
|
|
||
|
|
def test_rule_rejects_bad_target(self):
|
||
|
|
with self.assertRaises(ValueError):
|
||
|
|
RouterRule.from_mapping(
|
||
|
|
{"name": "x", "pattern": "a", "target": "mars"})
|
||
|
|
|
||
|
|
def test_regex_rule_matches(self):
|
||
|
|
rule = RouterRule.from_mapping(
|
||
|
|
{"name": "r", "kind": "regex", "pattern": r"\d{4}",
|
||
|
|
"target": RouteTarget.LOCAL})
|
||
|
|
hits = rule.find("温度 1234 度")
|
||
|
|
self.assertEqual(len(hits), 1)
|
||
|
|
self.assertEqual(hits[0][0], "1234")
|
||
|
|
|
||
|
|
|
||
|
|
if __name__ == "__main__":
|
||
|
|
unittest.main()
|