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()
|