Files
iAOP/core/llm-gateway/tests/test_router.py
T

191 lines
7.6 KiB
Python
Raw 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 -*-
"""敏感度路由引擎(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()