feat(backtest): integrate SlippageModel + ExecutionModel into BacktestEngine

This commit is contained in:
GitHub
2026-06-12 21:04:24 +08:00
parent d18af98855
commit 0945e47990
2 changed files with 163 additions and 17 deletions
+76 -2
View File
@@ -6,7 +6,7 @@ Coordinates Strategy → OrderSimulator → PortfolioTracker → PerformanceAnal
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
from typing import TYPE_CHECKING, Any
import pandas as pd
@@ -16,6 +16,10 @@ from easy_tdx.backtest.portfolio import PortfolioTracker
from easy_tdx.backtest.strategy import Strategy
from easy_tdx.backtest.types import BacktestResult, Signal, Trade
if TYPE_CHECKING:
from easy_tdx.backtest.execution import ExecutionModel
from easy_tdx.backtest.slippage import SlippageModel
@dataclass
class _StopCondition:
@@ -57,6 +61,8 @@ class BacktestEngine:
reject_policy: str = "reduce",
benchmark: pd.DataFrame | None = None,
chanlun_level: str | None = None,
slippage_model: SlippageModel | None = None,
execution_model: ExecutionModel | None = None,
):
"""Initialize engine.
@@ -73,6 +79,10 @@ class BacktestEngine:
benchmark: Benchmark data for performance comparison
chanlun_level: Auto-compute chanlun analysis at this level
(e.g. 'DAILY', '30MIN'). Strategy accesses via self.chanlun.
slippage_model: Pluggable slippage model (overrides flat slippage
when provided).
execution_model: Pluggable execution model (bypasses OrderSimulator
when provided).
"""
self._strategy_cls = strategy if isinstance(strategy, type) else type(strategy)
self._strategy_instance = strategy if isinstance(strategy, Strategy) else None
@@ -87,6 +97,8 @@ class BacktestEngine:
self._reject_policy = reject_policy
self._benchmark = benchmark
self._chanlun_level = chanlun_level
self._slippage_model = slippage_model
self._execution_model = execution_model
def run(self, df: pd.DataFrame, chanlun_result: Any | None = None) -> BacktestResult:
"""Run backtest.
@@ -113,6 +125,12 @@ class BacktestEngine:
signals = self._generate_signals(df, chanlun_result)
# Step 2: Order simulation
if self._execution_model is not None:
# ExecutionModel path
trades = self._execute_with_model(signals, df)
future_leak = False
else:
# OrderSimulator path (default)
simulator = OrderSimulator(
df,
execution=self._execution,
@@ -122,12 +140,14 @@ class BacktestEngine:
min_commission=self._min_commission,
stamp_tax=self._stamp_tax,
slippage=self._slippage,
slippage_model=self._slippage_model,
)
trades = simulator.simulate(
signals=signals,
cash=self._cash,
position=0.0,
)
future_leak = simulator.future_leak_warning
# Step 3: Portfolio tracking
trades = self._compute_pnls(trades)
@@ -149,7 +169,7 @@ class BacktestEngine:
"execution": self._execution,
"position_mode": self._position_mode,
"reject_policy": self._reject_policy,
"future_leak_warning": simulator.future_leak_warning,
"future_leak_warning": future_leak,
}
return BacktestResult(
@@ -160,6 +180,57 @@ class BacktestEngine:
config=config,
)
def _execute_with_model(self, signals: list[Signal], df: pd.DataFrame) -> list[Trade]:
"""Use ExecutionModel to process signals."""
assert self._execution_model is not None
all_trades: list[Trade] = []
cash = self._cash
position = 0.0
for signal in signals:
bar_idx = self._find_bar_index(df, signal.datetime)
if bar_idx is None:
continue
sub_trades = self._execution_model.execute(
signal=signal,
df=df,
bar_idx=bar_idx,
cash=cash,
position=position,
position_mode=self._position_mode,
commission=self._commission,
min_commission=self._min_commission,
stamp_tax=self._stamp_tax,
slippage_model=self._slippage_model,
)
for t in sub_trades:
if not t.rejected:
if t.direction == "BUY":
cash -= t.size * t.price + t.commission + t.slippage
position += t.size
else:
cash += t.size * t.price - t.commission - t.slippage
position -= t.size
all_trades.extend(sub_trades)
return all_trades
@staticmethod
def _find_bar_index(df: pd.DataFrame, datetime_val: int) -> int | None:
"""Find bar index by datetime value."""
dt_col = df["datetime"]
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
if hasattr(dt_col, "dt"):
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
def _generate_signals(self, df: pd.DataFrame, chanlun_result: Any | None) -> list[Signal]:
"""Generate signals from strategy.
@@ -385,6 +456,7 @@ class BacktestEngine:
"size",
"price",
"commission",
"slippage",
"pnl",
"rejected",
]
@@ -397,6 +469,7 @@ class BacktestEngine:
"size": t.size,
"price": t.price,
"commission": t.commission,
"slippage": t.slippage,
"pnl": t.pnl,
"rejected": t.rejected,
}
@@ -427,6 +500,7 @@ class BacktestEngine:
"size",
"price",
"commission",
"slippage",
"pnl",
"rejected",
]
+72
View File
@@ -7,6 +7,8 @@ import pandas as pd
from easy_tdx import MyTT
from easy_tdx.backtest.engine import BacktestEngine
from easy_tdx.backtest.execution import TWAPExecution
from easy_tdx.backtest.slippage import FixedSlippage
from easy_tdx.backtest.strategy import Strategy
@@ -552,3 +554,73 @@ def test_chanlun_manual_result_overrides_auto():
# Strategy should have received the manual result, not auto-computed one
assert CheckerStrategy.received == manual_result
# ── SlippageModel + ExecutionModel Integration ───────────────────────────────
class TestEngineSlippageModel:
"""BacktestEngine with SlippageModel integration."""
def test_engine_with_slippage_model(self) -> None:
"""Engine uses SlippageModel."""
class SimpleBuy(Strategy):
def init(self) -> None:
pass
def next(self) -> None:
if self._bar_index == 0:
self.buy(size=100)
df = _make_df(20)
engine = BacktestEngine(
SimpleBuy,
cash=100000,
slippage_model=FixedSlippage(per_share=0.05),
)
result = engine.run(df)
buy_trades = result.trades[result.trades["direction"] == "BUY"]
if len(buy_trades) > 0:
assert buy_trades.iloc[0]["slippage"] > 0
class TestEngineExecutionModel:
"""BacktestEngine with ExecutionModel integration."""
def test_engine_with_twap(self) -> None:
"""Engine uses TWAP execution."""
class SimpleBuy(Strategy):
def init(self) -> None:
pass
def next(self) -> None:
if self._bar_index == 0:
self.buy(size=300)
df = _make_df(20)
engine = BacktestEngine(
SimpleBuy,
cash=100000,
execution_model=TWAPExecution(n_bars=3),
)
result = engine.run(df)
buy_trades = result.trades[result.trades["direction"] == "BUY"]
assert len(buy_trades) >= 1
def test_engine_backward_compatible(self) -> None:
"""No new params: behavior unchanged."""
class SimpleBuy(Strategy):
def init(self) -> None:
pass
def next(self) -> None:
if self._bar_index == 0:
self.buy(size=100)
df = _make_df(20)
engine = BacktestEngine(SimpleBuy, cash=100000)
result = engine.run(df)
assert len(result.trades) >= 1