# -*- 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") class EngineDiagnosticsTest(unittest.TestCase): """规则引擎诊断(issue #43):priority / validate / stats / describe。""" def test_priority_from_mapping(self): rule = RouterRule.from_mapping( {"name": "r", "pattern": "x", "target": "block", "priority": 1}) self.assertEqual(rule.priority, 1) self.assertEqual(RouterRule.from_mapping( {"name": "r2", "pattern": "x", "target": "local"}).priority, 0) def test_priority_orders_rules(self): # block 规则 priority=1 → 先于 local 规则(priority=10)命中 # 用“紧急停泵”避免命中内置 rt_emergency_cmd(pattern=停机) rules = [ RouterRule.from_mapping( {"name": "low", "pattern": "紧急停泵", "target": "local", "priority": 10}), RouterRule.from_mapping( {"name": "high", "pattern": "紧急停泵", "target": "block", "priority": 1}), ] router = SensitivityRouter(rules=rules) self.assertEqual(router.route("紧急停泵").target, "block") names = router.rule_names self.assertLess(names.index("high"), names.index("low")) # priority 升序 def test_validate_rules(self): # 合法规则集(含内置保底):无问题;重复名在合并时按 name 覆盖不产生冲突 router = SensitivityRouter(rules=[ RouterRule.from_mapping({"name": "a", "pattern": "x", "target": "local"}), ]) self.assertEqual(router.validate_rules(), []) def test_stats(self): router = SensitivityRouter(rules=[ RouterRule.from_mapping({"name": "a", "pattern": "x", "target": "local"}), RouterRule.from_mapping( {"name": "b", "kind": "regex", "pattern": r"\d+", "target": "cloud"}), ]) stats = router.stats() self.assertEqual(stats["by_target"]["local"], 3) # 内置 PII 2 条 + a self.assertEqual(stats["by_target"]["cloud"], 1) self.assertEqual(stats["by_kind"]["regex"], 3) # 内置 PII 2 条 + b def test_describe_hit_chain(self): router = SensitivityRouter(rules=[ RouterRule.from_mapping({"name": "a", "pattern": "紧急停泵", "target": "local"}), RouterRule.from_mapping( {"name": "b", "pattern": "紧急停泵", "target": "block", "priority": 1}), ]) desc = router.describe("紧急停泵") names = [h["name"] for h in desc["hits"]] self.assertEqual(names, ["a", "b"]) # priority 升序(a=0 先于 b=1) self.assertEqual(router.route("紧急停泵").target, "local") # 决策取首个命中(a) if __name__ == "__main__": unittest.main()