feat: 完成 issue #6 LLM 网关 + RAG 模板化(混合网关主编排)
This commit is contained in:
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user