Files
iAOP/core/model-framework/registry_api.py
T

342 lines
15 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- 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/<name>/<ver>/promote 阶段提升(写)→ 200
- POST /api/v1/registry/models/<name>/rollback 回滚(写) → 200
鉴权双轨:
- FBA 轨:`Authorization: Bearer <FBA JWT>`(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()