feat: 完成 issue #6 LLM 网关 + RAG 模板化(混合网关主编排)

This commit is contained in:
2026-08-04 18:02:00 +08:00
parent fa5523a37d
commit 5d937e8efd
12 changed files with 1707 additions and 20 deletions
+131
View File
@@ -0,0 +1,131 @@
# -*- coding: utf-8 -*-
"""混合网关主编排(gateway)端到端单元测试:路由 → 生成 → 校验 → DLP 防线。"""
import os
import sys
import unittest
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import _bootstrap # noqa: F401
from llm_gateway.dlp import DlpEngine # noqa: E402
from llm_gateway.gateway import ( # noqa: E402
CloudBackend,
LLMGateway,
LocalBackend,
)
from llm_gateway.prompts import PromptRegistry # noqa: E402
from llm_gateway.router import RouteTarget, SensitivityRouter # noqa: E402
ROUTER_CONFIG = os.path.join(
os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
"config", "router.template.yaml",
)
PROMPTS_CONFIG = os.path.join(
os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
"config", "prompts.template.yaml",
)
def make_gateway() -> LLMGateway:
return LLMGateway(
dlp=DlpEngine(),
router=SensitivityRouter.from_template_config(ROUTER_CONFIG),
prompts=PromptRegistry.from_template_config(PROMPTS_CONFIG),
local=LocalBackend(),
cloud=CloudBackend(),
high_stakes_names=["alarm_explain"],
)
class GatewayRoutingTest(unittest.TestCase):
"""路由目标决定后端选择。"""
def setUp(self):
self.gw = make_gateway()
def test_sensitive_query_uses_local(self):
result = self.gw.ask("炉温当前是多少", rag_context=["SOP-炉温"])
self.assertEqual(result.route.target, RouteTarget.LOCAL)
self.assertIn("本地70B占位", result.answer)
self.assertFalse(result.needs_human)
def test_common_query_uses_cloud(self):
result = self.gw.ask("海绵钛是什么", rag_context=["科普手册"])
self.assertEqual(result.route.target, RouteTarget.CLOUD)
self.assertIn("云端API占位", result.answer)
def test_blocked_query_needs_human(self):
result = self.gw.ask("现场出现紧急停机指令", rag_context=[])
self.assertEqual(result.route.target, RouteTarget.BLOCK)
self.assertTrue(result.needs_human)
self.assertIn("人工确认", result.answer)
def test_dlp_blocked_query_forces_block(self):
# 身份证号触发 DLP → 即使模板规则未覆盖也 block
result = self.gw.ask("员工 110101199003071234 的炉温查询",
rag_context=["SOP"])
self.assertEqual(result.route.target, RouteTarget.BLOCK)
self.assertEqual(result.route.reason, "dlp_blocked")
class GatewayVerificationTest(unittest.TestCase):
"""引用溯源 + 信度阈值(高利害)。"""
def setUp(self):
self.gw = make_gateway()
def test_unsupported_citation_flagged(self):
# 占位后端回显 [来源: rag_context],与 rag_context 一致 → 支持
result = self.gw.ask("炉温偏高怎么处理",
rag_context=["沸腾氯化炉异常处置SOP"],
confidence=0.9)
self.assertTrue(result.verdict.supported)
def test_high_stakes_low_confidence_human_review(self):
# alarm_explain 为高利害模板:低信度 → 人工确认
result = self.gw.ask("解释报警并给出处置建议",
rag_context=["报警SOP"],
confidence=0.4)
self.assertEqual(result.route.target, RouteTarget.LOCAL)
self.assertEqual(result.verdict.action, "pass") # qa 非高利害,不启用阈值
gw2 = LLMGateway(
dlp=DlpEngine(),
router=SensitivityRouter.from_template_config(ROUTER_CONFIG),
prompts=PromptRegistry.from_template_config(PROMPTS_CONFIG),
prompt_name="alarm_explain",
high_stakes_names=["alarm_explain"],
)
result2 = gw2.ask("解释报警并给出处置建议",
rag_context=["报警SOP"],
confidence=0.4)
self.assertEqual(result2.verdict.action, "human_review")
self.assertTrue(result2.needs_human)
def test_prompt_version_binding(self):
# 显式绑定 qa@1.0.0(默认)——当前注册表已按模板加载
pv = self.gw.prompts.get("qa", version="1.0.0")
self.assertEqual(pv.version, "1.0.0")
class GatewayAuditTest(unittest.TestCase):
def test_audits_collectable(self):
gw = make_gateway()
gw.ask("炉温当前是多少", rag_context=["SOP"])
audits = gw.drain_audits()
self.assertIn("router", audits)
self.assertIn("guard", audits)
self.assertGreaterEqual(len(audits["router"]), 1)
# drain 后清空
self.assertEqual(gw.drain_audits()["router"], [])
def test_result_to_dict(self):
gw = make_gateway()
result = gw.ask("炉温当前是多少", rag_context=["SOP"])
d = result.to_dict()
self.assertIn("answer_id", d)
self.assertIn("route", d)
self.assertIn("verdict", d)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,103 @@
# -*- coding: utf-8 -*-
"""幻觉/事实性校验(hallucination)单元测试:引用溯源 / 信度阈值 / 评测。"""
import os
import sys
import unittest
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import _bootstrap # noqa: F401
from llm_gateway.hallucination import HallucinationGuard # noqa: E402
class CitationCheckTest(unittest.TestCase):
def setUp(self):
self.guard = HallucinationGuard(default_threshold=0.8)
def test_declared_sources_are_verified(self):
verdict = self.guard.check(
answer="炉温偏高应降温。[来源: 沸腾氯化炉异常处置SOP]",
sources=["沸腾氯化炉异常处置SOP", "交接班规范"],
)
self.assertTrue(verdict.supported)
self.assertEqual(verdict.action, "pass")
def test_missing_source_is_unsupported(self):
verdict = self.guard.check(
answer="应停机。[来源: 不存在的文档]",
sources=["沸腾氯化炉异常处置SOP"],
)
self.assertFalse(verdict.supported)
self.assertEqual(verdict.action, "unsupported")
self.assertIn("不存在的文档", verdict.missing_sources)
def test_no_citation_is_supported(self):
# 无引用声明 = 不判幻觉(引用为强制项由 RAG 模板保证)
verdict = self.guard.check(answer="按操作规程执行。", sources=[])
self.assertTrue(verdict.supported)
self.assertEqual(verdict.action, "pass")
class ConfidenceThresholdTest(unittest.TestCase):
def setUp(self):
self.guard = HallucinationGuard(default_threshold=0.8)
def test_low_confidence_high_stakes_human_review(self):
verdict = self.guard.check(
answer="建议立即停机。[来源: SOP]",
sources=["SOP"], confidence=0.55, high_stakes=True,
)
self.assertEqual(verdict.action, "human_review")
def test_high_confidence_high_stakes_passes(self):
verdict = self.guard.check(
answer="建议观察并记录。[来源: SOP]",
sources=["SOP"], confidence=0.95, high_stakes=True,
)
self.assertEqual(verdict.action, "pass")
def test_low_confidence_non_stakes_passes(self):
# 非高利害场景不启用阈值
verdict = self.guard.check(
answer="一般说明。[来源: 手册]", sources=["手册"],
confidence=0.3, high_stakes=False,
)
self.assertEqual(verdict.action, "pass")
def test_custom_threshold(self):
verdict = self.guard.check(
answer="处置建议。[来源: 手册]", sources=["手册"],
confidence=0.7, high_stakes=True, threshold=0.6,
)
self.assertEqual(verdict.action, "pass")
class EvaluateTest(unittest.TestCase):
def setUp(self):
self.guard = HallucinationGuard()
def test_evaluate_rates(self):
samples = [
{"answer": "a[来源: X]", "sources": ["X"], "confidence": 0.9, "high_stakes": True},
{"answer": "b[来源: Y]", "sources": ["Z"], "confidence": 0.9},
{"answer": "c[来源: X]", "sources": ["X"], "confidence": 0.4, "high_stakes": True},
]
report = self.guard.evaluate(samples)
self.assertEqual(report["total"], 3)
# 支持 2 条(a/c),human_review 1 条(c)
self.assertEqual(report["supported_rate"], round(2 / 3, 4))
self.assertEqual(report["human_review_rate"], round(1 / 3, 4))
def test_empty_samples(self):
report = self.guard.evaluate([])
self.assertEqual(report["total"], 0)
def test_audit_records(self):
self.guard.check("x[来源: A]", sources=["A"])
records = self.guard.drain_audit()
self.assertEqual(len(records), 1)
self.assertEqual(records[0]["action"], "pass")
if __name__ == "__main__":
unittest.main()
+115
View File
@@ -0,0 +1,115 @@
# -*- coding: utf-8 -*-
"""Prompt 版本管理(prompts)单元测试:登记 / 绑定 / 晋升 / 回滚 / 审计。"""
import os
import sys
import unittest
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import _bootstrap # noqa: F401
from llm_gateway.prompts import ( # noqa: E402
PromptRegistry,
validate_semver,
)
CONFIG_PATH = os.path.join(
os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
"config", "prompts.template.yaml",
)
class SemverTest(unittest.TestCase):
def test_valid_versions(self):
for v in ("1.0.0", "0.1.2", "10.20.30"):
self.assertTrue(validate_semver(v), v)
def test_invalid_versions(self):
for v in ("1.0", "v1.0.0", "1.0.0-rc1", "1.0.0.1", ""):
self.assertFalse(validate_semver(v), v)
class RegistryCoreTest(unittest.TestCase):
def setUp(self):
self.reg = PromptRegistry()
self.reg.update("qa", "问题:{query}", "1.0.0")
self.reg.update("qa", "问题:{query} 请引用SOP", "1.0.1")
def test_first_version_is_current(self):
self.assertEqual(self.reg.current("qa").version, "1.0.0")
def test_promote_switches_current(self):
self.reg.promote("qa", "1.0.1")
self.assertEqual(self.reg.current("qa").version, "1.0.1")
def test_runtime_binding_is_reproducible(self):
# 显式绑定旧版本:即使 current 已变,行为可复现
self.reg.promote("qa", "1.0.1")
pv = self.reg.get("qa", version="1.0.0")
self.assertEqual(pv.version, "1.0.0")
self.assertNotIn("SOP", pv.text)
def test_duplicate_version_rejected(self):
with self.assertRaises(ValueError):
self.reg.update("qa", "覆盖", "1.0.0")
def test_render(self):
pv = self.reg.current("qa")
self.assertEqual(pv.render(query="炉温"), "问题:炉温")
def test_missing_template_raises(self):
with self.assertRaises(KeyError):
self.reg.get("not_exist")
class RollbackTest(unittest.TestCase):
def test_rollback_returns_previous(self):
reg = PromptRegistry()
reg.update("t", "v0", "1.0.0")
reg.update("t", "v1", "1.0.1")
reg.promote("t", "1.0.1")
self.assertEqual(reg.current("t").version, "1.0.1")
previous = reg.rollback("t")
self.assertEqual(previous, "1.0.0")
self.assertEqual(reg.current("t").version, "1.0.0")
def test_rollback_without_history_returns_none(self):
reg = PromptRegistry()
reg.update("t", "v0", "1.0.0")
self.assertIsNone(reg.rollback("t"))
class AuditTest(unittest.TestCase):
def test_actions_recorded(self):
reg = PromptRegistry()
reg.update("t", "v0", "1.0.0")
reg.update("t", "v1", "1.0.1")
reg.promote("t", "1.0.1")
reg.rollback("t")
audit = reg.drain_audit()
actions = [a["action"] for a in audit]
self.assertEqual(actions, ["add", "add", "promote", "rollback"])
def test_drain_clears(self):
reg = PromptRegistry()
reg.update("t", "v0", "1.0.0")
self.assertEqual(len(reg.drain_audit()), 1)
self.assertEqual(reg.drain_audit(), [])
class TemplateLoadTest(unittest.TestCase):
def test_template_config_load(self):
reg = PromptRegistry.from_template_config(CONFIG_PATH)
names = reg.template_names
self.assertIn("qa", names)
self.assertIn("alarm_explain", names)
# shift_handover 应有两版本且 current 为 1.0.1(current: true)
self.assertEqual(reg.versions("shift_handover"), ["1.0.0", "1.0.1"])
self.assertEqual(reg.current("shift_handover").version, "1.0.1")
def test_invalid_version_rejected(self):
with self.assertRaises(ValueError):
PromptRegistry().update("t", "x", "not-semver")
if __name__ == "__main__":
unittest.main()
+136
View File
@@ -0,0 +1,136 @@
# -*- 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()