Files

116 lines
3.8 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 -*-
"""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()