feat(ws): /ws/realtime 接通 RealtimeDataFeed — 按需轮询 hub + fan-out + 冒烟脚本

- web/realtime_hub.py:RealtimeStreamHub——订阅集合变化按需启停 RealtimeDataFeed,
  无人订阅完全停止轮询(对齐 QuoteStreamer 节能);EventBus → 每连接队列
  fan-out(丢最旧保最新);去重标的上限 80;lifespan 挂载/关闭
- routers/realtime.py 重写:连接即订阅、断开退订、subscribe/unsubscribe 控制帧、
  30s ping 心跳;单一写者泵模型(全部出站帧经队列串行,防并发 send 交错)
- 修复 RealtimeDataFeed stop-before-start 竞态(_stop_requested 标志),
  补 2 个回归用例;hub 单测 10 例(并发订阅/退订竞态、跨标的串扰、背压、
  TestClient 端到端)
- 前端接入选文档方案(api_reference.md 协议+重连/心跳骨架;README 撤「未联动」):
  SSE 已覆盖看板/自选实时刷新,WS 定位按需单标的 tick,双通道冗余无必要
- scripts/ws_smoke.py:手动冒烟(mock 模式实测 tick 帧与动态订阅确认可见)
This commit is contained in:
GitHub
2026-09-01 23:21:42 +08:00
parent bec231e8be
commit e58de789ab
11 changed files with 898 additions and 82 deletions
+5
View File
@@ -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
+24 -8
View File
@@ -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 文档
启动服务后访问:
+63
View File
@@ -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`。
+86
View File
@@ -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}sCtrl+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()))
+15 -1
View File
@@ -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
# ------------------------------------------------------------------ #
+25
View File
@@ -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),
+31
View File
@@ -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,
+247
View File
@@ -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
+90 -71
View File
@@ -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"}`` 确认;
去重后标的总数上限 80get_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:
def _enqueue(frame: dict[str, Any]) -> None:
"""控制帧也走队列(容量满时静默丢弃,不影响行情流)。"""
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)
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
_enqueue({"type": "ping"})
continue
try:
await websocket.send_json({"type": "ping"})
except Exception:
break
except WebSocketDisconnect:
break
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)
+26
View File
@@ -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 == []
+284
View File
@@ -0,0 +1,284 @@
"""RealtimeStreamHubWebSocket 实时推送枢纽)单元测试。
全程用假行情客户端(不打真实网络、不受交易时段限制——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)