2026-09-19 13:07:15 +08:00

453 lines
20 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 -*-
"""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}
抛 RtxErrorPBL-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 FAILclient_key 未原样透传")
else:
msgs.append("tx_write PASSclient_key 原样透传(重放同键)")
blocked = []
for op in ("UPDATE", "DELETE", "TRUNCATE"):
try:
assert_append_only(op, EVENT_TABLE)
ok = False
msgs.append("tx_write FAILappend-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 PASSpbl_runtime_event 禁 %s" % "/".join(blocked))
try:
assert_append_only("INSERT", EVENT_TABLE)
assert_append_only("SELECT", EVENT_TABLE)
msgs.append("tx_write PASSINSERT/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.jsonbuild_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