# -*- 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