116 lines
3.8 KiB
Python
116 lines
3.8 KiB
Python
# -*- 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()
|