mirror of
https://ghfast.top/https://github.com/aeroxw/easy-tdx.git
synced 2026-09-12 14:34:15 +08:00
feat(backtest): integrate SlippageModel into OrderSimulator
This commit is contained in:
@@ -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"],
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user