pbl_domain_ext/tests/fake_db.py
2026-09-19 00:05:32 +08:00

292 lines
10 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 -*-
"""测试用假 sqlor内存表只实现 sqlor 标准 API 子集C / U / D / R / I / sqlExe。
契约(与 pbl_domain_ext/db.py 完全对齐):
* ``C(table, row)`` 插入一行
* ``U(table, spec)`` spec = {"data": {...}, "where": {...}},返回受影响行数
* ``D(table, spec)`` spec = {"where": {...}},物理删除(本模块不使用)
* ``R(table, spec)`` spec = {"where": {...}},读取
* ``I(row)`` 单参数,表名取 context 默认表pbl_domain_ref
* ``sqlExe(sql, params)`` 先按序把 ``%s`` 插值为字面量,再用极简解析器执行
``SELECT ... FROM t WHERE a = 'x' AND b IN ('1','2') [ORDER BY ...] [LIMIT n] [OFFSET m]``
以及 ``SELECT COUNT(1) AS cnt ...````1 = 0`` 恒不匹配fail-closed 断言用)。
基表只读把守world / scene / entity 上任何 C/U/D/I 都会记入
``base_write_attempts`` 并抛 AssertionError用于机械验证「不改基表」铁律。
"""
import re
__all__ = ["FakeSor", "make_fake_sor", "BASE_FIXTURES", "DEFAULT_TABLE"]
DEFAULT_TABLE = "pbl_domain_ref"
BASE_TABLES = ("world", "scene", "entity")
BASE_FIXTURES = {
"world": [
{"id": "W1", "tenant_id": "T1", "code": "WLD-001", "name": "火星基地",
"status": "published", "created_at": "2026-09-01 10:00:00"},
{"id": "W2", "tenant_id": "T1", "code": "WLD-002", "name": "深海实验室",
"status": "draft", "created_at": "2026-09-02 10:00:00"},
{"id": "W3", "tenant_id": "T2", "code": "WLD-101", "name": "他租户世界",
"status": "draft", "created_at": "2026-09-03 10:00:00"},
],
"scene": [
{"id": "S1", "tenant_id": "T1", "world_id": "W1", "code": "SCN-001",
"name": "基地外围", "status": "published",
"created_at": "2026-09-01 11:00:00"},
{"id": "S2", "tenant_id": "T1", "world_id": "W1", "code": "SCN-002",
"name": "内部舱段", "status": "draft",
"created_at": "2026-09-01 11:05:00"},
{"id": "S3", "tenant_id": "T1", "world_id": "W2", "code": "SCN-003",
"name": "海沟观测站", "status": "draft",
"created_at": "2026-09-02 11:00:00"},
{"id": "S9", "tenant_id": "T2", "world_id": "W3", "code": "SCN-101",
"name": "他租户场景", "status": "draft",
"created_at": "2026-09-03 11:00:00"},
],
"entity": [
{"id": "E1", "tenant_id": "T1", "scene_id": "S1", "code": "ENT-001",
"name": "漫游车", "kind": "vehicle", "created_at": "2026-09-01 12:00:00"},
{"id": "E2", "tenant_id": "T1", "scene_id": "S1", "code": "ENT-002",
"name": "太阳能板", "kind": "device", "created_at": "2026-09-01 12:01:00"},
{"id": "E3", "tenant_id": "T1", "scene_id": "S2", "code": "ENT-003",
"name": "气闸", "kind": "structure", "created_at": "2026-09-01 12:02:00"},
{"id": "E9", "tenant_id": "T2", "scene_id": "S9", "code": "ENT-101",
"name": "他租户实体", "kind": "vehicle", "created_at": "2026-09-03 12:00:00"},
],
}
def _lit(value):
"""把参数转成 SQL 字面量(仅测试解析用,不参与真实 SQL"""
if value is None:
return "NULL"
if isinstance(value, (int, float)):
return str(value)
return "'" + str(value).replace("'", "''") + "'"
class FakeSor:
"""内存 sqlor 替身。"""
def __init__(self, tables=None, base_fixtures=None, default_table=DEFAULT_TABLE):
self.tables = dict(tables) if tables is not None else {DEFAULT_TABLE: []}
self.tables.setdefault(DEFAULT_TABLE, [])
self.base = {t: [dict(r) for r in (base_fixtures or BASE_FIXTURES)[t]]
for t in BASE_TABLES}
self.default_table = default_table
self.base_write_attempts = []
self.calls = []
self.closed = False
# -------------------------------------------------------- 标准 API 子集
def _guard(self, table, op):
self.calls.append((op, table))
if table in self.base:
self.base_write_attempts.append((op, table))
raise AssertionError(
f"base table {table!r} is read-only for pbl_domain_ext (op={op})")
def _rows(self, table):
if table in self.base:
return self.base[table]
return self.tables.setdefault(table, [])
async def C(self, table, row):
self._guard(table, "C")
self._rows(table).append(dict(row))
return 1
async def I(self, row):
table = (row or {}).get("__table__") or self.default_table
return await self.C(table, row)
async def U(self, table, spec):
self._guard(table, "U")
spec = spec or {}
data = spec.get("data") or {}
where = spec.get("where") or {}
hits = 0
for row in self._rows(table):
if _match(row, where):
row.update(dict(data))
hits += 1
return hits
async def D(self, table, spec):
self._guard(table, "D")
where = (spec or {}).get("where") or {}
rows = self._rows(table)
keep = [r for r in rows if not _match(r, where)]
removed = len(rows) - len(keep)
self.tables[table] = keep
return removed
async def R(self, table, spec):
where = (spec or {}).get("where") or {}
return [dict(r) for r in self._rows(table) if _match(r, where)]
async def sqlExe(self, sql, params=None):
params = list(params or [])
self.calls.append(("sqlExe", sql.split()[0] if sql else ""))
text = _bind(sql, params)
return self._run_select(text)
async def close(self):
self.closed = True
# -------------------------------------------------------- 极简 SELECT 执行
def _run_select(self, text):
text = text.strip().rstrip(";")
m = re.match(
r"(?is)^SELECT\s+(?P<sel>.*?)\s+FROM\s+`?(?P<tbl>[\w]+)`?"
r"(?:\s+WHERE\s+(?P<where>.*?))?"
r"(?:\s+ORDER\s+BY\s+(?P<order>[\w`.]+\s*(?:asc|desc)?))?"
r"(?:\s+LIMIT\s+(?P<limit>\d+))?"
r"(?:\s+OFFSET\s+(?P<offset>\d+))?\s*$", text)
if not m:
raise AssertionError(f"FakeSor cannot parse SQL: {text}")
table = m.group("tbl")
rows = [dict(r) for r in self._rows(table)]
where = (m.group("where") or "").strip()
if where:
rows = [r for r in rows if _match_expr(r, where)]
order = (m.group("order") or "").strip()
if order:
parts = order.split()
col = parts[0].strip("`")
desc = len(parts) > 1 and parts[1].lower() == "desc"
rows.sort(key=lambda r: str(r.get(col, "")), reverse=desc)
if m.group("offset"):
rows = rows[int(m.group("offset")):]
if m.group("limit"):
rows = rows[:int(m.group("limit"))]
sel = m.group("sel").strip()
if re.match(r"(?i)^COUNT\(\s*1\s*\)", sel):
return [{"cnt": len(rows)}]
if sel != "*":
cols = [c.strip().strip("`").split(".")[-1] for c in _split_top(sel)]
rows = [{c: r.get(c) for c in cols} for r in rows]
return rows
# ------------------------------------------------------------------ 条件求值
def _match(row, where: dict) -> bool:
for key, value in (where or {}).items():
if key in ("__table__",):
continue
if str(row.get(key, "")) != str(value):
return False
return True
def _match_expr(row, expr: str) -> bool:
for part in _split_top(expr, sep="AND"):
part = part.strip()
if not part:
continue
if not _match_one(row, part):
return False
return True
def _match_one(row, cond: str) -> bool:
cond = cond.strip()
if re.match(r"^1\s*=\s*0$", cond):
return False
m = re.match(r"^`?(\w+)`?\s+IN\s*\((?P<vals>.*)\)$", cond, re.IGNORECASE)
if m:
vals = [_unquote(v) for v in _split_top(m.group("vals"))]
return str(row.get(m.group(1), "")) in [str(v) for v in vals]
m = re.match(r"^`?(\w+)`?\s*=\s*(?P<val>.+)$", cond)
if m:
want = _unquote(m.group("val"))
got = row.get(m.group(1))
if want == "NULL":
return got in (None, "")
return str("" if got is None else got) == str(want)
raise AssertionError(f"FakeSor cannot parse condition: {cond}")
def _unquote(token):
token = token.strip()
if len(token) >= 2 and token[0] == "'" and token[-1] == "'":
return token[1:-1].replace("''", "'")
if token.upper() == "NULL":
return None
try:
return int(token)
except ValueError:
return token
def _split_top(text: str, sep: str = ","):
"""按顶层分隔符切分,忽略括号与引号内的分隔符。"""
out, buf, depth, quote = [], [], 0, None
upper = text.upper()
sepu = sep.upper()
i = 0
while i < len(text):
ch = text[i]
if quote:
buf.append(ch)
if ch == quote:
quote = None
i += 1
continue
if ch in "([{":
depth += 1
elif ch in ")]}":
depth -= 1
elif ch in "'\"":
quote = ch
if depth == 0 and sep != "," and upper[i:i + len(sepu)] == sepu:
out.append("".join(buf))
buf = []
i += len(sepu)
continue
if depth == 0 and sep == "," and ch == ",":
out.append("".join(buf))
buf = []
i += 1
continue
buf.append(ch)
i += 1
if buf:
out.append("".join(buf))
return [s for s in (x.strip() for x in out) if s]
def _bind(sql: str, params) -> str:
"""把 ``%s`` 占位按序替换为字面量(仅测试解析用)。"""
out, idx = [], 0
i = 0
while i < len(sql):
if sql[i] == "%" and i + 1 < len(sql) and sql[i + 1] == "s":
if idx >= len(params):
raise AssertionError(f"FakeSor: not enough params for {sql}")
out.append(_lit(params[idx]))
idx += 1
i += 2
continue
out.append(sql[i])
i += 1
if idx != len(params):
raise AssertionError(
f"FakeSor: param count mismatch ({idx} used, {len(params)} given): {sql}")
return "".join(out)
def make_fake_sor(with_ref_table=True, base_fixtures=None):
"""构造 FakeSorwith_ref_table=False 时模拟关联表缺失场景。"""
tables = {DEFAULT_TABLE: []} if with_ref_table else {}
return FakeSor(tables=tables, base_fixtures=base_fixtures)