# -*- coding: utf-8 -*- """ 批次6 单测: ws 直连通道纯逻辑 (零外部依赖, 不连库) ==================================================== 运行: python scripts/test_batch6_units.py 覆盖: A. 协议 §2.1.1 **测试向量** —— 逐字节核对 payload_json / sha256 / canonical / sig。 这一组是本批最重要的断言: 规范化串的拼接是双方最容易写不一致的地方, 对不上就 一定卡在联调的"签名验不过", 而且看不出差在哪一段。 B. 签名/验签往返、篡改检出、跨密钥拒绝。 C. 信封解析: 版本、缺字段、非法 JSON、签名不匹配。 D. place_order 参数校验: 整百规则、限价必填与精度、点式代码、方向。 E. seq 水位推进: 连续、乱序、重复、缺口。 F. 双层去重键与成交金额自洽性。 G. 单表访问守卫: upsert (ON DUPLICATE KEY UPDATE) 不得被误判成多表。 H. dispatcher 纯助手: valid_until 归一、子单→父指令反推。 """ import os import sys import traceback sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) PASS, FAIL = [], [] def case(name): def deco(fn): try: fn() PASS.append(name) print(f" ok {name}") except Exception as e: FAIL.append((name, f"{type(e).__name__}: {e}", traceback.format_exc())) print(f" FAIL {name} —— {type(e).__name__}: {e}") return fn return deco def eq(got, exp, what=""): if got != exp: raise AssertionError(f"{what}\n 实际: {got!r}\n 期望: {exp!r}") # ================================================================ A. 测试向量 # 协议 §2.1.1 原文。测试密钥仅供联调, 不得用于生产。 VEC_SEED = "00112233445566778899aabbccddeeff00112233445566778899aabbccddeeff" VEC_PUB = "PM0kHP/Js2GARLl9A22GFFk9iwF8NA8d7odzOFUXZUs=" VEC_PAYLOAD = {"instruction_id": "INS-20260728-a3f19c04", "ts_code": "600000.SH", "side": "sell", "qty": 2000, "limit_price": 12.35, "valid_until": 1769512500000, "intent": "TRIM", "note": "保垫减仓·垫厚先收"} VEC_MSG_ID, VEC_TS, VEC_NONCE = "m-20260728-000123", 1769500000123, "9f2c8a1b7d3e4056" VEC_PJ = ('{"instruction_id":"INS-20260728-a3f19c04","intent":"TRIM","limit_price":12.35,' '"note":"保垫减仓·垫厚先收","qty":2000,"side":"sell","ts_code":"600000.SH",' '"valid_until":1769512500000}') VEC_SHA = "9fa3f56b92ef9d79290b46416899bdbebd93e7fdddbfe5bee5b5e09380bb7c63" VEC_CANON = f"1\nplace_order\n{VEC_MSG_ID}\n{VEC_TS}\n{VEC_NONCE}\n{VEC_SHA}" VEC_SIG = ("NzCILYw7ampmMtg41EBMubtlnDwr/jjGO+1cMTIxVdqaSaT779Kop5DAdADHL2aSbqrw2sp8" "hCQb+apCACIYCw==") def run(): from app.core import ws_codec as wsc print("\n[A] 协议 §2.1.1 测试向量 (逐字节)") @case("payload_json: 键序/无空格/中文不转义") def _(): eq(wsc.payload_json(VEC_PAYLOAD), VEC_PJ, "紧凑 JSON 与协议向量不一致") @case("payload sha256") def _(): eq(wsc.payload_sha256(VEC_PAYLOAD), VEC_SHA) @case("canonical 六段拼接") def _(): got = wsc.canonical(v=1, type_="place_order", msg_id=VEC_MSG_ID, ts=VEC_TS, nonce=VEC_NONCE, payload_hash=VEC_SHA) eq(got, VEC_CANON) eq(got.count("\n"), 5, "必须是 5 个换行 (6 段), 串尾不能有换行") @case("Ed25519 签名值") def _(): eq(wsc.sign(VEC_CANON, VEC_SEED), VEC_SIG) @case("由 seed 推出的公钥") def _(): eq(wsc.public_key_b64(VEC_SEED), VEC_PUB) @case("limit_price 序列化为 12.35 而非二进制尾数") def _(): eq(wsc.payload_json({"p": wsc.q2(12.35)}), '{"p":12.35}') eq(wsc.payload_json({"p": wsc.q2(12.34 * 1.0008)}), '{"p":12.35}') print("\n[B] 签名往返与篡改检出") @case("自签自验通过") def _(): assert wsc.verify(VEC_CANON, VEC_SIG, VEC_PUB) @case("payload 改一个字 → 验签失败") def _(): bad = dict(VEC_PAYLOAD, qty=2100) canon = wsc.canonical(v=1, type_="place_order", msg_id=VEC_MSG_ID, ts=VEC_TS, nonce=VEC_NONCE, payload_hash=wsc.payload_sha256(bad)) assert not wsc.verify(canon, VEC_SIG, VEC_PUB), "改了数量还能验过就等于没签名" @case("换一把密钥签 → 验签失败") def _(): other = "11" * 32 sig = wsc.sign(VEC_CANON, other) assert not wsc.verify(VEC_CANON, sig, VEC_PUB) @case("垃圾签名/垃圾公钥 → False 而不是抛异常") def _(): assert not wsc.verify(VEC_CANON, "not-base64!!", VEC_PUB) assert not wsc.verify(VEC_CANON, VEC_SIG, "zzz") @case("非法 seed 报 SIG_INVALID") def _(): try: wsc.sign(VEC_CANON, "abcd") raise AssertionError("短 seed 应当报错") except wsc.CodecError as e: eq(e.code, "SIG_INVALID") print("\n[C] 信封组包与解析") @case("build → parse 往返一致") def _(): env = wsc.build("ack", {"instruction_id": "INS-1", "accepted": True}, seed_hex=VEC_SEED, msg_id="m-1", ts=VEC_TS, nonce=VEC_NONCE, corr_id="INS-1", seq=42) back = wsc.parse(wsc.dumps(env), peer_pubkey_b64=VEC_PUB) eq(back["type"], "ack") eq(back["seq"], 42) eq(back["corr_id"], "INS-1") eq(back["payload"]["accepted"], True) @case("下行不带 seq (协议 §3: seq 仅上行必填)") def _(): env = wsc.build("ping", {}, seed_hex=VEC_SEED, msg_id="m-2") assert "seq" not in env, "下行擅自编号会给对端补发逻辑增加歧义" @case("协议版本不认 → VERSION") def _(): env = wsc.build("ping", {}, seed_hex=VEC_SEED, msg_id="m-3") env["v"] = 2 try: wsc.parse(wsc.dumps(env), peer_pubkey_b64=VEC_PUB) raise AssertionError("应当拒绝未知版本") except wsc.CodecError as e: eq(e.code, "VERSION") @case("缺信封字段 → BAD_PARAM") def _(): env = wsc.build("ping", {}, seed_hex=VEC_SEED, msg_id="m-4") env.pop("nonce") try: wsc.parse(wsc.dumps(env), peer_pubkey_b64=VEC_PUB) raise AssertionError("应当拒绝缺字段") except wsc.CodecError as e: eq(e.code, "BAD_PARAM") @case("非法 JSON → BAD_PARAM") def _(): try: wsc.parse("{不是 json", peer_pubkey_b64=VEC_PUB) raise AssertionError("应当拒绝") except wsc.CodecError as e: eq(e.code, "BAD_PARAM") @case("未配对端公钥 → SIG_INVALID (不能当成验过)") def _(): env = wsc.build("pong", {}, seed_hex=VEC_SEED, msg_id="m-5") try: wsc.parse(wsc.dumps(env), peer_pubkey_b64="") raise AssertionError("没有公钥就必须拒绝, 不能放行") except wsc.CodecError as e: eq(e.code, "SIG_INVALID") @case("补发消息 ts 很旧仍可解析 (PMS 不做时间窗校验)") def _(): # 协议 §6.1 规定补发「同 seq、同 msg_id、同内容、同签名」—— ts 还是原来那条的。 # 若照抄 §2.2 的 30 秒时间窗, 断网十分钟后补发的消息会被整批丢掉。 env = wsc.build("trade", {"trade_no": "QMT-1#1"}, seed_hex=VEC_SEED, msg_id="m-old", ts=1_600_000_000_000, seq=7) back = wsc.parse(wsc.dumps(env), peer_pubkey_b64=VEC_PUB) eq(back["seq"], 7) print("\n[D] place_order 参数校验 (本地拦下, 不换对端一个 BAD_PARAM)") def _po(**kw): base = dict(instruction_id="INS-20260728-aabbccdd", ts_code="600000.SH", side="buy", qty=200, limit_price=12.345, valid_until=1769512500000, intent="OPEN", note="x") base.update(kw) return wsc.place_order_payload(**base) def _reject(what, **kw): try: _po(**kw) raise AssertionError(f"{what} 应当被拒") except wsc.CodecError as e: eq(e.code, "BAD_PARAM", what) @case("限价四舍五入到 2 位") def _(): eq(_po()["limit_price"], 12.35) eq(_po(limit_price=2.675)["limit_price"], 2.68) # 不能被银行家舍入吃掉 @case("买入必须整百") def _(): _reject("买入 150 股", qty=150) eq(_po(qty=300)["qty"], 300) @case("卖出允许零股尾数 (清仓)") def _(): eq(_po(side="sell", qty=2130)["qty"], 2130) @case("限价不接受 null / 非正数") def _(): _reject("限价为 None", limit_price=None) _reject("限价为 0", limit_price=0) @case("代码必须点式") def _(): _reject("前缀式代码", ts_code="SH600000") @case("方向与数量非法") def _(): _reject("方向 hold", side="hold") _reject("数量 0", qty=0) @case("note 截断到 200 字, intent 非法回落 OPEN") def _(): eq(len(_po(note="补" * 500)["note"]), 200) eq(_po(intent="WHATEVER")["intent"], "OPEN") print("\n[E] seq 水位推进 (连续前缀语义)") @case("按序到达: 逐格前进") def _(): last, pend = 10, set() for s in (11, 12, 13): last, pend = wsc.next_watermark(last, pend, s) eq((last, pend), (13, set())) @case("缺口: 水位卡住, 后到的先暂存") def _(): last, pend = 10, set() last, pend = wsc.next_watermark(last, pend, 12) eq(last, 10, "12 到了但 11 没到, 水位不能跨过去 —— 跨了 11 就永久丢了") eq(pend, {12}) @case("缺口补齐: 一次性追上") def _(): last, pend = 10, {12, 13} last, pend = wsc.next_watermark(last, pend, 11) eq((last, pend), (13, set())) @case("重复帧: 水位不动") def _(): last, pend = wsc.next_watermark(20, set(), 15) eq((last, pend), (20, set())) last, pend = wsc.next_watermark(20, set(), 20) eq((last, pend), (20, set())) # --- 冷启动基线。这一组是 2026-07-28 联调沙箱抓出来的真实缺陷的回归防线: # 对端 seq 跨重启不回退 (§6.1), 首次对接时它可能已经在 10001; 而我们按 §4.1 传 # last_seq=0。当时的实现会死等 seq 1、2、3…, 水位永远推不动、ack 永远发不出, # 且不抛任何异常 —— 日志只有「上行乱序」在刷屏, 成交一条都没确认。 @case("冷启动: 对端从 10001 起 → 直接以 10000 为基线, 不算缺口") def _(): eq(wsc.cold_start_baseline(0, 10001), (10000, False)) @case("已有水位且连续 → 基线不动") def _(): eq(wsc.cold_start_baseline(10230, 10231), (10230, False)) eq(wsc.cold_start_baseline(10230, 10200), (10230, False)) # 补发的旧消息 @case("已有水位却出现缺口 → 推基线 + 判定为真缺口 (置 resync)") def _(): eq(wsc.cold_start_baseline(10230, 10500), (10499, True)) @case("冷启动基线接上后水位能正常前进") def _(): base, gap = wsc.cold_start_baseline(0, 10001) assert not gap last, pend = base, set() for s in (10001, 10002, 10003): last, pend = wsc.next_watermark(last, pend, s) eq((last, pend), (10003, set())) print("\n[F] 去重键与成交自洽") @case("trade 的第二层去重键取 trade_no") def _(): eq(wsc.dedup_key("trade", {"trade_no": "QMT-88123#1"}), "trade:QMT-88123#1") eq(wsc.dedup_key("order_update", {"instruction_id": "INS-1"}), None) eq(wsc.dedup_key("trade", {}), None) @case("amount = price × qty 校验 (差 > 0.01 告警)") def _(): assert wsc.trade_amount_ok({"price": 12.34, "qty": 600, "amount": 7404.00}) assert not wsc.trade_amount_ok({"price": 12.34, "qty": 600, "amount": 7000.00}) assert not wsc.trade_amount_ok({"price": 12.34, "qty": 600}) @case("终态判定") def _(): for s in ("FILLED", "CANCELLED", "EXPIRED", "REJECTED"): assert wsc.is_final(s), s for s in ("ACCEPTED", "SUBMITTED", "PARTIAL", ""): assert not wsc.is_final(s), s @case("只有 trade 需要 worker 入账") def _(): eq(tuple(wsc.NEEDS_LEDGER), ("trade",)) assert wsc.T_ORDER_UPDATE not in wsc.NEEDS_LEDGER, \ "order_update 只作状态跟踪 —— 拿它入账会和逐笔 trade 重复计数" print("\n[G] 严格单表访问守卫 (upsert 不得被误判成多表)") @case("ON DUPLICATE KEY UPDATE 的四条现存 upsert 全部放行") def _(): from app.db.session import assert_single_table for sql in ( "INSERT INTO pms_runtime_param (param_key, param_value, updated_by, updated_at) " "VALUES (:k, :v, :by, :ts) ON DUPLICATE KEY UPDATE param_value = :v, " "updated_by = :by, updated_at = :ts", "INSERT INTO pms_position (ts_code, status, updated_at) VALUES " "(:code, 'PLANNED', :ts) ON DUPLICATE KEY UPDATE updated_at = :ts", "INSERT INTO pms_daily_report (ymd, report_json, created_at) VALUES " "(:y, :r, :ts) ON DUPLICATE KEY UPDATE report_json = :r, created_at = :ts", "INSERT INTO pms_industry_map (ts_code, industry, updated_at) VALUES " "(:code, :ind, :ts) ON DUPLICATE KEY UPDATE industry = :ind, updated_at = :ts", "INSERT INTO pms_qmt_inbox (seq, msg_id) VALUES (:s, :m) " "ON DUPLICATE KEY UPDATE seq = seq", ): assert_single_table(sql) @case("真的多表 / JOIN / 逗号连表 仍然拦得住") def _(): from app.db.session import MultiTableSQL, assert_single_table for sql in ("SELECT a.* FROM pms_position a JOIN pms_lot b ON a.ts_code = b.ts_code", "SELECT * FROM pms_position, pms_lot WHERE 1 = 1", "INSERT INTO pms_lot (qty) SELECT qty FROM pms_position"): try: assert_single_table(sql) raise AssertionError(f"应当拦截: {sql[:50]}") except MultiTableSQL: pass @case("通道三表的 SQL 全部单表合规") def _(): from app.db.session import assert_single_table for sql in ( "SELECT * FROM pms_qmt_order WHERE status = 'QUEUED' ORDER BY id ASC LIMIT :n", "UPDATE pms_qmt_order SET status = 'SENDING', updated_at = :ts " "WHERE instruction_id = :iid AND status = 'QUEUED'", "SELECT seq FROM pms_qmt_inbox WHERE seq > :n ORDER BY seq ASC LIMIT :m", "SELECT status, COUNT(*) AS n FROM pms_qmt_order GROUP BY status", "UPDATE pms_ws_state SET conn_state = 'STOPPED', heartbeat_at = NULL, " "last_error = :le, updated_at = :ts WHERE id = :i", ): assert_single_table(sql) print("\n[H] dispatcher 纯助手") @case("valid_until 归一成 epoch 毫秒") def _(): from datetime import datetime from app.services.dispatcher import _to_epoch_ms dt = datetime(2026, 7, 28, 14, 35, 0) eq(_to_epoch_ms(dt), int(dt.timestamp() * 1000)) eq(_to_epoch_ms(1769512500000), 1769512500000) eq(_to_epoch_ms(1769512500), 1769512500000, "传秒的也要认") assert _to_epoch_ms(None) > 0, "给不出有效期时要兜底, 不能是无限期挂单" @case("子单 id 反推父指令") def _(): from app.services.dispatcher import _parent_of eq(_parent_of("INS_20260728_600000SH_EXIT_01_D03"), "INS_20260728_600000SH_EXIT_01") eq(_parent_of("INS_20260728_600000SH_EXIT_01"), "INS_20260728_600000SH_EXIT_01") eq(_parent_of("whatever_D07", parent_id="显式优先"), "显式优先") @case("连接类异常判别 (决定要不要重连、算不算发送失败次数)") def _(): from app.ws.runner import _is_conn_error class ConnectionClosed(Exception): pass assert _is_conn_error(ConnectionError("boom")) assert _is_conn_error(OSError("boom")) assert _is_conn_error(ConnectionClosed("bye")) assert not _is_conn_error(ValueError("参数不对")) @case("密钥不进 ParamStore (协议 §10.1.1)") def _(): from app.services import param_store snap_keys = {p["key"] for p in [{"key": k} for k in param_store._editable_keys()]} for k in param_store.SECRET_KEYS: assert k not in snap_keys, f"{k} 不该出现在可调参数里" eq(param_store.get(k), "", f"{k} 必须读不到") r = param_store.set_param(k, "deadbeef") eq(r["ok"], False, f"{k} 必须改不了") def main(): print("=" * 62) print("批次6: ws 直连通道纯逻辑") print("=" * 62) run() print("\n" + "-" * 62) print(f"通过 {len(PASS)} 例, 失败 {len(FAIL)} 例") if FAIL: for name, msg, tb in FAIL: print(f"\n--- {name}\n{tb}") sys.exit(1) print("BATCH6 PASS") if __name__ == "__main__": main()