feat(#41): 模型模板注册/加载/版本机制(多版本+阶段dev/staging/prod+提升+回滚+审计+持久化,PRD 5.3 ③模型模板注册加载版本机制)
This commit is contained in:
@@ -0,0 +1 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user