feat(backtest): add Strategy base class with DataProxy and crossover

- Add _SeriesAccessor for relative indexed data access ([0] current, [-1] previous)
- Add StrategyDataProxy for efficient DataFrame column access via numpy arrays
- Add crossover() function for golden cross detection (fast line crosses above slow line)
- Add Strategy abstract base class with:
  - init() for indicator registration via self.I()
  - next() for signal generation via buy()/sell()
  - Internal engine hooks (_bind_data, _call_init, _set_bar_index, etc.)
- All code is mypy strict compliant with full type annotations
- 25 unit tests covering all components

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
GitHub
2026-06-09 16:53:57 +08:00
co-authored by Claude Opus 4.8
parent f37b75ea42
commit 687851fc67
2 changed files with 917 additions and 0 deletions
+450
View File
@@ -0,0 +1,450 @@
"""策略基类和数据代理。
提供策略编写的声明式 API
- StrategyDataProxy: K线数据访问层,支持 OHLCV + 自定义指标列
- Strategy: 策略基类,init() 注册指标,next() 生成信号
- crossover: 金叉检测辅助函数
"""
from __future__ import annotations
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any
import numpy as np
import numpy.typing as npt
import pandas as pd
from easy_tdx.backtest.types import Signal
if TYPE_CHECKING:
from collections.abc import Callable
NDArray = npt.NDArray[np.float64]
else:
NDArray = np.ndarray
# ── 数据序列访问器 ─────────────────────────────────────────────────────────────
class _SeriesAccessor:
"""数据序列访问器,支持相对索引和 numpy 数组转换。
Examples:
>>> close = data.close
>>> close[0] # 当前 bar 的收盘价
>>> close[-1] # 前一根 bar 的收盘价
>>> ma = MyTT.MA(close.raw, 20) # 传入 numpy 数组
"""
__slots__ = ("_series", "_bar_index")
def __init__(self, series: NDArray, bar_index: int) -> None:
self._series = series
self._bar_index = bar_index
def __getitem__(self, key: int) -> float:
"""获取相对索引的值。
Args:
key: 0=当前值, -1=前一根, -2=前两根,依此类推
Returns:
对应位置的 float 值
"""
idx = self._bar_index + key
if idx < 0:
raise IndexError(f"索引 {key} 超出范围(bar_index={self._bar_index}")
return float(self._series[idx])
def __len__(self) -> int:
"""返回数组长度。"""
return len(self._series)
def __array__(self) -> NDArray:
"""允许传入 MyTT 函数(自动解包为 numpy 数组)。"""
return self._series
@property
def raw(self) -> NDArray:
"""获取完整原始 numpy 数组。"""
return self._series
# ── K线数据代理 ────────────────────────────────────────────────────────────────
class StrategyDataProxy:
"""K线数据代理,将 DataFrame 转为高效的 numpy 数组访问。
内部将所有列(除 datetime)转为 numpy 数组,通过 _SeriesAccessor
提供相对索引访问([0] 当前, [-1] 前一根)。
支持标准 OHLCV 列和任意自定义列(如 MACD_DIF, BOLL_UPPER)。
"""
__slots__ = ("_arrays", "_bar_index")
def __init__(self, df: pd.DataFrame) -> None:
"""初始化数据代理。
Args:
df: K线 DataFrame,必须包含 datetime, open, close, high, low, vol, amount
可包含额外列(如 MACD_DIF, BOLL_UPPER
"""
self._arrays: dict[str, NDArray] = {}
self._bar_index = 0
# 将所有列(除 datetime)转为 numpy 数组
for col in df.columns:
if col == "datetime":
continue
arr = df[col].to_numpy()
if len(arr) > 0 and isinstance(arr[0], (np.datetime64, pd.Timestamp)):
# datetime 列转为 int (YYYYMMDD)
self._arrays[col] = _datetime_to_int(arr)
else:
self._arrays[col] = arr.astype(np.float64)
def _set_index(self, idx: int) -> None:
"""设置当前 bar 索引(引擎调用)。"""
self._bar_index = idx
@property
def open(self) -> _SeriesAccessor:
"""开盘价序列。"""
return _SeriesAccessor(self._arrays["open"], self._bar_index)
@property
def close(self) -> _SeriesAccessor:
"""收盘价序列。"""
return _SeriesAccessor(self._arrays["close"], self._bar_index)
@property
def high(self) -> _SeriesAccessor:
"""最高价序列。"""
return _SeriesAccessor(self._arrays["high"], self._bar_index)
@property
def low(self) -> _SeriesAccessor:
"""最低价序列。"""
return _SeriesAccessor(self._arrays["low"], self._bar_index)
@property
def vol(self) -> _SeriesAccessor:
"""成交量序列。"""
return _SeriesAccessor(self._arrays["vol"], self._bar_index)
@property
def amount(self) -> _SeriesAccessor:
"""成交额序列。"""
return _SeriesAccessor(self._arrays["amount"], self._bar_index)
def __getattr__(self, name: str) -> _SeriesAccessor:
"""访问额外列(如 MACD_DIF, BOLL_UPPER)。
Raises:
AttributeError: 列不存在
"""
if name not in self._arrays:
raise AttributeError(f"'{name}' 不存在于数据中")
return _SeriesAccessor(self._arrays[name], self._bar_index)
# ── 金叉检测 ─────────────────────────────────────────────────────────────────────
def crossover(
a: NDArray | pd.Series | _SeriesAccessor,
b: NDArray | pd.Series | _SeriesAccessor,
) -> NDArray:
"""检测 a 从下方穿越 b(金叉)。
Args:
a: 快线序列
b: 慢线序列
Returns:
bool 数组,True 表示发生金叉
Examples:
>>> ma5 = MyTT.MA(data.close.raw, 5)
>>> ma20 = MyTT.MA(data.close.raw, 20)
>>> cross = crossover(ma5, ma20) # ma5 上穿 ma20
>>> if cross[bar_index]:
... strategy.buy(size=100)
"""
# 解包 _SeriesAccessor
if isinstance(a, _SeriesAccessor):
a = a.raw
if isinstance(b, _SeriesAccessor):
b = b.raw
# pd.Series 转 numpy
if isinstance(a, pd.Series):
a = a.to_numpy().astype(np.float64)
if isinstance(b, pd.Series):
b = b.to_numpy().astype(np.float64)
# 金叉:前一根 a <= b,当前 a > b
mask = np.zeros(len(a), dtype=bool)
mask[1:] = (a[:-1] <= b[:-1]) & (a[1:] > b[1:])
return mask
# ── 策略基类 ───────────────────────────────────────────────────────────────────
class Strategy(ABC):
"""策略基类,提供声明式回测 API。
用户子类实现:
- init(): 注册指标(通过 self.I()
- next(): 生成交易信号(通过 self.buy()/self.sell()
内部状态:
- self.data: StrategyDataProxy,访问 K线数据
- self.position: {"size": float},当前持仓
- self.chanlun: 缠论分析结果(预留)
示例:
>>> class MyStrategy(Strategy):
... def init(self):
... self.ma5 = self.I(MyTT.MA, self.data.close, 5)
... self.ma20 = self.I(MyTT.MA, self.data.close, 20)
... self.cross = crossover(self.ma5, self.ma20)
...
... def next(self):
... if self.cross[self._bar_index]:
... self.buy(size=100)
... elif self.position["size"] > 0:
... self.sell(size=0)
"""
def __init__(self) -> None:
"""初始化策略内部状态。"""
self._data_proxy: StrategyDataProxy | None = None
self._bar_index = 0
self._signals: list[Signal] = []
self._indicators: dict[str, NDArray] = {}
self._chanlun_result: dict[str, Any] | None = None
self._position_size = 0.0
self._cash = 0.0
self._datetime_array: NDArray | None = None
# ── 用户实现方法 ─────────────────────────────────────────────────────────────
@abstractmethod
def init(self) -> None:
"""注册指标。
在回测开始前调用一次,用户通过 self.I() 注册技术指标。
"""
@abstractmethod
def next(self) -> None:
"""生成交易信号。
每根 bar 调用一次,用户通过 self.buy()/self.sell() 生成信号。
"""
# ── 指标注册 ───────────────────────────────────────────────────────────────
def I(
self, func: Callable[..., NDArray], *args: Any, **kwargs: Any
) -> NDArray:
"""注册指标。
将 _SeriesAccessor 参数自动解包为 numpy 数组,调用 func 计算指标。
返回的数组存入 self._indicators,可在 next() 中访问。
Args:
func: 指标函数(如 MyTT.MA
*args: 参数(可能是 _SeriesAccessor,会自动解包)
**kwargs: 关键字参数
Returns:
指标值数组(numpy ndarray
Examples:
>>> ma5 = self.I(MyTT.MA, self.data.close, 5)
>>> macd = self.I(MyTT.MACD, self.data.close, self.data.low, self.data.high)
"""
# 解包 _SeriesAccessor 参数
unpacked_args = []
for arg in args:
if isinstance(arg, _SeriesAccessor):
unpacked_args.append(arg.raw)
else:
unpacked_args.append(arg)
# 调用函数
result = func(*unpacked_args, **kwargs)
# 存储指标(用于调试/日志)
func_name = getattr(func, "__name__", str(func))
self._indicators[func_name] = result
return result
# ── 交易信号生成 ─────────────────────────────────────────────────────────────
def buy(
self,
size: float = 0,
price: float | None = None,
stop_loss: float | None = None,
take_profit: float | None = None,
) -> None:
"""生成买入信号。
Args:
size: 交易数量(0 = 全仓,由引擎计算)
price: 限价(None = 市价单)
stop_loss: 止损价(None = 不设置)
take_profit: 止盈价(None = 不设置)
"""
if self._data_proxy is None:
raise RuntimeError("策略未绑定数据,请先调用 _bind_data()")
if self._datetime_array is None:
raise RuntimeError("数据未正确初始化")
signal = Signal(
datetime=int(self._datetime_array[self._bar_index]),
direction="BUY",
size=size,
price=price,
stop_loss=stop_loss,
take_profit=take_profit,
)
self._signals.append(signal)
def sell(
self,
size: float = 0,
price: float | None = None,
stop_loss: float | None = None,
take_profit: float | None = None,
) -> None:
"""生成卖出信号。
Args:
size: 交易数量(0 = 全仓,由引擎计算)
price: 限价(None = 市价单)
stop_loss: 止损价(None = 不设置)
take_profit: 止盈价(None = 不设置)
"""
if self._data_proxy is None:
raise RuntimeError("策略未绑定数据,请先调用 _bind_data()")
if self._datetime_array is None:
raise RuntimeError("数据未正确初始化")
signal = Signal(
datetime=int(self._datetime_array[self._bar_index]),
direction="SELL",
size=size,
price=price,
stop_loss=stop_loss,
take_profit=take_profit,
)
self._signals.append(signal)
# ── 状态访问 ─────────────────────────────────────────────────────────────────
@property
def data(self) -> StrategyDataProxy:
"""K线数据代理。"""
if self._data_proxy is None:
raise RuntimeError("策略未绑定数据,请先调用 _bind_data()")
return self._data_proxy
@property
def position(self) -> dict[str, float]:
"""当前持仓(简化 dict 格式)。
Returns:
{"size": float},正=多头,负=空头,0=空仓
"""
return {"size": self._position_size}
@property
def chanlun(self) -> dict[str, Any] | None:
"""缠论分析结果(预留)。"""
return self._chanlun_result
# ── 内部方法(引擎调用) ─────────────────────────────────────────────────────
def _bind_data(self, df: pd.DataFrame) -> None:
"""绑定 K线数据(引擎调用)。
Args:
df: K线 DataFrame
"""
self._data_proxy = StrategyDataProxy(df)
# 提取 datetime 数组
if "datetime" in df.columns:
dt_arr = df["datetime"].to_numpy()
self._datetime_array = _datetime_to_int(dt_arr)
else:
raise ValueError("DataFrame 必须包含 datetime 列")
def _call_init(self) -> None:
"""调用用户 init() 方法(引擎调用)。"""
self.init()
def _set_bar_index(self, idx: int) -> None:
"""设置当前 bar 索引(引擎调用)。
Args:
idx: bar 索引
"""
self._bar_index = idx
if self._data_proxy is not None:
self._data_proxy._set_index(idx)
def _call_next(self) -> None:
"""调用用户 next() 方法(引擎调用)。"""
self.next()
def _get_datetime(self) -> int:
"""获取当前 bar 的 datetime(引擎调用)。
Returns:
datetime int (YYYYMMDD)
"""
if self._datetime_array is None:
raise RuntimeError("数据未正确初始化")
return int(self._datetime_array[self._bar_index])
def _clear_signals(self) -> list[Signal]:
"""清空并返回已生成的信号(引擎调用)。
Returns:
当前累积的信号列表
"""
signals = self._signals
self._signals = []
return signals
# ── 辅助函数 ─────────────────────────────────────────────────────────────────────
def _datetime_to_int(arr: NDArray) -> NDArray:
"""将 datetime 数组转为 int (YYYYMMDD)。
Args:
arr: datetime 数组(np.datetime64 或 pd.Timestamp
Returns:
int 数组,格式 YYYYMMDD
"""
result = np.zeros(len(arr), dtype=np.float64)
for i, val in enumerate(arr):
if isinstance(val, (np.datetime64, pd.Timestamp)):
ts = pd.Timestamp(val)
result[i] = float(ts.strftime("%Y%m%d"))
else:
# 已经是 int 或可转为 int
result[i] = float(val)
return result