mirror of
https://ghfast.top/https://github.com/aeroxw/easy_tdx_max.git
synced 2026-09-12 15:44:18 +08:00
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:
co-authored by
Claude Opus 4.8
parent
f37b75ea42
commit
687851fc67
@@ -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
|
||||
@@ -0,0 +1,467 @@
|
||||
"""测试策略基类和数据代理。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from easy_tdx.backtest.strategy import (
|
||||
Strategy,
|
||||
StrategyDataProxy,
|
||||
_SeriesAccessor,
|
||||
crossover,
|
||||
)
|
||||
from easy_tdx.backtest.types import Signal
|
||||
|
||||
|
||||
# ── 辅助函数 ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _make_df(n: int = 20, seed: int = 42) -> pd.DataFrame:
|
||||
"""构造随机 OHLCV DataFrame。
|
||||
|
||||
Args:
|
||||
n: K线数量
|
||||
seed: 随机种子
|
||||
|
||||
Returns:
|
||||
DataFrame with columns: datetime, open, close, high, low, vol, amount
|
||||
"""
|
||||
rng = np.random.default_rng(seed)
|
||||
base = 10.0
|
||||
prices = base + rng.random(n) * 5
|
||||
|
||||
df = pd.DataFrame(
|
||||
{
|
||||
"datetime": pd.date_range("2024-01-01", periods=n, freq="D"),
|
||||
"open": prices,
|
||||
"close": prices + rng.random(n) - 0.5,
|
||||
"high": prices + rng.random(n),
|
||||
"low": prices - rng.random(n),
|
||||
"vol": rng.integers(1000, 10000, n),
|
||||
"amount": rng.integers(100000, 1000000, n),
|
||||
}
|
||||
)
|
||||
# 确保 high >= max(open, close), low <= min(open, close)
|
||||
df["high"] = df[["open", "close", "high"]].max(axis=1)
|
||||
df["low"] = df[["open", "close", "low"]].min(axis=1)
|
||||
return df
|
||||
|
||||
|
||||
def _make_df_with_extras(n: int = 20) -> pd.DataFrame:
|
||||
"""构造带额外列的 DataFrame(MACD_DIF, BOLL_UPPER)。
|
||||
|
||||
Args:
|
||||
n: K线数量
|
||||
|
||||
Returns:
|
||||
DataFrame with standard OHLCV columns + MACD_DIF, BOLL_UPPER
|
||||
"""
|
||||
df = _make_df(n)
|
||||
rng = np.random.default_rng(42)
|
||||
|
||||
# 添加额外列
|
||||
df["MACD_DIF"] = rng.random(n) * 2 - 1
|
||||
df["BOLL_UPPER"] = df["close"] + rng.random(n) * 2
|
||||
|
||||
return df
|
||||
|
||||
|
||||
# ── TestSeriesAccessor ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSeriesAccessor:
|
||||
"""测试 _SeriesAccessor 数据访问器。"""
|
||||
|
||||
def test_current_value(self) -> None:
|
||||
"""测试获取当前值 [0]。"""
|
||||
arr = np.array([1.0, 2.0, 3.0, 4.0, 5.0])
|
||||
acc = _SeriesAccessor(arr, bar_index=2)
|
||||
assert acc[0] == 3.0
|
||||
|
||||
def test_previous_value(self) -> None:
|
||||
"""测试获取前一根 [-1]。"""
|
||||
arr = np.array([1.0, 2.0, 3.0, 4.0, 5.0])
|
||||
acc = _SeriesAccessor(arr, bar_index=2)
|
||||
assert acc[-1] == 2.0
|
||||
|
||||
def test_previous_two(self) -> None:
|
||||
"""测试获取前两根 [-2]。"""
|
||||
arr = np.array([1.0, 2.0, 3.0, 4.0, 5.0])
|
||||
acc = _SeriesAccessor(arr, bar_index=2)
|
||||
assert acc[-2] == 1.0
|
||||
|
||||
def test_index_out_of_bounds_negative(self) -> None:
|
||||
"""测试索引越界(负方向)。"""
|
||||
arr = np.array([1.0, 2.0, 3.0, 4.0, 5.0])
|
||||
acc = _SeriesAccessor(arr, bar_index=0)
|
||||
with pytest.raises(IndexError, match="索引 -1 超出范围"):
|
||||
_ = acc[-1]
|
||||
|
||||
def test_len(self) -> None:
|
||||
"""测试 __len__ 返回数组长度。"""
|
||||
arr = np.array([1.0, 2.0, 3.0])
|
||||
acc = _SeriesAccessor(arr, bar_index=1)
|
||||
assert len(acc) == 3
|
||||
|
||||
def test_array_conversion(self) -> None:
|
||||
"""测试 __array__ 转为 numpy 数组。"""
|
||||
arr = np.array([1.0, 2.0, 3.0])
|
||||
acc = _SeriesAccessor(arr, bar_index=1)
|
||||
result = np.asarray(acc)
|
||||
np.testing.assert_array_equal(result, arr)
|
||||
|
||||
def test_raw_property(self) -> None:
|
||||
"""测试 raw 属性返回原始数组。"""
|
||||
arr = np.array([1.0, 2.0, 3.0])
|
||||
acc = _SeriesAccessor(arr, bar_index=1)
|
||||
np.testing.assert_array_equal(acc.raw, arr)
|
||||
|
||||
|
||||
# ── TestStrategyDataProxy ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestStrategyDataProxy:
|
||||
"""测试 StrategyDataProxy 数据代理。"""
|
||||
|
||||
def test_basic_columns(self) -> None:
|
||||
"""测试标准 OHLCV 列访问。"""
|
||||
df = _make_df(n=10)
|
||||
proxy = StrategyDataProxy(df)
|
||||
proxy._set_index(5)
|
||||
|
||||
assert isinstance(proxy.open[0], float)
|
||||
assert isinstance(proxy.close[0], float)
|
||||
assert isinstance(proxy.high[0], float)
|
||||
assert isinstance(proxy.low[0], float)
|
||||
assert isinstance(proxy.vol[0], float)
|
||||
assert isinstance(proxy.amount[0], float)
|
||||
|
||||
def test_previous_bar(self) -> None:
|
||||
"""测试访问前一根 bar 的数据。"""
|
||||
df = _make_df(n=10)
|
||||
proxy = StrategyDataProxy(df)
|
||||
proxy._set_index(5)
|
||||
|
||||
# close[0] 应该等于 df["close"].iloc[5]
|
||||
assert proxy.close[0] == pytest.approx(df["close"].iloc[5])
|
||||
# close[-1] 应该等于 df["close"].iloc[4]
|
||||
assert proxy.close[-1] == pytest.approx(df["close"].iloc[4])
|
||||
|
||||
def test_extra_columns_via_getattr(self) -> None:
|
||||
"""测试通过 __getattr__ 访问额外列。"""
|
||||
df = _make_df_with_extras(n=10)
|
||||
proxy = StrategyDataProxy(df)
|
||||
proxy._set_index(5)
|
||||
|
||||
# MACD_DIF 和 BOLL_UPPER 应该可以访问
|
||||
assert isinstance(proxy.MACD_DIF[0], float)
|
||||
assert isinstance(proxy.BOLL_UPPER[0], float)
|
||||
|
||||
# 验证值正确
|
||||
assert proxy.MACD_DIF[0] == pytest.approx(df["MACD_DIF"].iloc[5])
|
||||
assert proxy.BOLL_UPPER[0] == pytest.approx(df["BOLL_UPPER"].iloc[5])
|
||||
|
||||
def test_missing_column_raises(self) -> None:
|
||||
"""测试访问不存在的列抛出 AttributeError。"""
|
||||
df = _make_df(n=10)
|
||||
proxy = StrategyDataProxy(df)
|
||||
proxy._set_index(5)
|
||||
|
||||
with pytest.raises(AttributeError, match="列 'NONEXISTENT' 不存在"):
|
||||
_ = proxy.NONEXISTENT[0]
|
||||
|
||||
|
||||
# ── TestCrossover ───────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCrossover:
|
||||
"""测试 crossover 金叉检测函数。"""
|
||||
|
||||
def test_crossover_true(self) -> None:
|
||||
"""测试金叉检测(a 上穿 b)。"""
|
||||
a = np.array([1, 2, 3, 4, 5])
|
||||
b = np.array([5, 4, 3, 2, 1])
|
||||
mask = crossover(a, b)
|
||||
|
||||
# 索引 3 处:a[2]=3 <= b[2]=3, a[3]=4 > b[3]=2 → 金叉
|
||||
assert mask[3] == True
|
||||
# 其他位置无金叉
|
||||
assert mask[0] == False
|
||||
assert mask[1] == False
|
||||
assert mask[2] == False
|
||||
assert mask[4] == False
|
||||
|
||||
def test_crossover_false_no_cross(self) -> None:
|
||||
"""测试无金叉情况(a 全在 b 下方)。"""
|
||||
a = np.array([1, 2, 3, 4, 5])
|
||||
b = np.array([6, 7, 8, 9, 10])
|
||||
mask = crossover(a, b)
|
||||
|
||||
# 全部 False
|
||||
assert not np.any(mask)
|
||||
|
||||
def test_crossover_series(self) -> None:
|
||||
"""测试接受 pd.Series 参数。"""
|
||||
a = pd.Series([1, 2, 3, 4, 5])
|
||||
b = pd.Series([5, 4, 3, 2, 1])
|
||||
mask = crossover(a, b)
|
||||
|
||||
assert mask[3] == True
|
||||
|
||||
def test_crossover_with_accessor(self) -> None:
|
||||
"""测试接受 _SeriesAccessor 参数。"""
|
||||
arr_a = np.array([1, 2, 3, 4, 5])
|
||||
arr_b = np.array([5, 4, 3, 2, 1])
|
||||
acc_a = _SeriesAccessor(arr_a, bar_index=4)
|
||||
acc_b = _SeriesAccessor(arr_b, bar_index=4)
|
||||
|
||||
mask = crossover(acc_a, acc_b)
|
||||
assert mask[3] == True
|
||||
|
||||
|
||||
# ── TestStrategyBase ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestStrategyBase:
|
||||
"""测试 Strategy 策略基类。"""
|
||||
|
||||
def test_subclass_init_and_next(self) -> None:
|
||||
"""测试子类 init() 和 next() 被调用。"""
|
||||
|
||||
class SimpleStrategy(Strategy):
|
||||
init_called = False
|
||||
next_called_count = 0
|
||||
|
||||
def init(self) -> None:
|
||||
SimpleStrategy.init_called = True
|
||||
|
||||
def next(self) -> None:
|
||||
SimpleStrategy.next_called_count += 1
|
||||
|
||||
df = _make_df(n=10)
|
||||
strategy = SimpleStrategy()
|
||||
strategy._bind_data(df)
|
||||
strategy._call_init()
|
||||
|
||||
assert SimpleStrategy.init_called is True
|
||||
|
||||
# 模拟引擎遍历所有 bar
|
||||
for i in range(len(df)):
|
||||
strategy._set_bar_index(i)
|
||||
strategy._call_next()
|
||||
|
||||
assert SimpleStrategy.next_called_count == 10
|
||||
|
||||
def test_buy_sell_recording(self) -> None:
|
||||
"""测试 buy() 和 sell() 创建 Signal。"""
|
||||
|
||||
class SignalStrategy(Strategy):
|
||||
def init(self) -> None:
|
||||
pass
|
||||
|
||||
def next(self) -> None:
|
||||
if self._bar_index == 5:
|
||||
self.buy(size=100, price=10.0, stop_loss=9.0, take_profit=11.0)
|
||||
elif self._bar_index == 8:
|
||||
self.sell(size=100, price=10.5)
|
||||
|
||||
df = _make_df(n=10)
|
||||
strategy = SignalStrategy()
|
||||
strategy._bind_data(df)
|
||||
strategy._call_init()
|
||||
|
||||
for i in range(len(df)):
|
||||
strategy._set_bar_index(i)
|
||||
strategy._call_next()
|
||||
|
||||
signals = strategy._clear_signals()
|
||||
|
||||
assert len(signals) == 2
|
||||
assert signals[0].direction == "BUY"
|
||||
assert signals[0].size == 100
|
||||
assert signals[0].price == 10.0
|
||||
assert signals[0].stop_loss == 9.0
|
||||
assert signals[0].take_profit == 11.0
|
||||
|
||||
assert signals[1].direction == "SELL"
|
||||
assert signals[1].size == 100
|
||||
assert signals[1].price == 10.5
|
||||
|
||||
def test_I_registers_indicator(self) -> None:
|
||||
"""测试 self.I() 注册指标。"""
|
||||
from easy_tdx.MyTT import MA
|
||||
|
||||
class IndicatorStrategy(Strategy):
|
||||
def init(self) -> None:
|
||||
self.ma5 = self.I(MA, self.data.close, 5)
|
||||
self.ma20 = self.I(MA, self.data.close, 20)
|
||||
|
||||
def next(self) -> None:
|
||||
pass
|
||||
|
||||
df = _make_df(n=50)
|
||||
strategy = IndicatorStrategy()
|
||||
strategy._bind_data(df)
|
||||
strategy._call_init()
|
||||
|
||||
# 验证指标长度正确(MA 会产生 nan,但长度不变)
|
||||
assert len(strategy.ma5) == 50
|
||||
assert len(strategy.ma20) == 50
|
||||
|
||||
# 验证指标已注册
|
||||
assert "MA" in strategy._indicators
|
||||
|
||||
def test_full_position_buy(self) -> None:
|
||||
"""测试全仓买入(size=0)。"""
|
||||
|
||||
class FullPositionStrategy(Strategy):
|
||||
def init(self) -> None:
|
||||
pass
|
||||
|
||||
def next(self) -> None:
|
||||
if self._bar_index == 5:
|
||||
self.buy(size=0) # 全仓
|
||||
|
||||
df = _make_df(n=10)
|
||||
strategy = FullPositionStrategy()
|
||||
strategy._bind_data(df)
|
||||
strategy._call_init()
|
||||
|
||||
for i in range(len(df)):
|
||||
strategy._set_bar_index(i)
|
||||
strategy._call_next()
|
||||
|
||||
signals = strategy._clear_signals()
|
||||
|
||||
assert len(signals) == 1
|
||||
assert signals[0].direction == "BUY"
|
||||
assert signals[0].size == 0 # 0 表示全仓
|
||||
|
||||
def test_data_property(self) -> None:
|
||||
"""测试 self.data 属性返回 StrategyDataProxy。"""
|
||||
df = _make_df(n=10)
|
||||
|
||||
class DataAccessStrategy(Strategy):
|
||||
def init(self) -> None:
|
||||
# 验证 data 可访问
|
||||
assert hasattr(self.data, "close")
|
||||
assert hasattr(self.data, "open")
|
||||
|
||||
def next(self) -> None:
|
||||
# 验证 next() 中也可访问
|
||||
assert self.data.close[0] > 0
|
||||
|
||||
strategy = DataAccessStrategy()
|
||||
strategy._bind_data(df)
|
||||
strategy._call_init()
|
||||
|
||||
for i in range(len(df)):
|
||||
strategy._set_bar_index(i)
|
||||
strategy._call_next()
|
||||
|
||||
def test_position_property(self) -> None:
|
||||
"""测试 self.position 属性。"""
|
||||
df = _make_df(n=10)
|
||||
|
||||
class PositionStrategy(Strategy):
|
||||
def init(self) -> None:
|
||||
pass
|
||||
|
||||
def next(self) -> None:
|
||||
# 初始状态无持仓
|
||||
assert self.position["size"] == 0.0
|
||||
|
||||
strategy = PositionStrategy()
|
||||
strategy._bind_data(df)
|
||||
strategy._call_init()
|
||||
|
||||
for i in range(len(df)):
|
||||
strategy._set_bar_index(i)
|
||||
strategy._call_next()
|
||||
|
||||
def test_get_datetime(self) -> None:
|
||||
"""测试 _get_datetime() 方法。"""
|
||||
df = _make_df(n=10)
|
||||
|
||||
class DateTimeStrategy(Strategy):
|
||||
def init(self) -> None:
|
||||
pass
|
||||
|
||||
def next(self) -> None:
|
||||
# 验证 datetime 正确(YYYYMMDD 格式)
|
||||
dt = self._get_datetime()
|
||||
assert isinstance(dt, int)
|
||||
assert 20240101 <= dt <= 20241231
|
||||
|
||||
strategy = DateTimeStrategy()
|
||||
strategy._bind_data(df)
|
||||
strategy._call_init()
|
||||
|
||||
for i in range(len(df)):
|
||||
strategy._set_bar_index(i)
|
||||
strategy._call_next()
|
||||
|
||||
def test_clear_signals(self) -> None:
|
||||
"""测试 _clear_signals() 清空信号列表。"""
|
||||
|
||||
class ClearSignalStrategy(Strategy):
|
||||
def init(self) -> None:
|
||||
pass
|
||||
|
||||
def next(self) -> None:
|
||||
self.buy(size=100)
|
||||
|
||||
df = _make_df(n=10)
|
||||
strategy = ClearSignalStrategy()
|
||||
strategy._bind_data(df)
|
||||
strategy._call_init()
|
||||
|
||||
# 第一轮
|
||||
for i in range(len(df)):
|
||||
strategy._set_bar_index(i)
|
||||
strategy._call_next()
|
||||
signals1 = strategy._clear_signals()
|
||||
assert len(signals1) == 10
|
||||
|
||||
# 第二轮(无新信号)
|
||||
signals2 = strategy._clear_signals()
|
||||
assert len(signals2) == 0
|
||||
|
||||
def test_buy_without_bind_raises(self) -> None:
|
||||
"""测试未绑定数据时调用 buy() 抛出错误。"""
|
||||
|
||||
class BadStrategy(Strategy):
|
||||
def init(self) -> None:
|
||||
pass
|
||||
|
||||
def next(self) -> None:
|
||||
self.buy(size=100)
|
||||
|
||||
strategy = BadStrategy()
|
||||
# 未调用 _bind_data()
|
||||
|
||||
with pytest.raises(RuntimeError, match="策略未绑定数据"):
|
||||
strategy._call_init()
|
||||
strategy._set_bar_index(0)
|
||||
strategy._call_next()
|
||||
|
||||
def test_chanlun_property(self) -> None:
|
||||
"""测试 chanlun 属性(预留)。"""
|
||||
df = _make_df(n=10)
|
||||
|
||||
class ChanlunStrategy(Strategy):
|
||||
def init(self) -> None:
|
||||
assert self.chanlun is None
|
||||
|
||||
def next(self) -> None:
|
||||
pass
|
||||
|
||||
strategy = ChanlunStrategy()
|
||||
strategy._bind_data(df)
|
||||
strategy._call_init()
|
||||
|
||||
for i in range(len(df)):
|
||||
strategy._set_bar_index(i)
|
||||
strategy._call_next()
|
||||
Reference in New Issue
Block a user