342 lines
15 KiB
Python
342 lines
15 KiB
Python
# -*- 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()
|