Files
easy_tdx_max/tests/unit/test_realtime_feed.py
T
GitHub e58de789ab 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 帧与动态订阅确认可见)
2026-09-01 23:21:42 +08:00

335 lines
12 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""单元测试:RealtimeDataFeed 轮询数据源."""
from __future__ import annotations
import asyncio
import time
from datetime import datetime
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()])
# 动态构造一个「永远在未来」的 1 分钟时段,保证当前时刻必在盘外。
# 此前用固定 23:00-23:59 模拟盘外——测试跑在当地 23 点档时必炸
# (时间相关 flaky)。
_now = datetime.now()
now_min = _now.hour * 60 + _now.minute
feed = RealtimeDataFeed(
bus=bus,
symbols=[(0, "000001")],
sessions=(((now_min + 5) % (24 * 60), (now_min + 6) % (24 * 60)),),
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
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 == []