diff --git a/CHANGELOG.md b/CHANGELOG.md index 143db25..1d53b63 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,11 @@ ### 新增 - **Playwright E2E 前端测试基建**(升级计划 P4-1)——web-ui 引入 `@playwright/test`(`e2e/` + `playwright.config.ts`,`npm run test:e2e`)。**mock 方案选后端合成数据而非 page.route 拦截**:`EASY_TDX_E2E_MOCK=1` 时 serve 的 lifespan 把 TDX/MAC 客户端替换为合成数据客户端(`web/e2e_mock.py`,按 (market, code) CRC32 播种的确定性随机游走,分页语义与真实 /bars 一致),回测/WF/一条龙评估/自选/策略库继续走**真实后端代码**(它们本就不依赖行情连接),SSE 由 QuoteStreamer 真轮询合成数据全链路覆盖(mock 模式下轮询降到 2s 一拍,不受交易时段限制)。用例覆盖:看板五大指数区块+SSE 价格渲染、自选增删、回测全流程(净值图/绩效表/成交记录)、「附加分析」开关(WF 逐窗柱状图+一条龙评估卡)、策略库保存;`EASY_TDX_CONFIG_DIR` 指向每轮独立临时目录(断言可写死、不污染真实 `~/.easy_tdx`)。CI frontend job 追加 E2E 步骤;`verify_ci.sh` 补 `--no-frontend` 与前端 typecheck+build+E2E 段。新增 `tests/unit/test_e2e_mock.py`(11 例)守护 mock 与真实客户端的契约。 +- **WebSocket 实时推送联动 EventBus**(升级计划 P4-2)——`/ws/realtime/{symbol}` 从「不推送数据」变为真链路:新增 `web/realtime_hub.py`(RealtimeStreamHub),订阅集合变化时按需启停 `RealtimeDataFeed`(轮询 `get_stock_quotes` → `EventBus` → 每连接独立队列 fan-out,丢最旧保最新);**无人订阅完全停止轮询**(对齐 QuoteStreamer 节能语义);去重后标的上限 80;推送帧 `{type:"tick", symbol, market, code, price, volume, ts, open, high, low, pre_close, amount, name}`,30s 空闲 `ping` 心跳,客户端可 `subscribe`/`unsubscribe` 动态增删。端点重写为「单一写者泵」模型(全部出站帧经队列串行,杜绝并发 send 交错)。**前端接入选择只写文档不上组件**:看板/自选实时刷新已由 SSE `/stream/quotes`(全量快照、单连接共享)承担,WS 定位是按需单标的 tick(实时策略信号预留口),双通道同时拉同样行情属冗余——协议 + 自动重连/心跳容忍代码骨架落 `docs/api_reference.md` 与 README(「未联动」警示已撤)。新增 `scripts/ws_smoke.py` 手动冒烟(mock 模式随时可跑,实测可见 tick 帧与动态订阅确认)。环境变量 `EASY_TDX_WS_INTERVAL` 可调轮询间隔。 + +### 修复 + +- `RealtimeDataFeed` stop-before-start 竞态:`run_async`/`run_sync` 首行会把 `_running` 重置为 True,若 `stop()` 在任务首次调度前调用,停止请求被覆盖、任务永不退出(RealtimeStreamHub 换标的重建 feed 时必现死锁)。引入独立 `_stop_requested` 标志,启动前已请求停止则直接返回;`tests/unit/test_realtime_feed.py` 补 2 个回归用例。 ## [1.27.0] — 2026-09-01 diff --git a/README.md b/README.md index 107895d..a625193 100644 --- a/README.md +++ b/README.md @@ -1170,23 +1170,39 @@ curl -X POST "http://localhost:8000/api/v1/chanlun/analyze" \ ### WebSocket 实时行情 -> ⚠️ **当前未联动数据源**:`create_app()` 尚未创建 `EventBus`,此 WS 端点目前 -> 不会推送任何行情。计划在后续版本接入 `RealtimeDataFeed` 后打通。 -> 如需实时行情,请先用上文「实时行情轮询」的编程 API。 +`/api/v1/ws/realtime/{symbol}`(v1.28 起接通数据源):连接即订阅指定标的,服务端 +经 `RealtimeDataFeed`(按 `interval` 秒轮询五档快照 → `EventBus`)推送 tick 帧; +连接断开自动退订,无人订阅时完全停止轮询。盘外时段默认只睡不拉(交易时段过滤), +本地冒烟/演示可配合 `EASY_TDX_E2E_MOCK=1` 的合成行情随时验证(见 +`scripts/ws_smoke.py`)。 ```javascript -// JavaScript 示例 -const ws = new WebSocket("ws://localhost:8000/ws/realtime/SZ000001"); +const ws = new WebSocket("ws://localhost:8000/api/v1/ws/realtime/SZ000001"); ws.onmessage = (event) => { - const data = JSON.parse(event.data); - console.log(data); // {type: "tick", market: "SZ", code: "000001", price: 10.5, ...} + const frame = JSON.parse(event.data); + if (frame.type === "tick") { + // {type:"tick", symbol:"SZ000001", market:"SZ", code:"000001", + // price:10.5, volume:12345, ts:1760000000.0, + // open, high, low, pre_close, amount, name} + console.log(frame.symbol, frame.price, frame.ts); + } else if (frame.type === "ping") { + // 服务端 30s 空闲心跳,忽略即可(客户端无须回包) + } }; -// 动态订阅更多标的 +// 动态订阅更多标的(服务端回 {"type":"status","msg":"subscribed SH600000"}) ws.send(JSON.stringify({action: "subscribe", symbol: "SH600000"})); +// 退订 +ws.send(JSON.stringify({action: "unsubscribe", symbol: "SH600000"})); ``` +浏览器接入建议(自动重连 + 心跳容忍):`onclose` 后指数退避重连(参考 +`web-ui/src/stores/quotes.ts` 对 SSE 的同类处理);`{"type":"ping"}` 心跳帧直接 +忽略、不回包;连续 N 秒无任何帧(含 ping)再视为僵死连接主动重连。协议字段完整 +说明见 `docs/api_reference.md` 的 WebSocket 一节;单标的 WS 订阅与看板 SSE +(全量快照)并存不冲突,按需选用。 + ### API 文档 启动服务后访问: diff --git a/docs/api_reference.md b/docs/api_reference.md index e92c00c..2171671 100644 --- a/docs/api_reference.md +++ b/docs/api_reference.md @@ -639,3 +639,66 @@ compute_price_limits(market, code, name, pre_close, listed_days=None) | `KNOWN_HOSTS` | `list[str]` | A 股行情服务器列表 | | `KNOWN_EX_HOSTS` | `list[str]` | 扩展行情服务器列表 | | `XDXR_CATEGORY_NAMES` | `dict[int, str]` | 除权除息事件类型映射 | + +--- + +## WebSocket 实时行情(serve /ws/realtime/*) + +`easy-tdx serve` 后可建立 WebSocket 连接(v1.28 起联动 `RealtimeDataFeed`,此前 +该端点不推送数据): + +``` +ws://127.0.0.1:8000/api/v1/ws/realtime/{symbol} # symbol 如 SZ000001 / SH600519 +``` + +### 服务端推送帧(JSON) + +| type | 触发 | 字段 | +|------|------|------| +| `tick` | 轮询到标的的最新快照(价格/量变化才推,约 `interval` 秒一拍) | `symbol`、`market`、`code`、`price`、`volume`、`ts`(epoch 秒)、`open`、`high`、`low`、`pre_close`、`amount`、`name` | +| `ping` | 连续 30s 未收到客户端消息的心跳 | —(客户端忽略即可,无须回包) | +| `status` | 客户端 subscribe/unsubscribe 的确认 | `msg`(如 `subscribed SH600000`) | +| `error` | 非法 JSON / 未知 action / 超出订阅上限 | `msg` | + +### 客户端控制消息(JSON 文本帧) + +```json +{"action": "subscribe", "symbol": "SH600000"} +{"action": "unsubscribe", "symbol": "SH600000"} +``` + +### 行为约定 + +- **连接即订阅** path 上的 symbol;断开自动退订全部标的。 +- **按需轮询**:订阅集合为空时服务端不产生任何行情请求;去重后标的总数上限 + 80(`get_stock_quotes` 协议约束)。 +- **交易时段**:默认 A 股时段外只睡不拉(无 tick 帧,心跳照发);mock 模式 + (`EASY_TDX_E2E_MOCK=1`)不受限制。 +- **背压**:消费过慢时丢最旧快照保最新,不积压。 +- 环境变量:`EASY_TDX_WS_INTERVAL`(轮询间隔秒数,默认 3.0)。 + +### 前端接入方式(自动重连 + 心跳容忍) + +```typescript +function connectRealtime(symbol: string, onTick: (f: TickFrame) => void) { + let retry = 0 + let ws: WebSocket | null = null + const open = () => { + ws = new WebSocket(`ws://${location.host}/api/v1/ws/realtime/${symbol}`) + ws.onmessage = (e) => { + const frame = JSON.parse(e.data) + if (frame.type === 'tick') { retry = 0; onTick(frame) } // ping/status 忽略 + } + ws.onclose = () => { + retry += 1 + setTimeout(open, Math.min(1000 * 2 ** (retry - 1), 30_000)) // 指数退避 + } + } + open() + return () => ws?.close() +} +``` + +> 说明:看板/自选页的实时刷新已由 SSE `/stream/quotes`(全量快照、单连接共享) +> 承担;WS 通道定位是**按需订阅单标的 tick 事件**(后续实时策略信号的接入点), +> 两条链路按场景选用,不要求同时连接。手动冒烟见 `scripts/ws_smoke.py`。 diff --git a/scripts/ws_smoke.py b/scripts/ws_smoke.py new file mode 100644 index 0000000..ea54cbd --- /dev/null +++ b/scripts/ws_smoke.py @@ -0,0 +1,86 @@ +"""ws_smoke — WebSocket 实时行情推送手动冒烟脚本(v1.28)。 + +连 ``/api/v1/ws/realtime/{symbol}``,打印收到的每一帧,验证 +RealtimeDataFeed → EventBus → RealtimeStreamHub 的推送链路。 + +用法:: + + # 真实服务器(需 MAC 行情可达且在交易时段,或标的盘外有静止快照) + .venv/Scripts/python.exe scripts/ws_smoke.py --symbol SZ000001 + + # mock 模式(推荐本地冒烟:合成行情、不受交易时段限制) + # 终端 1: + EASY_TDX_E2E_MOCK=1 .venv/Scripts/python.exe -m easy_tdx serve --port 8000 --no-open-browser + # 终端 2: + .venv/Scripts/python.exe scripts/ws_smoke.py --url ws://127.0.0.1:8000/api/v1/ws/realtime/SZ000001 + +可选参数:--duration 秒数(默认 15)、--extra-symbol 运行中追加订阅的标的。 +退出码:收到至少一帧 tick = 0,否则 1。 +""" + +from __future__ import annotations + +import argparse +import asyncio +import json +import sys +import time + +import websockets + + +async def main() -> int: + parser = argparse.ArgumentParser(description="WebSocket 实时行情冒烟") + parser.add_argument( + "--url", + default="ws://127.0.0.1:8000/api/v1/ws/realtime/SZ000001", + help="WS 端点(默认本地 serve 的 SZ000001)", + ) + parser.add_argument("--duration", type=float, default=15.0, help="冒烟时长(秒)") + parser.add_argument( + "--extra-symbol", default="SH600519", help="运行中追加订阅的标的(演示动态订阅)" + ) + args = parser.parse_args() + + ticks = 0 + start = time.perf_counter() + + try: + async with websockets.connect(args.url) as ws: + print(f"已连接 {args.url},监听 {args.duration:.0f}s(Ctrl+C 提前退出)") + # 3 秒后演示动态订阅(应收到 status 确认 + 新标的 tick) + subscribe_at = start + 3.0 + + async def _maybe_subscribe() -> None: + if time.perf_counter() < subscribe_at: + await asyncio.sleep(subscribe_at - time.perf_counter()) + await ws.send(json.dumps({"action": "subscribe", "symbol": args.extra_symbol})) + print(f"→ 已发送动态订阅:{args.extra_symbol}") + + sub_task = asyncio.create_task(_maybe_subscribe()) + try: + while time.perf_counter() - start < args.duration: + raw = await asyncio.wait_for(ws.recv(), timeout=5.0) + frame = json.loads(raw) + if frame.get("type") == "tick": + ticks += 1 + print( + f"[tick] {frame['symbol']} 价={frame['price']:.2f} " + f"量={frame.get('volume', 0):.0f} 名={frame.get('name', '')} " + f"ts={frame['ts']:.0f}" + ) + else: + print(f"[{frame.get('type')}] {frame}") + finally: + sub_task.cancel() + await asyncio.gather(sub_task, return_exceptions=True) + except (OSError, TimeoutError) as exc: + print(f"连接失败:{exc}\n请确认 serve 已启动(easy-tdx serve)", file=sys.stderr) + return 1 + + print(f"—— 冒烟结束:共收到 {ticks} 帧 tick ——") + return 0 if ticks > 0 else 1 + + +if __name__ == "__main__": + sys.exit(asyncio.run(main())) diff --git a/src/easy_tdx/realtime/feed.py b/src/easy_tdx/realtime/feed.py index bfe9b30..72d8f17 100644 --- a/src/easy_tdx/realtime/feed.py +++ b/src/easy_tdx/realtime/feed.py @@ -186,6 +186,11 @@ class RealtimeDataFeed: self._fields = fields self._state = _FeedState() self._running = False + # stop() 在 run_async/run_sync 首次调度前调用的竞态防护: + # run_* 的首行会把 _running 重置为 True,若不另立 stop 标志, + # 「先 stop 后 start」的停止请求会被覆盖,任务永不退出 + # (RealtimeStreamHub 换标的重建 feed 时必现,v1.28 修复)。 + self._stop_requested = False @property def running(self) -> bool: @@ -206,6 +211,8 @@ class RealtimeDataFeed: max_iterations: 最多轮询多少轮(测试用);None 表示无限循环直到 :meth:`stop`。 """ + if self._stop_requested: + return # 启动前已请求停止(见 __init__ 的竞态说明) self._running = True try: count = 0 @@ -254,6 +261,8 @@ class RealtimeDataFeed: async def _run_sync_loop(self, client: Any, max_iterations: int | None) -> None: """同步客户端的轮询循环:阻塞调用丢到 executor。""" + if self._stop_requested: + return # 启动前已请求停止(见 __init__ 的竞态说明) self._running = True try: count = 0 @@ -267,7 +276,12 @@ class RealtimeDataFeed: self._running = False def stop(self) -> None: - """请求停止轮询(下一轮 sleep 结束后生效)。""" + """请求停止轮询(下一轮 sleep 结束后生效)。 + + 在 ``run_async`` / ``run_sync`` 首次获得调度之前调用同样有效 + (启动即退出),见 ``__init__`` 的竞态说明。 + """ + self._stop_requested = True self._running = False # ------------------------------------------------------------------ # diff --git a/src/easy_tdx/web/app.py b/src/easy_tdx/web/app.py index a717696..0a89219 100644 --- a/src/easy_tdx/web/app.py +++ b/src/easy_tdx/web/app.py @@ -141,6 +141,23 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: mac_client = None app.state.mac_client = mac_client + # --- WebSocket 实时推送枢纽(/ws/realtime/*,按需轮询 + fan-out) --- + # 与 MAC 客户端共用一条连接;无订阅时不轮询(见 realtime_hub.py)。 + # mock 模式关闭交易时段过滤并支持更密轮询,供 E2E / 冒烟脚本随时可推。 + app.state.realtime_hub = None + if mac_client is not None: + try: + from easy_tdx.web.realtime_hub import RealtimeStreamHub + + app.state.realtime_hub = RealtimeStreamHub( + mac_client, + interval=float(os.environ.get("EASY_TDX_WS_INTERVAL", "3.0")), + sessions=() if mock_mode else None, + ) + logger.info("RealtimeStreamHub 已挂载(/ws/realtime/* 就绪)") + except Exception: + logger.warning("RealtimeStreamHub 挂载失败 — WS 实时推送不可用", exc_info=True) + # --- 扩展市场客户端(可选) --- ex_client = None enable_ex = getattr(app.state, "enable_ex", False) @@ -166,6 +183,14 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: except Exception: logger.warning("QuoteStreamer stop failed", exc_info=True) + # --- 关闭 WebSocket 实时推送枢纽(停止按需轮询任务) --- + rt_hub = getattr(app.state, "realtime_hub", None) + if rt_hub is not None: + try: + await rt_hub.shutdown() + except Exception: + logger.warning("RealtimeStreamHub shutdown failed", exc_info=True) + # --- 依次关闭 --- for name, cli in [ ("Ex market client", ex_client), diff --git a/src/easy_tdx/web/e2e_mock.py b/src/easy_tdx/web/e2e_mock.py index 89b14ec..11c3a32 100644 --- a/src/easy_tdx/web/e2e_mock.py +++ b/src/easy_tdx/web/e2e_mock.py @@ -27,6 +27,7 @@ from __future__ import annotations import logging +import time import zlib from datetime import datetime, timedelta from typing import Any @@ -277,6 +278,36 @@ class MockMacClient: async def close(self) -> None: """lifespan 关闭时调用(无真实连接,空操作)。""" + async def get_stock_quotes( + self, stocks: list[tuple[int, str]], fields: Any = None + ) -> pd.DataFrame: + """批量五档快照(RealtimeStreamHub → RealtimeDataFeed 的轮询入口)。 + + 与 feed._row_to_event 的取数口径对齐:最新价落在 close 列、 + market 为 int;价格在合成序列上随每次调用轻微漂移(绕过 feed 去重, + 让 WS/E2E 在盘外也能持续看到推送帧)。 + """ + rows = [] + for market, code in stocks: + closes = _synth_closes(market_enum_key(_int_market_to_enum(int(market))), code, 30) + jitter = 1.0 + ((time.time() % 60.0) / 60.0 - 0.5) * 0.002 + price = float(closes[-1]) * jitter + rows.append( + { + "market": int(market), + "code": code, + "close": price, + "vol": 100_000.0 + (time.time() % 60.0) * 10.0, + "open": float(closes[-1]), + "high": price * 1.005, + "low": price * 0.995, + "pre_close": float(closes[-2]), + "amount": price * 100_000.0, + "name": _display_name(market, code), + } + ) + return pd.DataFrame(rows) + async def get_stock_kline( self, market: Any, diff --git a/src/easy_tdx/web/realtime_hub.py b/src/easy_tdx/web/realtime_hub.py new file mode 100644 index 0000000..51651f8 --- /dev/null +++ b/src/easy_tdx/web/realtime_hub.py @@ -0,0 +1,247 @@ +"""WebSocket 实时推送枢纽:RealtimeDataFeed(轮询)→ EventBus → 每连接队列 fan-out。 + +架构对齐 :class:`easy_tdx.web.quote_streamer.QuoteStreamer`(SSE 推送器)的 +生命周期与节能语义,差异在数据通路: + +- SSE ``/stream/quotes``:固定集合(指数 + 全部自选)全量快照,所有页面共享; +- WS ``/ws/realtime/{symbol}``:**按需订阅**的逐标的 tick 事件,走 + :class:`~easy_tdx.realtime.feed.RealtimeDataFeed`(轮询 ``get_stock_quotes`` + 五档快照 → :class:`~easy_tdx.realtime.engine.EventBus` 发布 MarketEvent)。 + +节能语义: + +- 订阅集合为空时不轮询(feed 任务不启动 / 已启动则停止)——与 QuoteStreamer + 「无人订阅时循环休眠」一致; +- 轮询集合 = 各连接订阅的去重并集(引用计数);集合变化时重启 feed + (feed 的 symbols 在构造时固定,重启代价 ≤ 一个 sleep 步长 0.5s); +- 盘外时段由 RealtimeDataFeed 自带的交易时段过滤兜底(只睡不拉)。 + +背压:队列满(连接消费慢)时丢最旧保最新——行情快照旧帧无价值。 +""" + +from __future__ import annotations + +import asyncio +import itertools +import logging +import time +from typing import Any + +from easy_tdx.realtime.engine import EventBus, MarketEvent +from easy_tdx.realtime.feed import RealtimeDataFeed + +logger = logging.getLogger(__name__) + +__all__ = ["RealtimeStreamHub", "parse_realtime_symbol", "MAX_WS_SYMBOLS"] + +#: 单 hub 订阅标的上限(get_stock_quotes 单次 80 只,协议约束)。 +MAX_WS_SYMBOLS = 80 + +_MARKET_STR_TO_INT: dict[str, int] = {"SZ": 0, "SH": 1, "BJ": 2} + + +def parse_realtime_symbol(symbol: str) -> tuple[int, str]: + """``"SZ000001"`` → ``(0, "000001")``(MAC 客户端的 int 市场约定)。 + + Raises: + ValueError: 格式非法(非 市场前缀 + 6 位数字)。 + """ + s = symbol.strip().upper() + if len(s) != 8 or s[:2] not in _MARKET_STR_TO_INT or not s[2:].isdigit(): + raise ValueError( + f"标的格式应为 市场前缀+6位代码(如 SZ000001 / SH600519),得到 '{symbol}'" + ) + return _MARKET_STR_TO_INT[s[:2]], s[2:] + + +def _event_to_frame(event: MarketEvent) -> dict[str, Any]: + """MarketEvent → 前端推送帧 ``{type, symbol, price, ts, ...}``。""" + symbol = f"{event.market}{event.code}" + return { + "type": "tick", + "symbol": symbol, + "market": event.market, + "code": event.code, + "price": event.price, + "volume": event.volume, + "ts": event.timestamp if event.timestamp > 0 else time.time(), + **event.data, # open/high/low/pre_close/amount/name 等(feed 附加字段) + } + + +class RealtimeStreamHub: + """按需轮询 + 每连接独立队列 fan-out。由 FastAPI lifespan 挂载到 + ``app.state.realtime_hub``,``/ws/realtime/*`` 端点调用。 + + Args: + quote_client: 拥有 ``async def get_stock_quotes(stocks, fields=None)`` + 的客户端(如 :class:`~easy_tdx.mac.client.AsyncMacClient`)。 + stocks 为 ``[(market_int, code), ...]``,返回列含 + market/code/close/vol/open/high/low/pre_close/amount/name。 + interval: 轮询间隔(秒),透传给 RealtimeDataFeed。 + sessions: 交易时段过滤,None = feed 默认(A 股时段);``()`` = 全天轮询 + (演示/E2E 用,配合 mock 数据不受交易时段限制)。 + queue_size: 每连接队列容量(满则丢最旧保最新)。 + """ + + def __init__( + self, + quote_client: Any, + *, + interval: float = 3.0, + sessions: tuple[tuple[int, int], ...] | None = None, + queue_size: int = 8, + ) -> None: + self._client = quote_client + self._interval = interval + self._sessions = sessions + self._queue_size = queue_size + self._bus = EventBus() + self._bus.subscribe_all(self._on_event) + self._queues: dict[int, asyncio.Queue[dict[str, Any]]] = {} + self._subs: dict[int, set[str]] = {} + self._refcount: dict[str, int] = {} + self._ids = itertools.count(1) + self._feed: RealtimeDataFeed | None = None + self._feed_task: asyncio.Task[None] | None = None + self._polling: list[tuple[int, str]] = [] # 当前 feed 正在轮询的集合 + self._lock = asyncio.Lock() + + # ── 订阅管理(WS 端点调用) ────────────────────────────────────────── + + def connect(self) -> tuple[int, asyncio.Queue[dict[str, Any]]]: + """注册一个连接;返回 (cid, queue)。queue 是该连接唯一的推送出口。""" + cid = next(self._ids) + q: asyncio.Queue[dict[str, Any]] = asyncio.Queue(maxsize=self._queue_size) + self._queues[cid] = q + self._subs[cid] = set() + return cid, q + + async def subscribe(self, cid: int, symbol: str) -> None: + """连接 cid 订阅 symbol(引用计数 +1;新标的进入轮询集合)。 + + Raises: + ValueError: symbol 格式非法,或去重后标的数超过 :data:`MAX_WS_SYMBOLS`。 + """ + parse_realtime_symbol(symbol) # 先做格式校验(不持锁) + key = symbol.strip().upper() + async with self._lock: + if key in self._subs[cid]: + return # 幂等:同连接重复订阅同一标的 + if self._refcount.get(key, 0) == 0 and len(self._active_symbols()) >= MAX_WS_SYMBOLS: + raise ValueError( + f"订阅标的总数已达上限 {MAX_WS_SYMBOLS}(get_stock_quotes 单次上限)" + ) + self._subs[cid].add(key) + self._refcount[key] = self._refcount.get(key, 0) + 1 + await self._sync_feed() + + async def unsubscribe(self, cid: int, symbol: str) -> None: + """连接 cid 退订 symbol(引用计数 -1;归零则移出轮询集合)。""" + key = symbol.strip().upper() + async with self._lock: + if key not in self._subs[cid]: + return + self._subs[cid].discard(key) + self._dec_ref(key) + await self._sync_feed() + + async def disconnect(self, cid: int) -> None: + """连接断开:退订其全部标的并同步轮询集合。""" + async with self._lock: + for key in self._subs.pop(cid, set()): + self._dec_ref(key) + self._queues.pop(cid, None) + await self._sync_feed() + + # ── 生命周期 ────────────────────────────────────────────────────────── + + async def shutdown(self) -> None: + """lifespan 关闭时调用:停止轮询任务(feed 正常退出,非 cancel)。""" + async with self._lock: + self._refcount.clear() + await self._stop_feed() + + @property + def poll_symbols(self) -> list[str]: + """当前去重后的订阅集合(诊断/测试用)。""" + return self._active_symbols() + + @property + def subscriber_count(self) -> int: + """当前连接数。""" + return len(self._queues) + + @property + def polling(self) -> bool: + """是否正在轮询(有订阅时 True)。""" + return self._feed_task is not None and not self._feed_task.done() + + # ── 内部 ────────────────────────────────────────────────────────────── + + def _active_symbols(self) -> list[str]: + """引用计数 > 0 的标的(排序保证 feed 重启的参数确定性)。""" + return sorted(s for s, c in self._refcount.items() if c > 0) + + def _dec_ref(self, key: str) -> None: + count = self._refcount.get(key, 0) - 1 + if count > 0: + self._refcount[key] = count + else: + self._refcount.pop(key, None) + + async def _sync_feed(self) -> None: + """把轮询集合同步成订阅并集(持锁调用)。 + + 集合未变直接返回;变化则停旧 feed、按新集合启动(空集合 = 完全停止)。 + """ + desired = [parse_realtime_symbol(s) for s in self._active_symbols()] + if desired == self._polling: + return + await self._stop_feed() + if desired: + self._feed = RealtimeDataFeed( + self._bus, desired, interval=self._interval, sessions=self._sessions + ) + self._feed_task = asyncio.get_running_loop().create_task( + self._feed.run_async(self._client) + ) + self._polling = desired + logger.info("RealtimeStreamHub 开始轮询 %d 只标的", len(desired)) + else: + logger.info("RealtimeStreamHub 无订阅,停止轮询") + + async def _stop_feed(self) -> None: + """优雅停止当前 feed 任务(持锁调用)。""" + if self._feed_task is None: + return + if self._feed is not None: + self._feed.stop() # 下一 sleep 步长(≤0.5s)内退出 + try: + await self._feed_task + except Exception: + logger.warning("RealtimeDataFeed 任务异常退出", exc_info=True) + self._feed_task = None + self._feed = None + self._polling = [] + + def _on_event(self, event: MarketEvent) -> None: + """EventBus 回调:按各连接的订阅集合 fan-out(丢最旧保最新)。""" + if not self._queues: + return + key = f"{event.market}{event.code}" + frame = _event_to_frame(event) + for cid, q in list(self._queues.items()): + if key not in self._subs.get(cid, set()): + continue + try: + q.put_nowait(frame) + except asyncio.QueueFull: + try: + q.get_nowait() + except asyncio.QueueEmpty: + pass + try: + q.put_nowait(frame) + except asyncio.QueueFull: + pass diff --git a/src/easy_tdx/web/routers/realtime.py b/src/easy_tdx/web/routers/realtime.py index 2e519d1..25bc4dd 100644 --- a/src/easy_tdx/web/routers/realtime.py +++ b/src/easy_tdx/web/routers/realtime.py @@ -1,4 +1,26 @@ -"""实时数据 WebSocket 路由。""" +"""实时数据 WebSocket 路由(v1.28 起联动 RealtimeDataFeed)。 + +``GET /api/v1/ws/realtime/{symbol}``(WebSocket): + +- 连接即订阅 path 上的 symbol(如 ``SZ000001``),服务端开始按需轮询并推送 + tick 帧;连接断开自动退订,无人订阅时完全停止轮询(节能语义与 + ``/stream/quotes`` 的 QuoteStreamer 一致,见 :mod:`easy_tdx.web.realtime_hub`)。 +- 服务端推送帧格式(JSON):: + + {"type": "tick", "symbol": "SZ000001", "market": "SZ", "code": "000001", + "price": 10.5, "volume": 12345.0, "ts": 1760000000.0, + "open": ..., "high": ..., "low": ..., "pre_close": ..., "amount": ..., "name": "平安银行"} + + 空闲时每 30s 一条 ``{"type": "ping"}`` 心跳;连接级错误发 + ``{"type": "error", "msg": ...}``。 +- 客户端控制消息(JSON 文本帧):: + + {"action": "subscribe", "symbol": "SH600000"} + {"action": "unsubscribe", "symbol": "SH600000"} + + 服务端回 ``{"type": "status", "msg": "subscribed SH600000"}`` 确认; + 去重后标的总数上限 80(get_stock_quotes 协议约束),超限回错误帧。 +""" from __future__ import annotations @@ -9,95 +31,92 @@ from typing import Any from fastapi import APIRouter, WebSocket, WebSocketDisconnect +from easy_tdx.web.realtime_hub import RealtimeStreamHub + logger = logging.getLogger(__name__) router = APIRouter(tags=["realtime"]) +#: 空闲心跳间隔(秒)。receive 超时未收到客户端消息即发一条 ping。 +_HEARTBEAT_SECONDS = 30.0 + + +async def _pump(websocket: WebSocket, queue: asyncio.Queue[dict[str, Any]]) -> None: + """唯一的发送协程:把 hub 队列里的帧逐条写给客户端。 + + 单一写者模型:tick / status / error / ping 全部经队列串行发出, + 避免多协程并发 send_json 的帧交错风险。 + """ + while True: + frame = await queue.get() + await websocket.send_json(frame) + @router.websocket("/ws/realtime/{symbol}") async def realtime_websocket(websocket: WebSocket, symbol: str) -> None: - """WebSocket 实时行情订阅。 - - 连接后自动订阅指定标的的实时事件。 - symbol 格式: SZ000001, SH600000 等。 - - 客户端可发送 JSON 消息来控制订阅: - - {"action": "subscribe", "symbol": "SZ000001"} - - {"action": "unsubscribe", "symbol": "SZ000001"} - - 服务端推送消息格式: - - {"type": "tick", "market": "SZ", "code": "000001", "price": 10.5, ...} - - {"type": "signal", "direction": "BUY", ...} - """ + """WebSocket 实时行情订阅(协议见模块 docstring)。""" await websocket.accept() - logger.info("WebSocket client connected for symbol: %s", symbol) + hub: RealtimeStreamHub | None = getattr(websocket.app.state, "realtime_hub", None) + if hub is None: + await websocket.send_json( + {"type": "error", "msg": "实时数据源不可用(MAC 行情客户端未连接)"} + ) + await websocket.close() + return - # Try to get EventBus from app state - event_bus = getattr(websocket.app.state, "event_bus", None) + cid, queue = hub.connect() + logger.info("WS realtime 连接:symbol=%s cid=%s", symbol, cid) - subscribed_symbols: set[str] = {symbol.upper()} - - async def _on_event(event: Any) -> None: - """EventBus 回调 → 推送 WebSocket 消息。""" - event_symbol = f"{event.market}{event.code}" - if event_symbol in subscribed_symbols: - try: - msg = { - "type": event.event_type.value, - "market": event.market, - "code": event.code, - "price": event.price, - "volume": event.volume, - "timestamp": event.timestamp, - "data": event.data, - } - await websocket.send_json(msg) - except Exception: - logger.warning("Failed to send WebSocket message") - - # Subscribe to event bus if available - if event_bus is not None: - event_bus.subscribe_all(_on_event) + def _enqueue(frame: dict[str, Any]) -> None: + """控制帧也走队列(容量满时静默丢弃,不影响行情流)。""" + try: + queue.put_nowait(frame) + except asyncio.QueueFull: + pass + pump_task = asyncio.create_task(_pump(websocket, queue)) try: + try: + await hub.subscribe(cid, symbol) + except ValueError as exc: + _enqueue({"type": "error", "msg": str(exc)}) + await websocket.close() + return + while True: - # Receive client messages (subscribe/unsubscribe control) try: - raw = await asyncio.wait_for(websocket.receive_text(), timeout=30.0) - data = json.loads(raw) - action = data.get("action", "") - - if action == "subscribe": - new_symbol = data.get("symbol", "").upper() - if new_symbol: - subscribed_symbols.add(new_symbol) - await websocket.send_json( - {"type": "status", "msg": f"subscribed {new_symbol}"} - ) - - elif action == "unsubscribe": - old_symbol = data.get("symbol", "").upper() - subscribed_symbols.discard(old_symbol) - await websocket.send_json( - {"type": "status", "msg": f"unsubscribed {old_symbol}"} - ) - + raw = await asyncio.wait_for(websocket.receive_text(), timeout=_HEARTBEAT_SECONDS) except asyncio.TimeoutError: - # Send heartbeat ping - try: - await websocket.send_json({"type": "ping"}) - except Exception: - break - except WebSocketDisconnect: - break + _enqueue({"type": "ping"}) + continue + + try: + data = json.loads(raw) except json.JSONDecodeError: - await websocket.send_json({"type": "error", "msg": "invalid JSON"}) + _enqueue({"type": "error", "msg": "invalid JSON"}) + continue + + action = str(data.get("action", "")) + target = str(data.get("symbol", "")) + if action == "subscribe" and target: + try: + await hub.subscribe(cid, target) + except ValueError as exc: + _enqueue({"type": "error", "msg": str(exc)}) + continue + _enqueue({"type": "status", "msg": f"subscribed {target.upper()}"}) + elif action == "unsubscribe" and target: + await hub.unsubscribe(cid, target) + _enqueue({"type": "status", "msg": f"unsubscribed {target.upper()}"}) + else: + _enqueue({"type": "error", "msg": f"unknown action: {raw[:80]}"}) except WebSocketDisconnect: - logger.info("WebSocket client disconnected: %s", symbol) + pass except Exception: - logger.exception("WebSocket error for %s", symbol) + logger.exception("WS realtime 连接异常:symbol=%s", symbol) finally: - if event_bus is not None: - event_bus.unsubscribe(symbol.upper(), _on_event) - logger.info("WebSocket connection closed: %s", symbol) + pump_task.cancel() + await asyncio.gather(pump_task, return_exceptions=True) + await hub.disconnect(cid) + logger.info("WS realtime 连接关闭:symbol=%s cid=%s", symbol, cid) diff --git a/tests/unit/test_realtime_feed.py b/tests/unit/test_realtime_feed.py index 9b76d5a..7cdeb5e 100644 --- a/tests/unit/test_realtime_feed.py +++ b/tests/unit/test_realtime_feed.py @@ -306,3 +306,29 @@ class TestStopFlag: await asyncio.gather(feed.run_async(client), _stop_soon()) assert feed.running is False + + async def test_stop_before_start_exits_immediately(self) -> None: + """回归(v1.28):stop 在 run_async 首次调度前调用也必须生效。 + + run_async 首行会把 _running 重置为 True,若无独立停止标志, + RealtimeStreamHub 这类「先 stop 再等任务退出」的用法会永久挂起。 + """ + bus = EventBus() + client = AsyncMockClient([_sample_quotes_df()]) + feed = RealtimeDataFeed(bus=bus, symbols=[(0, "000001")], sessions=(), interval=0.2) + feed.stop() + + task = asyncio.get_running_loop().create_task(feed.run_async(client)) + await asyncio.wait_for(task, timeout=2.0) # 启动即退出,不发轮询请求 + assert client.calls == [] + assert feed.running is False + + async def test_stop_before_start_sync_loop(self) -> None: + """run_sync 路径同样的竞态防护。""" + bus = EventBus() + client = SyncMockClient([_sample_quotes_df()]) + feed = RealtimeDataFeed(bus=bus, symbols=[(0, "000001")], sessions=(), interval=0.2) + feed.stop() + + await asyncio.wait_for(feed._run_sync_loop(client, None), timeout=2.0) + assert client.calls == [] diff --git a/tests/unit/test_realtime_ws.py b/tests/unit/test_realtime_ws.py new file mode 100644 index 0000000..d1d0960 --- /dev/null +++ b/tests/unit/test_realtime_ws.py @@ -0,0 +1,284 @@ +"""RealtimeStreamHub(WebSocket 实时推送枢纽)单元测试。 + +全程用假行情客户端(不打真实网络、不受交易时段限制——hub 传 sessions=()), +覆盖:订阅触发轮询、多客户端并发订阅同/不同标的、退订引用计数与竞态、 +无人订阅完全停止轮询、背压丢旧、以及 WS 端点端到端帧格式。 +""" + +from __future__ import annotations + +import asyncio +from typing import Any + +import pandas as pd +import pytest + +from easy_tdx.web.realtime_hub import ( + MAX_WS_SYMBOLS, + RealtimeStreamHub, + parse_realtime_symbol, +) + +pytest.importorskip("fastapi") + + +class FakeQuoteClient: + """假 MAC 行情客户端:每次调用价格/量递增(绕过 feed 的 (price, vol) 去重)。""" + + def __init__(self) -> None: + self.calls = 0 + self.requested: list[list[tuple[int, str]]] = [] + self._seq = 0 + + async def get_stock_quotes( + self, stocks: list[tuple[int, str]], fields: Any = None + ) -> pd.DataFrame: + self.calls += 1 + self.requested.append(list(stocks)) + self._seq += 1 + rows = [] + for market, code in stocks: + base = 10.0 + (int(code) % 100) * 0.1 + rows.append( + { + "market": int(market), + "code": code, + "close": base + self._seq * 0.01, + "vol": 1000.0 + self._seq, + "open": base, + "high": base + 1.0, + "low": base - 1.0, + "pre_close": base, + "amount": (base + self._seq * 0.01) * 1000.0, + "name": f"股票{code}", + } + ) + return pd.DataFrame(rows) + + +def _make_hub(**kwargs: Any) -> tuple[RealtimeStreamHub, FakeQuoteClient]: + client = FakeQuoteClient() + hub = RealtimeStreamHub( + client, + interval=0.1, # feed 内部 clamp 下限 + sessions=(), # 关闭交易时段过滤,测试不受运行时刻影响 + **kwargs, + ) + return hub, client + + +async def _next_tick(queue: asyncio.Queue[dict[str, Any]], timeout: float = 5.0) -> dict[str, Any]: + """取队列里的下一帧(跳过 status/ping 等控制帧)。""" + deadline = asyncio.get_event_loop().time() + timeout + while True: + remaining = deadline - asyncio.get_event_loop().time() + assert remaining > 0, "等待 tick 帧超时" + frame = await asyncio.wait_for(queue.get(), timeout=remaining) + if frame.get("type") == "tick": + return frame + + +# ── 符号解析 ───────────────────────────────────────────────────────────────── + + +def test_parse_symbol() -> None: + assert parse_realtime_symbol("SZ000001") == (0, "000001") + assert parse_realtime_symbol("sh600519") == (1, "600519") + assert parse_realtime_symbol("BJ920002") == (2, "920002") + for bad in ("600519", "XX000001", "SZ00001a", "SZ0000001", ""): + with pytest.raises(ValueError, match="标的格式"): + parse_realtime_symbol(bad) + + +# ── 订阅 / 轮询生命周期 ────────────────────────────────────────────────────── + + +async def test_subscribe_starts_polling_and_pushes_frames() -> None: + """订阅 → feed 开始轮询 → 队列收到 {symbol, price, ts} tick 帧。""" + hub, client = _make_hub() + assert not hub.polling # 无人订阅不轮询 + + cid, queue = hub.connect() + await hub.subscribe(cid, "SZ000001") + assert hub.polling + assert hub.poll_symbols == ["SZ000001"] + + frame = await _next_tick(queue) + assert frame["type"] == "tick" + assert frame["symbol"] == "SZ000001" + assert frame["market"] == "SZ" + assert frame["code"] == "000001" + assert frame["price"] > 0 + assert frame["ts"] > 0 + assert frame["name"] == "股票000001" + + # 第二帧(假客户端价格递增,绕过去重) + frame2 = await _next_tick(queue) + assert frame2["price"] != frame["price"] + assert client.calls >= 2 + + await hub.shutdown() + + +async def test_unsubscribe_last_client_stops_polling() -> None: + """退订到无人订阅 → 轮询完全停止(节能语义)。""" + hub, client = _make_hub() + cid, queue = hub.connect() + await hub.subscribe(cid, "SH600519") + await _next_tick(queue) + calls_before = client.calls + + await hub.unsubscribe(cid, "SH600519") + await asyncio.sleep(0.35) # 越过若干个 interval,确认不再发起新轮询 + assert not hub.polling + assert hub.poll_symbols == [] + assert client.calls == calls_before + + await hub.shutdown() + + +# ── 多客户端并发 ───────────────────────────────────────────────────────────── + + +async def test_multiple_clients_same_symbol_all_receive() -> None: + """两个客户端订阅同一标的:都收到推送;退订一个后轮询继续(引用计数)。""" + hub, client = _make_hub() + cid1, q1 = hub.connect() + cid2, q2 = hub.connect() + await hub.subscribe(cid1, "SZ000001") + await hub.subscribe(cid2, "sz000001") # 大小写归一到同一标的 + + f1 = await _next_tick(q1) + f2 = await _next_tick(q2) + assert f1["symbol"] == f2["symbol"] == "SZ000001" + # 去重后只轮询一份 + assert hub.poll_symbols == ["SZ000001"] + + # 客户端 1 断开:客户端 2 仍在订阅,轮询不停止 + await hub.disconnect(cid1) + await asyncio.sleep(0.2) + assert hub.polling + await _next_tick(q2) + + await hub.shutdown() + + +async def test_clients_different_symbols_no_cross_delivery() -> None: + """不同客户端订阅不同标的:各自只收到自己标的的帧。""" + hub, _client = _make_hub() + cid_a, q_a = hub.connect() + cid_b, q_b = hub.connect() + await hub.subscribe(cid_a, "SZ000001") + await hub.subscribe(cid_b, "SH600519") + assert hub.poll_symbols == ["SH600519", "SZ000001"] # 排序确定 + + for _ in range(3): # 多收几帧确认无串扰 + fa = await _next_tick(q_a) + fb = await _next_tick(q_b) + assert fa["symbol"] == "SZ000001" + assert fb["symbol"] == "SH600519" + + await hub.shutdown() + + +async def test_dynamic_subscribe_extends_poll_set() -> None: + """运行中追加订阅:新标的进入轮询集合并开始推送。""" + hub, client = _make_hub() + cid, queue = hub.connect() + await hub.subscribe(cid, "SZ000001") + await _next_tick(queue) + + await hub.subscribe(cid, "SH600519") + assert hub.poll_symbols == ["SH600519", "SZ000001"] + seen: set[str] = set() + for _ in range(8): + frame = await _next_tick(queue) + seen.add(frame["symbol"]) + if seen == {"SZ000001", "SH600519"}: + break + assert seen == {"SZ000001", "SH600519"} + # 每轮轮询带全量订阅集合(requested 记录 (market_int, code) 元组) + requested_keys = {f"SZ{c}" if m == 0 else f"SH{c}" for m, c in client.requested[-1]} + assert {"SZ000001", "SH600519"}.issubset(requested_keys) + + await hub.shutdown() + + +# ── 竞态 ───────────────────────────────────────────────────────────────────── + + +async def test_concurrent_subscribe_unsubscribe_race() -> None: + """多客户端并发订阅/退订同一批标的:最终引用计数与轮询状态一致。""" + hub, _client = _make_hub() + cids = [hub.connect()[0] for _ in range(8)] + symbols = ["SZ000001", "SH600519", "SZ399006"] + + # 并发混订:偶数客户端先订后退,奇数只订 + await asyncio.gather( + *[hub.subscribe(cid, sym) for cid in cids for sym in symbols], + *[hub.unsubscribe(cid, sym) for cid in cids[::2] for sym in symbols], + ) + # 奇数客户端(索引 1,3,5,7)仍各持有 3 个标的的订阅 + assert hub.poll_symbols == sorted(symbols) + assert hub.subscriber_count == 8 + + # 全部断开(并发)→ 轮询停止 + await asyncio.gather(*[hub.disconnect(cid) for cid in cids]) + await asyncio.sleep(0.2) + assert not hub.polling + assert hub.poll_symbols == [] + + await hub.shutdown() + + +async def test_symbol_cap_rejected() -> None: + """去重后标的数超 80(协议上限)拒绝并抛 ValueError。""" + hub, _client = _make_hub() + cid, _queue = hub.connect() + for i in range(MAX_WS_SYMBOLS): + await hub.subscribe(cid, f"SH{i:06d}") + with pytest.raises(ValueError, match="上限"): + await hub.subscribe(cid, "SZ000001") + await hub.shutdown() + + +async def test_backpressure_drops_oldest() -> None: + """队列满时丢最旧保最新(行情快照旧帧无价值)。""" + hub, _client = _make_hub(queue_size=2) + cid, queue = hub.connect() + await hub.subscribe(cid, "SZ000001") + + # 不消费,等队列塞满并溢出 + await asyncio.sleep(0.8) + assert queue.qsize() <= 2 + prices = [queue.get_nowait()["price"] for _ in range(queue.qsize())] + # 留下的是最新帧(假客户端价格单调递增) + assert prices == sorted(prices) + + await hub.shutdown() + + +# ── WS 端点端到端(TestClient)─────────────────────────────────────────────── + + +def test_ws_endpoint_end_to_end(monkeypatch: pytest.MonkeyPatch) -> None: + """连接 → 收 tick 帧(真实路由 + hub + feed,仅行情源为假)。 + + hub 的 feed 任务跑在 TestClient 的 portal 事件循环里,客户端退出时随循环 + 一并销毁,无需(也无法)在测试协程里再 shutdown。 + """ + from fastapi import FastAPI + from fastapi.testclient import TestClient + + from easy_tdx.web.routers.realtime import router + + app = FastAPI() + app.include_router(router, prefix="/api/v1") + hub, _client = _make_hub() + app.state.realtime_hub = hub + + with TestClient(app) as client, client.websocket_connect("/api/v1/ws/realtime/SZ000001") as ws: + frame = ws.receive_json() + assert frame["type"] == "tick" + assert frame["symbol"] == "SZ000001" + assert {"price", "ts", "market", "code"} <= set(frame)