feat: 完成 issue #182 [E1] 模型注册表后端化——registry_api.py HTTP 服务(四接口+双轨鉴权+种子+持久化)+ 11 单测全过 + README 反代约定
This commit is contained in:
@@ -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)
|
||||
Reference in New Issue
Block a user