# -*- coding: utf-8 -*- """ ws 直连通道三表的数据访问 (pms_qmt_order / pms_qmt_inbox / pms_ws_state) ======================================================================= 与 pms_repo 同纪律: 每个函数只碰一张表, SQL 全部经 db.session 的单表守卫, 可更新列走白名单。表结构与"为什么是三张表"见 ddl_pms_v1.sql 尾部。 三张表分别对应协议里的三件事: pms_qmt_order §4.2 place_order / §4.3 cancel_order 的出口队列 + 委托状态跟踪 pms_qmt_inbox §4.5「确认前必须已持久化」的那个"库" + §6.1 双层去重 pms_ws_state §6.1 seq 水位 (跨重启不回退) + 常驻进程存活心跳 """ from __future__ import annotations import json from datetime import datetime from app.db.session import execute, fetch_all, fetch_one _NOW = lambda: datetime.now() # noqa: E731 (容器时区 Asia/Shanghai) STATE_ID = 1 # pms_ws_state 恒一行 # 本地出口状态 (还没进协议状态机) OS_QUEUED, OS_SENDING, OS_SENT = "QUEUED", "SENDING", "SENT" OS_SEND_FAILED, OS_ABORTED = "SEND_FAILED", "ABORTED" # 协议状态 (§7.1) OS_ACCEPTED, OS_SUBMITTED, OS_PARTIAL = "ACCEPTED", "SUBMITTED", "PARTIAL" OS_FILLED, OS_CANCELLED, OS_EXPIRED, OS_REJECTED = ("FILLED", "CANCELLED", "EXPIRED", "REJECTED") FINAL = (OS_FILLED, OS_CANCELLED, OS_EXPIRED, OS_REJECTED, OS_SEND_FAILED, OS_ABORTED) # 「在途」= 还可能成交或还能撤的。撤单只对这些有意义。 LIVE = (OS_QUEUED, OS_SENDING, OS_SENT, OS_ACCEPTED, OS_SUBMITTED, OS_PARTIAL) CANCEL_NONE, CANCEL_REQUESTED, CANCEL_SENT = "NONE", "REQUESTED", "SENT" ORDER_COLS = { "status", "broker_order_id", "cum_qty", "cum_avg_price", "leaves_qty", "cancel_state", "cancel_id", "cancel_req_at", "reject_code", "reject_reason", "send_attempts", "sent_at", "final_at", "note", } def _dumps(v): return json.dumps(v, ensure_ascii=False) if not isinstance(v, (str, type(None))) else v def _loads(v, default=None): if v in (None, ""): return default if isinstance(v, (dict, list)): return v try: return json.loads(v) except (ValueError, TypeError): return default def _in_clause(values, prefix: str, params: dict) -> str: keys = [] for i, v in enumerate(values): keys.append(f":{prefix}{i}") params[f"{prefix}{i}"] = v return ", ".join(keys) # ================================================================ pms_ws_state def get_state() -> dict: r = fetch_one("SELECT * FROM pms_ws_state WHERE id = :i", {"i": STATE_ID}) if not r: return {"id": STATE_ID, "last_seq": 0, "acked_seq": 0, "server_seq": 0, "conn_state": "INIT", "heartbeat_at": None, "resync_flag": 0, "connected_at": None, "last_error": None, "stat": {}} d = dict(r) d["stat"] = _loads(d.pop("stat_json", None), {}) return d def ensure_state() -> int: """建表脚本已插过一行; 这里兜底 (库是别人手工建的/被清过 也不至于全线报错)。""" return execute( "INSERT INTO pms_ws_state (id, last_seq, acked_seq, server_seq, conn_state, " "updated_at) VALUES (:i, 0, 0, 0, 'INIT', :ts) " "ON DUPLICATE KEY UPDATE updated_at = :ts", {"i": STATE_ID, "ts": _NOW()}) def save_watermark(last_seq: int, acked_seq=None) -> int: """水位只进不退 —— GREATEST 兜住并发/乱序写回, 回退一格就意味着重复入账。""" sets = ["last_seq = GREATEST(last_seq, :ls)"] p = {"ls": int(last_seq), "i": STATE_ID, "ts": _NOW()} if acked_seq is not None: sets.append("acked_seq = GREATEST(acked_seq, :as_)") p["as_"] = int(acked_seq) return execute(f"UPDATE pms_ws_state SET {', '.join(sets)}, updated_at = :ts " f"WHERE id = :i", p) def set_conn(conn_state: str, *, connected_at=None, last_error=None, server_seq=None, resync=None, stat=None, beat: bool = True) -> int: sets, p = ["conn_state = :cs"], {"cs": conn_state, "i": STATE_ID, "ts": _NOW()} if beat: sets.append("heartbeat_at = :ts") if connected_at is not None: sets.append("connected_at = :ca") p["ca"] = connected_at if last_error is not None: sets.append("last_error = :le") p["le"] = str(last_error)[:300] if server_seq is not None: sets.append("server_seq = :ss") p["ss"] = int(server_seq) if resync is not None: sets.append("resync_flag = :rf") p["rf"] = 1 if resync else 0 if stat is not None: sets.append("stat_json = :sj") p["sj"] = _dumps(stat) return execute(f"UPDATE pms_ws_state SET {', '.join(sets)}, updated_at = :ts " f"WHERE id = :i", p) def mark_stopped(note: str = "") -> int: """优雅退出: 置 STOPPED 并**清空心跳**。 清心跳是关键一步 —— 不清的话, 停机后的 stale 窗口 (默认 15 秒) 里 dispatcher 仍 认为进程活着, 卖出指令还会继续往队列里排, 而已经没人会发它们了。 """ return execute("UPDATE pms_ws_state SET conn_state = 'STOPPED', heartbeat_at = NULL, " "last_error = :le, updated_at = :ts WHERE id = :i", {"le": str(note)[:300], "i": STATE_ID, "ts": _NOW()}) def beat(stat=None) -> int: """常驻进程存活心跳。dispatcher 看这个时间戳判断「ws 进程还在不在」—— 连接断了是一回事 (还能重连), 进程没了是另一回事 (队列永远发不出去)。""" sets, p = ["heartbeat_at = :ts"], {"i": STATE_ID, "ts": _NOW()} if stat is not None: sets.append("stat_json = :sj") p["sj"] = _dumps(stat) return execute(f"UPDATE pms_ws_state SET {', '.join(sets)}, updated_at = :ts " f"WHERE id = :i", p) # ================================================================ pms_qmt_inbox PUT_NEW, PUT_DUP_SEQ, PUT_DUP_KEY = "NEW", "DUP_SEQ", "DUP_KEY" def inbox_put(*, seq: int, msg_id: str, msg_type: str, payload: dict, msg_ts: int, corr_id=None, dedup_key=None, processed: int = 2, note=None) -> str: """落一条上行消息。返回 NEW / DUP_SEQ / DUP_KEY。 双层去重 (§6.1 + §5.5): 主键 seq 是第一层, 唯一索引 dedup_key (trade_no) 是第二层。 **判重必须先 SELECT, 不能拿 rowcount 当判据。** 这条踩过: `INSERT ... ON DUPLICATE KEY UPDATE seq = seq` 在"行已存在、值没变"时, MySQL 手册说 affected-rows 是 0 —— 但那是**没开 CLIENT_FOUND_ROWS** 的前提。SQLAlchemy 的 MySQL 方言默认就开着这个标志 (它要让 rowcount 反映"匹配到几行"而不是"改了几行"), 于是重复 插入照样回 1。实测: 表里始终只有 1 行, rowcount 却三次都是 1。 后果不是报错而是**静默双记**: 断线重连后 QMT 按 §6.1 重发的成交会被当成新成交, 再入账一次 —— 持仓和摊薄成本直接算错。多一次 SELECT 换这个确定性, 非常值。 **DUP_KEY 这条分支容易漏**: trade_no 撞了但 seq 是新的 —— QMT 用新序号重推了一条我们 已入过账的成交。这一笔不能再入账, 但**这个 seq 仍必须占住一行**, 否则连续水位永远卡在 它前面, ack_seq 再也推不动, 对端的消息也永远清理不掉。故补插一行 dedup_key=NULL、 processed=2 的存档行。 """ now = _NOW() seq = int(seq) if fetch_one("SELECT seq FROM pms_qmt_inbox WHERE seq = :s", {"s": seq}): return PUT_DUP_SEQ verdict = PUT_NEW if dedup_key and fetch_one("SELECT seq FROM pms_qmt_inbox WHERE dedup_key = :k", {"k": str(dedup_key)[:80]}): verdict = PUT_DUP_KEY note = f"重复成交 (dedup_key={dedup_key}), 不入账, 仅占位以推进 seq 水位" dedup_key, processed = None, 2 p = {"s": seq, "mi": str(msg_id)[:64], "mt": str(msg_type)[:24], "ci": (str(corr_id)[:64] if corr_id else None), "dk": (str(dedup_key)[:80] if dedup_key else None), "pj": _dumps(payload or {}), "mts": int(msg_ts or 0), "ts": now, "pc": int(processed), "nt": (str(note)[:300] if note else None), "pa": now if int(processed) != 0 else None} # ON DUPLICATE KEY UPDATE 只作兜底 (万一并发/时序意外), 判重结论以上面的 SELECT 为准 execute("INSERT INTO pms_qmt_inbox (seq, msg_id, msg_type, corr_id, dedup_key, " "payload_json, msg_ts, received_at, processed, processed_at, process_note) " "VALUES (:s, :mi, :mt, :ci, :dk, :pj, :mts, :ts, :pc, :pa, :nt) " "ON DUPLICATE KEY UPDATE seq = seq", p) return verdict def inbox_recover_watermark(stored_last_seq: int, limit: int = 20000) -> int: """用 inbox 重算真正的连续水位。 last_seq 落库是按批次刷的 (每 20 条 / 2 秒), 崩溃时可能落后于实际已落库的消息。 inbox 行才是事实, 这里从 stored 往后走连续段 —— 少 ack 一点只是让 QMT 多留一会儿, 多 ack 一点会让数据永久丢失, 所以宁可从保守值往前推。 """ cur = int(stored_last_seq or 0) rows = fetch_all("SELECT seq FROM pms_qmt_inbox WHERE seq > :n ORDER BY seq ASC LIMIT :m", {"n": cur, "m": int(limit)}) for r in rows: if int(r["seq"]) == cur + 1: cur += 1 else: break return cur def inbox_pending(limit: int = 500) -> list: rows = fetch_all("SELECT * FROM pms_qmt_inbox WHERE processed = 0 ORDER BY seq ASC " "LIMIT :n", {"n": int(limit)}) for r in rows: r["payload"] = _loads(r.get("payload_json"), {}) return rows def inbox_mark(seqs: list, processed: int = 1, note=None) -> int: if not seqs: return 0 p = {"pc": int(processed), "ts": _NOW(), "nt": (str(note)[:300] if note else None)} return execute(f"UPDATE pms_qmt_inbox SET processed = :pc, processed_at = :ts, " f"process_note = :nt WHERE seq IN ({_in_clause(seqs, 's', p)})", p) def inbox_list(*, msg_type=None, corr_id=None, limit: int = 200) -> list: where, p = [], {"n": int(limit)} if msg_type: where.append("msg_type = :mt") p["mt"] = msg_type if corr_id: where.append("corr_id = :ci") p["ci"] = corr_id sql = "SELECT * FROM pms_qmt_inbox" if where: sql += " WHERE " + " AND ".join(where) sql += " ORDER BY seq DESC LIMIT :n" rows = fetch_all(sql, p) for r in rows: r["payload"] = _loads(r.get("payload_json"), {}) return rows def inbox_pending_count() -> int: r = fetch_one("SELECT COUNT(*) AS n FROM pms_qmt_inbox WHERE processed = 0") return int((r or {}).get("n") or 0) # ================================================================ pms_qmt_order def enqueue_order(*, instruction_id, parent_id, ts_code, side, qty, limit_price, valid_until, intent="OPEN", note=None) -> int: """把一张待发委托落进出口队列 —— 这一步就是「先记账」, ws 进程随后才「后动作」。""" now = _NOW() return execute( "INSERT INTO pms_qmt_order (instruction_id, parent_id, ts_code, side, qty, " "limit_price, valid_until, intent, note, status, cancel_state, cum_qty, " "send_attempts, created_at, updated_at) VALUES (:iid, :pid, :code, :side, :qty, " ":px, :vu, :it, :nt, 'QUEUED', 'NONE', 0, 0, :ts, :ts)", {"iid": instruction_id, "pid": parent_id, "code": ts_code, "side": side, "qty": int(qty), "px": float(limit_price), "vu": int(valid_until), "it": intent, "nt": (str(note)[:200] if note else None), "ts": now}) def get_order(instruction_id: str): return fetch_one("SELECT * FROM pms_qmt_order WHERE instruction_id = :iid", {"iid": instruction_id}) def list_orders(*, statuses=None, parent_id=None, ts_code=None, limit: int = 200) -> list: where, p = [], {"n": int(limit)} if statuses: where.append(f"status IN ({_in_clause(list(statuses), 'st', p)})") if parent_id: where.append("parent_id = :pid") p["pid"] = parent_id if ts_code: where.append("ts_code = :code") p["code"] = ts_code sql = "SELECT * FROM pms_qmt_order" if where: sql += " WHERE " + " AND ".join(where) sql += " ORDER BY id DESC LIMIT :n" return fetch_all(sql, p) def next_queued(limit: int = 20, side=None) -> list: """取待发委托。按 id 升序 = 先进先发, 保证同一只票的分笔不乱序 (§1「不开第二条连接」 是为了避免乱序, 出口这一端也得守住)。""" p = {"n": int(limit)} sql = "SELECT * FROM pms_qmt_order WHERE status = 'QUEUED'" if side: sql += " AND side = :side" p["side"] = side sql += " ORDER BY id ASC LIMIT :n" return fetch_all(sql, p) def claim_order(instruction_id: str) -> bool: """QUEUED → SENDING 的原子认领。返回 False 说明被别人抢走了/状态已变, **不要再发**。 正常只有一个 ws 进程, 但滚动重启会有两个进程短暂并存 —— 那一瞬间靠这条 CAS 兜住, 而不是靠"我们约定只起一个"。 """ n = execute("UPDATE pms_qmt_order SET status = 'SENDING', updated_at = :ts " "WHERE instruction_id = :iid AND status = 'QUEUED'", {"iid": instruction_id, "ts": _NOW()}) return bool(n) def mark_sent(instruction_id: str) -> int: return execute("UPDATE pms_qmt_order SET status = 'SENT', sent_at = :ts, " "send_attempts = send_attempts + 1, updated_at = :ts " "WHERE instruction_id = :iid", {"iid": instruction_id, "ts": _NOW()}) def requeue_order(instruction_id: str, error: str = "", max_attempts: int = 3, count_attempt: bool = True) -> int: """发送失败退回队列; 试满次数置 SEND_FAILED 等人工 —— 不无限重试 (设计 §13 「指令下发失败/超时 → 不自动重发」的折中: 网络抖动允许有限重试, 但不能一直撞)。 count_attempt=False 用于**连接断了**的场景: 那是通道的问题不是这张单的问题, 不该 记在它头上, 否则一次几秒的抖动就能把待发单全烧成 SEND_FAILED。 注: MySQL 的 UPDATE ... SET 按书写顺序求值且后项可见前项新值 —— 故 status 必须写在 send_attempts **之前** (用旧值 + inc 判断), 顺序不可调换。同 pms_repo.close_lot_qty。 """ return execute( "UPDATE pms_qmt_order SET " "status = CASE WHEN send_attempts + :inc >= :mx THEN 'SEND_FAILED' " " ELSE 'QUEUED' END, " "send_attempts = send_attempts + :inc, " "reject_reason = :err, updated_at = :ts WHERE instruction_id = :iid", {"iid": instruction_id, "mx": int(max_attempts), "err": str(error)[:300], "inc": 1 if count_attempt else 0, "ts": _NOW()}) def reset_stuck_sending() -> int: """进程启动时把 SENDING 退回 QUEUED。 SENDING 意味着"认领了但没记到 SENT" —— 可能已经发出去了, 也可能没有。重发是安全的: 协议 §2.2 规定重复 instruction_id 不会二次下单, 只回 ack{duplicate:true} 带当前状态。 这正是幂等键存在的意义, 该用就用, 别为了"怕重复"把单子丢在半路。 """ return execute("UPDATE pms_qmt_order SET status = 'QUEUED', updated_at = :ts " "WHERE status = 'SENDING'", {"ts": _NOW()}) def abort_order(instruction_id: str, note: str) -> int: """未发出即本地作废 (典型: 排队期间 valid_until 已过, 发出去也只会立刻 EXPIRED)。""" now = _NOW() return execute("UPDATE pms_qmt_order SET status = 'ABORTED', reject_code = 'LOCAL_ABORT', " "reject_reason = :nt, final_at = :ts, updated_at = :ts " "WHERE instruction_id = :iid", {"iid": instruction_id, "nt": str(note)[:300], "ts": now}) def update_order(instruction_id: str, **fields) -> int: cols = [c for c in fields if c in ORDER_COLS] if not cols: return 0 p = {c: fields[c] for c in cols} p.update({"iid": instruction_id, "ts": _NOW()}) clause = ", ".join(f"{c} = :{c}" for c in cols) return execute(f"UPDATE pms_qmt_order SET {clause}, updated_at = :ts " f"WHERE instruction_id = :iid", p) def request_cancel(*, parent_id: str, cancel_id: str) -> int: """把某父指令名下所有在途子单标为待撤。ws 进程扫到后发 cancel_order。 cancel_state 同时是**区分 CANCELLED 与 EXPIRED 的本地依据** (协议 §7.2): 下游落库层 两者都写 cancelled, 但 PMS 自己知道有没有发过撤单。 """ p = {"pid": parent_id, "cid": cancel_id, "ts": _NOW()} return execute( f"UPDATE pms_qmt_order SET cancel_state = 'REQUESTED', cancel_id = :cid, " f"cancel_req_at = :ts, updated_at = :ts WHERE parent_id = :pid " f"AND cancel_state = 'NONE' AND status IN ({_in_clause(list(LIVE), 'lv', p)})", p) def next_cancel_requests(limit: int = 20) -> list: p = {"n": int(limit)} return fetch_all( f"SELECT * FROM pms_qmt_order WHERE cancel_state = 'REQUESTED' " f"AND status IN ({_in_clause(list(LIVE), 'lv', p)}) ORDER BY id ASC LIMIT :n", p) def mark_cancel_sent(instruction_id: str) -> int: return execute("UPDATE pms_qmt_order SET cancel_state = 'SENT', updated_at = :ts " "WHERE instruction_id = :iid", {"iid": instruction_id, "ts": _NOW()}) def queue_depth() -> dict: rows = fetch_all("SELECT status, COUNT(*) AS n FROM pms_qmt_order GROUP BY status") return {r["status"]: int(r["n"]) for r in rows}