453 lines
20 KiB
Python
453 lines
20 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""tx_write —— M11b 核心:**单事务**写事件 + 实体状态 + 快照,**提交后**广播。
|
||
|
||
契约(与 init.py env 注册、__init__.py 导出三处同步,缺一即 NameError):
|
||
apply_runtime_event(...) 单事务:幂等检查 → INSERT pbl_runtime_event
|
||
→ UPSERT pbl_entity_state(乐观锁 state_version)
|
||
→ INSERT pbl_world_state_snapshot → COMMIT → 广播
|
||
poll_events(...) 3s 轮询兜底:先取进程内 hub 增量,未命中/有缺口退回查库
|
||
read_states(...) 读实体状态(tenant_id 强制打头)
|
||
assert_append_only(...) pbl_runtime_event 禁 UPDATE/DELETE + 基表写保护
|
||
|
||
关键保证(QC #2/#3 关注点):
|
||
* 原子性:三条 DML 全在同一个 ``rtx_db.transaction()`` 内;任一失败 → rollback,
|
||
事件行**不存在**(真回滚,不是「标 rolled_back」),且不广播。
|
||
* 无幽灵更新:广播只在 ``await tx.commit()`` 成功之后(with 块退出后)发出。
|
||
* 幂等:同 idem_key 命中已有事件 → dedup=True,不重复写、不重复广播。
|
||
* 乐观锁:state_version 服务端递增,客户端禁写;base_version 不匹配 → 冲突回滚。
|
||
* fail-closed:库名/事务原语/租户任一不可得 → 抛错,绝不降级为逐条提交。
|
||
"""
|
||
|
||
import hashlib
|
||
import inspect
|
||
import json
|
||
import os
|
||
import time
|
||
|
||
from . import broadcast as _hub
|
||
from . import latency as _lat
|
||
from . import rtx_db
|
||
from .m11b_config import (ERR_APPEND_ONLY, ERR_BAD_PARAM, ERR_BROADCAST_AFTER_COMMIT,
|
||
ERR_STATE_CONFLICT, ERR_TENANT_MISSING, EVENT_TABLE,
|
||
PROTECTED_BASE_TABLES, SNAPSHOT_TABLE, STATE_TABLE)
|
||
|
||
__all__ = ["apply_runtime_event", "poll_events", "read_states", "assert_append_only",
|
||
"idem_key_of", "table_columns", "self_check"]
|
||
|
||
_BASE_TABLES_LC = tuple(t.lower() for t in PROTECTED_BASE_TABLES)
|
||
|
||
|
||
# ---------------------------------------------------------------- 列名自省
|
||
_COL_CACHE = {}
|
||
|
||
|
||
def table_columns(table):
|
||
"""从 models/{table}.json 读真实列名(禁止猜列)。读不到返回 None。"""
|
||
if table in _COL_CACHE:
|
||
return _COL_CACHE[table]
|
||
path = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
|
||
"models", "%s.json" % table)
|
||
cols = None
|
||
try:
|
||
with open(path, "r", encoding="utf-8") as f:
|
||
spec = json.load(f)
|
||
names = []
|
||
for fd in spec.get("fields") or []:
|
||
if isinstance(fd, dict) and fd.get("name"):
|
||
names.append(fd["name"])
|
||
elif isinstance(fd, str):
|
||
names.append(fd.split()[0])
|
||
cols = names or None
|
||
except Exception: # noqa: BLE001
|
||
cols = None
|
||
_COL_CACHE[table] = cols
|
||
return cols
|
||
|
||
|
||
def build_row(table, wanted):
|
||
"""按表真实列过滤字段(避免 Unknown column);列名未知时原样通过。"""
|
||
cols = table_columns(table)
|
||
if not cols:
|
||
return dict(wanted)
|
||
return dict((k, v) for k, v in wanted.items() if k in cols)
|
||
|
||
|
||
def _dumps(value):
|
||
if value is None:
|
||
return None
|
||
if isinstance(value, str):
|
||
return value
|
||
try:
|
||
return json.dumps(value, ensure_ascii=False, sort_keys=True)
|
||
except Exception: # noqa: BLE001
|
||
return json.dumps(str(value), ensure_ascii=False)
|
||
|
||
|
||
def _loads(value, default=None):
|
||
if value is None or value == "":
|
||
return default
|
||
if isinstance(value, (dict, list)):
|
||
return value
|
||
try:
|
||
return json.loads(value)
|
||
except Exception: # noqa: BLE001
|
||
return default
|
||
|
||
|
||
def _norm(rows):
|
||
return rtx_db._norm_rows(rows)
|
||
|
||
|
||
# ---------------------------------------------------------------- 幂等键
|
||
def idem_key_of(tenant_id, world_id, session_id, event_type, payload=None,
|
||
client_key=None):
|
||
"""幂等键:客户端给的优先(原样透传,保证重放同键),否则服务端派生。"""
|
||
if client_key:
|
||
return str(client_key)[:128]
|
||
basis = "|".join([str(tenant_id), str(world_id), str(session_id), str(event_type),
|
||
_dumps(payload or {})])
|
||
return "auto-" + hashlib.sha256(basis.encode("utf-8")).hexdigest()[:48]
|
||
|
||
|
||
# ---------------------------------------------------------------- 校验
|
||
def _require_tenant(tenant_id):
|
||
if tenant_id in (None, "", 0):
|
||
try:
|
||
from pbl_common.api import tenant_id as _tid # type: ignore
|
||
|
||
tenant_id = _tid()
|
||
except Exception: # noqa: BLE001
|
||
tenant_id = None
|
||
if tenant_id in (None, "", 0):
|
||
raise rtx_db.RtxError(ERR_TENANT_MISSING, "tenant_id 缺失,fail-closed 拒绝读写")
|
||
return tenant_id
|
||
|
||
|
||
def assert_append_only(operation, table, tenant_id=None):
|
||
"""守卫:pbl_runtime_event 只允许 INSERT/SELECT;基表禁止任何写。违规抛 RtxError。"""
|
||
op = str(operation or "").strip().upper()
|
||
tbl = str(table or "").strip().lower()
|
||
if tbl in _BASE_TABLES_LC and op in ("INSERT", "UPDATE", "DELETE", "REPLACE",
|
||
"TRUNCATE"):
|
||
raise rtx_db.RtxError(ERR_APPEND_ONLY, "基表 %s 只读(本模块禁改基表)" % tbl)
|
||
if tbl == EVENT_TABLE.lower() and op in ("UPDATE", "DELETE", "REPLACE", "TRUNCATE",
|
||
"ALTER", "DROP"):
|
||
raise rtx_db.RtxError(ERR_APPEND_ONLY,
|
||
"pbl_runtime_event 为 append-only,禁止 %s" % op)
|
||
return True
|
||
|
||
|
||
# ---------------------------------------------------------------- 事务内步骤
|
||
async def _find_idem(tx, tenant_id, key):
|
||
sql = ("SELECT id, seq_no, event_type FROM %s WHERE tenant_id=%%s AND idem_key=%%s "
|
||
"LIMIT 1" % EVENT_TABLE)
|
||
rows = _norm(await tx.execute(sql, (tenant_id, key)))
|
||
return rows[0] if rows else None
|
||
|
||
|
||
async def _next_seq(tx, tenant_id, world_id, session_id):
|
||
sql = ("SELECT COALESCE(MAX(seq_no), 0) AS max_seq FROM %s "
|
||
"WHERE tenant_id=%%s AND world_id=%%s AND session_id=%%s" % EVENT_TABLE)
|
||
rows = _norm(await tx.execute(sql, (tenant_id, world_id, session_id)))
|
||
try:
|
||
return int(rows[0].get("max_seq") or 0) + 1 if rows else 1
|
||
except (TypeError, ValueError):
|
||
return 1
|
||
|
||
|
||
async def _apply_state(tx, tenant_id, world_id, session_id, upd, seq_no,
|
||
base_version, now):
|
||
"""单条状态 UPSERT + 乐观锁。返回 {state_key, state_version, action, entity_id}。"""
|
||
if not isinstance(upd, dict) or not upd.get("state_key"):
|
||
raise rtx_db.RtxError(ERR_BAD_PARAM, "state_updates 每项必须含 state_key")
|
||
assert_append_only("UPDATE", STATE_TABLE, tenant_id)
|
||
state_key = str(upd["state_key"])[:128]
|
||
entity_id = upd.get("entity_id") or upd.get("entity_type") or ""
|
||
rows = _norm(await tx.execute(
|
||
"SELECT id, state_version FROM %s WHERE tenant_id=%%s AND world_id=%%s "
|
||
"AND session_id=%%s AND state_key=%%s LIMIT 1" % STATE_TABLE,
|
||
(tenant_id, world_id, session_id, state_key)))
|
||
if rows:
|
||
cur_ver = int(rows[0].get("state_version") or 0)
|
||
if base_version is not None and int(base_version) != cur_ver:
|
||
raise rtx_db.RtxError(
|
||
ERR_STATE_CONFLICT,
|
||
"state_key=%s 乐观锁冲突:base=%s 当前=%s(整体回滚)"
|
||
% (state_key, base_version, cur_ver))
|
||
new_ver = cur_ver + 1
|
||
vals = build_row(STATE_TABLE, {
|
||
"state_value": _dumps(upd.get("state_value")),
|
||
"state_version": new_ver, "seq_no": seq_no, "updated_at": now,
|
||
})
|
||
await tx.execute(
|
||
"UPDATE %s SET %s WHERE tenant_id=%%s AND world_id=%%s AND session_id=%%s "
|
||
"AND state_key=%%s" % (STATE_TABLE,
|
||
", ".join("%s=%%s" % c for c in vals.keys())),
|
||
tuple(list(vals.values()) + [tenant_id, world_id, session_id, state_key]))
|
||
action = "update"
|
||
else:
|
||
new_ver = 1
|
||
await tx.insert(STATE_TABLE, build_row(STATE_TABLE, {
|
||
"tenant_id": tenant_id, "world_id": world_id, "session_id": session_id,
|
||
"state_key": state_key, "entity_id": entity_id,
|
||
"state_value": _dumps(upd.get("state_value")),
|
||
"state_version": new_ver, "seq_no": seq_no,
|
||
"created_at": now, "updated_at": now,
|
||
}))
|
||
action = "insert"
|
||
return {"state_key": state_key, "state_version": new_ver, "action": action,
|
||
"entity_id": entity_id}
|
||
|
||
|
||
# ---------------------------------------------------------------- 主写入
|
||
async def apply_runtime_event(tenant_id=None, world_id=None, session_id=None,
|
||
event_type=None, payload=None, state_updates=None,
|
||
idem_key=None, client_key=None, snapshot=None,
|
||
base_version=None, env=None, occurred_at=None):
|
||
"""单事务写「事件 + 实体状态 + 快照」,提交后广播。
|
||
|
||
返回 {ok, event_id, seq_no, cursor, dedup, state_rows, latency_ms, channel}
|
||
抛 RtxError:PBL-TENANT-0001 / PBL-PARAM-0001 / PBL-STATE-CONFLICT / PBL-RTX-0002
|
||
"""
|
||
tenant_id = _require_tenant(tenant_id)
|
||
if world_id in (None, "") or session_id in (None, "") or not event_type:
|
||
raise rtx_db.RtxError(ERR_BAD_PARAM, "world_id/session_id/event_type 必填")
|
||
state_updates = state_updates or []
|
||
if not isinstance(state_updates, (list, tuple)):
|
||
raise rtx_db.RtxError(ERR_BAD_PARAM, "state_updates 必须是数组")
|
||
|
||
key = idem_key_of(tenant_id, world_id, session_id, event_type, payload,
|
||
client_key or idem_key)
|
||
hub = _hub.get_hub()
|
||
channel = _hub.channel_of(world_id, tenant_id)
|
||
now = occurred_at or time.strftime("%Y-%m-%d %H:%M:%S")
|
||
|
||
event_id = None
|
||
seq_no = None
|
||
state_events = []
|
||
dedup_result = None
|
||
stmt_count = 0
|
||
t0 = time.perf_counter() # 裁决计时起点(含事务全程)
|
||
async with rtx_db.transaction() as tx:
|
||
try:
|
||
# 1) 幂等:同键已存在 → 不重复写、不重复广播
|
||
dup = await _find_idem(tx, tenant_id, key)
|
||
if dup:
|
||
await tx.commit()
|
||
dedup_result = {
|
||
"ok": True, "dedup": True, "event_id": dup.get("id"),
|
||
"seq_no": dup.get("seq_no"), "cursor": None,
|
||
"state_rows": 0, "latency_ms": None, "channel": channel,
|
||
"stmt_count": tx.stmt_count,
|
||
}
|
||
else:
|
||
# 2) 会话内单调 seq_no
|
||
seq_no = await _next_seq(tx, tenant_id, world_id, session_id)
|
||
|
||
# 3) INSERT 事件(append-only)
|
||
assert_append_only("INSERT", EVENT_TABLE, tenant_id)
|
||
event_id = await tx.insert(EVENT_TABLE, build_row(EVENT_TABLE, {
|
||
"tenant_id": tenant_id, "world_id": world_id,
|
||
"session_id": session_id, "seq_no": seq_no,
|
||
"event_type": str(event_type)[:64],
|
||
"payload": _dumps(payload or {}), "idem_key": key,
|
||
"occurred_at": now, "created_at": now, "state": "applied",
|
||
}))
|
||
|
||
# 4) UPSERT 实体状态(乐观锁:state_version 服务端递增,客户端禁写)
|
||
for upd in state_updates:
|
||
state_events.append(await _apply_state(tx, tenant_id, world_id,
|
||
session_id, upd, seq_no,
|
||
base_version, now))
|
||
|
||
# 5) 快照(掉线补齐/离线兜底)
|
||
if snapshot is not None:
|
||
await tx.insert(SNAPSHOT_TABLE, build_row(SNAPSHOT_TABLE, {
|
||
"tenant_id": tenant_id, "world_id": world_id,
|
||
"session_id": session_id, "seq_no": seq_no,
|
||
"state": _dumps(snapshot),
|
||
"state_digest": _digest(snapshot), "created_at": now,
|
||
}))
|
||
|
||
# 6) 提交(任一 DML 失败 → 下面 rollback,事件行不存在)
|
||
await tx.commit()
|
||
stmt_count = tx.stmt_count
|
||
except Exception:
|
||
await tx.rollback()
|
||
raise
|
||
# ---- 事务已提交(with 块正常退出):此后才允许广播,杜绝幽灵更新 ----
|
||
ms = round((time.perf_counter() - t0) * 1000.0, 3)
|
||
if dedup_result is not None:
|
||
dedup_result["latency_ms"] = ms
|
||
_lat.record("adjudicate", ms)
|
||
return dedup_result
|
||
breach = _lat.record("adjudicate", ms)
|
||
evt = {
|
||
"tenant_id": tenant_id, "world_id": world_id, "session_id": session_id,
|
||
"event_id": event_id, "seq_no": seq_no, "event_type": str(event_type),
|
||
"payload": payload or {}, "states": state_events, "idem_key": key,
|
||
"occurred_at": now, "source": "apply_runtime_event",
|
||
}
|
||
try:
|
||
b_ms = hub.publish(channel, evt)
|
||
_lat.record("broadcast", b_ms)
|
||
_lat.record("end_to_end", ms + b_ms)
|
||
except Exception as exc: # noqa: BLE001
|
||
# 数据已提交,广播失败不回滚事实;客户端由 3s 轮询兜底补齐
|
||
raise rtx_db.RtxError(ERR_BROADCAST_AFTER_COMMIT,
|
||
"数据已提交但广播失败(poll 兜底):%r" % exc)
|
||
return {
|
||
"ok": True, "dedup": False, "event_id": event_id, "seq_no": seq_no,
|
||
"cursor": evt.get("cursor"), "state_rows": len(state_events),
|
||
"latency_ms": ms, "channel": channel, "sla_breach": bool(breach),
|
||
"stmt_count": stmt_count,
|
||
}
|
||
|
||
|
||
def _digest(obj):
|
||
return hashlib.sha1(_dumps(obj).encode("utf-8")).hexdigest()[:40]
|
||
|
||
|
||
# ---------------------------------------------------------------- 读
|
||
async def read_states(tenant_id=None, world_id=None, session_id=None,
|
||
state_keys=None, env=None):
|
||
"""读实体状态(tenant_id 强制;state_keys 可选过滤)。返回 list[dict]。"""
|
||
tenant_id = _require_tenant(tenant_id)
|
||
sql = ("SELECT state_key, entity_id, state_value, state_version, seq_no, updated_at "
|
||
"FROM %s WHERE tenant_id=%%s" % STATE_TABLE)
|
||
params = [tenant_id]
|
||
if world_id not in (None, ""):
|
||
sql += " AND world_id=%s"
|
||
params.append(world_id)
|
||
if session_id not in (None, ""):
|
||
sql += " AND session_id=%s"
|
||
params.append(session_id)
|
||
if state_keys:
|
||
sql += " AND state_key IN (%s)" % ",".join(["%s"] * len(state_keys))
|
||
params.extend([str(k) for k in state_keys])
|
||
sql += " ORDER BY state_key ASC LIMIT 500"
|
||
rows = await rtx_db.q_all(sql, tuple(params))
|
||
out = []
|
||
for r in rows:
|
||
r = dict(r)
|
||
r["state_value"] = _loads(r.get("state_value"), {})
|
||
out.append(r)
|
||
return out
|
||
|
||
|
||
async def poll_events(tenant_id=None, world_id=None, session_id=None,
|
||
after_cursor=None, after_seq=None, limit=200, env=None):
|
||
"""轮询兜底:先取 hub 增量;未命中/有缺口/带 after_seq 时退回查库(权威源)。
|
||
|
||
返回 {ok, events, next_cursor, next_seq, source: hub|db, truncated}
|
||
"""
|
||
tenant_id = _require_tenant(tenant_id)
|
||
hub = _hub.get_hub()
|
||
channel = _hub.channel_of(world_id, tenant_id)
|
||
limit = max(1, min(int(limit or 200), 500))
|
||
events, truncated = hub.pull(channel, after_cursor, limit)
|
||
if events and not truncated and after_seq in (None, ""):
|
||
return {"ok": True, "events": events, "next_cursor": events[-1]["cursor"],
|
||
"next_seq": None, "source": "hub", "truncated": False}
|
||
|
||
sql = ("SELECT id, seq_no, event_type, payload, occurred_at, created_at FROM %s "
|
||
"WHERE tenant_id=%%s" % EVENT_TABLE)
|
||
params = [tenant_id]
|
||
if world_id not in (None, ""):
|
||
sql += " AND world_id=%s"
|
||
params.append(world_id)
|
||
if session_id not in (None, ""):
|
||
sql += " AND session_id=%s"
|
||
params.append(session_id)
|
||
if after_seq not in (None, ""):
|
||
sql += " AND seq_no > %s"
|
||
params.append(int(after_seq))
|
||
elif truncated and events:
|
||
sql += " AND seq_no >= %s"
|
||
params.append(int(events[0].get("seq_no") or 0))
|
||
sql += " ORDER BY seq_no ASC LIMIT %s" % limit
|
||
rows = await rtx_db.q_all(sql, tuple(params))
|
||
out = []
|
||
for r in rows:
|
||
r = dict(r)
|
||
r["payload"] = _loads(r.get("payload"), {})
|
||
r["channel"] = channel
|
||
out.append(r)
|
||
next_seq = out[-1]["seq_no"] if out else (after_seq if after_seq is not None else 0)
|
||
return {"ok": True, "events": out, "next_cursor": hub.latest_cursor(),
|
||
"next_seq": next_seq, "source": "db", "truncated": bool(truncated)}
|
||
|
||
|
||
# ---------------------------------------------------------------- 自检
|
||
def self_check():
|
||
"""离线自检:幂等键派生、append-only 守卫、列名自省、契约协程签名。"""
|
||
msgs, ok = [], True
|
||
|
||
k1 = idem_key_of(7, 1, 11, "move", {"x": 1})
|
||
k2 = idem_key_of(7, 1, 11, "move", {"x": 1})
|
||
k3 = idem_key_of(7, 1, 11, "move", {"x": 2})
|
||
if k1 != k2 or k1 == k3:
|
||
ok = False
|
||
msgs.append("tx_write FAIL:幂等键不稳定/不敏感 %s %s %s" % (k1, k2, k3))
|
||
else:
|
||
msgs.append("tx_write PASS:幂等键同参同键、异参异键(%s…)" % k1[:12])
|
||
if idem_key_of(7, 1, 11, "move", {}, client_key="cli-9") != "cli-9":
|
||
ok = False
|
||
msgs.append("tx_write FAIL:client_key 未原样透传")
|
||
else:
|
||
msgs.append("tx_write PASS:client_key 原样透传(重放同键)")
|
||
|
||
blocked = []
|
||
for op in ("UPDATE", "DELETE", "TRUNCATE"):
|
||
try:
|
||
assert_append_only(op, EVENT_TABLE)
|
||
ok = False
|
||
msgs.append("tx_write FAIL:append-only 未拦 %s" % op)
|
||
except rtx_db.RtxError as exc:
|
||
if exc.code != ERR_APPEND_ONLY:
|
||
ok = False
|
||
msgs.append("tx_write FAIL:错误码 %s" % exc.code)
|
||
else:
|
||
blocked.append(op)
|
||
if blocked:
|
||
msgs.append("tx_write PASS:pbl_runtime_event 禁 %s" % "/".join(blocked))
|
||
try:
|
||
assert_append_only("INSERT", EVENT_TABLE)
|
||
assert_append_only("SELECT", EVENT_TABLE)
|
||
msgs.append("tx_write PASS:INSERT/SELECT 放行")
|
||
except rtx_db.RtxError as exc:
|
||
ok = False
|
||
msgs.append("tx_write FAIL:合法操作被拦 %s" % exc.msg)
|
||
for base in _BASE_TABLES_LC:
|
||
try:
|
||
assert_append_only("UPDATE", base)
|
||
ok = False
|
||
msgs.append("tx_write FAIL:基表 %s 未保护" % base)
|
||
except rtx_db.RtxError:
|
||
pass
|
||
msgs.append("tx_write PASS:基表写保护域 %d 张全覆盖" % len(_BASE_TABLES_LC))
|
||
|
||
cols = table_columns(EVENT_TABLE)
|
||
if cols:
|
||
row = build_row(EVENT_TABLE, {"tenant_id": 7, "world_id": 1, "seq_no": 1,
|
||
"not_a_column": "x"})
|
||
if "not_a_column" in row:
|
||
ok = False
|
||
msgs.append("tx_write FAIL:未知列未被过滤")
|
||
else:
|
||
msgs.append("tx_write PASS:列名自省生效(models/%s.json %d 列)"
|
||
% (EVENT_TABLE, len(cols)))
|
||
else:
|
||
msgs.append("tx_write WARN:离线未读到 models/%s.json,build_row 原样透传"
|
||
% EVENT_TABLE)
|
||
|
||
for fn, name in ((apply_runtime_event, "apply_runtime_event"),
|
||
(poll_events, "poll_events"), (read_states, "read_states")):
|
||
if not inspect.iscoroutinefunction(fn):
|
||
ok = False
|
||
msgs.append("tx_write FAIL:%s 非协程(dspy await 会拿到协程对象)" % name)
|
||
if ok:
|
||
msgs.append("tx_write PASS:三个写/读契约均为 async 协程")
|
||
msgs.append("tx_write.self_check PASS")
|
||
return ok, msgs
|