# -*- 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--<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