feat: 完成 issue #6 LLM 网关 + RAG 模板化(混合网关主编排)
This commit is contained in:
@@ -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,本版已提供核心)。
|
||||||
|
|||||||
@@ -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: 公开常识问答(脱敏/通用,可走云端)
|
||||||
@@ -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}>")
|
||||||
@@ -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}>"
|
||||||
@@ -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}>"
|
||||||
@@ -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}>"
|
||||||
@@ -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