diff --git a/core/model-framework/registry_api.py b/core/model-framework/registry_api.py new file mode 100644 index 0000000..b258fba --- /dev/null +++ b/core/model-framework/registry_api.py @@ -0,0 +1,341 @@ +# -*- 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() diff --git a/core/model-framework/tests/test_registry_api.py b/core/model-framework/tests/test_registry_api.py new file mode 100644 index 0000000..3d67880 --- /dev/null +++ b/core/model-framework/tests/test_registry_api.py @@ -0,0 +1,199 @@ +# -*- coding: utf-8 -*- +"""registry_api 单测 —— issue #182 [E1]:promote / rollback / 权限拒绝 / 持久化。 + +运行:python -m unittest discover -s core/model-framework/tests -p "test_registry_api.py" +或:python core/model-framework/tests/test_registry_api.py +""" +from __future__ import annotations + +import http.client +import json +import os +import sys +import tempfile +import threading +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 包(registry_api 自加载 auth 包) + +from model_framework.registry_api import make_server # noqa: E402 +from auth.session import issue_token # noqa: E402 +from auth.users import UserStore # noqa: E402 + + +class RegistryApiTest(unittest.TestCase): + @classmethod + def setUpClass(cls) -> None: + cls._tmp = tempfile.mkdtemp(prefix="iaop-registry-test-") + os.environ.pop("FBA_TOKEN_SECRET_KEY", None) # 强制降级轨 + # 用户:readonly(只读)/ engineer(写) / admin(写) + cls.store = UserStore() + cls.store.create("viewer1", "pass1234", role="readonly") + cls.store.create("engineer1", "pass1234", role="engineer") + cls.store.create("admin1", "pass1234", role="admin") + cls.tok_viewer = issue_token(1) # readonly(只读,写操作应 403) + cls.tok_engineer = issue_token(2) + cls.tok_admin = issue_token(3) + cls.server = make_server(port=0, data_dir=cls._tmp, user_store=cls.store) + cls.port = cls.server.server_address[1] + cls.thread = threading.Thread( + target=cls.server.serve_forever, daemon=True) + cls.thread.start() + # 用户:viewer(只读)/ engineer(写) / admin(写) + cls.store = UserStore() + cls.store.create("viewer1", "pass1234", role="readonly") + cls.store.create("engineer1", "pass1234", role="engineer") + cls.store.create("admin1", "pass1234", role="admin") + cls.tok_viewer = issue_token(1) # readonly(只读,写操作应 403) + cls.tok_engineer = issue_token(2) + cls.tok_admin = issue_token(3) + + @classmethod + def tearDownClass(cls) -> None: + cls.server.shutdown() + cls.server.server_close() + + # ---- 工具 ---------------------------------------------------------- + def _req(self, method, path, token=None, body=None): + conn = http.client.HTTPConnection("127.0.0.1", self.port, timeout=5) + headers = {} + if token: + headers["Authorization"] = f"Bearer {token}" + data = json.dumps(body).encode() if body is not None else None + if data is not None: + headers["Content-Type"] = "application/json" + conn.request(method, path, body=data, headers=headers) + resp = conn.getresponse() + raw = resp.read().decode("utf-8") + conn.close() + return resp.status, (json.loads(raw) if raw else {}) + + # ---- 鉴权拒绝路径 ---------------------------------------------------- + def test_401_without_token(self): + st, _ = self._req("GET", "/api/v1/registry/models") + self.assertEqual(st, 401) + + def test_403_viewer_cannot_write(self): + st, body = self._req("POST", "/api/v1/registry/models", + token=self.tok_viewer, + body={"name": "x", "version": "v1.0.0"}) + self.assertEqual(st, 403) + self.assertIn("权限不足", body.get("msg", "")) + + def test_403_viewer_cannot_promote(self): + st, _ = self._req("POST", "/api/v1/registry/models/quality_forecast/v1.0.0/promote", + token=self.tok_viewer) + self.assertEqual(st, 403) + + # ---- 读接口 --------------------------------------------------------- + def test_list_seeded(self): + st, body = self._req("GET", "/api/v1/registry/models", + token=self.tok_admin) + self.assertEqual(st, 200) + self.assertEqual(len(body.get("models", [])), 4) + names = {m["name"] for m in body["models"]} + self.assertEqual(names, + {"quality_forecast", "anomaly_detection", + "cross_process_optimizer"}) + # 种子阶段:quality_forecast 有 prod(v2.1.0) 与 staging(v2.2.0) + qf = [m for m in body["models"] if m["name"] == "quality_forecast"] + self.assertEqual({m["stage"] for m in qf}, {"prod", "staging"}) + + def test_list_filter_by_stage(self): + st, body = self._req("GET", "/api/v1/registry/models?stage=prod", + token=self.tok_admin) + self.assertEqual(st, 200) + self.assertTrue(all(m["stage"] == "prod" for m in body["models"])) + + # ---- 注册 ----------------------------------------------------------- + def test_register_new(self): + st, body = self._req("POST", "/api/v1/registry/models", + token=self.tok_engineer, + body={"name": "temp_control", + "version": "v1.0.0", + "backbone": "generic", + "description": "炉温控制模板", + "stage": "dev"}) + self.assertEqual(st, 201) + self.assertTrue(body.get("ok")) + # 已在列表 + st2, b2 = self._req("GET", "/api/v1/registry/models", + token=self.tok_admin) + self.assertIn("temp_control", + {m["name"] for m in b2["models"]}) + + def test_register_duplicate_conflict(self): + st, _ = self._req("POST", "/api/v1/registry/models", + token=self.tok_engineer, + body={"name": "quality_forecast", + "version": "v2.1.0"}) + self.assertEqual(st, 409) + + # ---- promote / rollback ---------------------------------------------- + def test_promote_flow(self): + # 注册 v9.9.9 → dev;逐级提升 → staging → prod + st, _ = self._req("POST", "/api/v1/registry/models", + token=self.tok_admin, + body={"name": "flow_test", "version": "v9.9.9", + "stage": "dev"}) + self.assertEqual(st, 201) + for _ in range(2): # dev→staging→prod + st, body = self._req( + "POST", "/api/v1/registry/models/flow_test/v9.9.9/promote", + token=self.tok_admin) + self.assertEqual(st, 200, body) + self.assertTrue(body.get("ok")) + st, body = self._req("GET", "/api/v1/registry/models?stage=prod", + token=self.tok_admin) + ft = [m for m in body["models"] + if m["name"] == "flow_test" and m["version"] == "v9.9.9"] + self.assertTrue(ft and ft[0]["stage"] == "prod") + + def test_rollback(self): + # 把 quality_forecast 的 prod 指针回滚到 v2.2.0(原 staging 版本) + st, body = self._req("POST", "/api/v1/registry/models/quality_forecast/rollback", + token=self.tok_admin, + body={"stage": "prod", "version": "v2.2.0"}) + self.assertEqual(st, 200, body) + self.assertEqual(body["stage"], "prod") + self.assertEqual(body["pointers"].get("prod"), "v2.2.0") + + def test_promote_unknown_model_404(self): + st, body = self._req("POST", "/api/v1/registry/models/nope/v1.0.0/promote", + token=self.tok_admin) + # TemplateRegistry 对未知模型抛 TemplateRegistryError → 409 + self.assertEqual(st, 409) + + # ---- 持久化 --------------------------------------------------------- + def test_persistence_after_restart(self): + # 自包含:先注册 persist_check,再重启读同一数据文件 + st, _ = self._req("POST", "/api/v1/registry/models", + token=self.tok_admin, + body={"name": "persist_check", "version": "v1.0.0", + "stage": "dev"}) + self.assertEqual(st, 201) + server2 = make_server(port=0, data_dir=self._tmp, user_store=self.store) + port2 = server2.server_address[1] + t2 = threading.Thread(target=server2.serve_forever, daemon=True) + t2.start() + try: + conn = http.client.HTTPConnection("127.0.0.1", port2, timeout=5) + conn.request("GET", "/api/v1/registry/models", + headers={"Authorization": f"Bearer {self.tok_admin}"}) + resp = conn.getresponse() + body = json.loads(resp.read().decode("utf-8")) + conn.close() + self.assertEqual(resp.status, 200) + names = {m["name"] for m in body["models"]} + self.assertIn("persist_check", names) + finally: + server2.shutdown() + server2.server_close() + + +if __name__ == "__main__": + unittest.main(verbosity=2) diff --git a/deploy/fba/README.md b/deploy/fba/README.md index d17cabe..2e92e00 100644 --- a/deploy/fba/README.md +++ b/deploy/fba/README.md @@ -150,3 +150,17 @@ docker exec -it fba_postgres psql -U postgres -d fba # 进数据库 ### 旧架构存档 - `web/` 静态演示外壳保留(FBA 停掉时兜底;/index.html、/auth/login.html 已加 C1 自动收口跳转) - 旧登录页三轨会话已精简为 FBA 单轨(session.js,FBA 接入版) + +### 模型注册表服务(issue #182 [E1],PRD 5.3 服务化) +- 代码:`core/model-framework/registry_api.py`(标准库 http.server,零依赖) +- 运行:`REGISTRY_DATA_DIR=<可挂卷目录> python3 core/model-framework/registry_api.py`(默认 :8002,绑定 127.0.0.1) +- 数据:`REGISTRY_DATA_DIR/registry.json`(JSON 持久化;首次启动种子 4 条演示模型) +- 鉴权双轨:FBA JWT(`FBA_TOKEN_SECRET_KEY` 配置时启用,写操作要求 `iaop:admin`/`iaop:studio` 权限码);未配置时降级 core/auth 会话(engineer/admin 角色可写) +- **nginx 反代约定**(加到宿主机 nginx server): + ```nginx + location /api/v1/registry/ { + proxy_pass http://127.0.0.1:8002; + proxy_set_header Host $host; + } + ``` +- 单测:`python core/model-framework/tests/test_registry_api.py`(11 用例:promote/rollback/403 权限拒绝/401/持久化)