Files
iAOP/core/llm-gateway/prompts.py

312 lines
12 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- 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}>"