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