pbl_blueprint/tests/fakedb.py
2026-09-16 12:49:06 +08:00

134 lines
4.4 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 / ServerEnv 替身,供 test_contract.py 做契约单测。
只实现本模块用到的 6 个 APIC / U / D / R / I / sqlExe语义与 sqlor 对齐:
where 支持 'col' 等值、'col!' 不等、'col~' LIKE、'col in' 列表)。
禁止在此文件写业务逻辑——它只模拟 DB 行为。
"""
import fnmatch
import os
import sys
REPO_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if REPO_ROOT not in sys.path:
sys.path.insert(0, REPO_ROOT)
def _match(row, where):
for key, want in (where or {}).items():
if key.endswith('!'):
col = key[:-1]
if str(row.get(col, '')) == str(want):
return False
continue
if key.endswith('~'):
col = key[:-1]
if not fnmatch.fnmatch(str(row.get(col, '')), str(want).replace('%', '*')):
return False
continue
if key.endswith(' in') and isinstance(want, (list, tuple)):
if row.get(key[:-4].strip()) not in want:
return False
continue
if str(row.get(key, '')) != str(want):
return False
return True
def _order_key(rows, orderby):
if not orderby:
return rows
parts = [p.strip() for p in orderby.split(',')]
result = list(rows)
for part in reversed(parts):
tok = part.split()
col = tok[0]
desc = len(tok) > 1 and tok[1].lower() == 'desc'
result.sort(key=lambda r: ('' if r.get(col) is None else r.get(col)), reverse=desc)
return result
class FakeSqlor(object):
def __init__(self, store):
self.store = store
self.calls = []
def C(self, table, data, dbname=None, **kw):
self.calls.append(('C', table, dict(data)))
self.store.setdefault(table, {})[data['id']] = dict(data)
return 1
def I(self, table, rows, dbname=None, **kw):
n = 0
for r in rows or []:
self.calls.append(('I', table, dict(r)))
self.store.setdefault(table, {})[r['id']] = dict(r)
n += 1
return n
def U(self, table, data, where, dbname=None, **kw):
self.calls.append(('U', table, dict(data), dict(where)))
n = 0
for rid, row in list(self.store.get(table, {}).items()):
if _match(row, where):
row.update(data)
n += 1
return n
def D(self, table, where, dbname=None, **kw):
self.calls.append(('D', table, dict(where)))
drop = [rid for rid, row in list(self.store.get(table, {}).items())
if _match(row, where)]
for rid in drop:
self.store[table].pop(rid, None)
return len(drop)
def R(self, table, where=None, fields='*', orderby=None, limit=None,
dbname=None, **kw):
self.calls.append(('R', table, dict(where or {})))
rows = [dict(r) for r in self.store.get(table, {}).values() if _match(r, where)]
rows = _order_key(rows, orderby)
if limit:
rows = rows[:int(limit)]
if fields and fields != '*':
keep = [f.strip() for f in fields.split(',')]
rows = [dict((k, r.get(k)) for k in keep if k in r) for r in rows]
return rows
def sqlExe(self, sql, params=None, dbname=None, **kw):
self.calls.append(('sqlExe', sql[:60]))
return []
class FakeServerEnv(object):
def __init__(self, dbname=None):
self._dbname = {'pbl_blueprint': dbname}
self.pbl_tenant_ctx = {'tenant_id': 'T1'}
self.pbl_current_user = 'tester'
def set_module_dbname(self, mod, dbname):
self._dbname[mod] = dbname
def get_module_dbname(self, mod):
return self._dbname.get(mod)
def install(store=None, tenant_id='T1', dbname=None):
"""把 FakeSqlor 注入 sys.modules['sqlor'],返回 (sor, env, store)。"""
import types
store = store if store is not None else {}
sor = FakeSqlor(store)
mod = types.ModuleType('sqlor')
mod.C, mod.U, mod.D, mod.R, mod.I, mod.sqlExe = (
sor.C, sor.U, sor.D, sor.R, sor.I, sor.sqlExe)
sys.modules['sqlor'] = mod
ah = types.ModuleType('ahserver')
env = FakeServerEnv(dbname)
env.pbl_tenant_ctx = {'tenant_id': tenant_id}
ah.ServerEnv = lambda *a, **kw: env
sys.modules.setdefault('ahserver', ah)
if 'ahserver' in sys.modules:
sys.modules['ahserver'].ServerEnv = lambda *a, **kw: env
return sor, env, store