tradingSystem/app/core/ws_codec.py

392 lines
18 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 -*-
"""
QMT WebSocket 协议编解码 (纯逻辑, 零 IO, 可单测)
================================================
对应 `QMT_WS_PROTOCOL.md` V1.0 的 §2 (签名与幂等) 与 §3 (消息信封)。
本模块只回答一件事: **一条消息怎么变成字节、字节怎么变回一条消息**。
不碰连接 (在 app/ws/runner.py)、不碰库 (在 app/repo/qmt_repo.py)、不碰业务
(在 app/services/dispatcher.py)。
三处最容易两边写不一致的地方, 在这里钉死:
1. payload 的紧凑 JSON —— 分隔符 "," ":" 无空格、键按 Unicode 码点升序、中文不转义。
2. 规范化串 —— 六段用**单个** \\n (0x0A) 连接, 串尾无换行; v/ts 以十进制整数字符串参与。
3. 价格 —— 上线前必须已经是 2 位小数, 且序列化结果得是 "12.35" 而不是
"12.350000000000001"。协议 §8 规定 QMT 收到超精度价格直接 BAD_PARAM 且不替我们
四舍五入 —— 定价权在 PMS, 就得由 PMS 在出栈这一步收干净。
协议 §2.1.1 给了一组测试向量, `scripts/test_batch6_units.py` 逐字节核对上述四个中间
结果 (payload_json / sha256 / canonical / sig)。**改动本模块后必须先跑通那一组**, 否则
联调时会卡在"签名验不过"而看不出差在哪一段。
--------------------------------------------------------------------------------
为什么 PMS **不**校验上行消息的时间窗 (这条容易踩, 单列出来)
--------------------------------------------------------------------------------
§2.2 的「ts 偏差 > 30 秒即拒绝」是 **QMT 校验 PMS** 的规则。反过来 PMS 不能照抄:
§6.1 规定断线补发的消息与首次推送**完全一致** —— 同 seq、同 msg_id、同内容、**同签名**,
也就是补发时 ts 仍是原来那条的发送时刻。断网十分钟后重连, 补发的第一批消息 ts 已经差了
十分钟, 若按时间窗校验会被整批丢弃, 补发机制当场失效, 而且失效得很安静。
PMS 侧防重放改靠两层: **seq 水位** (seq ≤ last_seq 一律丢弃, 见 §6.1) + **trade_no 去重**
(§5.5)。签名保证内容没被篡改, 序号保证同一条不会入账两次 —— 时间窗在这里是多余且有害的。
"""
from __future__ import annotations
import base64
import hashlib
import json
import secrets
import time
from decimal import ROUND_HALF_UP, Decimal
try: # cryptography 在 requirements.txt 里, 正常一定有
from cryptography.exceptions import InvalidSignature
from cryptography.hazmat.primitives.asymmetric.ed25519 import (Ed25519PrivateKey,
Ed25519PublicKey)
CRYPTO_AVAILABLE = True
except ImportError: # 缺库时不让 import 就炸 —— 由调用方拿到明确报错
CRYPTO_AVAILABLE = False
InvalidSignature = Exception # type: ignore
Ed25519PrivateKey = Ed25519PublicKey = None # type: ignore
PROTOCOL_VERSION = 1
# ---- 消息类型 (§4 下行 / §5 上行) ----
T_HELLO, T_PLACE, T_CANCEL = "hello", "place_order", "cancel_order"
T_Q_POS, T_Q_FUNDS, T_Q_ORDERS = "query_positions", "query_funds", "query_orders"
T_ACK_SEQ, T_PING = "ack_seq", "ping"
DOWNSTREAM_TYPES = (T_HELLO, T_PLACE, T_CANCEL, T_Q_POS, T_Q_FUNDS, T_Q_ORDERS,
T_ACK_SEQ, T_PING)
T_HELLO_ACK, T_ACK, T_REJECT = "hello_ack", "ack", "reject"
T_ORDER_UPDATE, T_TRADE, T_PERSIST = "order_update", "trade", "persist_result"
T_SNAPSHOT, T_POS_UPDATE, T_FUNDS_UPDATE, T_PONG = ("snapshot", "position_update",
"funds_update", "pong")
UPSTREAM_TYPES = (T_HELLO_ACK, T_ACK, T_REJECT, T_ORDER_UPDATE, T_TRADE, T_PERSIST,
T_SNAPSHOT, T_POS_UPDATE, T_FUNDS_UPDATE, T_PONG)
# 只有 trade 需要 worker 入账 (§5.5「账本以 trade 为唯一入账依据」), 其余上行消息
# 由常驻进程自己消化完就落库存档。这个集合是 inbox 里 processed 初值的判据。
NEEDS_LEDGER = (T_TRADE,)
# 指令状态 (§7.1)。终态唯一、不可再变。
ST_ACCEPTED, ST_SUBMITTED, ST_PARTIAL = "ACCEPTED", "SUBMITTED", "PARTIAL"
ST_FILLED, ST_CANCELLED, ST_EXPIRED, ST_REJECTED = ("FILLED", "CANCELLED", "EXPIRED",
"REJECTED")
FINAL_STATUSES = (ST_FILLED, ST_CANCELLED, ST_EXPIRED, ST_REJECTED)
PROTOCOL_STATUSES = (ST_ACCEPTED, ST_SUBMITTED, ST_PARTIAL) + FINAL_STATUSES
# 拒绝码 (§7.3)。retryable=True 的换个价/等一等还能再来, False 的重发多少次都一样。
REJECT_RETRYABLE = {
"SIG_INVALID": False, "TS_SKEW": True, "VERSION": False, "DUP_INSTRUCTION": False,
"BAD_PARAM": False, "UNKNOWN_CODE": False, "SUSPENDED": False,
"LIMIT_UP": True, "LIMIT_DOWN": True, "INSUFFICIENT_CASH": True,
"INSUFFICIENT_POSITION": True, "NOT_TRADING_TIME": True, "ALREADY_FINAL": False,
"BROKER_ERROR": True, "INTERNAL": True,
}
INTENTS = ("OPEN", "FILL", "ADD", "DCA", "TRIM", "EXIT", "T0")
class CodecError(ValueError):
"""编解码/校验失败。带 code 便于直接映射到协议 §7.3 的拒绝码。"""
def __init__(self, code: str, message: str = ""):
super().__init__(f"{code}: {message}" if message else code)
self.code = code
self.message = message or code
# ================================================================ 序列化
def payload_json(payload: dict) -> str:
"""payload 的紧凑 JSON —— 签名对象的第 6 段就是它的 sha256。
三个选项一个都不能改: sort_keys 按 Unicode 码点升序、separators 无空格、
ensure_ascii=False 让中文以 UTF-8 原文出现 (对方 Java 侧是默认行为)。
"""
return json.dumps(payload, sort_keys=True, separators=(",", ":"), ensure_ascii=False)
def payload_sha256(payload: dict) -> str:
return hashlib.sha256(payload_json(payload).encode("utf-8")).hexdigest()
def canonical(*, v: int, type_: str, msg_id: str, ts: int, nonce: str,
payload_hash: str) -> str:
"""规范化串 (§2.1)。六段, 单个换行连接, 串尾无换行。"""
return f"{int(v)}\n{type_}\n{msg_id}\n{int(ts)}\n{nonce}\n{payload_hash}"
def q2(x) -> float:
"""价格收成 2 位小数 (四舍五入), 且保证 json 序列化出来是 "12.35" 这种短表示。
择时模块算出来的限价常带一长串尾数 (现价 × 0.998 之类), 协议 §8 要求 2 位小数且
QMT 收到超精度**直接 BAD_PARAM 不替我们舍入** —— 定价权在 PMS, 就得在出栈这一步收
干净。走 Decimal 而不是内置 round(), 是为了避开 round(2.675, 2) = 2.67 那类
银行家舍入 + 二进制表示的双重意外。
"""
d = Decimal(str(x)).quantize(Decimal("0.01"), rounding=ROUND_HALF_UP)
return float(d)
def jsonable(v):
"""把 Decimal / date / 自定义对象收敛成 JSON 原生类型 (payload 里不许出现别的)。"""
if isinstance(v, dict):
return {str(k): jsonable(x) for k, x in v.items()}
if isinstance(v, (list, tuple)):
return [jsonable(x) for x in v]
if isinstance(v, bool) or v is None or isinstance(v, (int, str)):
return v
if isinstance(v, Decimal):
return float(v)
if isinstance(v, float):
return v
return str(v)
# ================================================================ 签名
def _seed_bytes(seed_hex: str) -> bytes:
s = (seed_hex or "").strip().lower().replace(" ", "")
if len(s) != 64:
raise CodecError("SIG_INVALID", f"私钥 seed 应为 32 字节 hex (64 字符), 收到 {len(s)} 字符")
try:
return bytes.fromhex(s)
except ValueError as e:
raise CodecError("SIG_INVALID", f"私钥 seed 非法 hex: {e}") from e
def sign(canon: str, seed_hex: str) -> str:
if not CRYPTO_AVAILABLE:
raise CodecError("INTERNAL", "缺少 cryptography 依赖, 无法签名 (requirements.txt)")
sk = Ed25519PrivateKey.from_private_bytes(_seed_bytes(seed_hex))
return base64.b64encode(sk.sign(canon.encode("utf-8"))).decode()
def public_key_b64(seed_hex: str) -> str:
"""由私钥 seed 推出公钥 —— 部署时把它带外送给 QMT 侧 (§10.1.1 第 2 条)。"""
if not CRYPTO_AVAILABLE:
raise CodecError("INTERNAL", "缺少 cryptography 依赖")
sk = Ed25519PrivateKey.from_private_bytes(_seed_bytes(seed_hex))
return base64.b64encode(sk.public_key().public_bytes_raw()).decode()
def verify(canon: str, sig_b64: str, pubkey_b64: str) -> bool:
"""验签。任何异常 (公钥非法/签名非法/长度不对/base64 烂) 一律 False —— 验不过就是验不过,
不区分「为什么验不过」, 免得给探测方留下侧信道。"""
if not CRYPTO_AVAILABLE:
raise CodecError("INTERNAL", "缺少 cryptography 依赖, 无法验签")
try:
pk = Ed25519PublicKey.from_public_bytes(base64.b64decode(pubkey_b64))
pk.verify(base64.b64decode(sig_b64), canon.encode("utf-8"))
return True
except Exception: # noqa: BLE001 —— 见上, 故意不细分
return False
# ================================================================ 标识符
def new_nonce() -> str:
"""每条消息随机 16 字节 hex (§2.2)。仅参与签名, 对端不必存。"""
return secrets.token_hex(16)
def new_msg_id(ymd: int, seq: int) -> str:
"""单条消息唯一, 仅供日志追踪 —— **不是**业务幂等键 (§3)。"""
return f"m-{int(ymd)}-{int(seq) % 1000000:06d}"
def new_instruction_id(ymd: int) -> str:
"""幂等键 (§2.2)。格式 INS-<yyyymmdd>-<8位随机>, 长度 21 ≤ 64。"""
return f"INS-{int(ymd)}-{secrets.token_hex(4)}"
def new_cancel_id(ymd: int) -> str:
return f"CXL-{int(ymd)}-{secrets.token_hex(4)}"
def now_ms(clock=None) -> int:
"""epoch 毫秒 (§1: 时间字段一律整数毫秒)。clock 可注入, 便于单测。"""
return int(clock) if clock is not None else int(time.time() * 1000)
# ================================================================ 组包 / 拆包
def build(type_: str, payload: dict, *, seed_hex: str, msg_id: str, ts: int = None,
nonce: str = None, corr_id: str = None, seq: int = None) -> dict:
"""组一条下行消息 (含签名)。返回 dict, 由调用方 json.dumps 后 send。
seq 只在上行必填 (§3), 下行留 None 即不带该字段 —— 别为了"对称"给下行编号,
那会让对方的补发逻辑多一个歧义来源。
"""
body = jsonable(payload or {})
ts = now_ms() if ts is None else int(ts)
nonce = nonce or new_nonce()
canon = canonical(v=PROTOCOL_VERSION, type_=type_, msg_id=msg_id, ts=ts, nonce=nonce,
payload_hash=payload_sha256(body))
env = {"v": PROTOCOL_VERSION, "type": type_, "msg_id": msg_id, "ts": ts,
"nonce": nonce, "payload": body, "sig": sign(canon, seed_hex)}
if corr_id:
env["corr_id"] = corr_id
if seq is not None:
env["seq"] = int(seq)
return env
def dumps(env: dict) -> str:
"""整条消息落到线上的字节。只有 payload 需要规范化, 信封本身怎么排都行 ——
但仍用同一套紧凑参数, 免得日志里两种风格混着看。"""
return json.dumps(env, separators=(",", ":"), ensure_ascii=False)
def parse(raw, *, peer_pubkey_b64: str) -> dict:
"""拆一条上行消息并验签。验不过抛 CodecError, **调用方直接丢弃且不做任何业务动作**
(§2.1: 验签失败 → 丢弃并回 reject{SIG_INVALID})。
这里不做时间窗校验 —— 原因见模块头部那段。
"""
if isinstance(raw, (bytes, bytearray)):
try:
raw = raw.decode("utf-8")
except UnicodeDecodeError as e:
raise CodecError("BAD_PARAM", f"非 UTF-8 帧: {e}") from e
try:
env = json.loads(raw)
except (ValueError, TypeError) as e:
raise CodecError("BAD_PARAM", f"非法 JSON: {e}") from e
if not isinstance(env, dict):
raise CodecError("BAD_PARAM", "消息顶层不是对象")
v = env.get("v")
if v != PROTOCOL_VERSION:
raise CodecError("VERSION", f"不支持的协议版本 {v!r} (本端 {PROTOCOL_VERSION})")
type_ = env.get("type")
if not type_ or not isinstance(type_, str):
raise CodecError("BAD_PARAM", "缺少 type")
for k in ("msg_id", "ts", "nonce", "sig"):
if env.get(k) in (None, ""):
raise CodecError("BAD_PARAM", f"缺少信封字段 {k}")
payload = env.get("payload")
if payload is None:
payload = {}
if not isinstance(payload, dict):
raise CodecError("BAD_PARAM", "payload 不是对象")
canon = canonical(v=v, type_=type_, msg_id=str(env["msg_id"]), ts=int(env["ts"]),
nonce=str(env["nonce"]), payload_hash=payload_sha256(payload))
if not peer_pubkey_b64:
raise CodecError("SIG_INVALID", "未配置对端公钥 (PMS_QMT_PEER_PUBKEY_B64), 无法验签")
if not verify(canon, str(env["sig"]), peer_pubkey_b64):
raise CodecError("SIG_INVALID", f"上行消息验签失败 type={type_} msg_id={env['msg_id']}")
env["payload"] = payload
if env.get("seq") is not None:
env["seq"] = int(env["seq"])
return env
# ================================================================ 各类下行 payload
def place_order_payload(*, instruction_id: str, ts_code: str, side: str, qty: int,
limit_price, valid_until: int, intent: str = "OPEN",
note: str = "") -> dict:
"""§4.2。落地前把易错项一次性校验干净 —— 宁可在本地报错, 不要换 QMT 一个 BAD_PARAM。"""
side = str(side or "").lower()
if side not in ("buy", "sell"):
raise CodecError("BAD_PARAM", f"side 只能是 buy/sell, 收到 {side!r}")
qty = int(qty or 0)
if qty <= 0:
raise CodecError("BAD_PARAM", f"qty 必须为正整数股, 收到 {qty}")
# 整百规则 (§8): 买入必整百; 卖出通常整百, 清仓允许零股尾数 —— 故只拦买入。
if side == "buy" and qty % 100:
raise CodecError("BAD_PARAM", f"买入数量必须整百, 收到 {qty}")
if limit_price in (None, ""):
raise CodecError("BAD_PARAM", "limit_price 必填 (协议不接受 null, 定价权在 PMS)")
px = q2(limit_price)
if px <= 0:
raise CodecError("BAD_PARAM", f"limit_price 必须为正, 收到 {limit_price!r}")
if not str(ts_code or "").strip():
raise CodecError("BAD_PARAM", "缺少 ts_code")
if "." not in str(ts_code):
raise CodecError("BAD_PARAM", f"ts_code 必须是点式 (600000.SH), 收到 {ts_code!r}")
intent = str(intent or "OPEN").upper()
if intent not in INTENTS:
intent = "OPEN"
if len(str(instruction_id or "")) > 64 or not instruction_id:
raise CodecError("BAD_PARAM", f"instruction_id 长度须在 1~64, 收到 {instruction_id!r}")
return {"instruction_id": str(instruction_id), "ts_code": str(ts_code), "side": side,
"qty": qty, "limit_price": px, "valid_until": int(valid_until),
"intent": intent, "note": str(note or "")[:200]}
def cancel_order_payload(*, cancel_id: str, instruction_id: str) -> dict:
return {"cancel_id": str(cancel_id), "instruction_id": str(instruction_id)}
def hello_payload(last_seq: int) -> dict:
"""§4.1。首次连接或本地无记录传 0。"""
return {"client": "pms", "last_seq": int(last_seq or 0), "protocol": PROTOCOL_VERSION}
def ack_seq_payload(seq: int) -> dict:
return {"seq": int(seq)}
# ================================================================ 上行 payload 读取
def dedup_key(type_: str, payload: dict):
"""第二层去重键 (§5.5)。目前只有 trade 有 —— trade_no 唯一, 命中即丢弃。
第一层是 seq (补发时同一条消息 seq 不变)。两层都设是为了防住 seq 实现出 bug 的场景。
"""
if type_ == T_TRADE:
tn = (payload or {}).get("trade_no")
return f"trade:{tn}" if tn else None
return None
def trade_amount_ok(payload: dict, tol: float = 0.01) -> bool:
"""§5.5: amount 应等于 price × qty, 差异 > 0.01 元告警。"""
try:
px, qty, amt = float(payload["price"]), int(payload["qty"]), float(payload["amount"])
except (KeyError, TypeError, ValueError):
return False
return abs(px * qty - amt) <= tol
def is_final(status: str) -> bool:
return str(status or "").upper() in FINAL_STATUSES
def cold_start_baseline(last_seq: int, first_seq: int) -> tuple:
"""对端序号不从 1 开始时, 本端水位该从哪里起算。返回 (基线, 是否真缺口)。
**这是本协议最安静的一种死法, 值得多写几行。** QMT 的 seq 是跨重启、跨交易日都不回退
的全局计数器 (§6.1), 我们接上去的时候它可能早就跑到几万了。而 PMS 首次连接按 §4.1 传
`last_seq=0` —— 若照着「水位必须连续」的规矩死等 seq 1、2、3…, 水位就永远推不动:
ack_seq 发不出去 → 对端的消息永远清理不掉 → 重连时又从头补发一遍。整个过程**不抛任何
异常**, 日志上只有一行「上行乱序」在刷屏, 而成交其实一条都没确认。
分两种情形:
last_seq == 0 冷启动。本端从没收过任何消息, 谈不上"" —— 以 first_seq-1 为起点。
这不是缺口, 这是起点。
last_seq > 0 真缺口, 中间那段再也拿不到了。基线同样要推上去 (卡死只会更糟),
但必须置 resync 标记走全量对账 —— §6.2。
"""
if int(first_seq) <= int(last_seq) + 1:
return int(last_seq), False
return int(first_seq) - 1, int(last_seq) > 0
def next_watermark(last_seq: int, arrived: set, seq: int) -> tuple:
"""收到 seq 后推进连续水位。返回 (新水位, 新的乱序暂存集合)。
水位必须是**连续前缀**的末位: 确认 10450 隐含确认之前所有 (§4.5 累积确认语义),
中间缺一条就不能往前跨 —— 跨过去那条就永久丢了, 而且丢得毫无痕迹。
正常情况下 QMT 按序补发, 这个集合始终是空的; 它存在是为了「万一乱序」不出错。
"""
arrived = set(arrived or ())
seq = int(seq)
if seq <= last_seq:
return last_seq, arrived # 已确认过, 重复帧
arrived.add(seq)
while (last_seq + 1) in arrived:
last_seq += 1
arrived.discard(last_seq)
return last_seq, arrived