feat(#41): 模型模板注册/加载/版本机制(多版本+阶段dev/staging/prod+提升+回滚+审计+持久化,PRD 5.3 ③模型模板注册加载版本机制)

This commit is contained in:
2026-08-05 01:19:01 +08:00
parent 793dd0a3b8
commit b6392525ce
7 changed files with 837 additions and 0 deletions
+1
View File
@@ -0,0 +1 @@
# -*- coding: utf-8 -*-
+23
View File
@@ -0,0 +1,23 @@
# -*- coding: utf-8 -*-
"""测试引导:把连字符目录 ``core/model-framework`` 加载为可导入包
``model_framework``,使测试可 ``from model_framework import ...``。
"""
import importlib.util
import os
import sys
PKG_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
def _load_package(name: str, path: str) -> None:
if name in sys.modules:
return
init_py = os.path.join(path, "__init__.py")
spec = importlib.util.spec_from_file_location(
name, init_py, submodule_search_locations=[path])
module = importlib.util.module_from_spec(spec)
sys.modules[name] = module
spec.loader.exec_module(module)
_load_package("model_framework", PKG_DIR)
@@ -0,0 +1,235 @@
# -*- coding: utf-8 -*-
"""模型模板注册 / 加载 / 版本机制单元测试(issue #41)。
覆盖:
- 版本号语义校验 ``is_valid_version``;
- Stage 枚举与 ``next_stage`` 阶段提升顺序;
- ModelTemplate 构造校验(name/version/backbone/stage)+ 序列化往返;
- TemplateRegistry:注册(拒重复 / force 覆盖)、加载(version/stage/默认)、
阶段提升 promote、回滚 rollback、set_stage、查询(list_*)、审计日志、
JSON 持久化 save/load 往返一致性。
"""
import json
import os
import sys
import tempfile
import time
import unittest
HERE = os.path.dirname(os.path.abspath(__file__))
if HERE not in sys.path:
sys.path.insert(0, HERE)
import _bootstrap # noqa: E402 加载 model_framework 包
from model_framework import ( # noqa: E402
ModelTemplate,
Stage,
TemplateRegistry,
TemplateRegistryError,
is_valid_version,
next_stage,
)
class TestVersionValidation(unittest.TestCase):
def test_valid_versions(self):
for v in ["v1", "1.0", "v1.2.3", "1.2.3", "v1.0-rc1", "v2.0.0+build5"]:
self.assertTrue(is_valid_version(v), f"应合法:{v}")
def test_invalid_versions(self):
for v in ["", "v", "abc", "v1.x", "1..2", None, "v 1"]:
self.assertFalse(is_valid_version(v), f"应非法:{v!r}")
class TestStage(unittest.TestCase):
def test_from_str(self):
self.assertEqual(Stage.from_str("dev"), Stage.DEV)
self.assertEqual(Stage.from_str("PROD"), Stage.PROD)
def test_from_str_invalid(self):
with self.assertRaises(TemplateRegistryError):
Stage.from_str("qa")
def test_next_stage(self):
self.assertEqual(next_stage(Stage.DEV), Stage.STAGING)
self.assertEqual(next_stage(Stage.STAGING), Stage.PROD)
self.assertIsNone(next_stage(Stage.PROD))
class TestModelTemplate(unittest.TestCase):
def test_construct_minimal(self):
t = ModelTemplate(name="m", version="v1")
self.assertEqual(t.backbone, "generic")
self.assertEqual(t.stage, Stage.DEV)
def test_rejects_empty_name(self):
with self.assertRaises(TemplateRegistryError):
ModelTemplate(name="", version="v1")
def test_rejects_bad_version(self):
with self.assertRaises(TemplateRegistryError):
ModelTemplate(name="m", version="abc")
def test_rejects_bad_backbone(self):
with self.assertRaises(TemplateRegistryError):
ModelTemplate(name="m", version="v1", backbone="magic")
def test_accepts_known_backbones(self):
for b in ("quality_forecast", "anomaly_detection",
"cross_process_opt", "recipe_opt", "generic"):
ModelTemplate(name="m", version="v1", backbone=b)
def test_roundtrip(self):
t = ModelTemplate(
name="qa-model", version="v1.2.0", backbone="quality_forecast",
hyperparams={"lr": 0.1}, feature_columns=("a", "b"),
target_column="y", metrics={"accuracy": 0.93},
stage="prod", description="d", extra={"k": "v"})
t2 = ModelTemplate.from_dict(json.loads(json.dumps(t.to_dict())))
self.assertEqual(t, t2)
self.assertEqual(t2.stage, Stage.PROD)
self.assertEqual(t2.metrics["accuracy"], 0.93)
class TestRegistryRegisterLoad(unittest.TestCase):
def setUp(self):
self.reg = TemplateRegistry()
self.t1 = ModelTemplate(name="m", version="v1", backbone="quality_forecast")
self.t2 = ModelTemplate(name="m", version="v2", backbone="quality_forecast")
def test_register_and_get_by_version(self):
self.reg.register(self.t1)
self.assertEqual(self.reg.get("m", "v1").version, "v1")
def test_register_duplicate_rejected(self):
self.reg.register(self.t1)
with self.assertRaises(TemplateRegistryError):
self.reg.register(self.t1)
def test_register_force_overwrites(self):
self.reg.register(self.t1)
t1_updated = ModelTemplate(
name="m", version="v1", description="updated")
self.reg.register(t1_updated, force=True)
self.assertEqual(self.reg.get("m", "v1").description, "updated")
def test_get_missing_name(self):
with self.assertRaises(TemplateRegistryError):
self.reg.get("nope")
def test_get_missing_version(self):
self.reg.register(self.t1)
with self.assertRaises(TemplateRegistryError):
self.reg.get("m", "v99")
def test_get_default_latest(self):
self.reg.register(self.t1)
time.sleep(0.01)
self.reg.register(self.t2)
self.assertEqual(self.reg.get("m").version, "v2")
def test_get_by_stage_pointer(self):
self.reg.register(self.t1)
# 新注册默认进 dev
self.assertEqual(self.reg.get("m", stage=Stage.DEV).version, "v1")
with self.assertRaises(TemplateRegistryError):
self.reg.get("m", stage=Stage.PROD)
class TestPromoteRollback(unittest.TestCase):
def setUp(self):
self.reg = TemplateRegistry()
self.reg.register(ModelTemplate(name="m", version="v1"))
self.reg.register(ModelTemplate(name="m", version="v2"))
def test_promote_chain(self):
self.reg.promote("m", "v1") # dev -> staging
self.assertEqual(self.reg.stage_pointer("m", Stage.STAGING), "v1")
self.reg.promote("m", "v1") # staging -> prod
self.assertEqual(self.reg.stage_pointer("m", Stage.PROD), "v1")
def test_promote_prod_raises(self):
self.reg.promote("m", "v1")
self.reg.promote("m", "v1") # 到 prod
with self.assertRaises(TemplateRegistryError):
self.reg.promote("m", "v1") # prod 无法继续
def test_rollback_stage_pointer(self):
self.reg.promote("m", "v2") # v2 -> staging
self.reg.promote("m", "v2") # v2 -> prod
# 回滚 prod 到 v1
self.reg.rollback("m", Stage.PROD, "v1")
self.assertEqual(self.reg.stage_pointer("m", Stage.PROD), "v1")
# v2 版本本身仍在(可审计)
self.assertIn("v2", self.reg.list_versions("m"))
def test_set_stage_direct(self):
self.reg.set_stage("m", "v1", Stage.PROD)
self.assertEqual(self.reg.get("m", "v1").stage, Stage.PROD)
self.assertEqual(self.reg.stage_pointer("m", Stage.PROD), "v1")
class TestQueries(unittest.TestCase):
def test_list_names_and_versions(self):
reg = TemplateRegistry()
reg.register(ModelTemplate(name="a", version="v1"))
reg.register(ModelTemplate(name="a", version="v2"))
reg.register(ModelTemplate(name="b", version="v1"))
self.assertEqual(reg.list_names(), ["a", "b"])
self.assertEqual(reg.list_versions("a"), ["v1", "v2"])
self.assertIn("a", reg)
self.assertNotIn("c", reg)
self.assertEqual(len(reg), 3)
def test_list_by_stage(self):
reg = TemplateRegistry()
reg.register(ModelTemplate(name="m", version="v1"))
reg.register(ModelTemplate(name="m", version="v2"))
reg.promote("m", "v2") # v2 -> staging
self.assertEqual(reg.list_by_stage("m", Stage.DEV), ["v1"])
self.assertEqual(reg.list_by_stage("m", Stage.STAGING), ["v2"])
class TestHistoryAndPersist(unittest.TestCase):
def test_history_logged(self):
reg = TemplateRegistry()
reg.register(ModelTemplate(name="m", version="v1"))
reg.promote("m", "v1")
h = reg.history("m")
actions = [e["action"] for e in h]
self.assertIn("register", actions)
self.assertIn("promote", actions)
def test_history_filter_by_name(self):
reg = TemplateRegistry()
reg.register(ModelTemplate(name="a", version="v1"))
reg.register(ModelTemplate(name="b", version="v1"))
self.assertEqual(len(reg.history("a")), 1)
self.assertEqual(len(reg.history("b")), 1)
def test_save_load_roundtrip(self):
reg = TemplateRegistry()
reg.register(ModelTemplate(
name="m", version="v1", backbone="quality_forecast",
metrics={"accuracy": 0.9}, stage="dev"))
reg.promote("m", "v1")
with tempfile.NamedTemporaryFile(
mode="w", suffix=".json", delete=False, encoding="utf-8") as fh:
path = fh.name
try:
reg.save(path)
reg2 = TemplateRegistry.load(path)
self.assertEqual(reg2.list_names(), ["m"])
self.assertEqual(reg2.get("m", "v1").backbone, "quality_forecast")
self.assertEqual(reg2.get("m", "v1").metrics["accuracy"], 0.9)
# 阶段指针恢复
self.assertEqual(reg2.stage_pointer("m", Stage.STAGING), "v1")
# 审计日志恢复
self.assertTrue(len(reg2.history("m")) >= 2)
finally:
os.unlink(path)
if __name__ == "__main__":
unittest.main()