180 lines
7.4 KiB
Python
180 lines
7.4 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""一次性修补:fake_db 真实 2 参签名 / api 契约修正 / 陈旧测试列名对齐。"""
|
||
import io
|
||
import os
|
||
|
||
ROOT = os.path.dirname(os.path.abspath(__file__))
|
||
|
||
|
||
def patch(path, pairs, required=True):
|
||
p = os.path.join(ROOT, path)
|
||
src = io.open(p, encoding="utf-8").read()
|
||
for old, new in pairs:
|
||
if old not in src:
|
||
if required:
|
||
raise SystemExit("NOT FOUND in %s:\n%s" % (path, old[:200]))
|
||
continue
|
||
src = src.replace(old, new, 1)
|
||
io.open(p, "w", encoding="utf-8").write(src)
|
||
print("patched", path)
|
||
|
||
|
||
# ---------------------------------------------------------------- fake_db
|
||
# QC#3:fake_db 的 U 曾是 3 参 (tbl,row,where),掩盖了真实 sqlor 2 参签名。
|
||
# 现改为真实签名 U(tbl, ns):ns 同时含 SET 值与 WHERE 条件(条件键 = id/tenant_id)。
|
||
patch("tests/fake_db.py", [
|
||
(""" def U(self, tbl, row, where):
|
||
self._guard_readonly(tbl, "U")
|
||
hits = 0
|
||
for exist in self.tables.get(tbl, []):
|
||
if self._match(exist, where):
|
||
exist.update(row)
|
||
hits += 1
|
||
self.sql_log.append(("U", tbl, dict(row), dict(where)))
|
||
return hits""",
|
||
""" #: 真实 sqlor U(table, ns) 只收 2 参——ns 同时携带 SET 值与 WHERE 条件。
|
||
#: 条件键约定为主键/租户键(本模块 update_ref 的 row 从不含这两列)。
|
||
WHERE_KEYS = ("id", "tenant_id")
|
||
|
||
def U(self, tbl, ns):
|
||
\"\"\"真实 sqlor 签名:U(table, ns),2 参(crud-spec Pitfall 22)。\"\"\"
|
||
self._guard_readonly(tbl, "U")
|
||
ns = dict(ns or {})
|
||
where = {k: ns.pop(k) for k in self.WHERE_KEYS if k in ns}
|
||
if not where:
|
||
raise AssertionError(
|
||
"sqlor U(table, ns) 的 ns 必须含 WHERE 条件键 %s,实得 %s"
|
||
% (list(self.WHERE_KEYS), sorted(ns)))
|
||
hits = 0
|
||
for exist in self.tables.get(tbl, []):
|
||
if self._match(exist, where):
|
||
exist.update(ns)
|
||
hits += 1
|
||
self.sql_log.append(("U", tbl, dict(ns), where))
|
||
return hits"""),
|
||
])
|
||
|
||
# ---------------------------------------------------------------- api.py
|
||
# ① world_get_context:基表未命中/跨租户必须报错,不得静默返回空上下文
|
||
patch("pbl_domain_ext/api.py", [
|
||
(""" base = db.read_base_row("world", world_id, tenant_id)
|
||
world_view = _base_view(base or {})
|
||
ref = db.select_ref_by_key(tenant_id, "world", world_id)""",
|
||
""" base = db.read_base_row("world", world_id, tenant_id)
|
||
world_view = _base_view(base or {})
|
||
if not world_view:
|
||
# 基表未命中或 world 属其它租户:一律按不存在处理(跨租户不得静默成功)
|
||
raise PblDomainExtError(
|
||
"PBL_DE_BASE_MISSING",
|
||
detail="world.%s 在 tenant_id=%s 下未命中" % (world_id, tenant_id))
|
||
ref = db.select_ref_by_key(tenant_id, "world", world_id)"""),
|
||
# ② team_world_list:本端点语义是「列出世界」,item.id = world id;
|
||
# 关联表主键另置 ref_pk,避免调用方把关联主键当世界 ID。
|
||
(""" for r in rows:
|
||
item = _ref_out(r)
|
||
item["world_id"] = item.get("ref_id")
|
||
item["world_code"] = item.get("ref_code")
|
||
item["world_name"] = item.get("ref_name")
|
||
items.append(item)""",
|
||
""" for r in rows:
|
||
item = _ref_out(r)
|
||
item["ref_pk"] = item.get("id") # 关联表主键
|
||
item["id"] = item.get("ref_id") # 对外身份 = world.id
|
||
item["world_id"] = item.get("ref_id")
|
||
item["world_code"] = item.get("ref_code")
|
||
item["world_name"] = item.get("ref_name")
|
||
items.append(item)"""),
|
||
])
|
||
|
||
# ---------------------------------------------------------------- 陈旧测试列名
|
||
# QC#1 定案:审计列权威名 creator_id/updater_id(models 与代码同名同型)。
|
||
patch("tests/test_domain_ref.py", [
|
||
('self.assertEqual(data["created_by"], OP)',
|
||
'self.assertEqual(data["creator_id"], OP) # QC#1 权威列名 creator_id'),
|
||
('self.assertEqual(self.rows()[0]["updated_by"], OP)',
|
||
'self.assertEqual(self.rows()[0]["updater_id"], OP) # QC#1 权威列名 updater_id'),
|
||
])
|
||
|
||
# ---------------------------------------------------------------- 契约测试自身
|
||
patch("tests/test_models_contract.py", [
|
||
(""" 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 不应被调用")
|
||
|
||
import asyncio
|
||
from pbl_domain_ext import db as dbm
|
||
old = dbm._SOR_HOLDER.get("sor")
|
||
dbm.set_sor(FakeSor())
|
||
try:
|
||
asyncio.get_event_loop().run_until_complete(
|
||
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(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")""",
|
||
""" calls = []
|
||
|
||
class AsyncSor(object):
|
||
\"\"\"真实 sqlor 形态:C/U/D/R 均为 async 且只收 2 参。\"\"\"
|
||
|
||
async def U(self, table, ns):
|
||
calls.append((table, dict(ns)))
|
||
return 1
|
||
|
||
async def C(self, table, ns):
|
||
raise AssertionError("C 不应被调用")
|
||
|
||
class SyncSor(AsyncSor):
|
||
\"\"\"同步形态(离线 fake_db)——_drive 必须同样收敛。\"\"\"
|
||
|
||
def U(self, table, ns):
|
||
calls.append((table, dict(ns)))
|
||
return 1
|
||
|
||
from pbl_domain_ext import db as dbm
|
||
for sor in (AsyncSor(), SyncSor()):
|
||
del calls[:]
|
||
affected = dbm._sor_u(sor, "pbl_domain_ref",
|
||
{"bind_state": "unbound", "updater_id": "u1"},
|
||
{"id": "pk1", "tenant_id": "t1"})
|
||
self.assertEqual(affected, 1, "返回值必须收敛为普通值,不得是 coroutine")
|
||
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")
|
||
|
||
def test_drive_converges_awaitable(self):
|
||
\"\"\"QC#3 回归:宿主 async sqlor 下不得把 coroutine 泄漏给 int()/调用方。\"\"\"
|
||
from pbl_domain_ext.db import _drive
|
||
|
||
async def _coro():
|
||
return 7
|
||
|
||
self.assertEqual(_drive(_coro()), 7)
|
||
self.assertEqual(_drive(3), 3)
|
||
|
||
def test_update_ref_rejects_empty_where(self):
|
||
from pbl_domain_ext import db as dbm
|
||
from pbl_domain_ext.errors import PblDomainExtError
|
||
with self.assertRaises(PblDomainExtError):
|
||
dbm.update_ref({}, {"bind_state": "unbound"})"""),
|
||
])
|
||
|
||
print("all patches applied")
|