pbl_domain_ext/tests/fake_db.py

273 lines
9.7 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。
要点:
* ``sqlExe`` 先把 ``%s`` 占位按序插值为字面量,再用极简解析器执行
SELECT / SELECT COUNT(1) / LIMIT / OFFSET / IN / AND 等值条件;
* world / scene / entity 三张基础域基表**只读**:任何 C/U/D/I 写入都会
记录到 ``base_write_attempts`` 并抛 AssertionError用于验证「不改基表」铁律。
"""
import re
__all__ = ["FakeSor", "make_fake_sor", "BASE_FIXTURES"]
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": "active"},
{"id": "S2", "tenant_id": "T1", "world_id": "W1", "code": "SCN-002",
"name": "实验舱", "status": "active"},
{"id": "S3", "tenant_id": "T1", "world_id": "W2", "code": "SCN-003",
"name": "深水区", "status": "active"},
],
"entity": [
{"id": "E1", "tenant_id": "T1", "scene_id": "S1", "code": "ENT-001",
"name": "探测车", "status": "idle"},
{"id": "E2", "tenant_id": "T1", "scene_id": "S1", "code": "ENT-002",
"name": "宇航员A", "status": "idle"},
{"id": "E3", "tenant_id": "T1", "scene_id": "S2", "code": "ENT-003",
"name": "培养皿", "status": "idle"},
],
}
_READONLY_TABLES = ("world", "scene", "entity")
def _literal(value):
"""把绑定参数转成 SQL 字面量(仅测试用)。"""
if value is None:
return "NULL"
if isinstance(value, bool):
return "1" if value else "0"
if isinstance(value, (int, float)):
return str(value)
return "'%s'" % str(value).replace("\\", "\\\\").replace("'", "''")
def _coerce(token):
"""把字面量 token 还原成 Python 值。"""
token = token.strip()
if token.upper() == "NULL":
return None
if len(token) >= 2 and token[0] == "'" and token[-1] == "'":
return token[1:-1].replace("''", "'").replace("\\\\", "\\")
try:
return int(token)
except ValueError:
pass
try:
return float(token)
except ValueError:
return token
def _loose_eq(left, right):
"""宽松相等:兼容 '0' vs 0、str vs int 的列值比较。"""
if left == right:
return True
if left is None or right is None:
return False
try:
return float(left) == float(right)
except (TypeError, ValueError):
return str(left) == str(right)
class FakeSor(object):
"""内存版 sqlor。"""
def __init__(self, tables=None, base_fixtures=None):
self.tables = {}
for name, rows in (base_fixtures if base_fixtures is not None
else BASE_FIXTURES).items():
self.tables[name] = [dict(r) for r in rows]
for name, rows in (tables or {}).items():
self.tables[name] = [dict(r) for r in rows]
self.tables.setdefault("pbl_domain_ref", [])
self.sql_log = []
self.base_write_attempts = []
# ---- sqlor 标准 API -------------------------------------------------
def C(self, tbl, row):
self._guard_readonly(tbl, "C")
self.tables.setdefault(tbl, []).append(dict(row))
self.sql_log.append(("C", tbl, dict(row)))
return 1
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
def D(self, tbl, where):
self._guard_readonly(tbl, "D")
keep, hits = [], 0
for exist in self.tables.get(tbl, []):
if self._match(exist, where):
hits += 1
else:
keep.append(exist)
self.tables[tbl] = keep
self.sql_log.append(("D", tbl, dict(where)))
return hits
def R(self, tbl, where=None, fields=None, order_by=None, limit=None):
rows = [dict(r) for r in self.tables.get(tbl, [])
if self._match(r, where or {})]
if fields:
rows = [{k: r.get(k) for k in fields} for r in rows]
if limit:
rows = rows[:int(limit)]
return rows
def I(self, tbl, rows):
self._guard_readonly(tbl, "I")
for row in rows or []:
self.tables.setdefault(tbl, []).append(dict(row))
return len(rows or [])
def sqlExe(self, sql, args=None):
args = tuple(args or ())
self.sql_log.append(("sqlExe", " ".join(sql.split()), args))
text = self._interpolate(" ".join(sql.split()), args)
upper = text.upper()
if upper.startswith("SELECT COUNT(1)"):
return [{"cnt": len(self._rows(text))}]
if upper.startswith("SELECT"):
return self._select(text)
if upper.startswith("ALTER") or upper.startswith("CREATE"):
return 0
raise AssertionError("FakeSor.sqlExe 不支持的语句: %s" % text)
# ---- 内部 -----------------------------------------------------------
def _guard_readonly(self, tbl, op):
if tbl in _READONLY_TABLES:
self.base_write_attempts.append((op, tbl))
raise AssertionError("铁律违规:禁止写基础域基表 %sop=%s" % (tbl, op))
@staticmethod
def _interpolate(text, args):
out, idx, buf = [], 0, []
i = 0
while i < len(text):
if text[i] == "%" and text[i:i + 2] == "%s":
buf.append(_literal(args[idx]) if idx < len(args) else "NULL")
idx += 1
i += 2
continue
buf.append(text[i])
i += 1
out.append("".join(buf))
return "".join(out)
@staticmethod
def _match(row, where):
for key, val in (where or {}).items():
col = key.strip("`")
if isinstance(val, (list, tuple, set)):
if not any(_loose_eq(row.get(col), v) for v in val):
return False
elif not _loose_eq(row.get(col), val):
return False
return True
@staticmethod
def _table_of(text):
m = re.search(r"FROM\s+`?(\w+)`?", text, re.IGNORECASE)
return m.group(1) if m else ""
def _conds(self, text):
m = re.search(r"WHERE\s+(.*?)(?:\s+ORDER\s+BY|\s+LIMIT|\s+GROUP\s+BY|\s*$)",
text, re.IGNORECASE | re.DOTALL)
if not m:
return []
conds = []
for part in re.split(r"\s+AND\s+", m.group(1), flags=re.IGNORECASE):
part = part.strip()
if not part:
continue
in_m = re.match(r"`?(\w+)`?\s+IN\s*\((.*)\)\s*$", part, re.IGNORECASE)
if in_m:
vals = [_coerce(v) for v in self._split_top(in_m.group(2))]
conds.append((in_m.group(1), vals))
continue
eq_m = re.match(r"`?(\w+)`?\s*=\s*(.+)$", part)
if eq_m:
conds.append((eq_m.group(1), _coerce(eq_m.group(2))))
return conds
@staticmethod
def _split_top(inner):
parts, depth, cur, quote = [], 0, [], False
for ch in inner:
if ch == "'":
quote = not quote
if not quote:
if ch == "(":
depth += 1
elif ch == ")":
depth -= 1
elif ch == "," and depth == 0:
parts.append("".join(cur))
cur = []
continue
cur.append(ch)
if cur:
parts.append("".join(cur))
return [p for p in parts if p.strip() != ""]
def _rows(self, text):
tbl = self._table_of(text)
if not tbl:
return []
conds = self._conds(text)
out = []
for row in self.tables.get(tbl, []):
hit = True
for col, val in conds:
if isinstance(val, list):
if not any(_loose_eq(row.get(col), v) for v in val):
hit = False
break
elif not _loose_eq(row.get(col), val):
hit = False
break
if hit:
out.append(dict(row))
return out
def _select(self, text):
rows = self._rows(text)
limit_m = re.search(r"LIMIT\s+(\d+)", text, re.IGNORECASE)
offset_m = re.search(r"OFFSET\s+(\d+)", text, re.IGNORECASE)
if offset_m:
rows = rows[int(offset_m.group(1)):]
if limit_m:
rows = rows[:int(limit_m.group(1))]
sel_m = re.match(r"SELECT\s+(.*?)\s+FROM\s", text, re.IGNORECASE | re.DOTALL)
if sel_m and sel_m.group(1).strip() != "*":
cols = [c.strip().strip("`").split(".")[-1]
for c in self._split_top(sel_m.group(1))]
rows = [{c: r.get(c) for c in cols} for r in rows]
return rows
def make_fake_sor(with_ref_table=True, base_fixtures=None):
"""构造 FakeSorwith_ref_table=False 时模拟关联表缺失场景。"""
tables = {"pbl_domain_ref": []} if with_ref_table else {}
return FakeSor(tables=tables, base_fixtures=base_fixtures)