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
+52 -6
View File
@@ -1,6 +1,6 @@
# iAOP-Core · LLM 网关(LLM Gateway) # iAOP-Core · LLM 网关(LLM Gateway)
对应 PRD 5.4「④ LLM 网关 + RAG」与 EPIC #6(Issue #48 等子任务): 对应 PRD 5.4「④ LLM 网关 + RAG」与 EPIC #6:
本地 70B(敏感/核心)+ 云端 API(脱敏/通用)**混合**,安全分级路由, 本地 70B(敏感/核心)+ 云端 API(脱敏/通用)**混合**,安全分级路由,
**敏感数据本地闭环,仅脱敏/公开内容可走云端 API**(数据不出厂)。 **敏感数据本地闭环,仅脱敏/公开内容可走云端 API**(数据不出厂)。
@@ -8,13 +8,23 @@
``` ```
core/llm-gateway/ core/llm-gateway/
├── __init__.py 包入口(导出 DLP 引擎 API) ├── __init__.py 包入口(导出各模块 API)
├── dlp.py DLP 敏感数据拦截引擎(Issue #48) ├── dlp.py DLP 敏感数据拦截引擎(Issue #48)
├── router.py 敏感度路由规则引擎(Issue #43 雏形)
├── prompts.py Prompt 版本管理(Issue #47 雏形)
├── hallucination.py 幻觉/事实性校验中间件(Issue #47 雏形)
├── gateway.py 混合网关主编排(EPIC #6 主体交付)
├── config/ ├── config/
│ └── dlp.template.yaml 模板 DLP 规则资产(ti-cl4 示例,换行业只改它) │ ├── dlp.template.yaml 模板 DLP 规则资产(ti-cl4 示例)
│ ├── router.template.yaml 模板敏感度路由规则资产
│ └── prompts.template.yaml 模板提示词版本库资产
└── tests/ └── tests/
├── _bootstrap.py 测试引导(目录含连字符,挂载包名 llm_gateway) ├── _bootstrap.py 测试引导(目录含连字符,挂载包名 llm_gateway)
└── test_dlp.py DLP 引擎单元测试 ├── test_dlp.py DLP 引擎单元测试
├── test_router.py 路由引擎单元测试
├── test_prompts.py Prompt 版本库单元测试
├── test_hallucination.py 幻觉/事实性校验单元测试
└── test_gateway.py 网关主编排端到端单元测试
``` ```
## DLP 敏感数据拦截(Issue #48) ## DLP 敏感数据拦截(Issue #48)
@@ -70,7 +80,43 @@ cd core/llm-gateway
python -m unittest discover -s tests -v python -m unittest discover -s tests -v
``` ```
## 混合网关主编排(EPIC #6 主体)
`LLMGateway.ask()` 串起完整闭环:**敏感度路由 → 本地/云端生成 → 引用溯源校验
→ DLP 出站防线**,覆盖 PRD 5.4 用户操作流程(提问 → 路由判断敏感级 →
本地/云端生成 → RAG 溯源校验 → 返回带引用的答案;异常转人工)。
```python
from llm_gateway import LLMGateway
from llm_gateway.dlp import DlpEngine
from llm_gateway.router import SensitivityRouter
from llm_gateway.prompts import PromptRegistry
gw = LLMGateway(
dlp=DlpEngine(),
router=SensitivityRouter.from_template_config("config/router.template.yaml"),
prompts=PromptRegistry.from_template_config("config/prompts.template.yaml"),
high_stakes_names=["alarm_explain"], # 高利害模板启用信度阈值
)
result = gw.ask("炉温偏高怎么处理", rag_context=["沸腾氯化炉异常处置SOP"])
print(result.route.target) # local / cloud / block
print(result.verdict.action) # pass / human_review / unsupported
if result.needs_human:
... # 转人工确认(PRD 5.4 异常时转人工)
```
- **敏感度路由**(router.py,Issue #43 雏形):模板配置驱动,DLP 拦截
fail-closed 强制 block,未知内容保守走本地(数据不出厂);
- **Prompt 版本管理**(prompts.py,Issue #47 雏形):semver 版本库、
运行时绑定(可复现)、一键回滚、变更审计;
- **幻觉/事实性校验**(hallucination.py,Issue #47 雏形):`[来源: X]`
引用溯源强制校验 + 高利害信度阈值 → 人工确认;
- **推理后端抽象**(gateway.py 内 `InferenceBackend`):业务代码只依赖
接口,本地 70B / 云端 API 具体接入由子任务 #44 / #45 实现。
## 后续子任务(EPIC #6 拆分,待扩展) ## 后续子任务(EPIC #6 拆分,待扩展)
- 敏感度路由(准确率 ≥ 96.5%)与路由准确率评估脚本(Issue #49); - 敏感度路由调优与路由准确率评估脚本(Issue #49,本版已提供评估入口);
- Prompt 版本管理与幻觉/事实性校验中间件(Issue #47)。 - 本地 70B 模型接入与推理封装(Issue #44);
- 云端 API(Qwen/DeepSeek)接入与安全网关(Issue #45);
- Prompt 版本管理 + 幻觉校验中间件完善(Issue #47,本版已提供核心)。
+48 -14
View File
@@ -1,19 +1,25 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
"""iAOP-Core · LLM 网关(LLM Gateway)—— 混合 LLM 的安全出站防线。 """iAOP-Core · LLM 网关(LLM Gateway)—— 混合 LLM 的安全出站防线与编排。
对应 PRD 5.4「④ LLM 网关 + RAG」与 EPIC #6: 对应 PRD 5.4「④ LLM 网关 + RAG」与 EPIC #6:
本地 70B(敏感/核心)+ 云端 API(脱敏/通用)混合,敏感数据**本地闭环**。 本地 70B(敏感/核心)+ 云端 API(脱敏/通用)混合,敏感数据**本地闭环**。
当前子模块(Issue #48): 模块组成:
- dlp DLP 敏感数据拦截引擎:出站内容(query / RAG context / 模型输出) - dlp DLP 敏感数据拦截引擎(Issue #48):出站内容(query / RAG context /
发往云端前做敏感规则检查,命中即拦截(目标 100% 拦截),全量审计。 模型输出)发往云端前做敏感规则检查,命中即拦截(目标 100% 拦截),
全量审计。
后续子任务(EPIC #6 拆分,将在本包扩展): - router 敏感度路由规则引擎(Issue #43 雏形):敏感度分级路由(local/cloud/
- 敏感度路由(router)、Prompt 版本管理与幻觉校验中间件、路由准确率评估。 block),模板配置驱动,DLP 拦截即 fail-closed 转 block。
- prompts Prompt 版本管理(Issue #47 雏形):semver 版本库、运行时绑定、
一键回滚、变更审计。
- hallucination 幻觉/事实性校验中间件(Issue #47 雏形):引用溯源 +
高利害信度阈值 → 人工确认。
- gateway 混合网关主编排(EPIC #6 主体):路由 → 生成 → 溯源校验 →
DLP 出站防线,端到端闭环。
测试:`python -m unittest discover -s tests -v`(在 core/llm-gateway 目录下执行)。 测试:`python -m unittest discover -s tests -v`(在 core/llm-gateway 目录下执行)。
""" """
__version__ = "0.1.0" __version__ = "0.2.0"
from .dlp import ( from .dlp import (
DLP_DEFAULT_RULES, DLP_DEFAULT_RULES,
@@ -23,12 +29,40 @@ from .dlp import (
DlpRule, DlpRule,
DlpRuleKind, DlpRuleKind,
) )
from .router import (
RouteDecision,
RouteTarget,
RouterRule,
SensitivityRouter,
)
from .prompts import (
PromptChange,
PromptRegistry,
PromptVersion,
validate_semver,
)
from .hallucination import (
GuardVerdict,
HallucinationGuard,
)
from .gateway import (
CloudBackend,
GatewayResult,
InferenceBackend,
LLMGateway,
LocalBackend,
)
__all__ = [ __all__ = [
"DlpRuleKind", # dlp
"DlpRule", "DlpRuleKind", "DlpRule", "DlpHit", "DlpResult", "DlpEngine", "DLP_DEFAULT_RULES",
"DlpHit", # router
"DlpResult", "RouteTarget", "RouterRule", "RouteDecision", "SensitivityRouter",
"DlpEngine", # prompts
"DLP_DEFAULT_RULES", "PromptVersion", "PromptChange", "PromptRegistry", "validate_semver",
# hallucination
"GuardVerdict", "HallucinationGuard",
# gateway
"InferenceBackend", "LocalBackend", "CloudBackend",
"GatewayResult", "LLMGateway",
] ]
@@ -0,0 +1,31 @@
# -*- coding: utf-8 -*-
# 模板「提示词版本库」资产示例:ti-cl4(氯化车间/海绵钛,Template-Ti 一期)。
#
# 说明:
# - 这是「提示词模板(版本化)」配置点(PRD 5.4):换行业只改本文件;
# - version 必须为 semver(主.次.补丁);变更须评审并记录,支持一键回滚;
# - current: true 的版本在加载后自动晋升为当前默认版本;
# - 模板正文用 {query} 等占位符(单行文本),运行时绑定版本渲染(可复现)。
template: ti-cl4
version: 1.0.0
templates:
- name: qa
version: 1.0.0
current: true
description: 通用工艺问答模板(v1 基线)
text: 你是氯化车间工艺助手。请基于给定资料回答问题,并标注来源。问题:{query}
- name: alarm_explain
version: 1.0.0
current: true
description: 报警解释模板(高利害,启用信度阈值)
text: 请解释以下报警的可能原因与处置建议,必须引用SOP来源:报警:{query}
- name: shift_handover
version: 1.0.1
current: true
description: 交接班摘要模板(v1.0.1:补充安全注意事项章节)
text: 生成交接班摘要,包含:生产概况、异常事项、安全注意事项。班次:{query}
- name: shift_handover
version: 1.0.0
current: false
description: 交接班摘要模板(v1 基线,无安全注意事项章节)
text: 生成交接班摘要,包含:生产概况、异常事项。班次:{query}
@@ -0,0 +1,53 @@
# -*- coding: utf-8 -*-
# 模板「敏感度路由规则」资产示例:ti-cl4(氯化车间/海绵钛,Template-Ti 一期)。
#
# 说明:
# - 这是「路由策略」配置点(PRD 5.4):换行业只改本文件,内核零改动;
# - target: local(敏感/核心 → 本地70B,数据不出厂)
# cloud(脱敏/通用 → 云端API)
# block(高危 → 直接拦截,转人工)
# - kind: keyword 大小写不敏感子串匹配;regex 正则匹配;
# - 通用 PII/高危规则由内核内置保底(ROUTER_DEFAULT_RULES),无需重复配置;
# 本文件只补充**行业路由语义**(工艺敏感 → 本地,公开常识 → 云端)。
template: ti-cl4
version: 1.0.0
rules:
# ---- 工艺敏感(必须走本地,数据不出厂) ----
- name: rt_proc_cl2_flow
category: process-parameter
kind: keyword
pattern: 氯气流量
target: local
description: 氯气流量参数(工艺敏感,本地闭环)
- name: rt_proc_furnace_temp
category: process-parameter
kind: keyword
pattern: 炉温
target: local
description: 炉温参数(工艺敏感,本地闭环)
- name: rt_proc_feeding_ratio
category: process-parameter
kind: keyword
pattern: 加料比
target: local
description: 加料配比参数(工艺敏感,本地闭环)
- name: rt_proc_ti_purity
category: process-parameter
kind: keyword
pattern: 钛纯度
target: local
description: 产品质量指标(钛纯度,本地闭环)
# ---- 高危(直接拦截转人工) ----
- name: rt_safety_emergency
category: safety
kind: keyword
pattern: 紧急停机
target: block
description: 紧急停机指令(高危,转人工确认)
# ---- 通用常识(可走云端,仅脱敏/公开内容) ----
- name: rt_common_knowledge
category: general
kind: keyword
pattern: 海绵钛是什么
target: cloud
description: 公开常识问答(脱敏/通用,可走云端)
+220
View File
@@ -0,0 +1,220 @@
# -*- coding: utf-8 -*-
"""iAOP-Core · LLM 网关 —— 混合网关主编排(EPIC #6 主体交付)。
对应 PRD 5.4「④ LLM 网关 + RAG」与 EPIC #6(Issue #48 DLP / #46 RAG 模板化
已完成,本模块为其上层编排):
用户提问 → 敏感度路由 → 本地/云端生成 → RAG 溯源校验 → 返回带引用答案
└──── DLP 出站检查(fail-closed,拦截即转本地/人工)────┘
`LLMGateway.ask()` 串起四个可插拔组件:
- `dlp`(DlpEngine):出站防线,云端通道必经检查;
- `router`(SensitivityRouter):敏感度分级路由(local / cloud / block);
- `prompts`(PromptRegistry):提示词模板版本绑定(可复现);
- `guard`(HallucinationGuard):引用溯源 + 信度阈值 → 人工确认;
- `backends`(LocalBackend / CloudBackend):推理后端抽象(可注入)。
设计说明:
- 本版提供**编排闭环 + 后端抽象接口**,本地 70B / 云端 API 的具体接入
由子任务 #44 / #45 实现;`LocalBackend` / `CloudBackend` 默认内置一个
最小实现(返回固定占位答案 + 回显引用),供端到端测试与演示。
测试:`python -m unittest discover -s tests -v`(在 core/llm-gateway 目录下执行)。
"""
from __future__ import annotations
import uuid
from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import Callable, Dict, List, Optional, Sequence
from .dlp import DlpEngine
from .router import RouteDecision, RouteTarget, SensitivityRouter
from .prompts import PromptRegistry
from .hallucination import GuardVerdict, HallucinationGuard
# ---------------------------------------------------------------------------
# 推理后端抽象(Issue #44 / #45 将实现具体后端,业务代码只依赖本接口)
# ---------------------------------------------------------------------------
class InferenceBackend:
"""推理后端接口抽象(对齐 PRD 5.6 InferenceBackend 思想)。
业务代码只依赖本接口,不感知具体硬件/厂商;切换后端 = 换实现。
子任务 #44(本地 70B)、#45(云端 Qwen/DeepSeek)将各自实现本接口。
"""
name: str = "base"
def generate(self, prompt: str, context: Sequence[str]) -> str:
"""根据 prompt 与 RAG 上下文生成回答。子类实现。"""
raise NotImplementedError
class LocalBackend(InferenceBackend):
"""本地 70B 后端占位实现:数据不出厂(敏感/核心走此通道)。
子任务 #44 将替换为真实本地模型推理封装(vLLM/TGI 等)。
"""
name = "local-70b"
def __init__(self, echo_context: bool = True) -> None:
self.echo_context = echo_context
def generate(self, prompt: str, context: Sequence[str]) -> str:
head = f"[本地70B占位] {prompt[:40]}"
refs = ""
if self.echo_context:
for i, src in enumerate(context[:3], 1):
refs += f"\n[来源: {src}]"
return head + refs
class CloudBackend(InferenceBackend):
"""云端 API 后端占位实现:仅接收 DLP 放行的脱敏/通用内容。
子任务 #45 将替换为 Qwen/DeepSeek API 接入 + 安全网关。
"""
name = "cloud-api"
def __init__(self, echo_context: bool = True) -> None:
self.echo_context = echo_context
def generate(self, prompt: str, context: Sequence[str]) -> str:
head = f"[云端API占位] {prompt[:40]}"
refs = ""
if self.echo_context:
for i, src in enumerate(context[:3], 1):
refs += f"\n[来源: {src}]"
return head + refs
# ---------------------------------------------------------------------------
# 网关输出
# ---------------------------------------------------------------------------
@dataclass(frozen=True)
class GatewayResult:
"""一次 ask() 的完整结果(含中间决策,便于审计与验收)。"""
query: str
answer: str
route: RouteDecision
verdict: GuardVerdict
answer_id: str = field(default_factory=lambda: uuid.uuid4().hex[:12])
created_at: str = field(default_factory=lambda: datetime.now(timezone.utc).isoformat())
@property
def needs_human(self) -> bool:
"""是否需转人工确认(路由 block 或校验 human_review/unsupported)。"""
return (self.route.target == RouteTarget.BLOCK
or self.verdict.action in ("human_review", "unsupported"))
def to_dict(self) -> Dict[str, object]:
return {
"answer_id": self.answer_id,
"created_at": self.created_at,
"query": self.query,
"answer": self.answer,
"route": self.route.to_dict(),
"verdict": self.verdict.to_dict(),
"needs_human": self.needs_human,
}
# ---------------------------------------------------------------------------
# 混合网关主编排
# ---------------------------------------------------------------------------
class LLMGateway:
"""混合 LLM 网关主编排:路由 → 生成 → 溯源校验 → DLP 出站防线。
构造参数均可注入(默认带内置 DLP 保底规则 + 空 router/prompts/guard)。
"""
def __init__(self,
dlp: Optional[DlpEngine] = None,
router: Optional[SensitivityRouter] = None,
prompts: Optional[PromptRegistry] = None,
guard: Optional[HallucinationGuard] = None,
local: Optional[InferenceBackend] = None,
cloud: Optional[InferenceBackend] = None,
prompt_name: str = "qa",
prompt_version: Optional[str] = None,
high_stakes_names: Optional[List[str]] = None) -> None:
self.dlp = dlp or DlpEngine()
self.router = router or SensitivityRouter()
self.prompts = prompts or PromptRegistry()
self.guard = guard or HallucinationGuard()
self.local = local or LocalBackend()
self.cloud = cloud or CloudBackend()
self.prompt_name = prompt_name
self.prompt_version = prompt_version
# 高利害提示词:命中即启用信度阈值(处置建议 / 报警解释等)
self.high_stakes_names = set(high_stakes_names or [])
# -- 主编排入口 --------------------------------------------------------
def ask(self, query: str,
rag_context: Optional[Sequence[str]] = None,
confidence: float = 1.0) -> GatewayResult:
"""完整处理一次用户提问。
`rag_context`:RAG 检索命中的文档标题列表(溯源校验用);
`confidence`:模型输出信度(0~1,高利害场景低于阈值转人工)。
"""
rag_context = list(rag_context or [])
# 1) DLP 出站检查:query 敏感即拦截(fail-closed,云端不可达)
dlp_result = self.dlp.check_outbound({"query": query})
# 2) 敏感度路由(含 DLP 结果 → block)
decision = self.router.route(query, dlp_blocked=dlp_result.blocked)
# 3) 选择后端与提示词版本(运行时绑定,可复现)
prompt = self.prompts.get(self.prompt_name, self.prompt_version)
backend = self.local if decision.target != RouteTarget.CLOUD else self.cloud
# 4) 生成(block 时也不调用后端,直接给出人工确认占位答案)
if decision.target == RouteTarget.BLOCK:
answer = "该请求已拦截(敏感度路由/规则触发),请转人工确认处理。"
backend_name = "none"
else:
rendered = prompt.render(query=query)
answer = backend.generate(rendered, rag_context)
backend_name = backend.name
# 5) 幻觉/事实性校验(引用溯源 + 高利害信度阈值)
high_stakes = prompt.name in self.high_stakes_names
verdict = self.guard.check(
answer=answer, sources=rag_context,
confidence=confidence, high_stakes=high_stakes,
)
# 6) 出站前最终 DLP 防线(模型输出若含敏感内容:云端通道拦截)
if decision.target == RouteTarget.CLOUD:
outbound = self.dlp.check_outbound({"output": answer})
if outbound.blocked:
answer = "输出经 DLP 复查拦截,已转本地/人工处理。"
return GatewayResult(
query=query, answer=answer, route=decision, verdict=verdict,
)
# -- 审计汇总 ----------------------------------------------------------
def drain_audits(self) -> Dict[str, List[Dict[str, object]]]:
"""取走各组件审计记录(DLP / 路由 / Prompt / 幻觉校验)。"""
return {
"dlp": self.dlp.drain_audit(),
"router": self.router.drain_audit(),
"prompts": self.prompts.drain_audit(),
"guard": self.guard.drain_audit(),
}
def __repr__(self) -> str: # pragma: no cover - 调试辅助
return (f"<LLMGateway router={self.router!r} prompts={self.prompts!r} "
f"guard={self.guard!r}>")
+142
View File
@@ -0,0 +1,142 @@
# -*- coding: utf-8 -*-
"""iAOP-Core · LLM 网关 —— 幻觉/事实性校验中间件(EPIC #6 主体,Issue #47 雏形)。
对应 PRD 5.4「④ LLM 网关 + RAG」:
- **事实性校验**:RAG 答案强制**引用溯源**(返回命中文档片段+来源);
对高利害输出(如处置建议)设置信度阈值,低于阈值触发"人工确认";
定期用评测集检验事实一致性。
本模块实现 `HallucinationGuard`:
- **引用溯源校验**:模型输出中声称引用的片段(`[来源: <doc>]`)必须能在
RAG 检索命中的文档片段中找到对应来源,找不到即判定 `unsupported`
(无源引用 = 幻觉嫌疑);
- **信度阈值**:对高利害输出(处置建议 / 报警解释)要求信度 ≥ 阈值,
低于阈值返回 `human_review`(转人工确认,PRD 5.4 异常时转人工);
- **评测集检验**:`evaluate()` 对 (prompt, answer, expected_sources) 样本
批量评估事实一致性(供"定期评测"脚本调用)。
设计说明(供子任务 #47 继续细化):
- 本版实现校验核心(溯源 + 信度阈值 + 评测入口);
- 子任务 #47 将在此基础上补齐与 Prompt 版本库的联动与评测报告脚本。
测试:`python -m unittest discover -s tests -v`(在 core/llm-gateway 目录下执行)。
"""
from __future__ import annotations
import re
import uuid
from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import Dict, List, Optional, Sequence
# 输出中引用声明的格式:`[来源: 文档标题]` 或 `[src: doc_id]`
_SOURCE_REF_RE = re.compile(r"\[来源[::]\s*([^\]]+)\]", re.IGNORECASE)
@dataclass(frozen=True)
class GuardVerdict:
"""一次事实性校验的结论。"""
answer: str
supported: bool # 所有引用声明均有真实来源
confidence: float # 调用方给出的信度(0~1)
threshold: float # 本次校验使用的信度阈值
action: str # pass / human_review / unsupported
missing_sources: List[str] = field(default_factory=list)
verdict_id: str = field(default_factory=lambda: uuid.uuid4().hex[:12])
created_at: str = field(default_factory=lambda: datetime.now(timezone.utc).isoformat())
def to_dict(self) -> Dict[str, object]:
return {
"verdict_id": self.verdict_id,
"created_at": self.created_at,
"supported": self.supported,
"confidence": self.confidence,
"threshold": self.threshold,
"action": self.action,
"missing_sources": self.missing_sources,
"answer": self.answer,
}
class HallucinationGuard:
"""幻觉/事实性校验中间件。
`check(answer, sources, confidence, high_stakes=False)`:
- `sources`:本次 RAG 检索实际命中的文档标题列表;
- `high_stakes=True`:启用信度阈值(处置建议 / 报警解释等),
低于阈值 → `human_review`;
- 输出中所有 `[来源: X]` 声明必须出现在 `sources` 中,
否则 → `unsupported`(缺失引用列表随结论返回)。
"""
def __init__(self, default_threshold: float = 0.8) -> None:
self.default_threshold = default_threshold
self._audit: List[Dict[str, object]] = []
def check(self, answer: str, sources: Sequence[str],
confidence: float = 1.0,
high_stakes: bool = False,
threshold: Optional[float] = None) -> GuardVerdict:
"""校验一条模型输出。返回结论(不修改输出,由调用方决定如何处置)。"""
th = threshold if threshold is not None else self.default_threshold
# 1) 引用溯源:输出中声明的来源必须真实存在
declared = _SOURCE_REF_RE.findall(answer)
available = set(sources)
missing = [s.strip() for s in declared if s.strip() not in available]
supported = not missing
# 2) 高利害 → 信度阈值
if high_stakes and confidence < th:
action = "human_review"
elif not supported:
action = "unsupported"
else:
action = "pass"
verdict = GuardVerdict(
answer=answer, supported=supported, confidence=confidence,
threshold=th, action=action, missing_sources=missing,
)
self._audit.append(verdict.to_dict())
return verdict
# -- 评测集检验(定期事实一致性评测入口) ------------------------------
def evaluate(self, samples: List[Dict[str, object]]) -> Dict[str, object]:
"""批量评估事实一致性。
`samples`:`[{"answer", "sources", "confidence", "high_stakes"}, ...]`。
返回支持率 / 人工复核率 / 未支持率。子任务 #47 将扩展为评测报告。
"""
total = len(samples)
if total == 0:
return {"supported_rate": 0.0, "human_review_rate": 0.0, "total": 0}
supported = 0
human = 0
for s in samples:
v = self.check(
answer=str(s.get("answer", "")),
sources=[str(x) for x in s.get("sources", [])],
confidence=float(s.get("confidence", 1.0)),
high_stakes=bool(s.get("high_stakes", False)),
)
if v.supported:
supported += 1
if v.action == "human_review":
human += 1
return {
"supported_rate": round(supported / total, 4),
"human_review_rate": round(human / total, 4),
"unsupported_rate": round((total - supported) / total, 4),
"total": total,
}
# -- 审计 --------------------------------------------------------------
def drain_audit(self) -> List[Dict[str, object]]:
out, self._audit = self._audit, []
return out
def __repr__(self) -> str: # pragma: no cover - 调试辅助
return f"<HallucinationGuard threshold={self.default_threshold}>"
+311
View File
@@ -0,0 +1,311 @@
# -*- coding: utf-8 -*-
"""iAOP-Core · LLM 网关 —— Prompt 版本管理(EPIC #6 主体,Issue #47 雏形)。
对应 PRD 5.4「④ LLM 网关 + RAG」:
- **Prompt 版本管理**:所有提示词模板纳入版本库(semver),变更须评审并记录,
支持一键回滚;运行时绑定模板版本,确保可复现。
本模块实现 `PromptRegistry`:
- 模板资产加载(`config/prompts.template.yaml`):每个提示词有 name / version
(semver)/ text / description;
- **运行时按 (name, version) 绑定**:生产流程显式声明使用的模板版本,
即使模板后续变更,已绑定版本行为不变(可复现);
- **版本历史**:同名的多个版本并存,`promote(name, version)` 设定当前默认版本,
`rollback(name)` 回滚到上一版本(一键回滚);
- **变更审计**:`update()` / `promote()` / `rollback()` 均落结构化变更记录。
设计说明(供子任务 #47 继续细化):
- 本版实现版本库核心(绑定 / 回滚 / 审计);
- 子任务 #47 将在此基础上补齐幻觉/事实性校验中间件(见 hallucination.py)。
测试:`python -m unittest discover -s tests -v`(在 core/llm-gateway 目录下执行)。
"""
from __future__ import annotations
import re
import uuid
from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import Dict, List, Optional, Tuple
# ---------------------------------------------------------------------------
# 轻量 YAML 子集解析(与 dlp.py / router.py 同款,模块内自持保持零耦合)。
# ---------------------------------------------------------------------------
def _parse_scalar(text: str) -> str:
t = text.split(" #", 1)[0].strip()
if len(t) >= 2 and t[0] == t[-1] and t[0] in ("'", '"'):
return t[1:-1]
return t
def _strip_comments(lines: List[str]) -> List[Tuple[str, int]]:
out = []
for i, ln in enumerate(lines):
s = ln.strip()
if not s or s.startswith("#"):
continue
out.append((ln, i + 1))
return out
def _parse_node(lines: List[Tuple[str, int]], i: int, indent: int):
text, no = lines[i]
if text.lstrip(" ").startswith("- "):
items: List[object] = []
while i < len(lines):
t, no2 = lines[i]
stripped = t.lstrip(" ")
if not stripped.startswith("- "):
break
lead_j = len(t) - len(t.lstrip(" "))
if lead_j != indent:
break
item_text = stripped[2:].strip()
if not item_text:
raise ValueError(f"prompts.yaml 第 {no2} 行:list 项为空")
if ":" in item_text:
map_indent = len(t) - len(t.lstrip(" ")) + 2
lines[i] = (" " * map_indent + item_text, no2)
v, i = _parse_node(lines, i, map_indent)
items.append(v)
else:
items.append(_parse_scalar(item_text))
i += 1
return items, i
result: Dict[str, object] = {}
while i < len(lines):
t, no = lines[i]
lead_j = len(t) - len(t.lstrip(" "))
if lead_j < indent or t.lstrip(" ").startswith("- "):
break
if lead_j > indent:
raise ValueError(f"prompts.yaml 第 {no} 行缩进异常(期望 {indent},实际 {lead_j})")
if ":" not in t:
raise ValueError(f"prompts.yaml 第 {no} 行不是合法键值对:{t!r}")
key, _, rest = t.partition(":")
key = key.strip()
rest = rest.strip()
if rest:
result[key] = _parse_scalar(rest)
i += 1
continue
if i + 1 >= len(lines):
raise ValueError(f"prompts.yaml 第 {no} 行 {key!r} 缺少值")
sub_indent = len(lines[i + 1][0]) - len(lines[i + 1][0].lstrip(" "))
if sub_indent <= indent:
raise ValueError(f"prompts.yaml 第 {no} 行 {key!r} 缺少值(无嵌套内容)")
v, i = _parse_node(lines, i + 1, sub_indent)
result[key] = v
return result, i
def _load_yaml_text(text: str) -> Dict[str, object]:
lines = _strip_comments(text.splitlines())
if not lines:
return {}
top_indent = len(lines[0][0]) - len(lines[0][0].lstrip(" "))
value, next_i = _parse_node(lines, 0, top_indent)
if not isinstance(value, dict):
raise ValueError("prompts.yaml 顶层必须是 map")
if next_i < len(lines):
raise ValueError(
f"prompts.yaml 第 {lines[next_i][1]} 行:顶层存在多个节点(缩进不一致)"
)
return value
# ---------------------------------------------------------------------------
# 版本模型
# ---------------------------------------------------------------------------
_SEMVER_RE = re.compile(r"^(0|[1-9]\d*)\.(0|[1-9]\d*)\.(0|[1-9]\d*)$")
def validate_semver(version: str) -> bool:
"""校验 semver 主.次.补丁格式(不含预发布后缀,够用且严格)。"""
return bool(_SEMVER_RE.match(version))
def _cmp_semver(a: str, b: str) -> int:
"""按 semver 比较:a < b 返回负数,相等 0,a > b 正数。"""
pa, pb = (tuple(int(x) for x in v.split(".")) for v in (a, b))
return (pa > pb) - (pa < pb)
@dataclass(frozen=True)
class PromptVersion:
"""一个不可变的 Prompt 模板版本。"""
name: str
version: str
text: str
description: str = ""
created_at: str = field(default_factory=lambda: datetime.now(timezone.utc).isoformat())
def render(self, **kwargs: object) -> str:
"""用 `{key}` 占位符渲染模板(缺参保持原样,供调用方校验)。"""
return self.text.format(**kwargs)
def to_dict(self) -> Dict[str, object]:
return {
"name": self.name,
"version": self.version,
"text": self.text,
"description": self.description,
"created_at": self.created_at,
}
@dataclass(frozen=True)
class PromptChange:
"""一次模板变更/晋升/回滚的审计记录。"""
name: str
action: str # add / update / promote / rollback
version: str
previous_version: Optional[str] = None
change_id: str = field(default_factory=lambda: uuid.uuid4().hex[:12])
created_at: str = field(default_factory=lambda: datetime.now(timezone.utc).isoformat())
def to_dict(self) -> Dict[str, object]:
return {
"change_id": self.change_id,
"created_at": self.created_at,
"name": self.name,
"action": self.action,
"version": self.version,
"previous_version": self.previous_version,
}
# ---------------------------------------------------------------------------
# Prompt 版本库
# ---------------------------------------------------------------------------
class PromptRegistry:
"""提示词模板版本库:多版本并存、当前默认版本、一键回滚、变更审计。
- `current(name)`:返回当前默认版本(最新 promote 的版本);
- `get(name, version=None)`:运行时绑定指定版本(可复现);
- `update(name, text, version, ...)`:登记新版本(同版本号覆盖报错,
防止无评审覆盖——变更须评审并记录,PRD 5.4);
- `promote(name, version)`:设定当前默认版本;
- `rollback(name)`:回滚到 promote 前的版本(一键回滚)。
"""
def __init__(self) -> None:
self._versions: Dict[str, List[PromptVersion]] = {} # name -> 版本列表(升序)
self._current: Dict[str, str] = {} # name -> 当前默认版本
self._history: Dict[str, List[str]] = {} # name -> 默认版本历史
self._audit: List[Dict[str, object]] = []
@classmethod
def from_template_config(cls, path: str) -> "PromptRegistry":
"""从模板资产加载(`config/prompts.template.yaml`)。"""
with open(path, "r", encoding="utf-8") as fh:
raw = _load_yaml_text(fh.read())
registry = cls()
for m in raw.get("templates", []):
if not isinstance(m, dict):
continue
name = str(m.get("name", ""))
version = str(m.get("version", ""))
text = str(m.get("text", ""))
if not name or not version or not text:
raise ValueError(f"prompts.yaml 模板缺少 name/version/text:{m!r}")
if not validate_semver(version):
raise ValueError(f"prompts.yaml 模板 {name} 版本非法(须 semver):{version!r}")
registry.update(name, text, version,
description=str(m.get("description", "")))
if str(m.get("current", "false")).lower() == "true":
registry.promote(name, version)
return registry
# -- 版本登记 ----------------------------------------------------------
def update(self, name: str, text: str, version: str,
description: str = "") -> PromptVersion:
"""登记(或覆盖同版本)一个模板版本。变更须显式记录(审计)。"""
if not validate_semver(version):
raise ValueError(f"版本非法(须 semver):{version!r}")
existing = self._versions.setdefault(name, [])
for pv in existing:
if pv.version == version:
raise ValueError(
f"模板 {name}@{version} 已存在,不允许无评审覆盖(PRD 5.4 变更须评审)"
)
pv = PromptVersion(name=name, version=version, text=text, description=description)
existing.append(pv)
existing.sort(key=lambda v: tuple(int(x) for x in v.version.split(".")))
if name not in self._current:
self._current[name] = version
self._history[name] = [version]
self._audit.append(PromptChange(
name=name, action="add", version=version,
).to_dict())
return pv
# -- 读取 / 绑定 ------------------------------------------------------
def get(self, name: str, version: Optional[str] = None) -> PromptVersion:
"""运行时绑定:未指定版本时返回当前默认版本(可复现:显式传版本)。"""
ver = version or self._current.get(name)
if ver is None:
raise KeyError(f"模板不存在:{name}")
for pv in self._versions.get(name, []):
if pv.version == ver:
return pv
raise KeyError(f"模板 {name}@{ver} 不存在")
def current(self, name: str) -> PromptVersion:
"""返回当前默认版本(不存在则 KeyError)。"""
return self.get(name)
def versions(self, name: str) -> List[str]:
"""该模板的全部可用版本(升序)。"""
return [pv.version for pv in self._versions.get(name, [])]
# -- 晋升 / 回滚 ------------------------------------------------------
def promote(self, name: str, version: str) -> str:
"""设定当前默认版本。返回生效的版本号。"""
if not any(pv.version == version for pv in self._versions.get(name, [])):
raise KeyError(f"模板 {name}@{version} 不存在,无法晋升")
previous = self._current.get(name)
self._current[name] = version
self._history.setdefault(name, []).append(version)
self._audit.append(PromptChange(
name=name, action="promote", version=version,
previous_version=previous,
).to_dict())
return version
def rollback(self, name: str) -> Optional[str]:
"""一键回滚到 promote 前的默认版本;无历史则返回 None。"""
hist = self._history.get(name, [])
if len(hist) < 2:
return None
previous = hist[-2]
self._current[name] = previous
hist.append(previous)
self._audit.append(PromptChange(
name=name, action="rollback", version=previous,
).to_dict())
return previous
# -- 审计 / 只读 ------------------------------------------------------
def drain_audit(self) -> List[Dict[str, object]]:
out, self._audit = self._audit, []
return out
@property
def template_names(self) -> List[str]:
return sorted(self._versions.keys())
def __repr__(self) -> str: # pragma: no cover - 调试辅助
return f"<PromptRegistry templates={self.template_names}>"
+365
View File
@@ -0,0 +1,365 @@
# -*- coding: utf-8 -*-
"""iAOP-Core · LLM 网关 —— 敏感度路由规则引擎(EPIC #6 主体,Issue #43 雏形)。
对应 PRD 5.4「④ LLM 网关 + RAG」与 EPIC #6:
本地 70B(敏感/核心)+ 云端 API(脱敏/通用)**混合**,安全分级路由。
本模块实现**路由决策层**:
- 依据「敏感度路由规则」(模板配置资产)对用户 query 做**敏感度分级**,
输出路由目标:`local`(敏感/核心,数据不出厂)/ `cloud`(脱敏/通用)/
`block`(触发高危规则,直接拦截,转人工)。
- 分级规则为**配置点**:`config/router.template.yaml`,换行业只改资产,
内核零改动(对齐 dlp / rag-kb 模板化思想)。
- 路由决策前**强制先过 DLP 出站检查**:query 若命中 DLP block 规则,
一律走本地(fail-closed),云端仅在 DLP 放行时允许(PRD 5.4 数据不出厂)。
设计说明(供子任务 #43 继续细化):
- 本版实现规则匹配与分级、模板加载、评估准确率的离线脚本接口;
- 子任务 #43 将在此基础上补齐敏感度词库覆盖与准确率 ≥ 96.5% 的调优基线。
测试:`python -m unittest discover -s tests -v`(在 core/llm-gateway 目录下执行)。
"""
from __future__ import annotations
import re
import uuid
from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import Dict, List, Optional, Tuple
# ---------------------------------------------------------------------------
# 轻量 YAML 子集解析(零第三方依赖,递归下降):与 dlp.py 同款(模块内自持,
# 保持模块零耦合)。足以解析 `config/router.template.yaml` 模板资产。
# ---------------------------------------------------------------------------
def _parse_scalar(text: str) -> str:
"""去掉标量两侧引号与行内注释(`key: value # comment`)。"""
t = text.split(" #", 1)[0].strip()
if len(t) >= 2 and t[0] == t[-1] and t[0] in ("'", '"'):
return t[1:-1]
return t
def _strip_comments(lines: List[str]) -> List[Tuple[str, int]]:
"""剔除空行与整行注释,保留行号(1 起)用于报错定位。"""
out = []
for i, ln in enumerate(lines):
s = ln.strip()
if not s or s.startswith("#"):
continue
out.append((ln, i + 1))
return out
def _parse_node(lines: List[Tuple[str, int]], i: int, indent: int):
"""递归解析从 lines[i] 开始、缩进为 `indent` 的一个节点。
返回 `(value, next_i)`:value 为 dict / list / str,next_i 为下一个
未消费行的下标。
"""
text, no = lines[i]
lead = len(text) - len(text.lstrip(" "))
# ---- list 节点:`- item` 或 `- key: val`(map 项) ----
if text.lstrip(" ").startswith("- "):
items: List[object] = []
while i < len(lines):
t, no2 = lines[i]
stripped = t.lstrip(" ")
if not stripped.startswith("- "):
break
lead_j = len(t) - len(t.lstrip(" "))
if lead_j != indent:
break
item_text = stripped[2:].strip()
if not item_text:
raise ValueError(f"router.yaml 第 {no2} 行:list 项为空")
if ":" in item_text:
map_indent = len(t) - len(t.lstrip(" ")) + 2
lines[i] = (" " * map_indent + item_text, no2)
v, i = _parse_node(lines, i, map_indent)
items.append(v)
else:
items.append(_parse_scalar(item_text))
i += 1
return items, i
# ---- map 节点:`key: value` / `key:`(嵌套值) ----
result: Dict[str, object] = {}
while i < len(lines):
t, no = lines[i]
lead_j = len(t) - len(t.lstrip(" "))
if lead_j < indent or t.lstrip(" ").startswith("- "):
break
if lead_j > indent:
raise ValueError(f"router.yaml 第 {no} 行缩进异常(期望 {indent},实际 {lead_j})")
if ":" not in t:
raise ValueError(f"router.yaml 第 {no} 行不是合法键值对:{t!r}")
key, _, rest = t.partition(":")
key = key.strip()
rest = rest.strip()
if rest:
result[key] = _parse_scalar(rest)
i += 1
continue
if i + 1 >= len(lines):
raise ValueError(f"router.yaml 第 {no} 行 {key!r} 缺少值")
sub_indent = len(lines[i + 1][0]) - len(lines[i + 1][0].lstrip(" "))
if sub_indent <= indent:
raise ValueError(f"router.yaml 第 {no} 行 {key!r} 缺少值(无嵌套内容)")
v, i = _parse_node(lines, i + 1, sub_indent)
result[key] = v
return result, i
def _load_yaml_text(text: str) -> Dict[str, object]:
"""解析 YAML 子集 → 嵌套 dict/list。顶层必须为 map。"""
lines = _strip_comments(text.splitlines())
if not lines:
return {}
top_indent = len(lines[0][0]) - len(lines[0][0].lstrip(" "))
value, next_i = _parse_node(lines, 0, top_indent)
if not isinstance(value, dict):
raise ValueError("router.yaml 顶层必须是 map")
if next_i < len(lines):
raise ValueError(
f"router.yaml 第 {lines[next_i][1]} 行:顶层存在多个节点(缩进不一致)"
)
return value
# ---------------------------------------------------------------------------
# 路由目标与规则模型
# ---------------------------------------------------------------------------
class RouteTarget:
"""路由目标常量。"""
LOCAL = "local" # 敏感/核心 → 本地 70B(数据不出厂)
CLOUD = "cloud" # 脱敏/通用 → 云端 API
BLOCK = "block" # 高危 → 直接拦截,转人工确认
@dataclass(frozen=True)
class RouterRule:
"""一条敏感度路由规则。
- `target`:命中后路由到哪(local / cloud / block);
- `kind`:`keyword`(大小写不敏感子串)或 `regex`(正则);
- `category`:敏感类别(工艺参数 / 个人信息 / 高危指令等),用于审计分组。
"""
name: str
category: str
kind: str
pattern: str
target: str
description: str = ""
_compiled: Optional["re.Pattern[str]"] = field(default=None, repr=False, compare=False)
@classmethod
def from_mapping(cls, m: Dict[str, object]) -> "RouterRule":
name = str(m.get("name", ""))
if not name:
raise ValueError("router 规则缺少 name")
kind = str(m.get("kind", "keyword"))
pattern = str(m.get("pattern", ""))
if not pattern:
raise ValueError(f"router 规则 {name} 缺少 pattern")
target = str(m.get("target", RouteTarget.LOCAL))
if target not in (RouteTarget.LOCAL, RouteTarget.CLOUD, RouteTarget.BLOCK):
raise ValueError(f"router 规则 {name} 的 target 非法:{target!r}")
if kind not in ("keyword", "regex"):
raise ValueError(f"router 规则 {name} 的 kind 非法:{kind!r}")
return cls(
name=name,
category=str(m.get("category", "general")),
kind=kind,
pattern=pattern,
target=target,
description=str(m.get("description", "")),
)
def _compiled_regex(self) -> "re.Pattern[str]":
if self.kind == "regex":
return re.compile(self.pattern, re.IGNORECASE)
return re.compile(re.escape(self.pattern), re.IGNORECASE)
def find(self, text: str) -> List[Tuple[str, int, int]]:
"""返回 (匹配文本, 起始, 结束) 列表;空串 pattern 返回空。"""
if not self.pattern:
return []
return [(m.group(0), m.start(), m.end()) for m in self._compiled_regex().finditer(text)]
@dataclass(frozen=True)
class RouteDecision:
"""一次路由决策结果(含审计所需上下文)。"""
query: str
target: str
reason: str # rule_hit / no_rule / dlp_blocked
rule_name: Optional[str] = None # 命中的规则(rule_hit 时)
category: Optional[str] = None
decision_id: str = field(default_factory=lambda: uuid.uuid4().hex[:12])
created_at: str = field(default_factory=lambda: datetime.now(timezone.utc).isoformat())
def to_dict(self) -> Dict[str, object]:
return {
"decision_id": self.decision_id,
"created_at": self.created_at,
"query": self.query,
"target": self.target,
"reason": self.reason,
"rule_name": self.rule_name,
"category": self.category,
}
# ---------------------------------------------------------------------------
# 敏感度路由引擎
# ---------------------------------------------------------------------------
class SensitivityRouter:
"""敏感度路由引擎:对 query 做分级路由(local / cloud / block)。
路由优先级(fail-closed):
1. DLP 出站检查拦截 → `block`(转发人工,绝不发云端);
2. 命中 `block` 路由规则 → `block`;
3. 命中 `local` 路由规则 → `local`(敏感优先本地,规则可覆盖 cloud);
4. 命中 `cloud` 规则 → `cloud`;
5. 未命中任何规则 → 默认 `local`(保守:未知 = 敏感,数据不出厂)。
"""
# 内置保底规则:即使未加载任何配置,通用高危/敏感内容默认生效。
ROUTER_DEFAULT_RULES: Tuple[RouterRule, ...] = (
RouterRule(
name="rt_id_card", category="pii", kind="regex",
pattern=r"\d{17}[\dXx]", target=RouteTarget.LOCAL,
description="身份证号(敏感,走本地)",
),
RouterRule(
name="rt_mobile", category="pii", kind="regex",
pattern=r"1[3-9]\d{9}", target=RouteTarget.LOCAL,
description="手机号(敏感,走本地)",
),
RouterRule(
name="rt_emergency_cmd", category="safety", kind="keyword",
pattern="停机", target=RouteTarget.BLOCK,
description="停机等安全指令(高危,转人工)",
),
)
def __init__(self, rules: Optional[List[RouterRule]] = None,
default_target: str = RouteTarget.LOCAL,
audit: bool = True) -> None:
# 内置保底 + 模板规则(同名覆盖内置:模板定制优先)
merged: Dict[str, RouterRule] = {r.name: r for r in self.ROUTER_DEFAULT_RULES}
for r in (rules or []):
merged[r.name] = r
self._rules: List[RouterRule] = list(merged.values())
self.default_target = default_target
self.audit = audit
self._audit_log: List[Dict[str, object]] = []
@classmethod
def from_template_config(cls, path: str,
default_target: str = RouteTarget.LOCAL) -> "SensitivityRouter":
"""从模板资产加载路由规则(`config/router.template.yaml`)。"""
with open(path, "r", encoding="utf-8") as fh:
raw = _load_yaml_text(fh.read())
rules = []
for m in raw.get("rules", []):
if isinstance(m, dict):
rules.append(RouterRule.from_mapping(m))
return cls(rules=rules, default_target=default_target)
# -- 决策 --------------------------------------------------------------
def route(self, query: str, dlp_blocked: bool = False) -> RouteDecision:
"""对单条 query 做路由决策。
`dlp_blocked`:上游 DLP 出站检查结果(true = 已拦截)。
命中 block 或 DLP 拦截时返回 `block`(fail-closed)。
"""
# 1) DLP 已拦截 → 直接 block
if dlp_blocked:
decision = RouteDecision(
query=query, target=RouteTarget.BLOCK,
reason="dlp_blocked", category="dlp",
)
self._record(decision)
return decision
# 2) 逐条规则(模板配置顺序 = 优先级)
for rule in self._rules:
if rule.find(query):
decision = RouteDecision(
query=query, target=rule.target,
reason="rule_hit", rule_name=rule.name,
category=rule.category,
)
self._record(decision)
return decision
# 3) 无规则命中 → 保守默认
decision = RouteDecision(
query=query, target=self.default_target, reason="no_rule",
)
self._record(decision)
return decision
# -- 评估(Issue #49 雏形:路由准确率离线评估脚本入口) ----------------
def evaluate(self, samples: List[Dict[str, object]]) -> Dict[str, object]:
"""离线评估路由准确率(目标 ≥ 96.5%)。
`samples`:`[{"query": str, "expected": "local"|"cloud"|"block"}, ...]`。
返回总体准确率 + 每类明细。子任务 #49 将扩展为评测集与报表脚本。
"""
total = len(samples)
if total == 0:
return {"accuracy": 0.0, "correct": 0, "total": 0, "by_target": {}}
correct = 0
by_target: Dict[str, Dict[str, int]] = {}
for s in samples:
expected = str(s["expected"])
got = self.route(str(s["query"]), dlp_blocked=bool(s.get("dlp_blocked", False)))
ok = got.target == expected
if ok:
correct += 1
agg = by_target.setdefault(expected, {"correct": 0, "total": 0})
agg["total"] += 1
if ok:
agg["correct"] += 1
return {
"accuracy": round(correct / total, 4),
"correct": correct,
"total": total,
"by_target": by_target,
}
# -- 审计 --------------------------------------------------------------
def _record(self, decision: RouteDecision) -> None:
if self.audit:
self._audit_log.append(decision.to_dict())
def drain_audit(self) -> List[Dict[str, object]]:
"""取走并清空审计记录(对接外部审计管道)。"""
out, self._audit_log = self._audit_log, []
return out
# -- 只读属性 ----------------------------------------------------------
@property
def rule_count(self) -> int:
return len(self._rules)
@property
def rule_names(self) -> List[str]:
return [r.name for r in self._rules]
def __repr__(self) -> str: # pragma: no cover - 调试辅助
return f"<SensitivityRouter rules={self.rule_count} default={self.default_target}>"
+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()