From 687851fc67e78f5eef89732b5c3251cb1f08b99c Mon Sep 17 00:00:00 2001 From: GitHub Date: Tue, 9 Jun 2026 16:53:57 +0800 Subject: [PATCH] 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 --- src/easy_tdx/backtest/strategy.py | 450 ++++++++++++++++++++++++++ tests/unit/test_backtest_strategy.py | 467 +++++++++++++++++++++++++++ 2 files changed, 917 insertions(+) create mode 100644 src/easy_tdx/backtest/strategy.py create mode 100644 tests/unit/test_backtest_strategy.py diff --git a/src/easy_tdx/backtest/strategy.py b/src/easy_tdx/backtest/strategy.py new file mode 100644 index 0000000..be03e10 --- /dev/null +++ b/src/easy_tdx/backtest/strategy.py @@ -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 diff --git a/tests/unit/test_backtest_strategy.py b/tests/unit/test_backtest_strategy.py new file mode 100644 index 0000000..6ba14b7 --- /dev/null +++ b/tests/unit/test_backtest_strategy.py @@ -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()