feat(backtest): add LimitExecution

This commit is contained in:
GitHub
2026-06-12 20:59:59 +08:00
parent fe68d9da95
commit d18af98855
2 changed files with 212 additions and 0 deletions
+113
View File
@@ -521,3 +521,116 @@ class VWAPExecution(ExecutionModel):
)
)
return trades
class LimitExecution(ExecutionModel):
"""限价单执行。
在目标价位挂单,仅当 bar_low <= price(买入)或 bar_high >= price(卖出)时成交。
无限价时退化为 ImmediateExecution。
"""
def __init__(self, ttl_bars: int = 5) -> None:
self._ttl_bars = max(1, ttl_bars)
self._fallback = ImmediateExecution()
def execute(
self,
signal: Signal,
df: pd.DataFrame,
bar_idx: int,
cash: float,
position: float,
position_mode: str,
commission: float,
min_commission: float,
stamp_tax: float,
slippage_model: SlippageModel | None,
) -> list[Trade]:
if signal.price is None:
return self._fallback.execute(
signal,
df,
bar_idx,
cash,
position,
position_mode,
commission,
min_commission,
stamp_tax,
slippage_model,
)
target_price = signal.price
for i in range(self._ttl_bars):
exec_idx = bar_idx + 1 + i
if exec_idx >= len(df):
break
row = df.iloc[exec_idx]
triggered = False
if signal.direction == "BUY" and float(row["low"]) <= target_price:
triggered = True
elif signal.direction == "SELL" and float(row["high"]) >= target_price:
triggered = True
if triggered:
if signal.direction == "BUY":
size = self._calc_buy_size(
signal.size,
target_price,
cash,
position_mode,
commission,
)
if size <= 0:
return []
comm = self._calc_commission(
size,
target_price,
False,
commission,
min_commission,
stamp_tax,
)
slip = self._calc_slippage(
size,
target_price,
False,
slippage_model,
df,
)
else:
size = signal.size if signal.size > 0 else position
if size <= 0:
return []
if size > position:
size = position
comm = self._calc_commission(
size,
target_price,
True,
commission,
min_commission,
stamp_tax,
)
slip = self._calc_slippage(
size,
target_price,
True,
slippage_model,
df,
)
return [
Trade(
datetime=self._get_datetime_int(df, exec_idx),
direction=signal.direction,
size=float(size),
price=target_price,
commission=comm,
slippage=slip,
)
]
return []
+99
View File
@@ -8,6 +8,7 @@ import pytest
from easy_tdx.backtest.execution import (
ExecutionModel,
ImmediateExecution,
LimitExecution,
TWAPExecution,
VWAPExecution,
)
@@ -319,3 +320,101 @@ class TestVWAPExecution:
slippage_model=None,
)
assert len(trades) <= 4
class TestLimitExecution:
"""限价单执行。"""
def test_buy_limit_filled(self) -> None:
df = _make_df(20)
model = LimitExecution(ttl_bars=5)
signal = Signal(datetime=20240101, direction="BUY", size=100, price=100.0)
trades = model.execute(
signal=signal,
df=df,
bar_idx=0,
cash=20000,
position=0,
position_mode="fixed",
commission=0.0003,
min_commission=5.0,
stamp_tax=0.001,
slippage_model=None,
)
assert len(trades) == 1
assert trades[0].price == 100.0
assert trades[0].direction == "BUY"
def test_sell_limit_filled(self) -> None:
df = _make_df(20)
model = LimitExecution(ttl_bars=5)
signal = Signal(datetime=20240101, direction="SELL", size=100, price=105.0)
trades = model.execute(
signal=signal,
df=df,
bar_idx=0,
cash=0,
position=200,
position_mode="fixed",
commission=0.0003,
min_commission=5.0,
stamp_tax=0.001,
slippage_model=None,
)
assert len(trades) == 1
assert trades[0].price == 105.0
def test_limit_not_triggered(self) -> None:
df = _make_df(10)
model = LimitExecution(ttl_bars=3)
signal = Signal(datetime=20240101, direction="BUY", size=100, price=50.0)
trades = model.execute(
signal=signal,
df=df,
bar_idx=0,
cash=20000,
position=0,
position_mode="fixed",
commission=0.0003,
min_commission=5.0,
stamp_tax=0.001,
slippage_model=None,
)
assert len(trades) == 0
def test_no_price_falls_back_to_immediate(self) -> None:
df = _make_df(10)
model = LimitExecution(ttl_bars=5)
signal = Signal(datetime=20240101, direction="BUY", size=100, price=None)
trades = model.execute(
signal=signal,
df=df,
bar_idx=0,
cash=20000,
position=0,
position_mode="fixed",
commission=0.0003,
min_commission=5.0,
stamp_tax=0.001,
slippage_model=None,
)
assert len(trades) == 1
assert trades[0].price == 101.0
def test_ttl_expires(self) -> None:
df = _make_df(20)
model = LimitExecution(ttl_bars=2)
signal = Signal(datetime=20240101, direction="BUY", size=100, price=98.0)
trades = model.execute(
signal=signal,
df=df,
bar_idx=0,
cash=20000,
position=0,
position_mode="fixed",
commission=0.0003,
min_commission=5.0,
stamp_tax=0.001,
slippage_model=None,
)
assert len(trades) == 0