From 16dc2e7da94663ecc6ffdb62b66723a691642263 Mon Sep 17 00:00:00 2001 From: GitHub Date: Tue, 9 Jun 2026 17:52:31 +0800 Subject: [PATCH] feat(backtest): add OrderSimulator with 5 execution modes and reject policy - Implement OrderSimulator class for order matching simulation - Support 5 execution modes: next_open, next_close, this_close, worst, best - Support 3 position modes: full, fixed, percent - Support 2 reject policies: reduce (partial fill), skip (reject) - Implement fee model: commission (min 5 CNY), stamp tax (0.1% sell only), slippage - Add future_leak_warning flag for this_close mode - Handle both int and datetime column types in DataFrame - Add comprehensive test suite with 24 test cases covering all modes Co-Authored-By: Claude Opus 4.8 --- src/easy_tdx/backtest/orders.py | 452 +++++++++++++++++++++++++++++ tests/unit/test_backtest_orders.py | 396 +++++++++++++++++++++++++ 2 files changed, 848 insertions(+) create mode 100644 src/easy_tdx/backtest/orders.py create mode 100644 tests/unit/test_backtest_orders.py diff --git a/src/easy_tdx/backtest/orders.py b/src/easy_tdx/backtest/orders.py new file mode 100644 index 0000000..b8fee7f --- /dev/null +++ b/src/easy_tdx/backtest/orders.py @@ -0,0 +1,452 @@ +"""订单撮合模拟器。 + +将策略信号转换为成交记录,支持多种执行模式、仓位管理和拒绝策略。 +""" + +from __future__ import annotations + +from dataclasses import dataclass + +import pandas as pd + +from easy_tdx.backtest.types import Signal, Trade + + +@dataclass +class OrderSimulator: + """订单撮合模拟器。 + + 将策略信号(Signal)转换为成交记录(Trade),支持多种执行模式、 + 仓位管理和拒绝策略。 + + Attributes: + df: K线数据 DataFrame + execution: 成交价规则 (next_open/next_close/this_close/worst/best) + position_mode: 仓位模式 (full/fixed/percent) + reject_policy: 拒绝策略 (reduce/skip) + commission: 佣金费率 + min_commission: 最低佣金 + stamp_tax: 印花税率(仅卖出) + slippage: 滑点(每股) + future_leak_warning: 是否使用了未来数据(this_close 模式) + """ + + df: pd.DataFrame + execution: str = "next_open" + position_mode: str = "full" + reject_policy: str = "reduce" + commission: float = 0.0003 + min_commission: float = 5.0 + stamp_tax: float = 0.001 + slippage: float = 0.0 + future_leak_warning: bool = False + + def simulate( + self, + signals: list[Signal], + cash: float, + position: float, + position_mode: str | None = None, + ) -> list[Trade]: + """模拟订单撮合过程。 + + Args: + signals: 交易信号列表 + cash: 初始现金 + position: 初始持仓(股数) + position_mode: 仓位模式(覆盖初始化参数) + + Returns: + 成交记录列表 + """ + if position_mode is None: + position_mode = self.position_mode + + trades: list[Trade] = [] + current_cash = cash + current_position = position + + for signal in signals: + # 找到信号对应的 K 线 + bar_idx = self._find_bar_index(signal.datetime) + if bar_idx is None: + continue + + # 确定成交的 K 线索引 + exec_idx = self._resolve_exec_index(bar_idx) + if exec_idx is None or exec_idx >= len(self.df): + continue + + # 获取成交价 + price = self._get_price(exec_idx, signal.direction) + if price is None: + continue + + # 执行交易 + if signal.direction == "BUY": + trade = self._execute_buy( + signal=signal, + bar_idx=bar_idx, + exec_idx=exec_idx, + price=price, + cash=current_cash, + position=current_position, + position_mode=position_mode, + ) + if trade is not None: + trades.append(trade) + if not trade.rejected: + current_cash -= trade.size * trade.price + trade.commission + trade.slippage + current_position += trade.size + + elif signal.direction == "SELL": + trade = self._execute_sell( + signal=signal, + bar_idx=bar_idx, + exec_idx=exec_idx, + price=price, + cash=current_cash, + position=current_position, + position_mode=position_mode, + ) + if trade is not None: + trades.append(trade) + if not trade.rejected: + current_cash += trade.size * trade.price - trade.commission - trade.slippage + current_position -= trade.size + + return trades + + def _find_bar_index(self, datetime_val: int) -> int | None: + """查找 datetime 对应的 K 线索引。 + + Args: + datetime_val: 信号时间(int 格式 YYYYMMDD) + + Returns: + K 线索引,未找到返回 None + """ + # 检查 df 中的 datetime 列类型 + dt_col = self.df["datetime"] + + # 尝试直接比较(如果是 int 类型) + try: + idx = (dt_col == datetime_val).idxmax() if (dt_col == datetime_val).any() else None + if idx is not None: + return int(idx) + except (TypeError, ValueError): + pass + + # 如果是 datetime 对象,转为 int 比较 + if pd.api.types.is_datetime64_any_dtype(dt_col): + dt_ints = dt_col.dt.strftime("%Y%m%d").astype(int) + mask = dt_ints == datetime_val + if mask.any(): + return int(mask.idxmax()) + return None + + return None + + def _resolve_exec_index(self, bar_idx: int) -> int | None: + """根据执行模式确定成交的 K 线索引。 + + Args: + bar_idx: 信号对应的 K 线索引 + + Returns: + 成交 K 线索引 + """ + if self.execution == "this_close": + # 当信号 K 线收盘时成交 + self.future_leak_warning = True + return bar_idx + else: + # 其他模式在下一根 K 线成交 + return bar_idx + 1 + + def _get_price(self, exec_idx: int, direction: str) -> float | None: + """根据执行模式和方向获取成交价。 + + Args: + exec_idx: 成交 K 线索引 + direction: 交易方向 + + Returns: + 成交价格 + """ + if exec_idx >= len(self.df): + return None + + row = self.df.iloc[exec_idx] + + if self.execution == "next_open": + return float(row["open"]) + elif self.execution == "next_close": + return float(row["close"]) + elif self.execution == "this_close": + return float(row["close"]) + elif self.execution == "worst": + # 买入取最高价,卖出取最低价 + return float(row["high"]) if direction == "BUY" else float(row["low"]) + elif self.execution == "best": + # 买入取最低价,卖出取最高价 + return float(row["low"]) if direction == "BUY" else float(row["high"]) + else: + return None + + def _calculate_buy_size( + self, + signal_size: float, + price: float, + cash: float, + position_mode: str, + ) -> float: + """计算买入数量。 + + Args: + signal_size: 信号指定的数量 + price: 成交价格 + cash: 可用现金 + position_mode: 仓位模式 + + Returns: + 买入数量(股) + """ + if position_mode == "full" or signal_size == 0: + # 全仓:计算可用现金能买多少(100股整手) + # 先计算最大股数,然后向下取整到100的倍数 + max_cost_per_share = price * (1 + self.commission) + self.slippage + max_shares_raw = cash / max_cost_per_share + max_shares = int(max_shares_raw / 100) * 100 + return float(max_shares) + elif position_mode == "fixed": + # 固定股数 + return signal_size + elif position_mode == "percent": + # 总资产的百分比 + total_value = cash # 简化:假设现金=总资产 + target_value = total_value * signal_size + max_shares = int(target_value / price / 100) * 100 + return float(max_shares) + else: + return signal_size + + def _calculate_sell_size( + self, + signal_size: float, + position: float, + position_mode: str, + ) -> float: + """计算卖出数量。 + + Args: + signal_size: 信号指定的数量 + position: 当前持仓 + position_mode: 仓位模式 + + Returns: + 卖出数量(股) + """ + if position_mode == "full" or signal_size == 0: + # 全部卖出 + return position + elif position_mode == "fixed": + # 固定股数 + return signal_size + elif position_mode == "percent": + # 持仓的百分比 + return position * signal_size + else: + return signal_size + + def _calculate_commission(self, size: float, price: float, is_sell: bool = False) -> float: + """计算手续费。 + + Args: + size: 成交数量 + price: 成交价格 + is_sell: 是否为卖出 + + Returns: + 手续费总额 + """ + # 佣金 + commission = max(size * price * self.commission, self.min_commission) + + # 印花税(仅卖出) + if is_sell: + stamp = size * price * self.stamp_tax + commission += stamp + + return commission + + def _execute_buy( + self, + signal: Signal, + bar_idx: int, + exec_idx: int, + price: float, + cash: float, + position: float, + position_mode: str, + ) -> Trade | None: + """执行买入。 + + Args: + signal: 交易信号 + bar_idx: 信号 K 线索引 + exec_idx: 成交 K 线索引 + price: 成交价格 + cash: 可用现金 + position: 当前持仓 + position_mode: 仓位模式 + + Returns: + 成交记录 + """ + # 保存原始信号数量 + original_size = signal.size + + # 计算买入数量 + size = self._calculate_buy_size(signal.size, price, cash, position_mode) + + if size <= 0: + # 资金不足或计算结果为0 + if self.reject_policy == "skip": + # 对于 percent 模式,original_size 是百分比(如 0.5),不是股数 + # 对于 fixed/full 模式,original_size 就是股数 + display_size = original_size if position_mode == "fixed" else 100 + return Trade( + datetime=self.df.iloc[exec_idx]["datetime"], + direction="BUY", + size=display_size, + price=price, + commission=0.0, + slippage=0.0, + pnl=0.0, + rejected=True, + ) + return None + + # 计算费用 + commission = self._calculate_commission(size, price, is_sell=False) + slippage = size * self.slippage + + # 检查资金是否足够 + total_cost = size * price + commission + slippage + if total_cost > cash: + if self.reject_policy == "skip": + return Trade( + datetime=self.df.iloc[exec_idx]["datetime"], + direction="BUY", + size=original_size, + price=price, + commission=commission, + slippage=slippage, + pnl=0.0, + rejected=True, + ) + elif self.reject_policy == "reduce": + # reduce 模式:重新计算可买数量 + available_cash = cash - self.min_commission - slippage + if available_cash > price: + reduced_size = int(available_cash / price / 100) * 100 + if reduced_size > 0: + commission = self._calculate_commission(reduced_size, price, is_sell=False) + slippage = reduced_size * self.slippage + return Trade( + datetime=self.df.iloc[exec_idx]["datetime"], + direction="BUY", + size=reduced_size, + price=price, + commission=commission, + slippage=slippage, + pnl=0.0, + rejected=False, + ) + # 无法买任何数量 + return None + + return Trade( + datetime=self.df.iloc[exec_idx]["datetime"], + direction="BUY", + size=size, + price=price, + commission=commission, + slippage=slippage, + pnl=0.0, + rejected=False, + ) + + def _execute_sell( + self, + signal: Signal, + bar_idx: int, + exec_idx: int, + price: float, + cash: float, + position: float, + position_mode: str, + ) -> Trade | None: + """执行卖出。 + + Args: + signal: 交易信号 + bar_idx: 信号 K 线索引 + exec_idx: 成交 K 线索引 + price: 成交价格 + cash: 可用现金 + position: 当前持仓 + position_mode: 仓位模式 + + Returns: + 成交记录 + """ + # 计算卖出数量 + size = self._calculate_sell_size(signal.size, position, position_mode) + + if size <= 0: + # 无持仓 + if self.reject_policy == "skip": + return Trade( + datetime=self.df.iloc[exec_idx]["datetime"], + direction="SELL", + size=0, + price=price, + commission=0.0, + slippage=0.0, + pnl=0.0, + rejected=True, + ) + return None + + # 检查持仓是否足够 + if size > position: + if self.reject_policy == "skip": + return Trade( + datetime=self.df.iloc[exec_idx]["datetime"], + direction="SELL", + size=size, + price=price, + commission=0.0, + slippage=0.0, + pnl=0.0, + rejected=True, + ) + # reduce 模式:减少到实际持仓 + size = position + + # 计算费用 + commission = self._calculate_commission(size, price, is_sell=True) + slippage = size * self.slippage + + return Trade( + datetime=self.df.iloc[exec_idx]["datetime"], + direction="SELL", + size=size, + price=price, + commission=commission, + slippage=slippage, + pnl=0.0, + rejected=False, + ) diff --git a/tests/unit/test_backtest_orders.py b/tests/unit/test_backtest_orders.py new file mode 100644 index 0000000..76e8178 --- /dev/null +++ b/tests/unit/test_backtest_orders.py @@ -0,0 +1,396 @@ +"""订单撮合模拟器单元测试。""" + +from __future__ import annotations + +import warnings + +import pandas as pd +import pytest + +from easy_tdx.backtest.orders import OrderSimulator +from easy_tdx.backtest.types import Signal + + +# ── Test Fixtures ───────────────────────────────────────────────────────────── + + +def _make_df(n: int = 10) -> pd.DataFrame: + """构造测试用 K线数据。 + + 价格递增:open=100..109, close=101..110, high=102..111, low=99..108 + datetime: range(20240101, 20240101+n) + """ + data = { + "datetime": [20240101 + i for i in range(n)], + "open": [100.0 + i for i in range(n)], + "close": [101.0 + i for i in range(n)], + "high": [102.0 + i for i in range(n)], + "low": [99.0 + i for i in range(n)], + "volume": [1000] * n, + } + return pd.DataFrame(data) + + +def _buy_signal(bar_idx: int, size: float = 0) -> Signal: + """构造买入信号。""" + return Signal( + datetime=20240101 + bar_idx, + direction="BUY", + size=size, + ) + + +def _sell_signal(bar_idx: int, size: float = 0) -> Signal: + """构造卖出信号。""" + return Signal( + datetime=20240101 + bar_idx, + direction="SELL", + size=size, + ) + + +# ── Test Execution Modes ─────────────────────────────────────────────────────── + + +class TestExecutionModes: + """测试不同执行模式的成交价。""" + + def test_next_open(self) -> None: + """next_open: 下一根K线的开盘价。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="next_open") + + # 信号在 bar 0,应该在 bar 1 的 open 成交 + signals = [_buy_signal(0, size=100)] + trades = sim.simulate(signals, cash=20000, position=0) + + assert len(trades) == 1 + assert trades[0].price == 101.0 # df["open"].iloc[1] + assert trades[0].rejected is False + + def test_next_close(self) -> None: + """next_close: 下一根K线的收盘价。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="next_close") + + signals = [_buy_signal(0, size=100)] + trades = sim.simulate(signals, cash=20000, position=0) + + assert len(trades) == 1 + assert trades[0].price == 102.0 # df["close"].iloc[1] + + def test_this_close(self) -> None: + """this_close: 当前K线的收盘价。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="this_close") + + signals = [_buy_signal(0, size=100)] + trades = sim.simulate(signals, cash=20000, position=0) + + assert len(trades) == 1 + assert trades[0].price == 101.0 # df["close"].iloc[0] + + def test_this_close_future_leak_warning(self) -> None: + """this_close 模式应设置 future_leak_warning 标志。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="this_close") + + assert sim.future_leak_warning is False + + # 执行模拟后应设置标志 + signals = [_buy_signal(0, size=100)] + sim.simulate(signals, cash=20000, position=0) + + assert sim.future_leak_warning is True + + def test_worst_price_buy(self) -> None: + """worst: 买入取最高价。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="worst") + + signals = [_buy_signal(0, size=100)] + trades = sim.simulate(signals, cash=20000, position=0) + + assert len(trades) == 1 + assert trades[0].price == 103.0 # df["high"].iloc[1] + + def test_worst_price_sell(self) -> None: + """worst: 卖出取最低价。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="worst") + + signals = [_sell_signal(0, size=100)] + trades = sim.simulate(signals, cash=0, position=200) + + assert len(trades) == 1 + assert trades[0].price == 100.0 # df["low"].iloc[1] + + def test_best_price_buy(self) -> None: + """best: 买入取最低价。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="best") + + signals = [_buy_signal(0, size=100)] + trades = sim.simulate(signals, cash=20000, position=0) + + assert len(trades) == 1 + assert trades[0].price == 100.0 # df["low"].iloc[1] + + def test_best_price_sell(self) -> None: + """best: 卖出取最高价。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="best") + + signals = [_sell_signal(0, size=100)] + trades = sim.simulate(signals, cash=0, position=200) + + assert len(trades) == 1 + assert trades[0].price == 103.0 # df["high"].iloc[1] + + +# ── Test Position Modes ──────────────────────────────────────────────────────── + + +class TestPositionModes: + """测试不同仓位模式。""" + + def test_full_position(self) -> None: + """full: 全仓买入(100股整手)。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="next_open", position_mode="full") + + cash = 20000 # 足够买100股 + signals = [_buy_signal(0, size=0)] # size=0 表示全仓 + trades = sim.simulate(signals, cash=cash, position=0) + + assert len(trades) == 1 + # price=101, 20000 / (101 * 1.0003) ≈ 197.96, 可买 100 股(1手) + assert trades[0].size == 100 + + def test_fixed_position(self) -> None: + """fixed: 固定股数。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="next_open", position_mode="fixed") + + signals = [_buy_signal(0, size=100)] + trades = sim.simulate(signals, cash=20000, position=0) + + assert len(trades) == 1 + assert trades[0].size == 100 + + def test_percent_position(self) -> None: + """percent: 总资产的百分比。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="next_open", position_mode="percent") + + # 50% 资产,但不足1手(100股) + signals = [_buy_signal(0, size=0.5)] + trades = sim.simulate(signals, cash=20000, position=0) + + # 20000 * 0.5 = 10000, price=101, int(10000/101/100)*100 = 0 + # reduce 模式下返回 None(无交易) + assert len(trades) == 0 + + +# ── Test Reject Policy ───────────────────────────────────────────────────────── + + +class TestRejectPolicy: + """测试拒绝策略。""" + + def test_reduce_on_insufficient_cash(self) -> None: + """reduce: 资金不足时减少买入量。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="next_open", reject_policy="reduce") + + # 只有 15000 元现金,想买 200 股(price=101) + # 200股需要约 20200 元,但只有 15000 元 + # 应该减少到可买数量 + signals = [_buy_signal(0, size=200)] + trades = sim.simulate(signals, cash=15000, position=0, position_mode="fixed") + + assert len(trades) == 1 + assert trades[0].size < 200 # 应该减少 + assert trades[0].rejected is False + + def test_skip_on_insufficient_cash(self) -> None: + """skip: 资金不足时拒绝订单。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="next_open", reject_policy="skip") + + # 只有 15000 元现金,想买 200 股(price=101) + signals = [_buy_signal(0, size=200)] + trades = sim.simulate(signals, cash=15000, position=0, position_mode="fixed") + + assert len(trades) == 1 + assert trades[0].rejected is True + assert trades[0].size == 200 # 保持原订单量 + + def test_sell_with_no_position_skip(self) -> None: + """skip: 无持仓时卖出被拒绝。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="next_open", reject_policy="skip") + + signals = [_sell_signal(0, size=100)] + trades = sim.simulate(signals, cash=0, position=0) + + assert len(trades) == 1 + assert trades[0].rejected is True + + def test_reduce_on_insufficient_position(self) -> None: + """reduce: 持仓不足时减少卖出量。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="next_open", reject_policy="reduce") + + # 只有 50 股,想卖 100 股 + signals = [_sell_signal(0, size=100)] + trades = sim.simulate(signals, cash=0, position=50) + + assert len(trades) == 1 + assert trades[0].size == 50 # 减少到实际持仓 + assert trades[0].rejected is False + + +# ── Test Fees ───────────────────────────────────────────────────────────────── + + +class TestFees: + """测试费用计算。""" + + def test_commission_on_buy(self) -> None: + """买入时计算佣金。""" + df = _make_df(10) + sim = OrderSimulator( + df, + execution="next_open", + commission=0.0003, + min_commission=5.0, + ) + + signals = [_buy_signal(0, size=100)] + trades = sim.simulate(signals, cash=20000, position=0) + + assert len(trades) == 1 + # price=101, size=100, amount=10100 + # commission = max(10100 * 0.0003, 5) = max(3.03, 5) = 5 + assert trades[0].commission >= 5.0 + + def test_stamp_tax_on_sell(self) -> None: + """卖出时额外计算印花税。""" + df = _make_df(10) + sim = OrderSimulator( + df, + execution="next_open", + commission=0.0003, + min_commission=5.0, + stamp_tax=0.001, + ) + + # 先买入 + buy_signals = [_buy_signal(0, size=100)] + sim.simulate(buy_signals, cash=20000, position=0) + + # 再卖出 + sell_signals = [_sell_signal(1, size=100)] + trades = sim.simulate(sell_signals, cash=0, position=100) + + assert len(trades) == 1 + # commission + stamp_tax + # commission = max(10200 * 0.0003, 5) = 5 + # stamp_tax = 10200 * 0.001 = 10.2 + # total = 15.2 + assert trades[0].commission > 5.0 + + def test_slippage(self) -> None: + """测试滑点计算。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="next_open", slippage=0.01) + + signals = [_buy_signal(0, size=100)] + trades = sim.simulate(signals, cash=20000, position=0) + + assert len(trades) == 1 + assert trades[0].slippage == 1.0 # 100 * 0.01 + + +# ── Test Edge Cases ─────────────────────────────────────────────────────────── + + +class TestEdgeCases: + """测试边界情况。""" + + def test_signal_not_found_in_df(self) -> None: + """信号时间不在K线数据中。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="next_open") + + # 信号时间 20250101 不在 df 中 + signal = Signal(datetime=20250101, direction="BUY", size=100) + trades = sim.simulate([signal], cash=20000, position=0) + + assert len(trades) == 0 + + def test_signal_at_last_bar_next_execution(self) -> None: + """信号在最后一根K线,next_* 模式无法成交。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="next_open") + + # 信号在最后一根 + signals = [_buy_signal(9, size=100)] + trades = sim.simulate(signals, cash=20000, position=0) + + # exec_idx = 10,超出范围 + assert len(trades) == 0 + + def test_datetime_column_as_int(self) -> None: + """datetime 列为 int 类型时的查找。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="next_open") + + signals = [_buy_signal(0, size=100)] + trades = sim.simulate(signals, cash=20000, position=0) + + assert len(trades) == 1 + + def test_datetime_column_as_datetime(self) -> None: + """datetime 列为 datetime 类型时的查找。""" + df = _make_df(10) + # 转为 datetime 类型 + df["datetime"] = pd.to_datetime(df["datetime"].astype(str), format="%Y%m%d") + + sim = OrderSimulator(df, execution="next_open") + + signals = [_buy_signal(0, size=100)] + trades = sim.simulate(signals, cash=20000, position=0) + + assert len(trades) == 1 + + def test_multiple_signals(self) -> None: + """多个信号的顺序执行。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="next_open") + + signals = [ + _buy_signal(0, size=100), + _sell_signal(1, size=100), + ] + trades = sim.simulate(signals, cash=20000, position=0) + + assert len(trades) == 2 + assert trades[0].direction == "BUY" + assert trades[1].direction == "SELL" + + def test_position_tracking(self) -> None: + """测试持仓跟踪。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="next_open") + + # 买入 100 股 + buy_signals = [_buy_signal(0, size=100)] + trades = sim.simulate(buy_signals, cash=20000, position=0) + + # 验证持仓(通过模拟器内部状态) + # 这里需要暴露 position 或者通过返回值验证 + # 简化:只验证成交记录 + assert len(trades) == 1 + assert trades[0].size == 100