tradingSystem/scripts/test_batch6_units.py

473 lines
19 KiB
Python
Raw Permalink Normal View History

2026-07-28 15:48:57 +08:00
# -*- 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 归一子单父指令反推
"""
2026-07-28 17:12:49 +08:00
import base64
2026-07-28 15:48:57 +08:00
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")
2026-07-28 17:12:49 +08:00
# --- 公钥形态。2026-07-28 与 QMT 侧交换密钥时踩到: 对方按 PEM 交换 (其 crypto.py
# 用 load_pem_public_key), 而协议 §2.1.1 正文写的是裸 32 字节 base64。PEM 正文
# 解出来是 44 字节 SPKI DER = 12 字节固定头 + 32 字节裸公钥, 直接粘进 .env 会得到
# 「44 字节, 长度不对」这种看不出所以然的报错。现在四种形态全收。
VEC_SPKI = "MCowBQYDK2VwAyEA" + VEC_PUB
VEC_PEM = f"-----BEGIN PUBLIC KEY-----\n{VEC_SPKI}\n-----END PUBLIC KEY-----\n"
@case("公钥四种形态都归一到裸 32 字节 base64")
def _():
for form in (VEC_PUB, VEC_SPKI, VEC_PEM,
VEC_PEM.replace("\n", "\\n"), # .env 里的字面 \n
base64.b64decode(VEC_PUB).hex()): # 64 位 hex
eq(wsc.normalize_pubkey(form), VEC_PUB, f"归一失败: {form[:24]}...")
@case("公钥四种形态都能直接验签")
def _():
for form in (VEC_PUB, VEC_SPKI, VEC_PEM, VEC_PEM.replace("\n", "\\n")):
assert wsc.verify(VEC_CANON, VEC_SIG, form), f"验签失败: {form[:24]}..."
@case("公钥形态识别 (自检脚本给人话提示用)")
def _():
eq(wsc.pubkey_form(VEC_PUB), "裸 base64 (协议形态)")
eq(wsc.pubkey_form(VEC_SPKI), "SPKI base64 (PEM 正文)")
eq(wsc.pubkey_form(VEC_PEM), "PEM 全文")
eq(wsc.pubkey_form(""), "")
@case("公钥非法输入仍然拒绝, 不能因为放宽格式就什么都收")
def _():
for bad_ in ("", "zzz!!!", "aGVsbG8=", "MCowBQYDK2VwAyEA"): # 空/非法/太短/只有头
try:
wsc.normalize_pubkey(bad_)
raise AssertionError(f"应当拒绝: {bad_!r}")
except wsc.CodecError as e:
eq(e.code, "SIG_INVALID", f"输入 {bad_!r}")
@case("本端公钥的 PEM 形态可被对方 load_pem_public_key 读取")
def _():
pem = wsc.public_key_pem(VEC_SEED)
assert pem.startswith("-----BEGIN PUBLIC KEY-----"), pem[:40]
eq(wsc.normalize_pubkey(pem), VEC_PUB, "PEM 与裸 base64 必须是同一把钥匙")
from cryptography.hazmat.primitives import serialization
k = serialization.load_pem_public_key(pem.encode()) # 对方那条路径
eq(base64.b64encode(k.public_bytes_raw()).decode(), VEC_PUB)
2026-07-28 15:48:57 +08:00
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()