463 lines
22 KiB
Python
463 lines
22 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""机械契约测试:models/pbl_domain_ref.json ↔ pbl_domain_ext 代码列名/类型一致性。
|
||
|
||
覆盖 QC 退回意见 #1/#2/#5/#6 的回归防护:
|
||
#1 代码读写的 7 个列(ref_code/ref_name/bind_state/bind_at/is_deleted/creator_id/updater_id)
|
||
必须全部出现在 models fields 中,且审计列名不得写成 created_by/updated_by;
|
||
#2 ref_id/blueprint_id 必须是 str 类型(与基表 world/scene/entity 主键 str(32) 一致),
|
||
禁止 long/bigint;
|
||
#5 init/data.json 必须是合法 JSON,且按 Format B 注入 appcodes 两组
|
||
(pbl_domain_ref_type / pbl_bind_state),parentid 与 models codes cond 对得上;
|
||
#6 主键生成必须走 appPublic.uniqueID.getID(禁止 uuid4)。
|
||
"""
|
||
|
||
import json
|
||
import os
|
||
import re
|
||
import sys
|
||
import unittest
|
||
|
||
ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||
if ROOT not in sys.path:
|
||
sys.path.insert(0, ROOT) # 允许 python3 tests/test_models_contract.py 直接跑
|
||
MODEL_PATH = os.path.join(ROOT, "models", "pbl_domain_ref.json")
|
||
CRUD_PATH = os.path.join(ROOT, "json", "pbl_domain_ref.json")
|
||
INIT_DATA_PATH = os.path.join(ROOT, "init", "data.json")
|
||
SQL_PATH = os.path.join(ROOT, "sql", "pbl_domain_ext.sql")
|
||
|
||
#: api.py / db.py / base.py 实际读写的列(QC#1 清单)
|
||
CODE_USED_COLUMNS = [
|
||
"id", "tenant_id", "ref_type", "ref_id", "ref_code", "ref_name",
|
||
"blueprint_id", "class_id", "team_id", "bind_state", "bind_at",
|
||
"ext_json", "is_deleted", "creator_id", "created_at", "updater_id", "updated_at",
|
||
]
|
||
|
||
FORBIDDEN_AUDIT_NAMES = ["created_by", "updated_by", "creator", "updater", "ext"]
|
||
|
||
|
||
def _load(path):
|
||
with open(path, "r", encoding="utf-8") as fh:
|
||
return json.load(fh)
|
||
|
||
|
||
class TestModelContract(unittest.TestCase):
|
||
"""models/pbl_domain_ref.json 与代码契约同名同型。"""
|
||
|
||
@classmethod
|
||
def setUpClass(cls):
|
||
cls.model = _load(MODEL_PATH)
|
||
cls.fields = {f["name"]: f for f in cls.model["fields"]}
|
||
|
||
def test_summary_shape(self):
|
||
"""summary 必须恰一条、primary 必须是数组 ["id"](table-definition-spec)。"""
|
||
self.assertEqual(len(self.model["summary"]), 1)
|
||
s = self.model["summary"][0]
|
||
self.assertEqual(s["name"], "pbl_domain_ref")
|
||
self.assertIsInstance(s["primary"], list)
|
||
self.assertEqual(s["primary"], ["id"])
|
||
|
||
def test_only_one_table(self):
|
||
"""Q-OPEN-3 薄扩展:models/ 下只允许 pbl_domain_ref 一张表,
|
||
禁止把 world/scene/entity 基表重新定义进来。"""
|
||
names = sorted(f for f in os.listdir(os.path.join(ROOT, "models"))
|
||
if f.endswith(".json"))
|
||
self.assertEqual(names, ["pbl_domain_ref.json"])
|
||
|
||
def test_code_columns_all_declared(self):
|
||
"""QC#1:代码读写的列必须全部在 fields 里声明。"""
|
||
missing = [c for c in CODE_USED_COLUMNS if c not in self.fields]
|
||
self.assertEqual(missing, [], "models 缺列: %s" % missing)
|
||
|
||
def test_no_forbidden_column_names(self):
|
||
"""QC#1:审计列名统一 creator_id/updater_id,扩展列统一 ext_json。"""
|
||
bad = [n for n in self.fields if n in FORBIDDEN_AUDIT_NAMES]
|
||
self.assertEqual(bad, [], "存在旧名/禁用列名: %s" % bad)
|
||
|
||
def test_ref_id_and_blueprint_id_are_str(self):
|
||
"""QC#2:ref_id / blueprint_id 必须 str,长度 >= 32,禁止 long/bigint。"""
|
||
for col, minlen in (("ref_id", 32), ("blueprint_id", 32), ("id", 32)):
|
||
f = self.fields[col]
|
||
self.assertEqual(f["type"], "str", "%s 类型应为 str" % col)
|
||
self.assertIsInstance(f["length"], int)
|
||
self.assertGreaterEqual(f["length"], minlen)
|
||
for col in self.fields:
|
||
self.assertNotIn(self.fields[col]["type"], ("long", "bigint"),
|
||
"%s 不得使用 bigint" % col)
|
||
|
||
def test_str_fields_have_int_length(self):
|
||
"""str/char 必须带正整数 length(禁止 "15,2" 之类字符串写法)。"""
|
||
for name, f in self.fields.items():
|
||
if f["type"] in ("str", "char"):
|
||
self.assertIsInstance(f.get("length"), int, "%s 缺 int length" % name)
|
||
self.assertGreater(f["length"], 0)
|
||
if f["type"] in ("float", "double", "ddouble", "decimal"):
|
||
self.assertIsInstance(f.get("length"), int)
|
||
self.assertIsInstance(f.get("dec"), int)
|
||
|
||
def test_indexes_shape(self):
|
||
"""索引必须用 idxfields 数组;唯一键 = (tenant_id, ref_type, ref_id)。"""
|
||
names = set()
|
||
for idx in self.model["indexes"]:
|
||
self.assertIn("idxtype", idx)
|
||
self.assertIsInstance(idx["idxfields"], list)
|
||
self.assertNotIn(idx["name"], names)
|
||
names.add(idx["name"])
|
||
uniq = [i for i in self.model["indexes"] if i["idxtype"] == "unique"]
|
||
self.assertEqual(len(uniq), 1)
|
||
self.assertEqual(uniq[0]["idxfields"], ["tenant_id", "ref_type", "ref_id"])
|
||
|
||
def test_codes_use_parentid(self):
|
||
"""codes 引用 appcodes_kv 必须 cond parentid=,禁止 id=;禁止 module.table 点号。"""
|
||
for c in self.model["codes"]:
|
||
self.assertNotIn(".", c["table"])
|
||
if c["table"] == "appcodes_kv":
|
||
self.assertTrue(c["cond"].startswith("parentid="),
|
||
"codes cond 必须 parentid= : %s" % c)
|
||
fields = [c["field"] for c in self.model["codes"]]
|
||
self.assertEqual(len(fields), len(set(fields)), "codes 存在重复 field")
|
||
|
||
def test_list_fields_match_model(self):
|
||
"""base.LIST_FIELDS 必须与 models fields 完全一致(无多无少)。"""
|
||
from pbl_domain_ext.base import LIST_FIELDS, ALL_COLUMNS, AUDIT_FIELDS
|
||
self.assertEqual(sorted(LIST_FIELDS), sorted(self.fields.keys()))
|
||
self.assertEqual(sorted(ALL_COLUMNS), sorted(self.fields.keys()))
|
||
for a in AUDIT_FIELDS:
|
||
self.assertIn(a, self.fields)
|
||
|
||
def test_sql_ddl_columns_match_model(self):
|
||
"""sql/pbl_domain_ext.sql 的建表列必须覆盖 models 全部列。"""
|
||
if not os.path.exists(SQL_PATH):
|
||
self.skipTest("sql/pbl_domain_ext.sql 不存在(由 json2ddl 生成)")
|
||
with open(SQL_PATH, "r", encoding="utf-8") as fh:
|
||
raw_ddl = fh.read()
|
||
# 剥离 -- 行注释与 /* */ 块注释:迁移说明注释里合法提到旧列名,
|
||
# 只有**生效语句**里的列名才是契约。
|
||
ddl = re.sub(r"/\*.*?\*/", "", raw_ddl, flags=re.S)
|
||
ddl = "\n".join(ln for ln in ddl.splitlines()
|
||
if not ln.strip().startswith("--"))
|
||
for col in self.fields:
|
||
self.assertRegex(ddl, r"`%s`" % col, "DDL 缺列 %s" % col)
|
||
for bad in FORBIDDEN_AUDIT_NAMES:
|
||
self.assertNotRegex(ddl, r"`%s`" % bad, "DDL 含禁用旧列名 %s" % bad)
|
||
# 生效建表语句必须恰为 1 张表(Q-OPEN-3 薄扩展:只新增 pbl_domain_ref)
|
||
created = re.findall(r"CREATE TABLE(?: IF NOT EXISTS)?\s+`?(\w+)`?", ddl, re.I)
|
||
self.assertEqual(created, ["pbl_domain_ref"],
|
||
"薄扩展只允许建 1 张 pbl_domain_ref,实得 %s" % created)
|
||
|
||
|
||
class TestCrudJsonContract(unittest.TestCase):
|
||
"""json/pbl_domain_ref.json 符合 crud-definition-spec(QC#4)并与 models 字段名一致。"""
|
||
|
||
#: crud-definition-spec 里 params 不允许出现的自创键
|
||
FORBIDDEN_PARAM_KEYS = ("filters", "sortname", "sortorder", "rows", "fields",
|
||
"columns", "grid", "form", "components", "select_fields")
|
||
|
||
@classmethod
|
||
def setUpClass(cls):
|
||
cls.crud = _load(CRUD_PATH)
|
||
cls.model_fields = {f["name"] for f in _load(MODEL_PATH)["fields"]}
|
||
cls.params = cls.crud["params"]
|
||
|
||
def test_tblname_and_editable(self):
|
||
self.assertEqual(self.crud["tblname"], "pbl_domain_ref")
|
||
for key in ("new_data_url", "update_data_url", "delete_data_url"):
|
||
self.assertIn(key, self.params, "params 顶层缺 %s" % key)
|
||
self.assertIn("entire_url", self.params[key])
|
||
self.assertIn("editable", self.params)
|
||
|
||
def test_no_forbidden_root_keys(self):
|
||
"""根键只允许 tblname/alias/title/params;禁止 tablename/grid/form/name/type。"""
|
||
self.assertNotIn("tablename", self.crud)
|
||
for bad in ("tablename", "grid", "form", "name", "type", "components"):
|
||
self.assertNotIn(bad, self.crud)
|
||
self.assertNotIn(bad, self.params)
|
||
|
||
def test_no_self_invented_param_keys(self):
|
||
"""QC#4:filters/sortname/sortorder/rows 是自创键,规范用 sortby + data_filter。"""
|
||
for bad in self.FORBIDDEN_PARAM_KEYS:
|
||
self.assertNotIn(bad, self.params, "params 含自创键 %s" % bad)
|
||
self.assertIsInstance(self.params.get("sortby"), list, "排序必须用 sortby 数组")
|
||
self.assertIn("data_filter", self.params, "筛选必须用 data_filter")
|
||
|
||
def test_data_filter_shape(self):
|
||
"""data_filter 必须是 {AND:[{field,op,var|const}]},且 field 都在 models 中。"""
|
||
df = self.params["data_filter"]
|
||
self.assertIn("AND", df)
|
||
self.assertGreaterEqual(len(df["AND"]), 2, "AND 数组长度须 >= 2")
|
||
for cond in df["AND"]:
|
||
self.assertIn(cond["field"], self.model_fields,
|
||
"data_filter 引用未声明列 %s" % cond["field"])
|
||
self.assertIn(cond["op"], ("=", "!=", ">", ">=", "<", "<=", "IN",
|
||
"NOT IN", "LIKE", "NOT LIKE",
|
||
"IS NULL", "IS NOT NULL"))
|
||
self.assertTrue("var" in cond or "const" in cond,
|
||
"条件必须带 var 或 const: %s" % cond)
|
||
|
||
def test_browserfields_structure(self):
|
||
"""QC#4:browserfields 结构为 {exclouded:[...], alters:{field:{uitype,...}}},
|
||
不是平铺字段定义。"""
|
||
bf = self.params["browserfields"]
|
||
self.assertIsInstance(bf, dict)
|
||
self.assertIn("alters", bf, "browserfields 缺 alters 段(平铺写法不符合规范)")
|
||
self.assertIn("exclouded", bf, "browserfields 缺 exclouded 段")
|
||
self.assertIsInstance(bf["exclouded"], list)
|
||
for col in bf["exclouded"]:
|
||
self.assertIn(col, self.model_fields, "exclouded 列 %s 未在 models 声明" % col)
|
||
for col, spec in bf["alters"].items():
|
||
self.assertIn(col, self.model_fields, "alters 列 %s 未在 models 声明" % col)
|
||
self.assertIsInstance(spec, dict)
|
||
self.assertIn("uitype", spec, "alters.%s 缺 uitype" % col)
|
||
if spec["uitype"] == "code":
|
||
self.assertTrue(spec.get("dataurl") or spec.get("data"),
|
||
"code 型字段 %s 必须给 dataurl 或 data" % col)
|
||
if spec.get("dataurl"):
|
||
self.assertIn("entire_url", spec["dataurl"])
|
||
|
||
def test_exclouded_not_at_params_top(self):
|
||
"""QC#4:exclouded 属于 browserfields,不得放在 params 顶层。"""
|
||
self.assertNotIn("exclouded", self.params)
|
||
|
||
def test_editexclouded_known_columns(self):
|
||
for col in self.params.get("editexclouded", []):
|
||
self.assertIn(col, self.model_fields,
|
||
"editexclouded 列 %s 未在 models 声明" % col)
|
||
|
||
def test_editexclouded_covers_not_null_defaults(self):
|
||
"""NOT NULL 且非用户可编辑的列必须进 editexclouded,否则提交报 cannot be null。"""
|
||
model = _load(MODEL_PATH)
|
||
auto = {"id", "tenant_id", "creator_id", "created_at", "updater_id",
|
||
"updated_at", "is_deleted", "bind_at", "ref_type", "ref_id",
|
||
"ref_code", "ref_name"}
|
||
edit = set(self.params.get("editexclouded", []))
|
||
for f in model["fields"]:
|
||
if f.get("nullable") == "no" and f["name"] in auto:
|
||
self.assertIn(f["name"], edit,
|
||
"NOT NULL 列 %s 未 editexclouded" % f["name"])
|
||
|
||
def test_data_url_not_duplicated_in_editable(self):
|
||
"""QC#12:三个写 URL 只在 params 顶层一份,禁止在 editable 内重复定义。
|
||
|
||
重复定义会让 xls2crud 模板读取歧义(历史上 discount 模块因嵌套导致
|
||
自定义 create 逻辑被绕过)。editable 只允许保留 get_data_url。
|
||
"""
|
||
editable = self.params.get("editable", {})
|
||
for key in ("new_data_url", "update_data_url", "delete_data_url"):
|
||
self.assertNotIn(key, editable,
|
||
"params.editable 内不得重复定义 %s(权威位置=params 顶层)" % key)
|
||
self.assertIn("get_data_url", editable,
|
||
"editable 缺 get_data_url(列表取数地址)")
|
||
|
||
def test_field_maxlen_matches_model(self):
|
||
"""QC#23:base.FIELD_MAXLEN 必须与 models/*.json 的 length 单一来源一致。"""
|
||
import sys
|
||
sys.path.insert(0, os.path.join(ROOT, "pbl_domain_ext"))
|
||
from pbl_domain_ext import base
|
||
for f in _load(MODEL_PATH)["fields"]:
|
||
name = f["name"]
|
||
self.assertIn(name, base.FIELD_MAXLEN, "FIELD_MAXLEN 缺列 %s" % name)
|
||
if f.get("type") in ("str", "char") and f.get("length"):
|
||
self.assertEqual(base.FIELD_MAXLEN[name], f["length"],
|
||
"FIELD_MAXLEN[%s]=%s != model length=%s"
|
||
% (name, base.FIELD_MAXLEN[name], f["length"]))
|
||
else:
|
||
self.assertIsNone(base.FIELD_MAXLEN[name],
|
||
"非字符列 %s 不应有截断宽度" % name)
|
||
|
||
def test_clip_respects_column_width(self):
|
||
"""QC#23:clip() 按列宽裁剪——blueprint_id=32、ref_code=64、ref_name=128。"""
|
||
from pbl_domain_ext import base
|
||
self.assertEqual(len(base.clip("blueprint_id", "b" * 80)), 32)
|
||
self.assertEqual(len(base.clip("ref_code", "c" * 80)), 64)
|
||
self.assertEqual(len(base.clip("class_id", "x" * 80)), 64)
|
||
self.assertEqual(len(base.clip("ref_name", "n" * 200)), 128)
|
||
self.assertEqual(base.clip("ref_code", " w1 "), "w1")
|
||
# 非字符列(text/timestamp)不截断
|
||
self.assertEqual(base.clip("ext_json", "e" * 5000)[:5], "eeeee")
|
||
|
||
def test_no_new_data_url_query_params(self):
|
||
"""new_data_url 不得带 query 参数(避免与 editexclouded 合并成 list)。"""
|
||
url = self.params["new_data_url"]
|
||
self.assertNotIn("?", url)
|
||
|
||
def test_no_jinja_control_blocks(self):
|
||
"""CRUD JSON 禁止 Jinja2 控制块,只允许 {{entire_url(...)}} 插值。"""
|
||
with open(CRUD_PATH, "r", encoding="utf-8") as fh:
|
||
raw = fh.read()
|
||
self.assertIsNone(re.search(r"\{%\s*(if|for|each|end)", raw),
|
||
"CRUD JSON 含 Jinja2 控制块")
|
||
|
||
def test_editable_urls_have_dspy_endpoints(self):
|
||
"""每个 editable/data_url 指向的 .dspy 必须真实存在于 wwwroot/api/。"""
|
||
names = []
|
||
for key in ("new_data_url", "update_data_url", "delete_data_url"):
|
||
names.append(self.params[key])
|
||
for spec in self.params["browserfields"]["alters"].values():
|
||
du = spec.get("dataurl", "")
|
||
if "pbl_domain_ext" in du:
|
||
names.append(du)
|
||
api_dir = os.path.join(ROOT, "wwwroot", "api")
|
||
existing = set(os.listdir(api_dir))
|
||
for url in names:
|
||
m = re.search(r"entire_url\('([^']+)'\)", url)
|
||
self.assertIsNotNone(m, "URL 未用 entire_url 包裹: %s" % url)
|
||
fn = m.group(1).rstrip("/").split("/")[-1]
|
||
if "." not in fn: # CRUD alias 目录,由 xls2ui 生成
|
||
continue
|
||
self.assertIn(fn, existing, "缺少 .dspy 端点文件: %s" % fn)
|
||
|
||
|
||
class TestInitDataContract(unittest.TestCase):
|
||
"""QC#5:init/data.json Format B 种子必须真实存在且与 codes cond 对齐。"""
|
||
|
||
@classmethod
|
||
def setUpClass(cls):
|
||
cls.data = _load(INIT_DATA_PATH)
|
||
cls.codes = _load(MODEL_PATH)["codes"]
|
||
|
||
def test_is_valid_json_format_b(self):
|
||
self.assertIn("appcodes", self.data)
|
||
self.assertIsInstance(self.data["appcodes"], list)
|
||
self.assertGreaterEqual(len(self.data["appcodes"]), 2)
|
||
|
||
def test_parentids_cover_codes_cond(self):
|
||
parents = {g["parentid"] for g in self.data["appcodes"]}
|
||
for c in self.codes:
|
||
if c["table"] != "appcodes_kv":
|
||
continue
|
||
m = re.search(r"parentid='([^']+)'", c["cond"])
|
||
self.assertIsNotNone(m, "codes cond 解析失败: %s" % c)
|
||
self.assertIn(m.group(1), parents,
|
||
"codes 引用的编码组 %s 未在 init/data.json 注入" % m.group(1))
|
||
|
||
def test_items_and_id_length(self):
|
||
"""每组必须有 items(k/v);parentid+k 生成的 id 不得超 VARCHAR(32)。"""
|
||
for g in self.data["appcodes"]:
|
||
self.assertLessEqual(len(g["parentid"]), 22,
|
||
"parentid 过长会导致 appcodes_kv.id 超 32: %s" % g["parentid"])
|
||
self.assertTrue(g.get("items"), "%s 无 items" % g["parentid"])
|
||
for it in g["items"]:
|
||
self.assertTrue(it.get("k") and it.get("v"))
|
||
self.assertLessEqual(len("%s_%s" % (g["parentid"], it["k"])), 32)
|
||
|
||
def test_ref_type_items_match_code_enum(self):
|
||
"""world/scene/entity 三值必须齐(与 base.REF_TYPES 一致)。"""
|
||
from pbl_domain_ext.base import REF_TYPES, BIND_STATES
|
||
groups = {g["parentid"]: {it["k"] for it in g["items"]} for g in self.data["appcodes"]}
|
||
self.assertEqual(groups["pbl_domain_ref_type"], set(REF_TYPES))
|
||
self.assertEqual(groups["pbl_bind_state"], set(BIND_STATES))
|
||
|
||
|
||
class TestIdGeneration(unittest.TestCase):
|
||
"""QC#6:主键必须走 appPublic.uniqueID.getID,禁止 uuid4。"""
|
||
|
||
def test_no_uuid_in_package(self):
|
||
pkg = os.path.join(ROOT, "pbl_domain_ext")
|
||
for fn in sorted(os.listdir(pkg)):
|
||
if not fn.endswith(".py"):
|
||
continue
|
||
with open(os.path.join(pkg, fn), "r", encoding="utf-8") as fh:
|
||
src = fh.read()
|
||
self.assertNotIn("uuid", src, "%s 不得使用 uuid 生成主键" % fn)
|
||
|
||
def test_base_uses_platform_getid(self):
|
||
with open(os.path.join(ROOT, "pbl_domain_ext", "base.py"), "r", encoding="utf-8") as fh:
|
||
src = fh.read()
|
||
self.assertIn("from appPublic.uniqueID import getID", src)
|
||
|
||
def test_gen_id_shape(self):
|
||
from pbl_domain_ext.base import gen_id
|
||
a, b = gen_id(), gen_id()
|
||
# 平台 appPublic.uniqueID.getID() 实际返回 21 位;离线降级路径 32 位 hex。
|
||
# 两者都在 models 列宽 str(32) 内,断言按「唯一 + 不超列宽」而非固定长度。
|
||
self.assertGreaterEqual(len(a), 16, "主键过短,碰撞风险不可接受: %r" % a)
|
||
self.assertLessEqual(len(a), 32, "主键超列宽 str(32): %r" % a)
|
||
self.assertNotEqual(a, b)
|
||
|
||
|
||
class TestHostAgnosticImports(unittest.TestCase):
|
||
"""QC#4:模块宿主无关——禁止 import 具体宿主应用(sage/pipeline_app 等)。"""
|
||
|
||
FORBIDDEN_HOSTS = ("sage", "pipeline_app", "hrs6", "hrs7")
|
||
ALLOWED_ROOTS = ("ahserver", "sqlor", "appPublic", "apppublic", "pbl_domain_ext",
|
||
"json", "os", "re", "sys", "time", "hashlib", "unittest")
|
||
|
||
def test_no_host_import(self):
|
||
pkg = os.path.join(ROOT, "pbl_domain_ext")
|
||
pat = re.compile(r"^\s*(?:from|import)\s+([A-Za-z_][\w\.]*)", re.M)
|
||
for fn in sorted(os.listdir(pkg)):
|
||
if not fn.endswith(".py"):
|
||
continue
|
||
with open(os.path.join(pkg, fn), "r", encoding="utf-8") as fh:
|
||
src = fh.read()
|
||
for mod in pat.findall(src):
|
||
top = mod.split(".")[0]
|
||
self.assertNotIn(top, self.FORBIDDEN_HOSTS,
|
||
"%s 违反宿主无关铁律: import %s" % (fn, mod))
|
||
|
||
def test_serverenv_from_ahserver(self):
|
||
with open(os.path.join(ROOT, "pbl_domain_ext", "db.py"), "r", encoding="utf-8") as fh:
|
||
src = fh.read()
|
||
self.assertIn("from ahserver.serverenv import ServerEnv", src)
|
||
self.assertNotIn("from sage", src)
|
||
|
||
|
||
class TestSqlorSignatures(unittest.TestCase):
|
||
"""QC#3:sor.C/U/D/R 只收 2 参——真实签名断言(不再被 fake_db 掩盖)。"""
|
||
|
||
def test_no_three_arg_crud_calls(self):
|
||
pkg = os.path.join(ROOT, "pbl_domain_ext")
|
||
bad = re.compile(r"sor\.[CUDR]\(\s*[^()]*?,\s*[^()]*?,\s*[^()]*?\)")
|
||
for fn in sorted(os.listdir(pkg)):
|
||
if not fn.endswith(".py"):
|
||
continue
|
||
with open(os.path.join(pkg, fn), "r", encoding="utf-8") as fh:
|
||
src = fh.read()
|
||
self.assertIsNone(bad.search(src), "%s 存在 3 参 sor.C/U/D/R 调用" % fn)
|
||
|
||
def test_sor_I_single_arg(self):
|
||
pkg = os.path.join(ROOT, "pbl_domain_ext")
|
||
bad = re.compile(r"sor\.I\(\s*[^()]*?,")
|
||
for fn in sorted(os.listdir(pkg)):
|
||
if not fn.endswith(".py"):
|
||
continue
|
||
with open(os.path.join(pkg, fn), "r", encoding="utf-8") as fh:
|
||
src = fh.read()
|
||
self.assertIsNone(bad.search(src), "%s sor.I 只能 1 参" % fn)
|
||
|
||
def test_update_ref_merges_where_into_ns(self):
|
||
"""update_ref 必须把 where 并入 ns 后两参调用 sor.U。"""
|
||
calls = []
|
||
|
||
class FakeSor(object):
|
||
async def U(self, table, ns):
|
||
calls.append((table, dict(ns)))
|
||
return 1
|
||
|
||
def C(self, *a, **k):
|
||
raise AssertionError("C 不应被调用")
|
||
|
||
from pbl_domain_ext import db as dbm
|
||
old = dbm._SOR_HOLDER.get("sor")
|
||
dbm.set_sor(FakeSor())
|
||
try:
|
||
# _sor_u 是同步契约函数:内部 _drive 负责把 async sor.U 的协程收敛掉,
|
||
# 测试不得再用 run_until_complete 包一层(协程已被 _drive 消费,
|
||
# run_until_complete(None) 会抛 TypeError —— 旧断言的错误写法)。
|
||
rc = dbm._sor_u(FakeSor(), "pbl_domain_ref",
|
||
{"bind_state": "unbound", "updater_id": "u1"},
|
||
{"id": "pk1", "tenant_id": "t1"})
|
||
finally:
|
||
dbm._SOR_HOLDER["sor"] = old
|
||
self.assertEqual(int(rc or 0), 1, "应返回受影响行数")
|
||
self.assertEqual(len(calls), 1)
|
||
table, ns = calls[0]
|
||
self.assertEqual(table, "pbl_domain_ref")
|
||
# where 条件必须在 ns 内,且不被 SET 覆盖
|
||
self.assertEqual(ns["id"], "pk1")
|
||
self.assertEqual(ns["tenant_id"], "t1")
|
||
self.assertEqual(ns["bind_state"], "unbound")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
unittest.main(verbosity=2)
|