mirror of
https://ghfast.top/https://github.com/aeroxw/easy-tdx.git
synced 2026-09-12 13:24:15 +08:00
feat(realtime): 新增 RealtimeDataFeed 轮询数据源(issue #34)
EventBus 原本是纯发布/订阅管道,库内没有任何代码调用 publish(),导致 RealtimeStrategy 订阅后无数据可收(issue #34 根因)。 新增 RealtimeDataFeed:轮询 get_stock_quotes 快照 → 封装 MarketEvent → bus.publish,让 RealtimeStrategy 拿来就能跑。 - feed.py: async(run_async) + sync(run_sync, executor 线程避免卡 loop) 支持 interval/dedup(去重)/sessions(交易时段过滤,默认 9:15-11:30/13:00-15:00) symbol key 自动拼 f"{market}{code}" 匹配 EventBus.publish - test_realtime_feed.py: 20 个单测覆盖 publish 路径/symbol key 匹配/ 去重开关/空 df/客户端异常不崩溃/盘外不拉取/stop 退出 - engine.py + __init__.py docstring: 「API 骨架待实现」→「需配合 feed 使用」 - README: 新增实时行情轮询 Python API 段落(含 issue #34 三个坑说明), WebSocket 段落标注「当前未联动数据源」现状 - examples/11_realtime/: 可直接运行的异步示例 通达信协议无服务端推送,本模块是约 3 秒延迟的快照轮询近似,非逐笔 tick。 验证:ruff + mypy(strict) + pytest 全量 956 passed(+20 新增),无回归。
This commit is contained in:
@@ -1112,6 +1112,10 @@ curl -X POST "http://localhost:8000/api/v1/chanlun/analyze" \
|
|||||||
|
|
||||||
### WebSocket 实时行情
|
### WebSocket 实时行情
|
||||||
|
|
||||||
|
> ⚠️ **当前未联动数据源**:`create_app()` 尚未创建 `EventBus`,此 WS 端点目前
|
||||||
|
> 不会推送任何行情。计划在后续版本接入 `RealtimeDataFeed` 后打通。
|
||||||
|
> 如需实时行情,请先用上文「实时行情轮询」的编程 API。
|
||||||
|
|
||||||
```javascript
|
```javascript
|
||||||
// JavaScript 示例
|
// JavaScript 示例
|
||||||
const ws = new WebSocket("ws://localhost:8000/ws/realtime/SZ000001");
|
const ws = new WebSocket("ws://localhost:8000/ws/realtime/SZ000001");
|
||||||
@@ -1615,6 +1619,65 @@ df = client.get_financial_report("600519", report_type="llb", num=4)
|
|||||||
> - 有同比的科目附加 ``{科目}_同比`` 列(float 比例,如 0.06336 = +6.3%)
|
> - 有同比的科目附加 ``{科目}_同比`` 列(float 比例,如 0.06336 = +6.3%)
|
||||||
> - 大类标题行(如 ``流动资产``,原 ``item_value=""``)保留为 None,反映报表结构
|
> - 大类标题行(如 ``流动资产``,原 ``item_value=""``)保留为 None,反映报表结构
|
||||||
|
|
||||||
|
### 实时行情轮询(RealtimeDataFeed)
|
||||||
|
|
||||||
|
> ⚠️ 通达信协议**没有服务端推送**,只有请求/响应。本模块的「实时」是 **轮询五档快照近似**
|
||||||
|
> (默认约 3 秒延迟),适合盘中信号提醒、轻量监控;**不适合高频 / 逐笔撮合**。
|
||||||
|
|
||||||
|
`EventBus` 自身是纯发布/订阅管道,不会产生数据。要让 `RealtimeStrategy` 跑起来,
|
||||||
|
需要配合 `RealtimeDataFeed`:它自动完成 `get_stock_quotes → MarketEvent → bus.publish`。
|
||||||
|
|
||||||
|
```python
|
||||||
|
import asyncio
|
||||||
|
from easy_tdx.mac.client import AsyncMacClient
|
||||||
|
from easy_tdx.realtime import (
|
||||||
|
EventBus,
|
||||||
|
RealtimeStrategy,
|
||||||
|
MarketEvent,
|
||||||
|
RealtimeDataFeed,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class MyStrategy(RealtimeStrategy):
|
||||||
|
def on_tick(self, event: MarketEvent) -> None:
|
||||||
|
print(f"{event.market}{event.code} price={event.price} vol={event.volume}")
|
||||||
|
|
||||||
|
|
||||||
|
async def main():
|
||||||
|
bus = EventBus()
|
||||||
|
strategy = MyStrategy()
|
||||||
|
bus.subscribe("SZ000001", strategy.on_tick) # 注意:key 必须带市场前缀
|
||||||
|
|
||||||
|
feed = RealtimeDataFeed(
|
||||||
|
bus=bus,
|
||||||
|
symbols=[(0, "000001"), (1, "600519")], # [(Market.SZ, code), ...],单批 ≤80 只
|
||||||
|
interval=3.0, # 轮询间隔(秒),下限 0.1
|
||||||
|
dedup=True, # (price, volume) 未变的标的跳过发布
|
||||||
|
# sessions=(), # 传空 tuple 表示全天轮询;默认仅 9:15-11:30 / 13:00-15:00
|
||||||
|
)
|
||||||
|
async with AsyncMacClient.from_best_host() as client:
|
||||||
|
await feed.run_async(client) # Ctrl+C 或 feed.stop() 退出
|
||||||
|
|
||||||
|
|
||||||
|
asyncio.run(main())
|
||||||
|
```
|
||||||
|
|
||||||
|
同步客户端(`MacClient`)用 `run_sync`,feed 会把阻塞调用丢到线程池,不卡事件循环:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from easy_tdx.mac.client import MacClient
|
||||||
|
|
||||||
|
with MacClient.from_best_host() as client:
|
||||||
|
feed.run_sync(client)
|
||||||
|
```
|
||||||
|
|
||||||
|
> **关键坑(issue #34)**:
|
||||||
|
> - 订阅 key 必须与 `publish` 内部拼的 `f"{market}{code}"` 一致,即 `"SZ000001"`
|
||||||
|
> 而不是 `"000001"`;否则事件分发匹配不到。
|
||||||
|
> - 传给 `subscribe` 的必须是**实例的绑定方法** `strategy.on_tick`,
|
||||||
|
> 不是未绑定的类方法 `MyStrategy.on_tick`。
|
||||||
|
> - 整个流程跑在 `asyncio.run()` 里,否则协程不会被调度。
|
||||||
|
|
||||||
## 枚举参考
|
## 枚举参考
|
||||||
|
|
||||||
### Period(K 线周期)
|
### Period(K 线周期)
|
||||||
|
|||||||
@@ -0,0 +1,73 @@
|
|||||||
|
"""演示:RealtimeDataFeed + RealtimeStrategy 实时行情轮询。
|
||||||
|
|
||||||
|
通达信协议没有服务端推送,只有请求/响应。本示例用 RealtimeDataFeed 轮询
|
||||||
|
``get_stock_quotes`` 的五档快照,封装成 MarketEvent 喂给 EventBus,驱动
|
||||||
|
RealtimeStrategy 的 on_tick 回调。
|
||||||
|
|
||||||
|
注意事项:
|
||||||
|
- 「实时」是轮询快照近似(默认约 3 秒延迟),不是逐笔 tick。
|
||||||
|
- 订阅 key 必须与 publish 内部拼的 f"{market}{code}" 一致,即 "SZ000001"。
|
||||||
|
- 传给 subscribe 的必须是实例的绑定方法 strategy.on_tick。
|
||||||
|
- 默认仅在 A 股交易时段(9:15-11:30 / 13:00-15:00)轮询;演示用 sessions=()
|
||||||
|
全天轮询。
|
||||||
|
|
||||||
|
使用客户端:AsyncMacClient(异步)
|
||||||
|
关键参数:
|
||||||
|
- symbols: [(Market.SZ, code), ...],单批最多 80 只(协议上限)
|
||||||
|
- interval: 轮询间隔(秒),默认 3.0
|
||||||
|
- dedup: (price, volume) 未变的标的跳过发布,默认 True
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
from easy_tdx import Market
|
||||||
|
from easy_tdx.mac.client import AsyncMacClient
|
||||||
|
from easy_tdx.realtime import EventBus, MarketEvent, RealtimeDataFeed, RealtimeStrategy
|
||||||
|
|
||||||
|
|
||||||
|
class MyStrategy(RealtimeStrategy):
|
||||||
|
"""示例策略:价格超过阈值时打印提醒。"""
|
||||||
|
|
||||||
|
def __init__(self, threshold: float = 10.0) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.threshold = threshold
|
||||||
|
|
||||||
|
def on_tick(self, event: MarketEvent) -> None:
|
||||||
|
flag = " ⚡ 超阈值!" if event.price > self.threshold else ""
|
||||||
|
print(
|
||||||
|
f"{event.market}{event.code} price={event.price:.2f} "
|
||||||
|
f"vol={int(event.volume)} name={event.data.get('name', '')}{flag}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def main() -> None:
|
||||||
|
bus = EventBus()
|
||||||
|
strategy = MyStrategy(threshold=10.5)
|
||||||
|
# 注意:key 带市场前缀,且传实例绑定方法
|
||||||
|
bus.subscribe("SZ000001", strategy.on_tick)
|
||||||
|
bus.subscribe("SH600519", strategy.on_tick)
|
||||||
|
|
||||||
|
feed = RealtimeDataFeed(
|
||||||
|
bus=bus,
|
||||||
|
symbols=[(Market.SZ, "000001"), (Market.SH, "600519")],
|
||||||
|
interval=3.0,
|
||||||
|
dedup=True,
|
||||||
|
sessions=(), # 演示用,全天轮询;实盘删掉此行用默认交易时段
|
||||||
|
)
|
||||||
|
|
||||||
|
print("开始轮询(Ctrl+C 退出)...")
|
||||||
|
async with AsyncMacClient.from_best_host() as client:
|
||||||
|
await feed.run_async(client)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
try:
|
||||||
|
asyncio.run(main())
|
||||||
|
except KeyboardInterrupt:
|
||||||
|
print("\n已停止")
|
||||||
|
|
||||||
|
# 运行结果(示例):
|
||||||
|
# 开始轮询(Ctrl+C 退出)...
|
||||||
|
# SZ000001 price=10.52 vol=102345 name=平安银行 ⚡ 超阈值!
|
||||||
|
# SH600519 price=1798.00 vol=5234 name=贵州茅台
|
||||||
|
# SZ000001 price=10.53 vol=102580 name=平安银行 ⚡ 超阈值!
|
||||||
@@ -1,10 +1,34 @@
|
|||||||
"""实时数据推送模块。
|
"""实时数据推送模块。
|
||||||
|
|
||||||
提供事件驱动的行情推送框架,基于 asyncio 实现。
|
提供事件驱动的行情推送框架,基于 asyncio 实现。
|
||||||
这是 API 骨架设计,实际协议层订阅功能待实现。
|
|
||||||
|
|
||||||
核心组件:
|
核心组件:
|
||||||
- EventBus: 事件总线,发布/订阅行情事件
|
- :class:`EventBus`: 事件总线,发布/订阅行情事件
|
||||||
- RealtimeStrategy: 实时策略基类
|
- :class:`RealtimeStrategy`: 实时策略基类
|
||||||
- MarketEvent: 行情事件数据结构
|
- :class:`MarketEvent`: 行情事件数据结构
|
||||||
|
- :class:`RealtimeDataFeed`: 轮询数据源,把 ``get_stock_quotes`` 快照喂给 EventBus
|
||||||
|
|
||||||
|
.. note::
|
||||||
|
|
||||||
|
通达信协议没有服务端推送,本模块的「实时」是 **轮询快照近似**(默认约 3 秒
|
||||||
|
延迟)。``EventBus`` 自身不会产生数据,必须配合 :class:`RealtimeDataFeed`
|
||||||
|
(或自行实现 ``bus.publish`` 的数据源)才能驱动 ``RealtimeStrategy``。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from easy_tdx.realtime.engine import (
|
||||||
|
EventBus,
|
||||||
|
EventHandler,
|
||||||
|
EventType,
|
||||||
|
MarketEvent,
|
||||||
|
RealtimeStrategy,
|
||||||
|
)
|
||||||
|
from easy_tdx.realtime.feed import RealtimeDataFeed
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"EventBus",
|
||||||
|
"EventHandler",
|
||||||
|
"EventType",
|
||||||
|
"MarketEvent",
|
||||||
|
"RealtimeDataFeed",
|
||||||
|
"RealtimeStrategy",
|
||||||
|
]
|
||||||
|
|||||||
@@ -5,7 +5,13 @@
|
|||||||
- 实时策略信号触发
|
- 实时策略信号触发
|
||||||
- 多标的并发监控
|
- 多标的并发监控
|
||||||
|
|
||||||
这是 API 骨架,transport 层的协议级订阅待实现。
|
.. note::
|
||||||
|
|
||||||
|
``EventBus`` 是纯发布/订阅管道,自身不会产生数据。要驱动
|
||||||
|
:class:`RealtimeStrategy`,需要配合 :mod:`easy_tdx.realtime.feed` 的
|
||||||
|
:class:`RealtimeDataFeed`(轮询 ``get_stock_quotes`` 快照)或自行实现
|
||||||
|
调用 ``bus.publish`` 的数据源。通达信协议无服务端推送,本模块的「实时」
|
||||||
|
是轮询快照近似(默认约 3 秒延迟)。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|||||||
@@ -0,0 +1,323 @@
|
|||||||
|
"""轮询数据源:把 ``get_stock_quotes`` 的快照喂给 :class:`EventBus`。
|
||||||
|
|
||||||
|
通达信协议是请求/响应,没有服务端推送,所以「实时」只能靠**轮询五档快照**
|
||||||
|
近似实现。本模块提供 :class:`RealtimeDataFeed`,自动完成::
|
||||||
|
|
||||||
|
get_stock_quotes → MarketEvent → bus.publish
|
||||||
|
|
||||||
|
让 :class:`~easy_tdx.realtime.engine.RealtimeStrategy` 拿来就能跑,不必每个
|
||||||
|
用户都手写一遍轮询循环、symbol key 拼接、asyncio 调度。
|
||||||
|
|
||||||
|
.. note::
|
||||||
|
|
||||||
|
这是 **约 ``interval`` 秒延迟的快照轮询**,不是逐笔 tick。适合盘中信号
|
||||||
|
提醒、轻量监控;不适合高频 / 逐笔撮合。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
import time
|
||||||
|
from collections.abc import Iterable
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import pandas as pd
|
||||||
|
|
||||||
|
from .engine import EventBus, EventType, MarketEvent
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Market 整数 → 字符串前缀。EventBus.publish 用 ``f"{market}{code}"`` 做 key,
|
||||||
|
# 所以这里必须产出与订阅 key 一致的字符串前缀。
|
||||||
|
_MARKET_INT_TO_STR: dict[int, str] = {0: "SZ", 1: "SH", 2: "BJ"}
|
||||||
|
|
||||||
|
# A 股常规交易时段(含集合竞价)。盘外时段默认不轮询,避免空转 + 被服务器限流。
|
||||||
|
_DEFAULT_SESSIONS: tuple[tuple[int, int], ...] = (
|
||||||
|
(9 * 60 + 15, 11 * 60 + 30), # 09:15 - 11:30(含集合竞价)
|
||||||
|
(13 * 60 + 0, 15 * 60 + 0), # 13:00 - 15:00
|
||||||
|
)
|
||||||
|
|
||||||
|
_BATCH_SIZE = 80 # get_stock_quotes 单次上限(协议约束)
|
||||||
|
|
||||||
|
|
||||||
|
def _market_int_to_str(market: int) -> str:
|
||||||
|
"""把 ``Market`` 整数转成 ``EventBus`` 用的字符串前缀。
|
||||||
|
|
||||||
|
未知市场回退为十进制字符串,保证 ``publish`` key 可拼接、不丢事件。
|
||||||
|
"""
|
||||||
|
return _MARKET_INT_TO_STR.get(market, str(market))
|
||||||
|
|
||||||
|
|
||||||
|
def _is_in_session(now: float, sessions: tuple[tuple[int, int], ...]) -> bool:
|
||||||
|
"""判断 ``now``(epoch 秒)的本地时分是否落在某个交易时段内。
|
||||||
|
|
||||||
|
空 ``sessions`` 表示不做时段过滤,始终返回 True。
|
||||||
|
"""
|
||||||
|
if not sessions:
|
||||||
|
return True
|
||||||
|
lt = time.localtime(now)
|
||||||
|
minute_of_day = lt.tm_hour * 60 + lt.tm_min
|
||||||
|
for start, end in sessions:
|
||||||
|
if start <= minute_of_day < end:
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _row_to_event(row: pd.Series, timestamp: float) -> MarketEvent | None:
|
||||||
|
"""把 quotes DataFrame 的一行转成 :class:`MarketEvent`。
|
||||||
|
|
||||||
|
- ``market`` 列是 int,转成字符串前缀以匹配 ``EventBus.publish`` 的 key。
|
||||||
|
- 最新价取 ``close`` 列(``get_stock_quotes`` 没有 ``price`` 列,最新成交价
|
||||||
|
落在 ``close``)。
|
||||||
|
- 缺列时回退为 0.0,保证事件仍可发布。
|
||||||
|
"""
|
||||||
|
code = str(row.get("code", ""))
|
||||||
|
if not code:
|
||||||
|
return None
|
||||||
|
market_raw = row.get("market")
|
||||||
|
try:
|
||||||
|
market_int = int(market_raw)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
market_int = -1
|
||||||
|
market_str = _market_int_to_str(market_int)
|
||||||
|
|
||||||
|
def _f(key: str) -> float:
|
||||||
|
try:
|
||||||
|
return float(row.get(key, 0.0) or 0.0)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return 0.0
|
||||||
|
|
||||||
|
return MarketEvent(
|
||||||
|
event_type=EventType.TICK,
|
||||||
|
code=code,
|
||||||
|
market=market_str,
|
||||||
|
price=_f("close"),
|
||||||
|
volume=_f("vol"),
|
||||||
|
timestamp=timestamp,
|
||||||
|
data={
|
||||||
|
"open": _f("open"),
|
||||||
|
"high": _f("high"),
|
||||||
|
"low": _f("low"),
|
||||||
|
"pre_close": _f("pre_close"),
|
||||||
|
"amount": _f("amount"),
|
||||||
|
"name": str(row.get("name", "")),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class _FeedState:
|
||||||
|
"""跨周期去重用的上次价格/成交量缓存。"""
|
||||||
|
|
||||||
|
last: dict[str, tuple[float, float]] = field(default_factory=dict)
|
||||||
|
|
||||||
|
def changed(self, key: str, price: float, volume: float) -> bool:
|
||||||
|
prev = self.last.get(key)
|
||||||
|
self.last[key] = (price, volume)
|
||||||
|
return prev != (price, volume)
|
||||||
|
|
||||||
|
|
||||||
|
class RealtimeDataFeed:
|
||||||
|
"""轮询 ``get_stock_quotes`` 并把快照发布到 :class:`EventBus`。
|
||||||
|
|
||||||
|
同时支持 **异步客户端**(推荐,:meth:`run_async`)和 **同步客户端**
|
||||||
|
(:meth:`run_sync`,把阻塞调用丢到 executor 线程,避免卡住事件循环)。
|
||||||
|
|
||||||
|
用法(异步,最常见)::
|
||||||
|
|
||||||
|
from easy_tdx.mac.client import AsyncMacClient
|
||||||
|
from easy_tdx.realtime.engine import EventBus, RealtimeStrategy
|
||||||
|
from easy_tdx.realtime.feed import RealtimeDataFeed
|
||||||
|
|
||||||
|
bus = EventBus()
|
||||||
|
bus.subscribe("SZ000001", MyStrategy().on_tick)
|
||||||
|
|
||||||
|
feed = RealtimeDataFeed(
|
||||||
|
bus=bus,
|
||||||
|
symbols=[(0, "000001"), (1, "600519")], # [(Market.SZ, code), ...]
|
||||||
|
)
|
||||||
|
async with AsyncMacClient.from_best_host() as client:
|
||||||
|
await feed.run_async(client)
|
||||||
|
|
||||||
|
用法(同步客户端)::
|
||||||
|
|
||||||
|
from easy_tdx.mac.client import MacClient
|
||||||
|
|
||||||
|
feed = RealtimeDataFeed(bus=bus, symbols=[(0, "000001")])
|
||||||
|
with MacClient.from_best_host() as client:
|
||||||
|
feed.run_sync(client) # 阻塞,直到 Ctrl+C 或 feed.stop()
|
||||||
|
|
||||||
|
Attributes:
|
||||||
|
bus: 目标事件总线。
|
||||||
|
symbols: ``[(market_int, code), ...]`` 列表,单次最多 80 只(协议上限)。
|
||||||
|
interval: 轮询间隔(秒),默认 3.0。
|
||||||
|
dedup: 是否对 ``(price, volume)`` 未变化的标的跳过发布,默认 True。
|
||||||
|
sessions: 交易时段(``[(start_min, end_min), ...]``,分钟数),
|
||||||
|
盘外时段只睡眠不拉取;传空 tuple 表示不做时段过滤(全天轮询)。
|
||||||
|
fields: 透传给 ``get_stock_quotes`` 的字段选择,默认 None 用客户端默认。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
bus: EventBus,
|
||||||
|
symbols: Iterable[tuple[int, str]],
|
||||||
|
*,
|
||||||
|
interval: float = 3.0,
|
||||||
|
dedup: bool = True,
|
||||||
|
sessions: tuple[tuple[int, int], ...] | None = None,
|
||||||
|
fields: object = None,
|
||||||
|
) -> None:
|
||||||
|
symbol_list = list(symbols)
|
||||||
|
if not symbol_list:
|
||||||
|
raise ValueError("symbols 不能为空")
|
||||||
|
if len(symbol_list) > _BATCH_SIZE:
|
||||||
|
raise ValueError(
|
||||||
|
f"symbols 最多 {_BATCH_SIZE} 只(get_stock_quotes 单次上限),"
|
||||||
|
f"当前传了 {len(symbol_list)} 只"
|
||||||
|
)
|
||||||
|
|
||||||
|
self._bus = bus
|
||||||
|
self._symbols = symbol_list
|
||||||
|
self._interval = max(0.1, interval)
|
||||||
|
self._dedup = dedup
|
||||||
|
self._sessions = _DEFAULT_SESSIONS if sessions is None else sessions
|
||||||
|
self._fields = fields
|
||||||
|
self._state = _FeedState()
|
||||||
|
self._running = False
|
||||||
|
|
||||||
|
@property
|
||||||
|
def running(self) -> bool:
|
||||||
|
"""是否正在轮询。"""
|
||||||
|
return self._running
|
||||||
|
|
||||||
|
async def run_async(
|
||||||
|
self,
|
||||||
|
client: Any,
|
||||||
|
*,
|
||||||
|
max_iterations: int | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""异步轮询循环(主推入口)。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
client: 拥有 ``async def get_stock_quotes`` 的客户端
|
||||||
|
(如 :class:`~easy_tdx.mac.client.AsyncMacClient`)。
|
||||||
|
max_iterations: 最多轮询多少轮(测试用);None 表示无限循环直到
|
||||||
|
:meth:`stop`。
|
||||||
|
"""
|
||||||
|
self._running = True
|
||||||
|
try:
|
||||||
|
count = 0
|
||||||
|
while self._running:
|
||||||
|
if max_iterations is not None and count >= max_iterations:
|
||||||
|
break
|
||||||
|
count += 1
|
||||||
|
await self._poll_once_async(client)
|
||||||
|
await self._sleep_or_stop()
|
||||||
|
finally:
|
||||||
|
self._running = False
|
||||||
|
|
||||||
|
def run_sync(
|
||||||
|
self,
|
||||||
|
client: Any,
|
||||||
|
*,
|
||||||
|
max_iterations: int | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""同步轮询循环(阻塞调用方)。
|
||||||
|
|
||||||
|
同步客户端的 ``get_stock_quotes`` 是阻塞 socket 调用,直接放进 asyncio
|
||||||
|
事件循环会卡死整个 loop。本方法把每次拉取丢到默认 executor 线程执行,
|
||||||
|
发布事件仍走事件循环,从而既不卡 loop、又不用换异步客户端。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
client: 拥有同步 ``get_stock_quotes`` 的客户端
|
||||||
|
(如 :class:`~easy_tdx.mac.client.MacClient`)。
|
||||||
|
max_iterations: 最多轮询多少轮(测试用)。
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
loop = asyncio.get_event_loop()
|
||||||
|
if loop.is_running():
|
||||||
|
raise RuntimeError(
|
||||||
|
"检测到正在运行的事件循环;请在循环内改用 "
|
||||||
|
"await feed.run_async(...) 而非 feed.run_sync(...)"
|
||||||
|
)
|
||||||
|
except RuntimeError:
|
||||||
|
loop = asyncio.new_event_loop()
|
||||||
|
asyncio.set_event_loop(loop)
|
||||||
|
|
||||||
|
try:
|
||||||
|
loop.run_until_complete(self._run_sync_loop(client, max_iterations))
|
||||||
|
finally:
|
||||||
|
if not loop.is_running():
|
||||||
|
loop.close()
|
||||||
|
|
||||||
|
async def _run_sync_loop(self, client: Any, max_iterations: int | None) -> None:
|
||||||
|
"""同步客户端的轮询循环:阻塞调用丢到 executor。"""
|
||||||
|
self._running = True
|
||||||
|
try:
|
||||||
|
count = 0
|
||||||
|
while self._running:
|
||||||
|
if max_iterations is not None and count >= max_iterations:
|
||||||
|
break
|
||||||
|
count += 1
|
||||||
|
await self._poll_once_sync(client)
|
||||||
|
await self._sleep_or_stop()
|
||||||
|
finally:
|
||||||
|
self._running = False
|
||||||
|
|
||||||
|
def stop(self) -> None:
|
||||||
|
"""请求停止轮询(下一轮 sleep 结束后生效)。"""
|
||||||
|
self._running = False
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------ #
|
||||||
|
# 内部
|
||||||
|
# ------------------------------------------------------------------ #
|
||||||
|
|
||||||
|
async def _poll_once_async(self, client: Any) -> None:
|
||||||
|
now = time.time()
|
||||||
|
if not self._in_session(now):
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
df = await client.get_stock_quotes(self._symbols, self._fields)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("async get_stock_quotes 失败,本轮跳过")
|
||||||
|
return
|
||||||
|
await self._publish_df(df, now)
|
||||||
|
|
||||||
|
async def _poll_once_sync(self, client: Any) -> None:
|
||||||
|
now = time.time()
|
||||||
|
if not self._in_session(now):
|
||||||
|
return
|
||||||
|
loop = asyncio.get_event_loop()
|
||||||
|
try:
|
||||||
|
df = await loop.run_in_executor(
|
||||||
|
None, client.get_stock_quotes, self._symbols, self._fields
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("sync get_stock_quotes 失败,本轮跳过")
|
||||||
|
return
|
||||||
|
await self._publish_df(df, now)
|
||||||
|
|
||||||
|
async def _publish_df(self, df: pd.DataFrame, timestamp: float) -> None:
|
||||||
|
if df is None or df.empty:
|
||||||
|
return
|
||||||
|
for _, row in df.iterrows():
|
||||||
|
event = _row_to_event(row, timestamp)
|
||||||
|
if event is None:
|
||||||
|
continue
|
||||||
|
key = f"{event.market}{event.code}"
|
||||||
|
if self._dedup and not self._state.changed(key, event.price, event.volume):
|
||||||
|
continue
|
||||||
|
await self._bus.publish(event)
|
||||||
|
|
||||||
|
def _in_session(self, now: float) -> bool:
|
||||||
|
return _is_in_session(now, self._sessions)
|
||||||
|
|
||||||
|
async def _sleep_or_stop(self) -> None:
|
||||||
|
"""按 interval 睡眠,但每 0.5s 检查一次 stop 标志,缩短退出延迟。"""
|
||||||
|
elapsed = 0.0
|
||||||
|
step = 0.5
|
||||||
|
while elapsed < self._interval and self._running:
|
||||||
|
await asyncio.sleep(min(step, self._interval - elapsed))
|
||||||
|
elapsed += step
|
||||||
@@ -0,0 +1,303 @@
|
|||||||
|
"""单元测试:RealtimeDataFeed 轮询数据源."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import time
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import pandas as pd
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from easy_tdx.realtime.engine import EventBus, EventType, MarketEvent
|
||||||
|
from easy_tdx.realtime.feed import (
|
||||||
|
RealtimeDataFeed,
|
||||||
|
_is_in_session,
|
||||||
|
_market_int_to_str,
|
||||||
|
_row_to_event,
|
||||||
|
)
|
||||||
|
|
||||||
|
# ── 测试数据 ────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def _sample_quotes_df() -> pd.DataFrame:
|
||||||
|
"""模拟 get_stock_quotes 返回的 DataFrame。
|
||||||
|
|
||||||
|
列结构与 MacClient._quotes_to_df 一致:market(int) / code / name + 字段。
|
||||||
|
"""
|
||||||
|
return pd.DataFrame(
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"market": 0, # SZ
|
||||||
|
"code": "000001",
|
||||||
|
"name": "平安银行",
|
||||||
|
"close": 10.50,
|
||||||
|
"vol": 100000,
|
||||||
|
"open": 10.30,
|
||||||
|
"high": 10.60,
|
||||||
|
"low": 10.20,
|
||||||
|
"pre_close": 10.40,
|
||||||
|
"amount": 1050000.0,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"market": 1, # SH
|
||||||
|
"code": "600519",
|
||||||
|
"name": "贵州茅台",
|
||||||
|
"close": 1800.0,
|
||||||
|
"vol": 5000,
|
||||||
|
"open": 1790.0,
|
||||||
|
"high": 1810.0,
|
||||||
|
"low": 1785.0,
|
||||||
|
"pre_close": 1795.0,
|
||||||
|
"amount": 9000000.0,
|
||||||
|
},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class AsyncMockClient:
|
||||||
|
"""模拟异步客户端:按预设序列返回 DataFrame。"""
|
||||||
|
|
||||||
|
def __init__(self, frames: list[pd.DataFrame]) -> None:
|
||||||
|
self._frames = frames
|
||||||
|
self._idx = 0
|
||||||
|
self.calls: list[list[tuple[int, str]]] = []
|
||||||
|
|
||||||
|
async def get_stock_quotes(
|
||||||
|
self, stocks: list[tuple[int, str]], fields: object = None
|
||||||
|
) -> pd.DataFrame:
|
||||||
|
self.calls.append(list(stocks))
|
||||||
|
if self._idx < len(self._frames):
|
||||||
|
df = self._frames[self._idx]
|
||||||
|
self._idx += 1
|
||||||
|
return df
|
||||||
|
return pd.DataFrame()
|
||||||
|
|
||||||
|
|
||||||
|
class SyncMockClient:
|
||||||
|
"""模拟同步客户端。"""
|
||||||
|
|
||||||
|
def __init__(self, frames: list[pd.DataFrame]) -> None:
|
||||||
|
self._frames = frames
|
||||||
|
self._idx = 0
|
||||||
|
self.calls: list[list[tuple[int, str]]] = []
|
||||||
|
|
||||||
|
def get_stock_quotes(
|
||||||
|
self, stocks: list[tuple[int, str]], fields: object = None
|
||||||
|
) -> pd.DataFrame:
|
||||||
|
self.calls.append(list(stocks))
|
||||||
|
if self._idx < len(self._frames):
|
||||||
|
df = self._frames[self._idx]
|
||||||
|
self._idx += 1
|
||||||
|
return df
|
||||||
|
return pd.DataFrame()
|
||||||
|
|
||||||
|
|
||||||
|
# ── 纯函数测试 ──────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
class TestMarketIntToStr:
|
||||||
|
def test_known_markets(self) -> None:
|
||||||
|
assert _market_int_to_str(0) == "SZ"
|
||||||
|
assert _market_int_to_str(1) == "SH"
|
||||||
|
assert _market_int_to_str(2) == "BJ"
|
||||||
|
|
||||||
|
def test_unknown_falls_back(self) -> None:
|
||||||
|
assert _market_int_to_str(99) == "99"
|
||||||
|
|
||||||
|
|
||||||
|
class TestIsInSession:
|
||||||
|
def test_in_morning_session(self) -> None:
|
||||||
|
# 构造一个 10:30 的 epoch(任意日期,tm_hour=10, tm_min=30)
|
||||||
|
t = time.mktime(time.strptime("2026-07-13 10:30:00", "%Y-%m-%d %H:%M:%S"))
|
||||||
|
sessions = ((9 * 60 + 15, 11 * 60 + 30), (13 * 60, 15 * 60))
|
||||||
|
assert _is_in_session(t, sessions) is True
|
||||||
|
|
||||||
|
def test_outside_session(self) -> None:
|
||||||
|
t = time.mktime(time.strptime("2026-07-13 08:00:00", "%Y-%m-%d %H:%M:%S"))
|
||||||
|
sessions = ((9 * 60 + 15, 11 * 60 + 30), (13 * 60, 15 * 60))
|
||||||
|
assert _is_in_session(t, sessions) is False
|
||||||
|
|
||||||
|
def test_empty_sessions_always_in(self) -> None:
|
||||||
|
t = time.mktime(time.strptime("2026-07-13 03:00:00", "%Y-%m-%d %H:%M:%S"))
|
||||||
|
assert _is_in_session(t, ()) is True
|
||||||
|
|
||||||
|
|
||||||
|
class TestRowToEvent:
|
||||||
|
def test_basic_conversion(self) -> None:
|
||||||
|
df = _sample_quotes_df()
|
||||||
|
event = _row_to_event(df.iloc[0], timestamp=1700000000.0)
|
||||||
|
assert event is not None
|
||||||
|
assert event.code == "000001"
|
||||||
|
assert event.market == "SZ"
|
||||||
|
assert event.price == 10.50
|
||||||
|
assert event.volume == 100000
|
||||||
|
assert event.event_type == EventType.TICK
|
||||||
|
assert event.data["name"] == "平安银行"
|
||||||
|
assert event.data["pre_close"] == 10.40
|
||||||
|
|
||||||
|
def test_sh_market_prefix(self) -> None:
|
||||||
|
df = _sample_quotes_df()
|
||||||
|
event = _row_to_event(df.iloc[1], timestamp=0.0)
|
||||||
|
assert event is not None
|
||||||
|
assert event.market == "SH"
|
||||||
|
assert event.code == "600519"
|
||||||
|
assert event.price == 1800.0
|
||||||
|
|
||||||
|
def test_missing_columns_default_zero(self) -> None:
|
||||||
|
row = pd.Series({"code": "000002", "market": 0, "name": "万科A"})
|
||||||
|
event = _row_to_event(row, timestamp=0.0)
|
||||||
|
assert event is not None
|
||||||
|
assert event.price == 0.0
|
||||||
|
assert event.volume == 0.0
|
||||||
|
|
||||||
|
def test_empty_code_returns_none(self) -> None:
|
||||||
|
row = pd.Series({"code": "", "market": 0})
|
||||||
|
assert _row_to_event(row, timestamp=0.0) is None
|
||||||
|
|
||||||
|
|
||||||
|
# ── Feed 构造测试 ───────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
class TestFeedConstruction:
|
||||||
|
def test_empty_symbols_raises(self) -> None:
|
||||||
|
with pytest.raises(ValueError, match="不能为空"):
|
||||||
|
RealtimeDataFeed(bus=EventBus(), symbols=[])
|
||||||
|
|
||||||
|
def test_too_many_symbols_raises(self) -> None:
|
||||||
|
symbols = [(0, f"{i:06d}") for i in range(81)]
|
||||||
|
with pytest.raises(ValueError, match="80"):
|
||||||
|
RealtimeDataFeed(bus=EventBus(), symbols=symbols)
|
||||||
|
|
||||||
|
def test_interval_clamped_to_minimum(self) -> None:
|
||||||
|
feed = RealtimeDataFeed(bus=EventBus(), symbols=[(0, "000001")], interval=0.01)
|
||||||
|
assert feed._interval == 0.1
|
||||||
|
|
||||||
|
def test_sessions_override_empty(self) -> None:
|
||||||
|
feed = RealtimeDataFeed(bus=EventBus(), symbols=[(0, "000001")], sessions=())
|
||||||
|
# 空 sessions → _in_session 始终 True
|
||||||
|
assert feed._in_session(0.0) is True
|
||||||
|
|
||||||
|
|
||||||
|
# ── 异步 publish 路径测试 ───────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
class TestAsyncPublishPath:
|
||||||
|
async def test_events_published_to_correct_keys(self) -> None:
|
||||||
|
"""关键测试:market int 0 → 'SZ',订阅 'SZ000001' 必须收到。"""
|
||||||
|
bus = EventBus()
|
||||||
|
received: list[MarketEvent] = []
|
||||||
|
bus.subscribe("SZ000001", lambda e: received.append(e))
|
||||||
|
bus.subscribe("SH600519", lambda e: received.append(e))
|
||||||
|
|
||||||
|
client = AsyncMockClient([_sample_quotes_df()])
|
||||||
|
feed = RealtimeDataFeed(
|
||||||
|
bus=bus,
|
||||||
|
symbols=[(0, "000001"), (1, "600519")],
|
||||||
|
sessions=(), # 测试不受时段限制
|
||||||
|
interval=0.1,
|
||||||
|
)
|
||||||
|
await feed.run_async(client, max_iterations=1)
|
||||||
|
|
||||||
|
assert len(received) == 2
|
||||||
|
codes = {e.code for e in received}
|
||||||
|
assert codes == {"000001", "600519"}
|
||||||
|
# symbol key 必须匹配:SZ 前缀
|
||||||
|
sz_event = next(e for e in received if e.code == "000001")
|
||||||
|
assert sz_event.market == "SZ"
|
||||||
|
assert sz_event.price == 10.50
|
||||||
|
|
||||||
|
async def test_dedup_skips_unchanged(self) -> None:
|
||||||
|
bus = EventBus()
|
||||||
|
received: list[MarketEvent] = []
|
||||||
|
bus.subscribe_all(lambda e: received.append(e))
|
||||||
|
|
||||||
|
same_df = _sample_quotes_df()
|
||||||
|
client = AsyncMockClient([same_df.copy(), same_df.copy()])
|
||||||
|
feed = RealtimeDataFeed(
|
||||||
|
bus=bus,
|
||||||
|
symbols=[(0, "000001"), (1, "600519")],
|
||||||
|
dedup=True,
|
||||||
|
sessions=(),
|
||||||
|
interval=0.1,
|
||||||
|
)
|
||||||
|
await feed.run_async(client, max_iterations=2)
|
||||||
|
|
||||||
|
# 第一轮 2 个事件,第二轮因 price/volume 不变被去重
|
||||||
|
assert len(received) == 2
|
||||||
|
|
||||||
|
async def test_dedup_disabled_publishes_all(self) -> None:
|
||||||
|
bus = EventBus()
|
||||||
|
received: list[MarketEvent] = []
|
||||||
|
bus.subscribe_all(lambda e: received.append(e))
|
||||||
|
|
||||||
|
same_df = _sample_quotes_df()
|
||||||
|
client = AsyncMockClient([same_df.copy(), same_df.copy()])
|
||||||
|
feed = RealtimeDataFeed(
|
||||||
|
bus=bus,
|
||||||
|
symbols=[(0, "000001"), (1, "600519")],
|
||||||
|
dedup=False,
|
||||||
|
sessions=(),
|
||||||
|
interval=0.1,
|
||||||
|
)
|
||||||
|
await feed.run_async(client, max_iterations=2)
|
||||||
|
|
||||||
|
assert len(received) == 4
|
||||||
|
|
||||||
|
async def test_empty_df_publishes_nothing(self) -> None:
|
||||||
|
bus = EventBus()
|
||||||
|
received: list[MarketEvent] = []
|
||||||
|
bus.subscribe_all(lambda e: received.append(e))
|
||||||
|
|
||||||
|
client = AsyncMockClient([pd.DataFrame()])
|
||||||
|
feed = RealtimeDataFeed(bus=EventBus(), symbols=[(0, "000001")], sessions=(), interval=0.1)
|
||||||
|
# 用 subscribe_all 的 bus
|
||||||
|
feed._bus = bus
|
||||||
|
await feed.run_async(client, max_iterations=1)
|
||||||
|
assert received == []
|
||||||
|
|
||||||
|
async def test_client_error_does_not_crash(self) -> None:
|
||||||
|
"""get_stock_quotes 抛异常时,feed 应跳过该轮,不崩溃。"""
|
||||||
|
bus = EventBus()
|
||||||
|
received: list[MarketEvent] = []
|
||||||
|
bus.subscribe_all(lambda e: received.append(e))
|
||||||
|
|
||||||
|
failing_client = MagicMock()
|
||||||
|
failing_client.get_stock_quotes = MagicMock(side_effect=ConnectionError("boom"))
|
||||||
|
|
||||||
|
feed = RealtimeDataFeed(bus=bus, symbols=[(0, "000001")], sessions=(), interval=0.1)
|
||||||
|
# 不应抛异常
|
||||||
|
await feed.run_async(failing_client, max_iterations=1)
|
||||||
|
assert received == []
|
||||||
|
|
||||||
|
|
||||||
|
class TestSessionGating:
|
||||||
|
async def test_outside_session_no_fetch(self) -> None:
|
||||||
|
"""盘外时段不应调用 get_stock_quotes。"""
|
||||||
|
bus = EventBus()
|
||||||
|
client = AsyncMockClient([_sample_quotes_df()])
|
||||||
|
|
||||||
|
# 用一个不可能命中的时段(如 23:00-23:59)模拟盘外
|
||||||
|
feed = RealtimeDataFeed(
|
||||||
|
bus=bus,
|
||||||
|
symbols=[(0, "000001")],
|
||||||
|
sessions=((23 * 60, 23 * 60 + 59),),
|
||||||
|
interval=0.1,
|
||||||
|
)
|
||||||
|
await feed.run_async(client, max_iterations=1)
|
||||||
|
assert len(client.calls) == 0 # 盘外,没拉数据
|
||||||
|
|
||||||
|
|
||||||
|
class TestStopFlag:
|
||||||
|
async def test_stop_terminates_loop(self) -> None:
|
||||||
|
bus = EventBus()
|
||||||
|
client = AsyncMockClient([_sample_quotes_df()])
|
||||||
|
feed = RealtimeDataFeed(bus=bus, symbols=[(0, "000001")], sessions=(), interval=0.2)
|
||||||
|
|
||||||
|
# 在短延迟后请求停止
|
||||||
|
async def _stop_soon() -> None:
|
||||||
|
await asyncio.sleep(0.15)
|
||||||
|
feed.stop()
|
||||||
|
|
||||||
|
await asyncio.gather(feed.run_async(client), _stop_soon())
|
||||||
|
assert feed.running is False
|
||||||
Reference in New Issue
Block a user