# -*- coding: utf-8 -*- """模型注册表 HTTP 服务 —— issue #182 [E1](PRD 5.3 服务化)。 标准库 http.server(零依赖,对齐 core/auth/auth_api.py 风格),暴露: - GET /api/v1/registry/models 列表(?stage= 过滤)→ 200 - POST /api/v1/registry/models 注册新版本(写) → 201 - POST /api/v1/registry/models///promote 阶段提升(写)→ 200 - POST /api/v1/registry/models//rollback 回滚(写) → 200 鉴权双轨: - FBA 轨:`Authorization: Bearer `(FBA_TOKEN_SECRET_KEY 配置时启用), 写操作要求权限码 `iaop:admin` 或 `iaop:studio`(回源 FBA /auth/codes); - 降级轨:core/auth 会话 token(iaop_session),写操作要求 engineer/admin 角色。 数据:服务端 JSON 持久化(REGISTRY_DATA_DIR 可挂卷);首次启动种子 4 条演示模型。 运行:python -m core.model_framework.registry_api(默认 :8002 绑定 127.0.0.1) nginx 反代约定见 deploy/fba/README.md(/api/v1/registry/ → 本服务)。 """ from __future__ import annotations import importlib.util import json import os import sys from http import HTTPStatus from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from typing import Any, Dict, List, Optional, Tuple def _load_package(name: str, path: str) -> None: """连字符目录 → 可导入包(core/model-framework → model_framework;core/auth → auth)。""" 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) _PKG_DIR = os.path.dirname(os.path.abspath(__file__)) # core/model-framework _CORE_DIR = os.path.dirname(_PKG_DIR) # core _load_package("model_framework", _PKG_DIR) _load_package("auth", os.path.join(_CORE_DIR, "auth")) from model_framework.template_registry import ( # noqa: E402 ModelTemplate, Stage, TemplateRegistry, TemplateRegistryError, next_stage, ) from auth.fba_jwt import FbaAuth, FbaAuthError # noqa: E402 from auth.users import UserStore # noqa: E402 from auth.session import parse_token # noqa: E402 WRITE_CODES = {"iaop:admin", "iaop:studio"} # FBA 权限码(写操作) FALLBACK_WRITE_ROLES = ("engineer", "admin") # 降级轨角色 # 种子模型(对齐前端 models.vue 演示数据,首次启动初始化) SEED_MODELS: List[Dict[str, Any]] = [ {"name": "quality_forecast", "version": "v2.1.0", "stage": "prod", "backbone": "quality_forecast", "description": "TiCl4 纯度预测(GBDT)"}, {"name": "quality_forecast", "version": "v2.2.0", "stage": "staging", "backbone": "quality_forecast", "description": "TiCl4 纯度预测(DNN)"}, {"name": "anomaly_detection", "version": "v1.4.0", "stage": "prod", "backbone": "anomaly_detection", "description": "炉温异常检测(iForest)"}, {"name": "cross_process_optimizer", "version": "v1.1.0", "stage": "dev", "backbone": "cross_process_opt", "description": "跨工序寻优"}, ] class RegistryError(Exception): def __init__(self, status: int, message: str) -> None: super().__init__(message) self.status = status self.message = message # --------------------------------------------------------------------------- # 服务 # --------------------------------------------------------------------------- class RegistryService: """持有 TemplateRegistry 与数据文件路径(进程级单例,handler 共享)。""" def __init__(self, data_file: str, user_store: Optional[UserStore] = None) -> None: self.data_file = data_file os.makedirs(os.path.dirname(data_file) or ".", exist_ok=True) self.registry = self._load_or_seed() self.fba: Optional[FbaAuth] = FbaAuth.from_env() # 降级轨用户源:默认 bootstrap admin(admin/change-me-now,id=1); # 测试/生产可注入共享 UserStore(与 auth_api 同源)。 self._store = user_store if self._store is None: self._store = UserStore() self._store.ensure_bootstrap_admin("admin", "change-me-now") def _load_or_seed(self) -> TemplateRegistry: if os.path.exists(self.data_file): return TemplateRegistry.load(self.data_file) reg = TemplateRegistry() for m in SEED_MODELS: stage = Stage.from_str(m["stage"]) reg.register(ModelTemplate( name=m["name"], version=m["version"], backbone=m["backbone"], description=m.get("description", ""), stage=stage, )) reg.set_stage(m["name"], m["version"], stage) self._save(reg) return reg def _save(self, reg: Optional[TemplateRegistry] = None) -> None: (reg or self.registry).save(self.data_file) # ---- 鉴权 ---------------------------------------------------------- def auth(self, headers: Dict[str, str]) -> Tuple[str, Any]: """校验请求身份,返回 (mode, principal);失败抛 RegistryError(401)。""" authorization = (headers.get("Authorization") or "").strip() if self.fba is not None and authorization.startswith("Bearer "): try: claims = self.fba.verify(authorization) codes = self.fba.get_codes(authorization) return ("fba", {"user": claims.get("sub"), "codes": codes}) except FbaAuthError as exc: raise RegistryError(exc.status, exc.message) # 降级轨:core/auth 会话 token token = authorization[7:].strip() if authorization.startswith("Bearer ") else "" if not token: cookie = (headers.get("Cookie") or "").split(";") for part in cookie: k, _, v = part.strip().partition("=") if k == "iaop_session": token = v sess = parse_token(token) if token else None if sess is None: raise RegistryError(HTTPStatus.UNAUTHORIZED, "未登录或会话已过期") user = self._store.get(sess.user_id) if user is None or not user.active: raise RegistryError(HTTPStatus.UNAUTHORIZED, "账号不可用") return ("fallback", user) def require_write(self, mode: str, principal: Any) -> None: """写操作权限守卫:FBA 权限码或降级角色。""" if mode == "fba": codes = principal.get("codes") or set() if not codes.intersection(WRITE_CODES): raise RegistryError( HTTPStatus.FORBIDDEN, "权限不足:需要 iaop:admin 或 iaop:studio 权限码") return role = getattr(principal, "role", "") if role not in FALLBACK_WRITE_ROLES: raise RegistryError( HTTPStatus.FORBIDDEN, f"权限不足:当前角色 {role or 'unknown'} 不可执行写操作") # --------------------------------------------------------------------------- # HTTP Handler # --------------------------------------------------------------------------- class RegistryAPIHandler(BaseHTTPRequestHandler): server: "RegistryHTTPServer" def log_message(self, fmt: str, *args: Any) -> None: print(f"[registry] {self.address_string()} {fmt % args}", flush=True) # ---- 路由 ------------------------------------------------------------ def _route(self) -> None: parts = [p for p in (self.path.split("?", 1)[0]).split("/") if p] # 期望:api/v1/registry/models[...] if len(parts) < 4 or parts[:4] != ["api", "v1", "registry", "models"]: raise RegistryError(HTTPStatus.NOT_FOUND, "未找到路由") tail = parts[4:] if self.command == "GET" and not tail: self._list_models() elif self.command == "POST" and not tail: self._register() elif (self.command == "POST" and len(tail) == 3 and tail[2] == "promote"): self._promote(tail[0], tail[1]) elif (self.command == "POST" and len(tail) == 2 and tail[1] == "rollback"): self._rollback(tail[0]) else: raise RegistryError(HTTPStatus.NOT_FOUND, "未找到路由") def _read_body(self) -> Dict[str, Any]: try: length = int(self.headers.get("Content-Length") or 0) except ValueError: raise RegistryError(HTTPStatus.BAD_REQUEST, "Content-Length 非法") if length <= 0: return {} raw = self.rfile.read(length) try: return json.loads(raw.decode("utf-8") or "{}") except (ValueError, UnicodeDecodeError): raise RegistryError(HTTPStatus.BAD_REQUEST, "请求体不是合法 JSON") # ---- GET /api/v1/registry/models ------------------------------------- def _list_models(self) -> None: stage_filter = self._query("stage") mode, principal = self.server.svc.auth(self.headers) # 读接口也要求登录 reg = self.server.svc.registry out: List[Dict[str, Any]] = [] for name in sorted(reg.list_names()): for ver in reg.list_versions(name): tpl = reg.get(name, version=ver) if stage_filter and tpl.stage.value != stage_filter: continue out.append(tpl.to_dict()) self._json(HTTPStatus.OK, { "models": out, "auth_mode": mode, "user": principal.get("user") if mode == "fba" else getattr(principal, "username", None), }) # ---- POST /api/v1/registry/models ------------------------------------ def _register(self) -> None: mode, principal = self.server.svc.auth(self.headers) self.server.svc.require_write(mode, principal) body = self._read_body() name = (body.get("name") or "").strip() version = (body.get("version") or "").strip() if not name or not version: raise RegistryError(HTTPStatus.BAD_REQUEST, "name / version 必填") try: stage = Stage.from_str((body.get("stage") or "dev")) except ValueError: raise RegistryError(HTTPStatus.BAD_REQUEST, "stage 非法(dev/staging/prod)") tpl = ModelTemplate( name=name, version=version, backbone=(body.get("backbone") or "generic"), description=(body.get("description") or ""), stage=stage, ) try: self.server.svc.registry.register(tpl) self.server.svc.registry.set_stage(name, version, stage) except TemplateRegistryError as exc: raise RegistryError(HTTPStatus.CONFLICT, str(exc)) self.server.svc._save() self._json(HTTPStatus.CREATED, {"ok": True, "model": tpl.to_dict()}) # ---- POST .../{name}/{version}/promote -------------------------------- def _promote(self, name: str, version: str) -> None: mode, principal = self.server.svc.auth(self.headers) self.server.svc.require_write(mode, principal) try: tpl = self.server.svc.registry.promote(name, version) except TemplateRegistryError as exc: raise RegistryError(HTTPStatus.CONFLICT, str(exc)) self.server.svc._save() self._json(HTTPStatus.OK, {"ok": True, "model": tpl.to_dict()}) # ---- POST .../{name}/rollback ----------------------------------------- def _rollback(self, name: str) -> None: mode, principal = self.server.svc.auth(self.headers) self.server.svc.require_write(mode, principal) body = self._read_body() stage_str = (body.get("stage") or "prod") version = (body.get("version") or "").strip() if not version: raise RegistryError(HTTPStatus.BAD_REQUEST, "version 必填") try: stage = Stage.from_str(stage_str) except ValueError: raise RegistryError(HTTPStatus.BAD_REQUEST, "stage 非法") try: tpl = self.server.svc.registry.rollback(name, stage, version) except TemplateRegistryError as exc: raise RegistryError(HTTPStatus.CONFLICT, str(exc)) self.server.svc._save() ptrs = self.server.svc.registry._stage_pointers.get(name, {}) self._json(HTTPStatus.OK, { "ok": True, "model": tpl.to_dict(), "stage": stage.value, "pointers": {st.value: ver for st, ver in ptrs.items()}, }) # ---- 辅助 ------------------------------------------------------------- def _query(self, key: str) -> Optional[str]: query = self.path.split("?", 1)[1] if "?" in self.path else "" for pair in query.split("&"): k, _, v = pair.partition("=") if k == key: return v return None def _json(self, status: int, payload: Dict[str, Any]) -> None: data = json.dumps(payload, ensure_ascii=False).encode("utf-8") self.send_response(status) self.send_header("Content-Type", "application/json; charset=utf-8") self.send_header("Content-Length", str(len(data))) self.end_headers() self.wfile.write(data) # ---- 入口 ------------------------------------------------------------- def do_GET(self) -> None: self._dispatch() def do_POST(self) -> None: self._dispatch() def _dispatch(self) -> None: try: self._route() except RegistryError as exc: self._json(exc.status, {"code": exc.status, "msg": exc.message}) except Exception as exc: # noqa: BLE001 self._json(HTTPStatus.INTERNAL_SERVER_ERROR, {"code": 500, "msg": f"服务器内部错误: {exc}"}) class RegistryHTTPServer(ThreadingHTTPServer): daemon_threads = True def __init__(self, addr: Tuple[str, int], svc: RegistryService) -> None: self.svc = svc super().__init__(addr, RegistryAPIHandler) def make_server(port: int = 8002, data_dir: str = "", user_store: Optional[UserStore] = None) -> RegistryHTTPServer: data_dir = data_dir or os.environ.get( "REGISTRY_DATA_DIR", os.path.join("deploy", "data", "registry")) data_file = os.path.join(data_dir, "registry.json") svc = RegistryService(data_file, user_store=user_store) return RegistryHTTPServer(("127.0.0.1", port), svc) if __name__ == "__main__": port = int(os.environ.get("REGISTRY_PORT", "8002")) server = make_server(port) print(f"[registry] 模型注册表服务监听 :{port} " f"(data={server.svc.data_file},fba_auth={server.svc.fba is not None})", flush=True) server.serve_forever()