From 6414c2cc119df22393fb1e8a40dcb08f1e3c50ad Mon Sep 17 00:00:00 2001 From: GitHub Date: Fri, 12 Jun 2026 20:50:42 +0800 Subject: [PATCH] feat(backtest): integrate SlippageModel into OrderSimulator --- src/easy_tdx/backtest/orders.py | 43 +++++++++++++++-- tests/unit/test_backtest_orders.py | 74 ++++++++++++++++++++++++++++++ 2 files changed, 114 insertions(+), 3 deletions(-) diff --git a/src/easy_tdx/backtest/orders.py b/src/easy_tdx/backtest/orders.py index 638e948..dfd2712 100644 --- a/src/easy_tdx/backtest/orders.py +++ b/src/easy_tdx/backtest/orders.py @@ -6,11 +6,16 @@ from __future__ import annotations from dataclasses import dataclass +from typing import TYPE_CHECKING +import numpy as np import pandas as pd from easy_tdx.backtest.types import Signal, Trade +if TYPE_CHECKING: + from easy_tdx.backtest.slippage import SlippageModel + @dataclass class OrderSimulator: @@ -39,6 +44,7 @@ class OrderSimulator: min_commission: float = 5.0 stamp_tax: float = 0.001 slippage: float = 0.0 + slippage_model: SlippageModel | None = None future_leak_warning: bool = False def simulate( @@ -288,6 +294,37 @@ class OrderSimulator: return commission + def _compute_slippage(self, size: float, price: float, is_sell: bool) -> float: + """计算滑点成本。""" + if self.slippage_model is not None: + volume = self._get_current_volume() + volatility = self._estimate_volatility() + return self.slippage_model.compute( + price=price, + size=size, + volume=volume, + volatility=volatility, + direction="SELL" if is_sell else "BUY", + ) + return size * self.slippage + + def _get_current_volume(self) -> float: + """获取最后一根K线的成交量。""" + if "volume" in self.df.columns and len(self.df) > 0: + return float(self.df["volume"].iloc[-1]) + return 0.0 + + def _estimate_volatility(self) -> float: + """从收盘价估计年化波动率。""" + if "close" not in self.df.columns or len(self.df) < 2: + return 0.0 + close = self.df["close"].to_numpy() + returns = np.diff(close) / close[:-1] + if len(returns) < 2: + return 0.0 + daily_vol = float(np.std(returns)) + return daily_vol * np.sqrt(252) + def _execute_buy( self, signal: Signal, @@ -338,7 +375,7 @@ class OrderSimulator: # 计算费用 commission = self._calculate_commission(size, price, is_sell=False) - slippage = size * self.slippage + slippage = self._compute_slippage(size, price, is_sell=False) # 检查资金是否足够 total_cost = size * price + commission + slippage @@ -361,7 +398,7 @@ class OrderSimulator: 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 + slippage = self._compute_slippage(reduced_size, price, is_sell=False) return Trade( datetime=self.df.iloc[exec_idx]["datetime"], direction="BUY", @@ -446,7 +483,7 @@ class OrderSimulator: # 计算费用 commission = self._calculate_commission(size, price, is_sell=True) - slippage = size * self.slippage + slippage = self._compute_slippage(size, price, is_sell=True) return Trade( datetime=self.df.iloc[exec_idx]["datetime"], diff --git a/tests/unit/test_backtest_orders.py b/tests/unit/test_backtest_orders.py index 7aca070..72879da 100644 --- a/tests/unit/test_backtest_orders.py +++ b/tests/unit/test_backtest_orders.py @@ -3,8 +3,10 @@ from __future__ import annotations import pandas as pd +import pytest from easy_tdx.backtest.orders import OrderSimulator +from easy_tdx.backtest.slippage import FixedSlippage, PercentSlippage from easy_tdx.backtest.types import Signal # ── Test Fixtures ───────────────────────────────────────────────────────────── @@ -390,3 +392,75 @@ class TestEdgeCases: # 简化:只验证成交记录 assert len(trades) == 1 assert trades[0].size == 100 + + +# ── Test SlippageModel Integration ───────────────────────────────────────────── + + +class TestSlippageModelIntegration: + """测试 OrderSimulator 与 SlippageModel 集成。""" + + def test_fixed_slippage_model(self) -> None: + """FixedSlippage 与旧 slippage 参数等价。""" + df = _make_df(10) + sim = OrderSimulator( + df, + execution="next_open", + slippage_model=FixedSlippage(per_share=0.01), + ) + signals = [_buy_signal(0, size=100)] + trades = sim.simulate(signals, cash=20000, position=0) + assert len(trades) == 1 + assert trades[0].slippage == pytest.approx(1.0) + + def test_percent_slippage_model(self) -> None: + """PercentSlippage 计算。""" + df = _make_df(10) + sim = OrderSimulator( + df, + execution="next_open", + slippage_model=PercentSlippage(rate=0.001), + ) + signals = [_buy_signal(0, size=100)] + trades = sim.simulate(signals, cash=20000, position=0) + assert len(trades) == 1 + # price=101 (next_open), 101 × 100 × 0.001 = 10.1 + assert trades[0].slippage == pytest.approx(10.1) + + def test_slippage_model_overrides_slippage_param(self) -> None: + """slippage_model 优先于 slippage 参数。""" + df = _make_df(10) + sim = OrderSimulator( + df, + execution="next_open", + position_mode="fixed", + slippage=999.0, + slippage_model=FixedSlippage(per_share=0.01), + ) + signals = [_buy_signal(0, size=100)] + trades = sim.simulate(signals, cash=20000, position=0) + assert len(trades) == 1 + assert trades[0].slippage == pytest.approx(1.0) + + def test_sell_with_slippage_model(self) -> None: + """卖出时也使用滑点模型。""" + df = _make_df(10) + sim = OrderSimulator( + df, + execution="next_open", + slippage_model=FixedSlippage(per_share=0.02), + ) + signals = [_sell_signal(0, size=100)] + trades = sim.simulate(signals, cash=0, position=100) + assert len(trades) == 1 + # position_mode=full, size=0 → sell all position=100 + assert trades[0].slippage == pytest.approx(2.0) + + def test_no_slippage_model_uses_old_param(self) -> None: + """不提供 model 时使用旧 slippage 参数。""" + df = _make_df(10) + sim = OrderSimulator(df, execution="next_open", slippage=0.05) + signals = [_buy_signal(0, size=100)] + trades = sim.simulate(signals, cash=20000, position=0) + assert len(trades) == 1 + assert trades[0].slippage == pytest.approx(5.0)