mirror of
https://ghfast.top/https://github.com/aeroxw/easy-tdx.git
synced 2026-09-12 15:44:15 +08:00
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
This commit is contained in:
@@ -0,0 +1,10 @@
|
|||||||
|
"""实时数据推送模块。
|
||||||
|
|
||||||
|
提供事件驱动的行情推送框架,基于 asyncio 实现。
|
||||||
|
这是 API 骨架设计,实际协议层订阅功能待实现。
|
||||||
|
|
||||||
|
核心组件:
|
||||||
|
- EventBus: 事件总线,发布/订阅行情事件
|
||||||
|
- RealtimeStrategy: 实时策略基类
|
||||||
|
- MarketEvent: 行情事件数据结构
|
||||||
|
"""
|
||||||
@@ -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))
|
||||||
@@ -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"
|
||||||
Reference in New Issue
Block a user