diff --git a/core/llm-gateway/README.md b/core/llm-gateway/README.md index 2625af3..a75e5ab 100644 --- a/core/llm-gateway/README.md +++ b/core/llm-gateway/README.md @@ -1,6 +1,6 @@ # iAOP-Core · LLM 网关(LLM Gateway) -对应 PRD 5.4「④ LLM 网关 + RAG」与 EPIC #6(Issue #48 等子任务): +对应 PRD 5.4「④ LLM 网关 + RAG」与 EPIC #6: 本地 70B(敏感/核心)+ 云端 API(脱敏/通用)**混合**,安全分级路由, **敏感数据本地闭环,仅脱敏/公开内容可走云端 API**(数据不出厂)。 @@ -8,13 +8,23 @@ ``` core/llm-gateway/ -├── __init__.py 包入口(导出 DLP 引擎 API) +├── __init__.py 包入口(导出各模块 API) ├── dlp.py DLP 敏感数据拦截引擎(Issue #48) +├── router.py 敏感度路由规则引擎(Issue #43 雏形) +├── prompts.py Prompt 版本管理(Issue #47 雏形) +├── hallucination.py 幻觉/事实性校验中间件(Issue #47 雏形) +├── gateway.py 混合网关主编排(EPIC #6 主体交付) ├── config/ -│ └── dlp.template.yaml 模板 DLP 规则资产(ti-cl4 示例,换行业只改它) +│ ├── dlp.template.yaml 模板 DLP 规则资产(ti-cl4 示例) +│ ├── router.template.yaml 模板敏感度路由规则资产 +│ └── prompts.template.yaml 模板提示词版本库资产 └── tests/ ├── _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) @@ -70,7 +80,43 @@ cd core/llm-gateway 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 拆分,待扩展) -- 敏感度路由(准确率 ≥ 96.5%)与路由准确率评估脚本(Issue #49); -- Prompt 版本管理与幻觉/事实性校验中间件(Issue #47)。 +- 敏感度路由调优与路由准确率评估脚本(Issue #49,本版已提供评估入口); +- 本地 70B 模型接入与推理封装(Issue #44); +- 云端 API(Qwen/DeepSeek)接入与安全网关(Issue #45); +- Prompt 版本管理 + 幻觉校验中间件完善(Issue #47,本版已提供核心)。 diff --git a/core/llm-gateway/__init__.py b/core/llm-gateway/__init__.py index b714c0c..09da2cc 100644 --- a/core/llm-gateway/__init__.py +++ b/core/llm-gateway/__init__.py @@ -1,19 +1,25 @@ # -*- coding: utf-8 -*- -"""iAOP-Core · LLM 网关(LLM Gateway)—— 混合 LLM 的安全出站防线。 +"""iAOP-Core · LLM 网关(LLM Gateway)—— 混合 LLM 的安全出站防线与编排。 对应 PRD 5.4「④ LLM 网关 + RAG」与 EPIC #6: 本地 70B(敏感/核心)+ 云端 API(脱敏/通用)混合,敏感数据**本地闭环**。 -当前子模块(Issue #48): -- dlp DLP 敏感数据拦截引擎:出站内容(query / RAG context / 模型输出) - 发往云端前做敏感规则检查,命中即拦截(目标 100% 拦截),全量审计。 - -后续子任务(EPIC #6 拆分,将在本包扩展): -- 敏感度路由(router)、Prompt 版本管理与幻觉校验中间件、路由准确率评估。 +模块组成: +- dlp DLP 敏感数据拦截引擎(Issue #48):出站内容(query / RAG context / + 模型输出)发往云端前做敏感规则检查,命中即拦截(目标 100% 拦截), + 全量审计。 +- router 敏感度路由规则引擎(Issue #43 雏形):敏感度分级路由(local/cloud/ + 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 目录下执行)。 """ -__version__ = "0.1.0" +__version__ = "0.2.0" from .dlp import ( DLP_DEFAULT_RULES, @@ -23,12 +29,40 @@ from .dlp import ( DlpRule, 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__ = [ - "DlpRuleKind", - "DlpRule", - "DlpHit", - "DlpResult", - "DlpEngine", - "DLP_DEFAULT_RULES", + # dlp + "DlpRuleKind", "DlpRule", "DlpHit", "DlpResult", "DlpEngine", "DLP_DEFAULT_RULES", + # router + "RouteTarget", "RouterRule", "RouteDecision", "SensitivityRouter", + # prompts + "PromptVersion", "PromptChange", "PromptRegistry", "validate_semver", + # hallucination + "GuardVerdict", "HallucinationGuard", + # gateway + "InferenceBackend", "LocalBackend", "CloudBackend", + "GatewayResult", "LLMGateway", ] diff --git a/core/llm-gateway/config/prompts.template.yaml b/core/llm-gateway/config/prompts.template.yaml new file mode 100644 index 0000000..d6b0100 --- /dev/null +++ b/core/llm-gateway/config/prompts.template.yaml @@ -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} diff --git a/core/llm-gateway/config/router.template.yaml b/core/llm-gateway/config/router.template.yaml new file mode 100644 index 0000000..3fae30b --- /dev/null +++ b/core/llm-gateway/config/router.template.yaml @@ -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: 公开常识问答(脱敏/通用,可走云端) diff --git a/core/llm-gateway/gateway.py b/core/llm-gateway/gateway.py new file mode 100644 index 0000000..d612dec --- /dev/null +++ b/core/llm-gateway/gateway.py @@ -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"") diff --git a/core/llm-gateway/hallucination.py b/core/llm-gateway/hallucination.py new file mode 100644 index 0000000..ff04aab --- /dev/null +++ b/core/llm-gateway/hallucination.py @@ -0,0 +1,142 @@ +# -*- coding: utf-8 -*- +"""iAOP-Core · LLM 网关 —— 幻觉/事实性校验中间件(EPIC #6 主体,Issue #47 雏形)。 + +对应 PRD 5.4「④ LLM 网关 + RAG」: +- **事实性校验**:RAG 答案强制**引用溯源**(返回命中文档片段+来源); + 对高利害输出(如处置建议)设置信度阈值,低于阈值触发"人工确认"; + 定期用评测集检验事实一致性。 + +本模块实现 `HallucinationGuard`: +- **引用溯源校验**:模型输出中声称引用的片段(`[来源: ]`)必须能在 + 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"" diff --git a/core/llm-gateway/prompts.py b/core/llm-gateway/prompts.py new file mode 100644 index 0000000..11ab605 --- /dev/null +++ b/core/llm-gateway/prompts.py @@ -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"" diff --git a/core/llm-gateway/router.py b/core/llm-gateway/router.py new file mode 100644 index 0000000..dc34b18 --- /dev/null +++ b/core/llm-gateway/router.py @@ -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"" diff --git a/core/llm-gateway/tests/test_gateway.py b/core/llm-gateway/tests/test_gateway.py new file mode 100644 index 0000000..da28145 --- /dev/null +++ b/core/llm-gateway/tests/test_gateway.py @@ -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() diff --git a/core/llm-gateway/tests/test_hallucination.py b/core/llm-gateway/tests/test_hallucination.py new file mode 100644 index 0000000..e8b70c7 --- /dev/null +++ b/core/llm-gateway/tests/test_hallucination.py @@ -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() diff --git a/core/llm-gateway/tests/test_prompts.py b/core/llm-gateway/tests/test_prompts.py new file mode 100644 index 0000000..267019e --- /dev/null +++ b/core/llm-gateway/tests/test_prompts.py @@ -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() diff --git a/core/llm-gateway/tests/test_router.py b/core/llm-gateway/tests/test_router.py new file mode 100644 index 0000000..0ecbc02 --- /dev/null +++ b/core/llm-gateway/tests/test_router.py @@ -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()