diff --git a/core/model-framework/README.md b/core/model-framework/README.md new file mode 100644 index 0000000..d1cee66 --- /dev/null +++ b/core/model-framework/README.md @@ -0,0 +1,71 @@ +# iAOP-Core · 模型框架层(AI Model Framework) + +对应 PRD 5.3「③ AI 模型框架」与 EPIC #5「内核平台化改造」。 + +本层把化工 AI 的「模型资产」统一注册、版本化、阶段化管理:模型 / 配方 / +估计器是资产,需要注册表统一管理——注册、加载、版本、阶段(灰度)、回滚、审计。 + +## 当前已交付 + +| 模块 | 对应 issue | PRD 5.3 模型 | 说明 | +|------|-----------|-------------|------| +| `template_registry` | #41 | ③ 模型模板注册/加载/版本机制 | 多版本 + 阶段(dev/staging/prod) + 提升 + 回滚 + 审计 + 持久化 | + +## 模板注册表(`template_registry.py`) + +模板化之后,模型 / 配方是「资产」,需要注册表统一管理: +- **注册**:登记一个模型模板(含主干、超参、特征列、指标、版本); +- **加载**:按 name + version/别名/stage 取出; +- **版本**:同模板多版本共存,可回滚、可审计; +- **阶段**:版本带 stage 标签(dev/staging/prod),灰度发布可控。 + +### 核心组件 + +- **`ModelTemplate`**:模型模板数据对象(name/version/backbone/hyperparams/ + feature_columns/metrics/stage),不可变、可序列化往返。 +- **`TemplateRegistry`**:注册表核心 API: + - `register`(校验完整性 + 同版本号拒重复,force 可覆盖) + - `get`(按 version / stage / 默认最新加载) + - `promote`(dev→staging→prod 逐级提升) + - `rollback`(stage 指针回退,保留历史可审计) + - `set_stage`(直接设 stage,紧急回滚) + - `list_versions` / `list_by_stage` / `stage_pointer` / `history` + - `save` / `load`(JSON 持久化,重启恢复) +- **`Stage`**:阶段枚举(DEV / STAGING / PROD),PRD 灰度三段制。 +- **`is_valid_version`**:语义化版本号校验(v1 / 1.0.0 / v1.2-rc1)。 + +### 快速开始 + +```python +from template_registry import ModelTemplate, Stage, TemplateRegistry + +reg = TemplateRegistry() +reg.register(ModelTemplate( + name="ti-quality", version="v1.0", backbone="quality_forecast", + feature_columns=("furnace_temp", "cl2_flow"), + metrics={"accuracy": 0.91}, stage="dev")) +reg.register(ModelTemplate(name="ti-quality", version="v1.1", + backbone="quality_forecast", metrics={"accuracy": 0.94})) + +reg.promote("ti-quality", "v1.1") # dev -> staging +reg.promote("ti-quality", "v1.1") # staging -> prod +reg.rollback("ti-quality", Stage.PROD, "v1.0") # 紧急回滚 + +serving = reg.get("ti-quality", stage=Stage.PROD) # 按阶段加载 +reg.save("registry.json") # 持久化审计 +``` + +## 测试 + +```bash +cd core/model-framework +python -m unittest discover -s tests -v +python _sanity_check.py +``` + +## 与规划模块的关系 + +接口风格对齐 #34(Model Recipe)、#36(quality_forecast)、#38 +(cross_process_optimizer)、#40(pipeline)。本模块是 #40 `ModelRegistry` +的完整版(多阶段 + 回滚 + 持久化),#40 的 `RegisterStep` 未来可直接对接本 +注册表,业务侧零改动。 diff --git a/core/model-framework/__init__.py b/core/model-framework/__init__.py new file mode 100644 index 0000000..7561660 --- /dev/null +++ b/core/model-framework/__init__.py @@ -0,0 +1,22 @@ +# -*- coding: utf-8 -*- +"""iAOP-Core · 模型框架层(AI Model Framework)。 + +对应 PRD 5.3「③ AI 模型框架」与 EPIC #5(内核平台化改造)。 + +当前已交付(自包含,不依赖未合并分支): +- ``template_registry``:模型模板注册 / 加载 / 版本机制(多版本 + 阶段 dev/ + staging/prod + 提升 + 回滚 + 审计 + JSON 持久化),issue #41。 + +规划(待相关 PR 合入后无缝对接,业务侧零改动): +- ``pipeline``(issue #40,PR #107)、``cross_process_optimizer``(issue #38, + PR #106)等可注册为本注册表的具名模板。 +""" +from model_framework.template_registry import ( # noqa: F401 + ALLOWED_BACKBONES, + ModelTemplate, + Stage, + TemplateRegistry, + TemplateRegistryError, + is_valid_version, + next_stage, +) diff --git a/core/model-framework/_sanity_check.py b/core/model-framework/_sanity_check.py new file mode 100644 index 0000000..6fe1c1b --- /dev/null +++ b/core/model-framework/_sanity_check.py @@ -0,0 +1,90 @@ +# -*- coding: utf-8 -*- +"""模型模板注册 / 加载 / 版本机制 sanity 检查(issue #41)。 + +验证 PRD 5.3 模板资产管理的端到端能力: +1. 注册多版本 → 阶段提升(dev→staging→prod)→ 回滚; +2. 按版本 / 阶段 / 默认加载; +3. JSON 持久化往返一致; +4. 审计日志完整。 +""" +import os +import sys +import tempfile + +HERE = os.path.dirname(os.path.abspath(__file__)) +if HERE not in sys.path: + sys.path.insert(0, HERE) + +from template_registry import ( # noqa: E402 + ModelTemplate, Stage, TemplateRegistry, is_valid_version, next_stage) + + +def main() -> int: + failures = [] + + # 版本号校验 + for v in ["v1", "1.0.0", "v2.1-rc3"]: + if not is_valid_version(v): + failures.append(f"版本号 {v!r} 应合法") + + # 注册 + 提升 + 回滚 + try: + reg = TemplateRegistry() + reg.register(ModelTemplate( + name="ti-quality", version="v1.0", backbone="quality_forecast", + feature_columns=("furnace_temp", "cl2_flow"), + metrics={"accuracy": 0.91}, stage="dev")) + reg.register(ModelTemplate( + name="ti-quality", version="v1.1", backbone="quality_forecast", + feature_columns=("furnace_temp", "cl2_flow", "impurity_fe"), + metrics={"accuracy": 0.94}, stage="dev")) + + assert reg.list_versions("ti-quality") == ["v1.0", "v1.1"] + + # v1.1 逐级提升到 prod + reg.promote("ti-quality", "v1.1") # dev->staging + reg.promote("ti-quality", "v1.1") # staging->prod + assert reg.stage_pointer("ti-quality", Stage.PROD) == "v1.1" + + # 紧急回滚 prod 到 v1.0 + reg.rollback("ti-quality", Stage.PROD, "v1.0") + assert reg.stage_pointer("ti-quality", Stage.PROD) == "v1.0" + assert "v1.1" in reg.list_versions("ti-quality") # 历史保留 + + # 按阶段加载 + serving = reg.get("ti-quality", stage=Stage.PROD) + assert serving.version == "v1.0", f"prod 应为 v1.0,实为 {serving.version}" + + print(f"[ti-quality] 版本={reg.list_versions('ti-quality')} " + f"prod={reg.stage_pointer('ti-quality', Stage.PROD)}") + print(f" v1.0 metrics={reg.get('ti-quality','v1.0').metrics}") + print(f" v1.1 metrics={reg.get('ti-quality','v1.1').metrics}") + print(f" 审计日志 {len(reg.history('ti-quality'))} 条") + except Exception as exc: # noqa: BLE001 + failures.append(f"注册/提升/回滚流程失败:{exc}") + + # 持久化往返 + try: + with tempfile.NamedTemporaryFile( + mode="w", suffix=".json", delete=False, encoding="utf-8") as fh: + path = fh.name + reg.save(path) + reg2 = TemplateRegistry.load(path) + assert reg2.list_versions("ti-quality") == ["v1.0", "v1.1"] + assert reg2.stage_pointer("ti-quality", Stage.PROD) == "v1.0" + print(f"[persist] save/load 往返一致,恢复 {len(reg2)} 个模板") + os.unlink(path) + except Exception as exc: # noqa: BLE001 + failures.append(f"持久化往返失败:{exc}") + + if failures: + print("\n失败项:") + for f in failures: + print(f" ✗ {f}") + return 1 + print("\n✓ template_registry sanity check 通过") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/core/model-framework/template_registry.py b/core/model-framework/template_registry.py new file mode 100644 index 0000000..6008de3 --- /dev/null +++ b/core/model-framework/template_registry.py @@ -0,0 +1,395 @@ +# -*- coding: utf-8 -*- +"""模型模板注册 / 加载 / 版本机制(PRD 5.3 ③ 模型框架)。 + +对应 issue #41(父 EPIC #5「③ AI 模型框架 配置化重构」、PRD 5.3 +「③ 模型模板注册 / 加载 / 版本机制」)。 + +PRD 5.3 的核心诉求 +------------------ + +模板化之后,模型 / 配方 / 估计器是「资产」,需要**注册表**统一管理: +- **注册**:把一个模型模板(含主干、超参、特征列、指标、版本)登记入库; +- **加载**:按名字 + 版本(或别名)取出可用的模板实例; +- **版本**:同一模板多版本共存,可回滚、可审计; +- **阶段**:版本带 stage 标签(dev / staging / prod),灰度发布可控; +- **校验**:注册时校验模板完整性,防止脏数据进库。 + +本模块交付一个独立的 ``TemplateRegistry``,自包含、不依赖未合并分支,是 +issue #40 ``ModelRegistry`` 的深化(完整版本 + 阶段 + 回滚 + 审计 + 持久化)。 + +本模块交付什么 +-------------- + +1. **``ModelTemplate``**:模型模板数据对象(name / version / backbone / + hyperparams / feature_columns / metrics / stage / extra),不可变、可序列化。 +2. **``TemplateRegistry``**:模板注册表核心 API: + - ``register``:注册一个版本(校验完整性 + 同版本号拒重复); + - ``get``:按 name + version/别名加载; + - ``promote``:把版本提升到下一 stage(dev→staging→prod); + - ``rollback``:把某 stage 回滚到指定版本; + - ``list_versions`` / ``list_by_stage`` / ``history``:查询; + - ``save`` / ``load``:JSON 持久化(文件 / 目录),便于审计与重启恢复。 +3. **``Stage``**:阶段枚举(DEV / STAGING / PROD),PRD 灰度发布三段制。 +4. **校验**:注册时校验 name/version/backbone 非空、版本号语义合法、 + stage 合法,拒绝脏模板。 + +零外部强依赖 +------------ + +纯 Python 实现,无 numpy/sklearn 依赖,CI 可加载与校验。 + +与 issue #34 / #36 / #38 / #40 的关系 +-------------------------------------- + +接口风格对齐 #34 声明式数据对象、#36 / #38 ``Recipe``、#40 ``ModelRegistry``。 +本模块是 #40 ``ModelRegistry`` 的完整版(多阶段 + 回滚 + 持久化),#40 的 +``RegisterStep`` 未来可直接对接本注册表,业务侧零改动。 +""" + +from __future__ import annotations + +import json +import os +import re +import time +from dataclasses import dataclass, field +from enum import Enum +from typing import Any, Dict, List, Optional, Tuple + +__all__ = [ + "ModelTemplate", + "TemplateRegistry", + "Stage", + "TemplateRegistryError", + "is_valid_version", + "next_stage", +] + + +class TemplateRegistryError(Exception): + """模板注册表统一异常(模板非法 / 版本冲突 / 版本不存在 / 阶段非法)。""" + + +# --------------------------------------------------------------------------- +# 阶段(Stage):灰度发布三段制 +# --------------------------------------------------------------------------- + +class Stage(str, Enum): + """模型版本阶段:DEV(开发)→ STAGING(预发)→ PROD(生产)。""" + + DEV = "dev" + STAGING = "staging" + PROD = "prod" + + @classmethod + def from_str(cls, s: str) -> "Stage": + try: + return cls(s.lower()) + except ValueError: + raise TemplateRegistryError( + f"非法阶段 {s!r},允许:{[s.value for s in Stage]}") + + +# 阶段提升顺序:dev → staging → prod +_STAGE_ORDER = [Stage.DEV, Stage.STAGING, Stage.PROD] + + +def next_stage(stage: Stage) -> Optional[Stage]: + """返回下一阶段;PROD 已是最高返回 None。""" + try: + idx = _STAGE_ORDER.index(stage) + except ValueError: + return None + if idx + 1 >= len(_STAGE_ORDER): + return None + return _STAGE_ORDER[idx + 1] + + +# --------------------------------------------------------------------------- +# 版本号语义校验(语义化版本 vMAJOR.MINOR.PATCH 或简单的 vN / N) +# --------------------------------------------------------------------------- + +_VERSION_RE = re.compile(r"^v?\d+(\.\d+)*([\-+][0-9A-Za-z.\-]+)?$") + + +def is_valid_version(version: str) -> bool: + """校验版本号是否合法(v1 / 1.0 / v1.2.3 / v1.0-rc1 等)。""" + if not isinstance(version, str) or not version.strip(): + return False + return bool(_VERSION_RE.match(version.strip())) + + +# --------------------------------------------------------------------------- +# 模型模板数据对象 +# --------------------------------------------------------------------------- + +#: 允许的主干类型(对齐 PRD 5.3 四类模型模板 + 通用) +ALLOWED_BACKBONES = ( + "quality_forecast", # ① 质量预测(issue #36) + "anomaly_detection", # 异常检测(issue #37) + "cross_process_opt", # ③ 跨工序寻优(issue #38) + "recipe_opt", # ② 配方优化 + "generic", # 通用 +) + + +@dataclass(frozen=True) +class ModelTemplate: + """模型模板:描述一个可加载的模型资产(主干 + 超参 + 特征 + 指标 + 版本)。 + + 不可变数据对象,``to_dict`` / ``from_dict`` 可序列化往返,便于持久化与 + 配置台展示。注册表以 (name, version) 为主键管理多个模板实例。 + """ + + name: str + version: str + backbone: str = "generic" + hyperparams: Dict[str, Any] = field(default_factory=dict) + feature_columns: Tuple[str, ...] = field(default_factory=tuple) + target_column: str = "" + metrics: Dict[str, float] = field(default_factory=dict) + stage: Stage = Stage.DEV + registered_at: float = field(default_factory=time.time) + description: str = "" + extra: Dict[str, Any] = field(default_factory=dict) + + def __post_init__(self) -> None: + if not self.name: + raise TemplateRegistryError("ModelTemplate 缺少 name") + if not is_valid_version(self.version): + raise TemplateRegistryError( + f"非法版本号 {self.version!r}(例:v1 / 1.0.0 / v1.2-rc1)") + if self.backbone not in ALLOWED_BACKBONES: + raise TemplateRegistryError( + f"非法主干 {self.backbone!r},允许:{ALLOWED_BACKBONES}") + if not isinstance(self.stage, Stage): + # 允许从字符串构造(dataclass frozen 用 object.__setattr__) + object.__setattr__(self, "stage", Stage.from_str(str(self.stage))) + + @property + def key(self) -> Tuple[str, str]: + return (self.name, self.version) + + def to_dict(self) -> Dict[str, Any]: + return { + "name": self.name, + "version": self.version, + "backbone": self.backbone, + "hyperparams": dict(self.hyperparams), + "feature_columns": list(self.feature_columns), + "target_column": self.target_column, + "metrics": dict(self.metrics), + "stage": self.stage.value, + "registered_at": self.registered_at, + "description": self.description, + "extra": dict(self.extra), + } + + @classmethod + def from_dict(cls, d: Dict[str, Any]) -> "ModelTemplate": + try: + return cls( + name=d["name"], + version=d["version"], + backbone=d.get("backbone", "generic"), + hyperparams=dict(d.get("hyperparams", {})), + feature_columns=tuple(d.get("feature_columns", [])), + target_column=d.get("target_column", ""), + metrics={k: float(v) for k, v in d.get("metrics", {}).items()}, + stage=Stage.from_str(d.get("stage", "dev")), + registered_at=float(d.get("registered_at", time.time())), + description=d.get("description", ""), + extra=dict(d.get("extra", {})), + ) + except KeyError as exc: # pragma: no cover + raise TemplateRegistryError(f"模板缺少必填字段:{exc}") from exc + + +# --------------------------------------------------------------------------- +# 模板注册表 +# --------------------------------------------------------------------------- + +class TemplateRegistry: + """模型模板注册表:注册 / 加载 / 版本 / 阶段提升 / 回滚 / 持久化。 + + * 以 ``name`` 维护多个 ``version``; + * 每个版本带 ``Stage``(dev/staging/prod),``promote`` 逐级提升; + * ``rollback`` 把某 stage 指针回退到指定版本(保留历史,可审计); + * ``save`` / ``load`` JSON 持久化(单文件或目录每模板一文件)。 + """ + + def __init__(self) -> None: + self._templates: Dict[str, Dict[str, ModelTemplate]] = {} + # 每个 name 的阶段指针:{stage: version} + self._stage_pointers: Dict[str, Dict[Stage, str]] = {} + # 审计日志:[(ts, action, name, version, detail)] + self._history: List[Tuple[float, str, str, str, str]] = [] + + # -- 注册 / 校验 -------------------------------------------------------- + + def register(self, template: ModelTemplate, *, force: bool = False) -> ModelTemplate: + """注册一个模板版本。 + + * 同 (name, version) 默认拒绝重复(``force=True`` 可覆盖); + * 注册后自动成为该模板 dev 阶段指针(若该 stage 无指针); + * 记入审计日志。 + """ + versions = self._templates.setdefault(template.name, {}) + if template.version in versions and not force: + raise TemplateRegistryError( + f"{template.name}@{template.version} 已存在(force=True 可覆盖)") + versions[template.version] = template + pointers = self._stage_pointers.setdefault(template.name, {}) + # 新注册默认进 dev;若 dev 无指针则指向它 + if Stage.DEV not in pointers: + pointers[Stage.DEV] = template.version + self._log("register", template.name, template.version, + f"stage={template.stage.value}") + return template + + # -- 加载 --------------------------------------------------------------- + + def get(self, name: str, + version: Optional[str] = None, + stage: Optional[Stage] = None) -> ModelTemplate: + """按 name + version 或 name + stage 加载模板。 + + * ``version`` 优先;其次 ``stage``(取该阶段指针); + * 都不提供则取该模板最新注册版本。 + """ + versions = self._templates.get(name) + if not versions: + raise TemplateRegistryError(f"模板 {name!r} 未注册") + if version is not None: + if version not in versions: + raise TemplateRegistryError( + f"{name!r} 无版本 {version!r}(可用:{sorted(versions)})") + return versions[version] + if stage is not None: + ptr = self._stage_pointers.get(name, {}).get(stage) + if ptr is None: + raise TemplateRegistryError( + f"{name!r} 无 {stage.value} 阶段指针") + return versions[ptr] + # 默认:按注册时间最新;时间相同时按版本号字典序最新(确定性) + latest = max(versions.values(), + key=lambda t: (t.registered_at, t.version)) + return latest + + # -- 阶段提升 / 回滚 ----------------------------------------------------- + + def promote(self, name: str, version: str) -> ModelTemplate: + """把指定版本提升到下一 stage(dev→staging→prod)。 + + 提升后该版本成为新 stage 的指针版本。 + """ + tpl = self.get(name, version) + nxt = next_stage(tpl.stage) + if nxt is None: + raise TemplateRegistryError( + f"{name}@{version} 已在 PROD,无法继续提升") + # 更新该模板版本的 stage(需重建不可变对象) + new_tpl = ModelTemplate(**{**tpl.to_dict(), "stage": nxt.value}) + self._templates[name][version] = new_tpl + self._stage_pointers.setdefault(name, {})[nxt] = version + self._log("promote", name, version, f"{tpl.stage.value}->{nxt.value}") + return new_tpl + + def rollback(self, name: str, stage: Stage, version: str) -> ModelTemplate: + """把某 stage 的指针回退到指定版本(保留历史版本,可审计)。""" + tpl = self.get(name, version) + if not isinstance(stage, Stage): + stage = Stage.from_str(str(stage)) + self._stage_pointers.setdefault(name, {})[stage] = version + self._log("rollback", name, version, f"stage={stage.value}") + return tpl + + def set_stage(self, name: str, version: str, stage: Stage) -> ModelTemplate: + """直接把某版本设到指定 stage(覆盖 promote 的逐级约束,用于紧急回滚)。""" + tpl = self.get(name, version) + if not isinstance(stage, Stage): + stage = Stage.from_str(str(stage)) + new_tpl = ModelTemplate(**{**tpl.to_dict(), "stage": stage.value}) + self._templates[name][version] = new_tpl + self._stage_pointers.setdefault(name, {})[stage] = version + self._log("set_stage", name, version, f"stage={stage.value}") + return new_tpl + + # -- 查询 --------------------------------------------------------------- + + def list_names(self) -> List[str]: + return sorted(self._templates.keys()) + + def list_versions(self, name: str) -> List[str]: + return sorted(self._templates.get(name, {}).keys()) + + def list_by_stage(self, name: str, stage: Stage) -> List[str]: + """列出某模板处于指定 stage 的所有版本号。""" + if not isinstance(stage, Stage): + stage = Stage.from_str(str(stage)) + return sorted(v for v, t in self._templates.get(name, {}).items() + if t.stage == stage) + + def stage_pointer(self, name: str, stage: Stage) -> Optional[str]: + if not isinstance(stage, Stage): + stage = Stage.from_str(str(stage)) + return self._stage_pointers.get(name, {}).get(stage) + + def history(self, name: Optional[str] = None) -> List[Dict[str, Any]]: + """返回审计日志(可按 name 过滤)。""" + out = [] + for ts, action, n, ver, detail in self._history: + if name is not None and n != name: + continue + out.append({"time": ts, "action": action, "name": n, + "version": ver, "detail": detail}) + return out + + # -- 持久化 -------------------------------------------------------------- + + def save(self, path: str) -> None: + """把整个注册表序列化到 JSON 文件(含模板 + 阶段指针 + 审计日志)。""" + data = { + "templates": { + name: {ver: t.to_dict() for ver, t in vers.items()} + for name, vers in self._templates.items() + }, + "stage_pointers": { + name: {s.value: v for s, v in ptrs.items()} + for name, ptrs in self._stage_pointers.items() + }, + "history": [ + {"time": ts, "action": a, "name": n, "version": v, "detail": d} + for ts, a, n, v, d in self._history + ], + } + with open(path, "w", encoding="utf-8") as fh: + json.dump(data, fh, ensure_ascii=False, indent=2) + + @classmethod + def load(cls, path: str) -> "TemplateRegistry": + """从 JSON 文件恢复注册表。""" + with open(path, "r", encoding="utf-8") as fh: + data = json.load(fh) + reg = cls() + for name, vers in data.get("templates", {}).items(): + for ver, td in vers.items(): + reg._templates.setdefault(name, {})[ver] = ModelTemplate.from_dict(td) + for name, ptrs in data.get("stage_pointers", {}).items(): + for s_str, v in ptrs.items(): + reg._stage_pointers.setdefault(name, {})[Stage.from_str(s_str)] = v + for h in data.get("history", []): + reg._history.append( + (h["time"], h["action"], h["name"], h["version"], h["detail"])) + return reg + + # -- 内部 ---------------------------------------------------------------- + + def _log(self, action: str, name: str, version: str, detail: str) -> None: + self._history.append((time.time(), action, name, version, detail)) + + def __len__(self) -> int: + return sum(len(v) for v in self._templates.values()) + + def __contains__(self, name: str) -> bool: + return name in self._templates diff --git a/core/model-framework/tests/__init__.py b/core/model-framework/tests/__init__.py new file mode 100644 index 0000000..40a96af --- /dev/null +++ b/core/model-framework/tests/__init__.py @@ -0,0 +1 @@ +# -*- coding: utf-8 -*- diff --git a/core/model-framework/tests/_bootstrap.py b/core/model-framework/tests/_bootstrap.py new file mode 100644 index 0000000..208a8a8 --- /dev/null +++ b/core/model-framework/tests/_bootstrap.py @@ -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) diff --git a/core/model-framework/tests/test_template_registry.py b/core/model-framework/tests/test_template_registry.py new file mode 100644 index 0000000..ecd9e90 --- /dev/null +++ b/core/model-framework/tests/test_template_registry.py @@ -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()