388 lines
17 KiB
Python
388 lines
17 KiB
Python
# -*- 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) 是第二层。
|
|
|
|
**DUP_KEY 这条分支容易漏, 单独说明**: trade_no 撞了但 seq 是新的 —— 说明 QMT 用
|
|
新序号重推了一条我们已入过账的成交。这一笔不能再入账, 但**这个 seq 仍必须占住一行**,
|
|
否则连续水位永远卡在它前面, ack_seq 再也推不动, QMT 那边的消息也就永远清理不掉。
|
|
所以这里补插一行 dedup_key=NULL、processed=2 的存档行。
|
|
"""
|
|
now = _NOW()
|
|
p = {"s": int(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}
|
|
sql = ("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")
|
|
if execute(sql, p):
|
|
return PUT_NEW
|
|
if fetch_one("SELECT seq FROM pms_qmt_inbox WHERE seq = :s", {"s": int(seq)}):
|
|
return PUT_DUP_SEQ
|
|
# 见 docstring: trade_no 重复但 seq 是新的 —— 占位存档, 让水位能继续往前推
|
|
p["dk"] = None
|
|
p["pc"] = 2
|
|
p["pa"] = now
|
|
p["nt"] = f"重复成交 (dedup_key={dedup_key}), 不入账, 仅占位以推进 seq 水位"
|
|
execute(sql, p)
|
|
return PUT_DUP_KEY
|
|
|
|
|
|
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}
|