tradingSystem/scripts/test_batch6_units.py

473 lines
19 KiB
Python
Raw Permalink 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 -*-
"""
批次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 base64
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 与 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)
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()