From 92edc189bb689abafee8b51297c9403f609c3b41 Mon Sep 17 00:00:00 2001 From: Justin Gu <97915@qq.com> Date: Thu, 11 Jun 2026 02:34:00 +0800 Subject: [PATCH] feat: realtime event-driven market data push framework - Add EventBus for async publish/subscribe market events - Add MarketEvent dataclass with tick/bar/signal/error types - Add RealtimeStrategy base class with on_tick/on_bar callbacks - Add emit_signal() for strategy-to-engine signal publishing - Support per-symbol and global subscriptions - API skeleton: transport-level subscription TBD - Add 10 tests covering events, bus, and strategy --- src/easy_tdx/realtime/__init__.py | 10 ++ src/easy_tdx/realtime/engine.py | 211 ++++++++++++++++++++++++++++++ tests/unit/test_realtime.py | 164 +++++++++++++++++++++++ 3 files changed, 385 insertions(+) create mode 100644 src/easy_tdx/realtime/__init__.py create mode 100644 src/easy_tdx/realtime/engine.py create mode 100644 tests/unit/test_realtime.py diff --git a/src/easy_tdx/realtime/__init__.py b/src/easy_tdx/realtime/__init__.py new file mode 100644 index 0000000..58c09f0 --- /dev/null +++ b/src/easy_tdx/realtime/__init__.py @@ -0,0 +1,10 @@ +"""实时数据推送模块。 + +提供事件驱动的行情推送框架,基于 asyncio 实现。 +这是 API 骨架设计,实际协议层订阅功能待实现。 + +核心组件: +- EventBus: 事件总线,发布/订阅行情事件 +- RealtimeStrategy: 实时策略基类 +- MarketEvent: 行情事件数据结构 +""" diff --git a/src/easy_tdx/realtime/engine.py b/src/easy_tdx/realtime/engine.py new file mode 100644 index 0000000..faaddd8 --- /dev/null +++ b/src/easy_tdx/realtime/engine.py @@ -0,0 +1,211 @@ +"""实时数据推送引擎。 + +基于 asyncio 的事件驱动架构,支持: +- 行情订阅与推送 +- 实时策略信号触发 +- 多标的并发监控 + +这是 API 骨架,transport 层的协议级订阅待实现。 +""" + +from __future__ import annotations + +import asyncio +import logging +from collections.abc import Callable +from dataclasses import dataclass, field +from enum import Enum +from typing import Any + +logger = logging.getLogger(__name__) + + +class EventType(Enum): + """事件类型。""" + + TICK = "tick" + BAR = "bar" + SIGNAL = "signal" + ERROR = "error" + + +@dataclass +class MarketEvent: + """行情事件。 + + Attributes: + event_type: 事件类型 + code: 股票代码(如 "000001") + market: 市场(如 "SZ") + price: 最新价格 + volume: 成交量 + timestamp: 事件时间戳 + data: 额外数据 + """ + + event_type: EventType + code: str + market: str + price: float = 0.0 + volume: float = 0.0 + timestamp: float = 0.0 + data: dict[str, Any] = field(default_factory=dict) + + +# 回调函数类型 +EventHandler = Callable[[MarketEvent], Any] + + +class EventBus: + """异步事件总线。 + + 发布/订阅模式,支持多个订阅者监听行情事件。 + + 用法:: + + bus = EventBus() + bus.subscribe("SZ000001", on_tick) + await bus.publish(event) + """ + + def __init__(self) -> None: + self._subscribers: dict[str, list[EventHandler]] = {} + self._global_subscribers: list[EventHandler] = [] + + def subscribe(self, symbol: str, handler: EventHandler) -> None: + """订阅指定标的的事件。 + + Args: + symbol: 标的标识(如 "SZ000001") + handler: 事件处理回调 + """ + if symbol not in self._subscribers: + self._subscribers[symbol] = [] + self._subscribers[symbol].append(handler) + + def subscribe_all(self, handler: EventHandler) -> None: + """订阅所有标的的事件。 + + Args: + handler: 事件处理回调 + """ + self._global_subscribers.append(handler) + + def unsubscribe(self, symbol: str, handler: EventHandler) -> None: + """取消订阅。 + + Args: + symbol: 标的标识 + handler: 要移除的回调 + """ + if symbol in self._subscribers: + self._subscribers[symbol] = [h for h in self._subscribers[symbol] if h != handler] + + async def publish(self, event: MarketEvent) -> None: + """发布事件到所有匹配的订阅者。 + + Args: + event: 行情事件 + """ + symbol = f"{event.market}{event.code}" + + # 通知特定标的的订阅者 + handlers = self._subscribers.get(symbol, []) + for handler in handlers: + try: + result = handler(event) + if asyncio.iscoroutine(result): + await result + except Exception: + logger.exception("Event handler error for %s", symbol) + + # 通知全局订阅者 + for handler in self._global_subscribers: + try: + result = handler(event) + if asyncio.iscoroutine(result): + await result + except Exception: + logger.exception("Global event handler error") + + @property + def subscriber_count(self) -> int: + """当前订阅者总数。""" + count = len(self._global_subscribers) + for handlers in self._subscribers.values(): + count += len(handlers) + return count + + +class RealtimeStrategy: + """实时策略基类。 + + 用户子类实现 on_tick / on_bar 方法,在收到实时行情时 + 自动触发,可通过 event_bus 发布交易信号。 + + 用法:: + + class MyRealtimeStrategy(RealtimeStrategy): + def on_tick(self, event: MarketEvent) -> None: + if event.price > self.threshold: + self.emit_signal("BUY", event) + + strategy = MyRealtimeStrategy(threshold=10.5) + bus = EventBus() + bus.subscribe_all(strategy.on_tick) + """ + + def __init__(self, event_bus: EventBus | None = None, **kwargs: Any) -> None: + """初始化实时策略。 + + Args: + event_bus: 事件总线(可选,用于发送信号) + **kwargs: 用户自定义参数 + """ + self._event_bus = event_bus + self._params = kwargs + + def on_tick(self, event: MarketEvent) -> None: + """处理 tick 事件(用户实现)。 + + Args: + event: 行情事件 + """ + + def on_bar(self, event: MarketEvent) -> None: + """处理 bar 完成事件(用户实现)。 + + Args: + event: 行情事件 + """ + + def emit_signal( + self, + direction: str, + event: MarketEvent, + size: float = 0, + price: float | None = None, + ) -> None: + """发送交易信号事件。 + + Args: + direction: 交易方向 ("BUY" / "SELL") + event: 触发信号的行情事件 + size: 交易数量(0 = 全仓) + price: 信号价格(None = 市价) + """ + if self._event_bus is None: + logger.warning("No event bus configured, signal not sent") + return + + signal_event = MarketEvent( + event_type=EventType.SIGNAL, + code=event.code, + market=event.market, + price=price or event.price, + volume=size, + timestamp=event.timestamp, + data={"direction": direction}, + ) + # fire-and-forget signal publish + asyncio.ensure_future(self._event_bus.publish(signal_event)) diff --git a/tests/unit/test_realtime.py b/tests/unit/test_realtime.py new file mode 100644 index 0000000..f5c069d --- /dev/null +++ b/tests/unit/test_realtime.py @@ -0,0 +1,164 @@ +"""单元测试:实时数据推送引擎.""" + +from __future__ import annotations + +import asyncio + +import pytest + +from easy_tdx.realtime.engine import ( + EventBus, + EventType, + MarketEvent, + RealtimeStrategy, +) + + +class TestMarketEvent: + """测试行情事件.""" + + def test_event_creation(self) -> None: + """创建事件应正确设置字段.""" + event = MarketEvent( + event_type=EventType.TICK, + code="000001", + market="SZ", + price=10.5, + volume=1000, + timestamp=1700000000.0, + ) + + assert event.event_type == EventType.TICK + assert event.code == "000001" + assert event.market == "SZ" + assert event.price == 10.5 + + def test_event_default_values(self) -> None: + """默认值应为零和空字典.""" + event = MarketEvent( + event_type=EventType.BAR, + code="600000", + market="SH", + ) + + assert event.price == 0.0 + assert event.volume == 0.0 + assert event.data == {} + + +class TestEventBus: + """测试事件总线.""" + + def test_subscribe_and_count(self) -> None: + """订阅后计数应正确.""" + bus = EventBus() + bus.subscribe("SZ000001", lambda e: None) + bus.subscribe("SZ000001", lambda e: None) + bus.subscribe("SH600000", lambda e: None) + + assert bus.subscriber_count == 3 + + def test_subscribe_global(self) -> None: + """全局订阅应被计入.""" + bus = EventBus() + bus.subscribe_all(lambda e: None) + + assert bus.subscriber_count == 1 + + def test_unsubscribe(self) -> None: + """取消订阅后计数应减少.""" + bus = EventBus() + handler = lambda e: None # noqa: E731 + bus.subscribe("SZ000001", handler) + bus.unsubscribe("SZ000001", handler) + + assert bus.subscriber_count == 0 + + @pytest.mark.asyncio + async def test_publish_to_specific_subscriber(self) -> None: + """发布事件应通知特定订阅者.""" + bus = EventBus() + received: list[MarketEvent] = [] + + async def handler(event: MarketEvent) -> None: + received.append(event) + + bus.subscribe("SZ000001", handler) + + event = MarketEvent( + event_type=EventType.TICK, + code="000001", + market="SZ", + price=10.5, + ) + await bus.publish(event) + + assert len(received) == 1 + assert received[0].price == 10.5 + + @pytest.mark.asyncio + async def test_publish_to_global_subscriber(self) -> None: + """全局订阅者应收到所有事件.""" + bus = EventBus() + received: list[MarketEvent] = [] + + bus.subscribe_all(lambda e: received.append(e)) + + event = MarketEvent( + event_type=EventType.TICK, + code="000001", + market="SZ", + ) + await bus.publish(event) + + assert len(received) == 1 + + +class TestRealtimeStrategy: + """测试实时策略基类.""" + + def test_on_tick_default_does_nothing(self) -> None: + """默认 on_tick 不应抛异常.""" + strategy = RealtimeStrategy() + event = MarketEvent( + event_type=EventType.TICK, + code="000001", + market="SZ", + ) + strategy.on_tick(event) # should not raise + + @pytest.mark.asyncio + async def test_emit_signal_without_bus(self) -> None: + """无 event_bus 时 emit_signal 不应抛异常.""" + strategy = RealtimeStrategy() + event = MarketEvent( + event_type=EventType.TICK, + code="000001", + market="SZ", + price=10.0, + ) + strategy.emit_signal("BUY", event) + + @pytest.mark.asyncio + async def test_emit_signal_publishes_to_bus(self) -> None: + """emit_signal 应发布 SIGNAL 类型事件.""" + bus = EventBus() + received: list[MarketEvent] = [] + bus.subscribe("SZ000001", lambda e: received.append(e)) + + strategy = RealtimeStrategy(event_bus=bus) + event = MarketEvent( + event_type=EventType.TICK, + code="000001", + market="SZ", + price=10.0, + timestamp=1700000000.0, + ) + + strategy.emit_signal("BUY", event) + # Give the asyncio task time to run + await asyncio.sleep(0.1) + + assert len(received) == 1 + assert received[0].event_type == EventType.SIGNAL + assert received[0].data["direction"] == "BUY"