From afcc100e3539c312cc779edb2226edb3d5450e4f Mon Sep 17 00:00:00 2001 From: Justin Gu <97915@qq.com> Date: Sun, 12 Jul 2026 16:56:41 +0800 Subject: [PATCH] =?UTF-8?q?feat(realtime):=20=E6=96=B0=E5=A2=9E=20Realtime?= =?UTF-8?q?DataFeed=20=E8=BD=AE=E8=AF=A2=E6=95=B0=E6=8D=AE=E6=BA=90?= =?UTF-8?q?=EF=BC=88issue=20#34=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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 新增),无回归。 --- README.md | 63 ++++ examples/11_realtime/realtime_feed_async.py | 73 +++++ src/easy_tdx/realtime/__init__.py | 32 +- src/easy_tdx/realtime/engine.py | 8 +- src/easy_tdx/realtime/feed.py | 323 ++++++++++++++++++++ tests/unit/test_realtime_feed.py | 303 ++++++++++++++++++ 6 files changed, 797 insertions(+), 5 deletions(-) create mode 100644 examples/11_realtime/realtime_feed_async.py create mode 100644 src/easy_tdx/realtime/feed.py create mode 100644 tests/unit/test_realtime_feed.py diff --git a/README.md b/README.md index e2a0408..6937a57 100644 --- a/README.md +++ b/README.md @@ -1112,6 +1112,10 @@ curl -X POST "http://localhost:8000/api/v1/chanlun/analyze" \ ### WebSocket 实时行情 +> ⚠️ **当前未联动数据源**:`create_app()` 尚未创建 `EventBus`,此 WS 端点目前 +> 不会推送任何行情。计划在后续版本接入 `RealtimeDataFeed` 后打通。 +> 如需实时行情,请先用上文「实时行情轮询」的编程 API。 + ```javascript // JavaScript 示例 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%) > - 大类标题行(如 ``流动资产``,原 ``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 线周期) diff --git a/examples/11_realtime/realtime_feed_async.py b/examples/11_realtime/realtime_feed_async.py new file mode 100644 index 0000000..5a53177 --- /dev/null +++ b/examples/11_realtime/realtime_feed_async.py @@ -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=平安银行 ⚡ 超阈值! diff --git a/src/easy_tdx/realtime/__init__.py b/src/easy_tdx/realtime/__init__.py index 58c09f0..e44e24e 100644 --- a/src/easy_tdx/realtime/__init__.py +++ b/src/easy_tdx/realtime/__init__.py @@ -1,10 +1,34 @@ """实时数据推送模块。 提供事件驱动的行情推送框架,基于 asyncio 实现。 -这是 API 骨架设计,实际协议层订阅功能待实现。 核心组件: -- EventBus: 事件总线,发布/订阅行情事件 -- RealtimeStrategy: 实时策略基类 -- MarketEvent: 行情事件数据结构 +- :class:`EventBus`: 事件总线,发布/订阅行情事件 +- :class:`RealtimeStrategy`: 实时策略基类 +- :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", +] diff --git a/src/easy_tdx/realtime/engine.py b/src/easy_tdx/realtime/engine.py index faaddd8..052be71 100644 --- a/src/easy_tdx/realtime/engine.py +++ b/src/easy_tdx/realtime/engine.py @@ -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 diff --git a/src/easy_tdx/realtime/feed.py b/src/easy_tdx/realtime/feed.py new file mode 100644 index 0000000..bfe9b30 --- /dev/null +++ b/src/easy_tdx/realtime/feed.py @@ -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 diff --git a/tests/unit/test_realtime_feed.py b/tests/unit/test_realtime_feed.py new file mode 100644 index 0000000..10b4b91 --- /dev/null +++ b/tests/unit/test_realtime_feed.py @@ -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