Files
easy_tdx_max/docs/superpowers/plans/2026-06-12-v1.14.0-slippage-execution.md

58 KiB
Raw Permalink Blame History

v1.14.0 滑点模型 + 执行仿真 实施计划

For agentic workers: REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task.

Goal: 为回测引擎新增可插拔的滑点模型和执行仿真引擎,同时保持完全向后兼容。

Architecture:backtest/ 下新增 slippage.pyexecution.py,修改 orders.pyengine.py 集成新模型。所有新功能通过可选参数启用,不破坏现有 API。

Tech Stack: 纯 numpy/pandas,无新外部依赖。


文件结构

文件 操作 职责
src/easy_tdx/backtest/slippage.py 新增 SlippageModel ABC + 4 种内置模型
src/easy_tdx/backtest/execution.py 新增 ExecutionModel ABC + 4 种内置模型
src/easy_tdx/backtest/orders.py 修改 集成 SlippageModel
src/easy_tdx/backtest/engine.py 修改 集成 SlippageModel + ExecutionModel
tests/unit/test_backtest_slippage.py 新增 滑点模型测试(~15 个)
tests/unit/test_backtest_execution.py 新增 执行仿真测试(~20 个)

Task 1: SlippageModel 基类 + FixedSlippage + PercentSlippage

Files:

  • Create: src/easy_tdx/backtest/slippage.py

  • Test: tests/unit/test_backtest_slippage.py

  • Step 1: Write tests

"""滑点模型单元测试。"""
from __future__ import annotations

import pytest

from easy_tdx.backtest.slippage import (
    FixedSlippage,
    PercentSlippage,
    SlippageModel,
    SquareRootSlippage,
    VolumeSlippage,
)


class TestSlippageBase:
    """基类验证。"""

    def test_cannot_instantiate_abc(self) -> None:
        """不能直接实例化 ABC。"""
        with pytest.raises(TypeError):
            SlippageModel()  # type: ignore[abstract]

    def test_subclass_must_implement_compute(self) -> None:
        """子类必须实现 compute。"""

        class BadModel(SlippageModel):
            pass

        with pytest.raises(TypeError):
            BadModel()  # type: ignore[abstract]


class TestFixedSlippage:
    """固定每股滑点。"""

    def test_zero_per_share(self) -> None:
        """per_share=0 时无滑点。"""
        model = FixedSlippage(per_share=0.0)
        cost = model.compute(price=10.0, size=100, volume=10000, volatility=0.3, direction="BUY")
        assert cost == 0.0

    def test_basic(self) -> None:
        """基本计算:100 股 × 0.01 元/股 = 1.0。"""
        model = FixedSlippage(per_share=0.01)
        cost = model.compute(price=10.0, size=100, volume=10000, volatility=0.3, direction="BUY")
        assert cost == pytest.approx(1.0)

    def test_large_size(self) -> None:
        """大单。"""
        model = FixedSlippage(per_share=0.05)
        cost = model.compute(price=50.0, size=10000, volume=500000, volatility=0.2, direction="SELL")
        assert cost == pytest.approx(500.0)

    def test_direction_irrelevant(self) -> None:
        """方向不影响固定滑点。"""
        model = FixedSlippage(per_share=0.01)
        buy_cost = model.compute(price=10.0, size=100, volume=10000, volatility=0.3, direction="BUY")
        sell_cost = model.compute(price=10.0, size=100, volume=10000, volatility=0.3, direction="SELL")
        assert buy_cost == sell_cost


class TestPercentSlippage:
    """按成交金额百分比滑点。"""

    def test_zero_rate(self) -> None:
        """rate=0 时无滑点。"""
        model = PercentSlippage(rate=0.0)
        cost = model.compute(price=10.0, size=100, volume=10000, volatility=0.3, direction="BUY")
        assert cost == 0.0

    def test_basic(self) -> None:
        """10元 × 100股 × 0.001 = 1.0。"""
        model = PercentSlippage(rate=0.001)
        cost = model.compute(price=10.0, size=100, volume=10000, volatility=0.3, direction="BUY")
        assert cost == pytest.approx(1.0)

    def test_high_price(self) -> None:
        """高价股。"""
        model = PercentSlippage(rate=0.002)
        cost = model.compute(price=100.0, size=500, volume=20000, volatility=0.25, direction="BUY")
        # 100 × 500 × 0.002 = 100.0
        assert cost == pytest.approx(100.0)
  • Step 2: Run tests to verify they fail

Run: python -m pytest tests/unit/test_backtest_slippage.py -v -k "TestSlippageBase or TestFixed or TestPercent" --no-header Expected: FAIL (import error)

  • Step 3: Write implementation
"""可插拔滑点模型。"""
from __future__ import annotations

from abc import ABC, abstractmethod

import numpy as np


class SlippageModel(ABC):
    """滑点模型基类。

    所有滑点模型必须实现 compute() 方法,返回总滑点成本(金额)。
    """

    @abstractmethod
    def compute(
        self,
        price: float,
        size: float,
        volume: float,
        volatility: float,
        direction: str,
    ) -> float:
        """计算滑点成本。

        Args:
            price: 成交价格
            size: 订单数量(股)
            volume: 当日成交量(股),0 表示无数据
            volatility: 近期年化波动率,0 表示无数据
            direction: 交易方向 BUY / SELL

        Returns:
            总滑点成本(金额,非比率)
        """
        ...


class FixedSlippage(SlippageModel):
    """固定每股滑点(向后兼容)。"""

    def __init__(self, per_share: float = 0.01) -> None:
        self._per_share = per_share

    def compute(
        self,
        price: float,
        size: float,
        volume: float,
        volatility: float,
        direction: str,
    ) -> float:
        return size * self._per_share


class PercentSlippage(SlippageModel):
    """按成交金额百分比滑点。"""

    def __init__(self, rate: float = 0.001) -> None:
        self._rate = rate

    def compute(
        self,
        price: float,
        size: float,
        volume: float,
        volatility: float,
        direction: str,
    ) -> float:
        return price * size * self._rate
  • Step 4: Run tests to verify they pass

Run: python -m pytest tests/unit/test_backtest_slippage.py -v -k "TestSlippageBase or TestFixed or TestPercent" --no-header Expected: 9 passed

  • Step 5: Commit
git add src/easy_tdx/backtest/slippage.py tests/unit/test_backtest_slippage.py
git commit -m "feat(backtest): add SlippageModel ABC + FixedSlippage + PercentSlippage"

Task 2: SquareRootSlippage + VolumeSlippage

Files:

  • Modify: src/easy_tdx/backtest/slippage.py

  • Modify: tests/unit/test_backtest_slippage.py

  • Step 1: Write tests

test_backtest_slippage.py 末尾追加:

class TestSquareRootSlippage:
    """方根市场冲击模型。"""

    def test_zero_size(self) -> None:
        """size=0 时无冲击。"""
        model = SquareRootSlippage(impact_coeff=0.1)
        cost = model.compute(price=10.0, size=0, volume=10000, volatility=0.3, direction="BUY")
        assert cost == 0.0

    def test_small_participation_rate(self) -> None:
        """低参与率(小单),冲击成本低。"""
        model = SquareRootSlippage(impact_coeff=0.1)
        # size=100, volume=1000000, participation_rate=0.0001
        cost = model.compute(
            price=10.0, size=100, volume=1_000_000, volatility=0.3, direction="BUY"
        )
        # σ=0.3, √(0.0001)=0.01, impact = 0.3 × 0.01 × 10 × 100 × 0.1 = 0.03
        assert cost == pytest.approx(0.03)

    def test_high_participation_rate(self) -> None:
        """高参与率(大单),冲击成本高。"""
        model = SquareRootSlippage(impact_coeff=0.1)
        # size=100000, volume=200000, participation_rate=0.5
        cost = model.compute(
            price=10.0, size=100_000, volume=200_000, volatility=0.3, direction="BUY"
        )
        # 应该显著大于小单
        small_cost = model.compute(
            price=10.0, size=100, volume=1_000_000, volatility=0.3, direction="BUY"
        )
        assert cost > small_cost * 10

    def test_zero_volume_fallback(self) -> None:
        """volume=0 时退化为 PercentSlippage(rate=0.001)。"""
        model = SquareRootSlippage(impact_coeff=0.1)
        cost = model.compute(price=10.0, size=100, volume=0, volatility=0.3, direction="BUY")
        # 退化为 10 × 100 × 0.001 = 1.0
        assert cost == pytest.approx(1.0)

    def test_zero_volatility_fallback(self) -> None:
        """volatility=0 时退化为 PercentSlippage(rate=0.001)。"""
        model = SquareRootSlippage(impact_coeff=0.1)
        cost = model.compute(price=10.0, size=100, volume=10000, volatility=0.0, direction="BUY")
        assert cost == pytest.approx(1.0)


class TestVolumeSlippage:
    """成交量比例滑点。"""

    def test_zero_size(self) -> None:
        """size=0 时无滑点。"""
        model = VolumeSlippage(base_bps=10.0)
        cost = model.compute(price=10.0, size=0, volume=10000, volatility=0.3, direction="BUY")
        assert cost == 0.0

    def test_basic(self) -> None:
        """基本计算。"""
        model = VolumeSlippage(base_bps=10.0)
        # base_bps=10 → rate=10/10000=0.001
        # participation = 100/10000 = 0.01
        # cost = 0.001 × 0.01 × 10 × 100 = 0.01
        cost = model.compute(price=10.0, size=100, volume=10000, volatility=0.3, direction="BUY")
        assert cost == pytest.approx(0.01)

    def test_high_participation(self) -> None:
        """高参与率时滑点高。"""
        model = VolumeSlippage(base_bps=10.0)
        cost_high = model.compute(price=10.0, size=5000, volume=10000, volatility=0.3, direction="BUY")
        cost_low = model.compute(price=10.0, size=100, volume=10000, volatility=0.3, direction="BUY")
        assert cost_high > cost_low

    def test_zero_volume_fallback(self) -> None:
        """volume=0 时退化为 PercentSlippage(rate=base_bps/10000)。"""
        model = VolumeSlippage(base_bps=10.0)
        cost = model.compute(price=10.0, size=100, volume=0, volatility=0.3, direction="BUY")
        # 退化为 10/10000 × 10 × 100 = 1.0
        assert cost == pytest.approx(1.0)
  • Step 2: Run tests to verify they fail

Run: python -m pytest tests/unit/test_backtest_slippage.py -v -k "TestSquareRoot or TestVolume" --no-header Expected: FAIL (import error)

  • Step 3: Write implementation

slippage.py 末尾追加:

class SquareRootSlippage(SlippageModel):
    """方根市场冲击模型(Almgren-Chriss 简化版)。

    impact = σ × √(participation_rate) × price × size × impact_coeff

    当 volume=0 或 volatility=0 时退化为 PercentSlippage(rate=0.001)。
    """

    def __init__(self, impact_coeff: float = 0.1) -> None:
        self._impact_coeff = impact_coeff
        self._fallback = PercentSlippage(rate=0.001)

    def compute(
        self,
        price: float,
        size: float,
        volume: float,
        volatility: float,
        direction: str,
    ) -> float:
        if size <= 0:
            return 0.0
        if volume <= 0 or volatility <= 0:
            return self._fallback.compute(price, size, volume, volatility, direction)
        participation_rate = min(size / volume, 1.0)
        impact = volatility * np.sqrt(participation_rate) * price * size * self._impact_coeff
        return float(impact)


class VolumeSlippage(SlippageModel):
    """成交量比例滑点。

    cost = (base_bps / 10000) × (size / volume) × price × size

    当 volume=0 时退化为 PercentSlippage(rate=base_bps/10000)。
    """

    def __init__(self, base_bps: float = 10.0) -> None:
        self._base_bps = base_bps
        self._fallback = PercentSlippage(rate=base_bps / 10000.0)

    def compute(
        self,
        price: float,
        size: float,
        volume: float,
        volatility: float,
        direction: str,
    ) -> float:
        if size <= 0:
            return 0.0
        if volume <= 0:
            return self._fallback.compute(price, size, volume, volatility, direction)
        rate = self._base_bps / 10000.0
        participation = min(size / volume, 1.0)
        return rate * participation * price * size
  • Step 4: Run tests

Run: python -m pytest tests/unit/test_backtest_slippage.py -v --no-header Expected: 19 passed

  • Step 5: Commit
git add src/easy_tdx/backtest/slippage.py tests/unit/test_backtest_slippage.py
git commit -m "feat(backtest): add SquareRootSlippage + VolumeSlippage"

Task 3: OrderSimulator 集成 SlippageModel

Files:

  • Modify: src/easy_tdx/backtest/orders.py

  • Step 1: Write tests

tests/unit/test_backtest_orders.py 末尾追加:

from easy_tdx.backtest.slippage import FixedSlippage, PercentSlippage, SquareRootSlippage


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)  # 100 × 0.01

    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",
            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)  # 用 model,不是 999*100

    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=200)
        assert len(trades) == 1
        assert trades[0].slippage == pytest.approx(2.0)  # 100 × 0.02

    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)  # 100 × 0.05
  • Step 2: Run tests to verify they fail

Run: python -m pytest tests/unit/test_backtest_orders.py -v -k "TestSlippageModel" --no-header Expected: FAIL (slippage_model param not accepted)

  • Step 3: Modify orders.py

Changes to src/easy_tdx/backtest/orders.py:

  1. Add import:
from __future__ import annotations

from dataclasses import dataclass, field

import pandas as pd

from easy_tdx.backtest.types import Signal, Trade
  1. Add slippage_model field to OrderSimulator:
@dataclass
class OrderSimulator:
    # ... existing fields ...
    slippage: float = 0.0
    slippage_model: SlippageModel | None = None  # NEW
    future_leak_warning: bool = False

But we can't use from __future__ import annotations with dataclass field type hints that reference imported types directly. Since SlippageModel is in the same package, use TYPE_CHECKING:

from __future__ import annotations

from dataclasses import dataclass
from typing import TYPE_CHECKING

import pandas as pd

from easy_tdx.backtest.types import Signal, Trade

if TYPE_CHECKING:
    from easy_tdx.backtest.slippage import SlippageModel
  1. Replace slippage calculations in _execute_buy and _execute_sell:

In _execute_buy, replace:

slippage = size * self.slippage

with:

slippage = self._compute_slippage(size, price, False)

Same in _execute_sell:

slippage = self._compute_slippage(size, price, True)
  1. Add helper method:
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)

Also need to add import numpy as np at top.

  • Step 4: Run all backtest tests

Run: python -m pytest tests/unit/test_backtest_orders.py -v --no-header Expected: All existing + 5 new tests pass

  • Step 5: Commit
git add src/easy_tdx/backtest/orders.py tests/unit/test_backtest_orders.py
git commit -m "feat(backtest): integrate SlippageModel into OrderSimulator"

Task 4: ExecutionModel 基类 + ImmediateExecution

Files:

  • Create: src/easy_tdx/backtest/execution.py

  • Create: tests/unit/test_backtest_execution.py

  • Step 1: Write tests

"""执行仿真引擎单元测试。"""
from __future__ import annotations

import pytest

from easy_tdx.backtest.execution import ExecutionModel, ImmediateExecution
from easy_tdx.backtest.types import Signal


def _make_df(n: int = 20) -> "pd.DataFrame":
    """构造测试用K线数据。"""
    import pandas as pd
    data = {
        "datetime": [20240101 + i for i in range(n)],
        "open": [100.0 + i for i in range(n)],
        "close": [101.0 + i for i in range(n)],
        "high": [102.0 + i for i in range(n)],
        "low": [99.0 + i for i in range(n)],
        "volume": [10000] * n,
    }
    return pd.DataFrame(data)


class TestExecutionBase:
    """基类验证。"""

    def test_cannot_instantiate_abc(self) -> None:
        with pytest.raises(TypeError):
            ExecutionModel()  # type: ignore[abstract]


class TestImmediateExecution:
    """即时成交(向后兼容)。"""

    def test_buy_signal(self) -> None:
        """买入信号在下一 bar 成交。"""
        df = _make_df(10)
        model = ImmediateExecution()
        signal = Signal(datetime=20240101, direction="BUY", size=100)
        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].direction == "BUY"
        assert trades[0].price == 101.0  # next bar open

    def test_sell_signal(self) -> None:
        """卖出信号在下一 bar 成交。"""
        df = _make_df(10)
        model = ImmediateExecution()
        signal = Signal(datetime=20240101, direction="SELL", size=100)
        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].direction == "SELL"

    def test_signal_at_last_bar(self) -> None:
        """信号在最后一根K线,无法成交。"""
        df = _make_df(10)
        model = ImmediateExecution()
        signal = Signal(datetime=20240109, direction="BUY", size=100)
        trades = model.execute(
            signal=signal, df=df, bar_idx=9,
            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_with_slippage_model(self) -> None:
        """使用滑点模型。"""
        from easy_tdx.backtest.slippage import FixedSlippage
        df = _make_df(10)
        model = ImmediateExecution()
        signal = Signal(datetime=20240101, direction="BUY", size=100)
        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=FixedSlippage(per_share=0.01),
        )
        assert len(trades) == 1
        assert trades[0].slippage == pytest.approx(1.0)

    def test_commission_on_buy(self) -> None:
        """买入佣金计算。"""
        df = _make_df(10)
        model = ImmediateExecution()
        signal = Signal(datetime=20240101, direction="BUY", size=100)
        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
        # 101 × 100 × 0.0003 = 3.03, min_commission=5 → 5.0
        assert trades[0].commission >= 5.0

    def test_stamp_tax_on_sell(self) -> None:
        """卖出印花税。"""
        df = _make_df(10)
        model = ImmediateExecution()
        signal = Signal(datetime=20240101, direction="SELL", size=100)
        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
        # commission = max(101×100×0.0003, 5) + 101×100×0.001 = 5 + 10.1 = 15.1
        assert trades[0].commission > 10.0

    def test_full_position_buy(self) -> None:
        """full 模式全仓买入(100股整手)。"""
        df = _make_df(10)
        model = ImmediateExecution()
        signal = Signal(datetime=20240101, direction="BUY", size=0)  # size=0 = full
        trades = model.execute(
            signal=signal, df=df, bar_idx=0,
            cash=20000, position=0, position_mode="full",
            commission=0.0003, min_commission=5.0, stamp_tax=0.001,
            slippage_model=None,
        )
        assert len(trades) == 1
        assert trades[0].size == 100  # 20000 / (101 × 1.0003) ≈ 197 → 100
  • Step 2: Run tests to verify they fail

Run: python -m pytest tests/unit/test_backtest_execution.py -v -k "TestExecutionBase or TestImmediate" --no-header Expected: FAIL (import error)

  • Step 3: Write implementation
"""可插拔执行仿真引擎。"""
from __future__ import annotations

from abc import ABC, abstractmethod
from typing import TYPE_CHECKING

import numpy as np
import pandas as pd

from easy_tdx.backtest.types import Trade

if TYPE_CHECKING:
    from easy_tdx.backtest.slippage import SlippageModel
    from easy_tdx.backtest.types import Signal


class ExecutionModel(ABC):
    """执行仿真基类。

    将信号转换为一笔或多笔成交记录。
    """

    @abstractmethod
    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]:
        """将信号转换为一笔或多笔成交。"""
        ...

    def _calc_commission(
        self, size: float, price: float, is_sell: bool,
        commission: float, min_commission: float, stamp_tax: float,
    ) -> float:
        """计算手续费。"""
        comm = max(size * price * commission, min_commission)
        if is_sell:
            comm += size * price * stamp_tax
        return comm

    def _calc_slippage(
        self, size: float, price: float, is_sell: bool,
        slippage_model: SlippageModel | None, df: pd.DataFrame,
    ) -> float:
        """计算滑点。"""
        if slippage_model is None:
            return 0.0
        volume = float(df["volume"].iloc[-1]) if "volume" in df.columns else 0.0
        volatility = self._estimate_volatility(df)
        return slippage_model.compute(
            price=price, size=size, volume=volume,
            volatility=volatility,
            direction="SELL" if is_sell else "BUY",
        )

    def _estimate_volatility(self, df: pd.DataFrame) -> float:
        """从收盘价估计近期年化波动率。"""
        if "close" not in df.columns or len(df) < 2:
            return 0.0
        close = df["close"].to_numpy()
        returns = np.diff(close) / close[:-1]
        if len(returns) < 2:
            return 0.0
        return float(np.std(returns)) * np.sqrt(252)

    def _calc_buy_size(
        self, signal_size: float, price: float, cash: float,
        position_mode: str, commission: float,
    ) -> float:
        """计算买入数量。"""
        if position_mode == "full" or signal_size == 0:
            max_cost = price * (1 + commission)
            max_shares = int(cash / max_cost / 100) * 100
            return float(max_shares)
        elif position_mode == "percent":
            target_value = cash * signal_size
            return float(int(target_value / price / 100) * 100)
        return signal_size


class ImmediateExecution(ExecutionModel):
    """即时成交(向后兼容,与现有 OrderSimulator 行为一致)。"""

    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]:
        exec_idx = bar_idx + 1
        if exec_idx >= len(df):
            return []

        price = float(df["open"].iloc[exec_idx])

        if signal.direction == "BUY":
            size = self._calc_buy_size(signal.size, price, cash, position_mode, commission)
            if size <= 0:
                return []
            comm = self._calc_commission(size, price, False, commission, min_commission, stamp_tax)
            slip = self._calc_slippage(size, price, False, slippage_model, df)
            return [Trade(
                datetime=int(df["datetime"].iloc[exec_idx])
                if hasattr(df["datetime"].iloc[exec_idx], "strftime")
                else int(df["datetime"].iloc[exec_idx]),
                direction="BUY",
                size=size,
                price=price,
                commission=comm,
                slippage=slip,
            )]

        elif signal.direction == "SELL":
            size = signal.size if signal.size > 0 else position
            if size <= 0:
                return []
            if size > position:
                size = position
            comm = self._calc_commission(size, price, True, commission, min_commission, stamp_tax)
            slip = self._calc_slippage(size, price, True, slippage_model, df)
            return [Trade(
                datetime=int(df["datetime"].iloc[exec_idx])
                if hasattr(df["datetime"].iloc[exec_idx], "strftime")
                else int(df["datetime"].iloc[exec_idx]),
                direction="SELL",
                size=size,
                price=price,
                commission=comm,
                slippage=slip,
            )]

        return []
  • Step 4: Run tests

Run: python -m pytest tests/unit/test_backtest_execution.py -v -k "TestExecutionBase or TestImmediate" --no-header Expected: 8 passed

  • Step 5: Commit
git add src/easy_tdx/backtest/execution.py tests/unit/test_backtest_execution.py
git commit -m "feat(backtest): add ExecutionModel ABC + ImmediateExecution"

Task 5: TWAPExecution + VWAPExecution

Files:

  • Modify: src/easy_tdx/backtest/execution.py

  • Modify: tests/unit/test_backtest_execution.py

  • Step 1: Write tests

test_backtest_execution.py 追加:

from easy_tdx.backtest.execution import TWAPExecution, VWAPExecution


class TestTWAPExecution:
    """时间加权平均价格执行。"""

    def test_split_buy_into_3_bars(self) -> None:
        """买入订单拆分为 3 个子订单。"""
        df = _make_df(20)
        model = TWAPExecution(n_bars=3)
        signal = Signal(datetime=20240101, direction="BUY", size=300)
        trades = model.execute(
            signal=signal, df=df, bar_idx=0,
            cash=100000, position=0, position_mode="fixed",
            commission=0.0003, min_commission=5.0, stamp_tax=0.001,
            slippage_model=None,
        )
        assert len(trades) == 3
        total_size = sum(t.size for t in trades)
        assert total_size <= 300  # 可能有100股取整损失
        # 每笔在不同 bar 执行
        prices = [t.price for t in trades]
        assert prices[0] != prices[1]  # 不同 bar 价格不同

    def test_split_sell_into_2_bars(self) -> None:
        """卖出订单拆分为 2 个子订单。"""
        df = _make_df(20)
        model = TWAPExecution(n_bars=2)
        signal = Signal(datetime=20240101, direction="SELL", size=200)
        trades = model.execute(
            signal=signal, df=df, bar_idx=0,
            cash=0, position=500, position_mode="fixed",
            commission=0.0003, min_commission=5.0, stamp_tax=0.001,
            slippage_model=None,
        )
        assert len(trades) == 2
        assert sum(t.size for t in trades) == 200.0

    def test_truncates_at_data_end(self) -> None:
        """数据不足 n_bars 时截断。"""
        df = _make_df(5)  # 只有 5 根 K 线
        model = TWAPExecution(n_bars=10)
        signal = Signal(datetime=20240101, direction="BUY", size=1000)
        trades = model.execute(
            signal=signal, df=df, bar_idx=0,
            cash=100000, position=0, position_mode="fixed",
            commission=0.0003, min_commission=5.0, stamp_tax=0.001,
            slippage_model=None,
        )
        # bar_idx=0, 可用 bar 1-4, 最多 4 笔
        assert len(trades) <= 4

    def test_full_position_mode(self) -> None:
        """full 模式下拆分全仓买入。"""
        df = _make_df(20)
        model = TWAPExecution(n_bars=3)
        signal = Signal(datetime=20240101, direction="BUY", size=0)
        trades = model.execute(
            signal=signal, df=df, bar_idx=0,
            cash=60000, position=0, position_mode="full",
            commission=0.0003, min_commission=5.0, stamp_tax=0.001,
            slippage_model=None,
        )
        assert len(trades) == 3
        assert all(t.size > 0 for t in trades)


class TestVWAPExecution:
    """成交量加权平均价格执行。"""

    def test_basic_buy(self) -> None:
        """基本买入执行。"""
        df = _make_df(20)
        model = VWAPExecution(n_bars=3, volume_lookback=10)
        signal = Signal(datetime=20240101, direction="BUY", size=300)
        trades = model.execute(
            signal=signal, df=df, bar_idx=5,
            cash=100000, position=0, position_mode="fixed",
            commission=0.0003, min_commission=5.0, stamp_tax=0.001,
            slippage_model=None,
        )
        assert len(trades) == 3
        total_size = sum(t.size for t in trades)
        assert total_size <= 300

    def test_volume_weighted_split(self) -> None:
        """成交量大的 bar 分配更多数量。"""
        import pandas as pd
        df = _make_df(20)
        # 让不同 bar 的 volume 不同
        df.loc[6, "volume"] = 50000  # bar 7 volume 很大
        df.loc[7, "volume"] = 50000
        df.loc[8, "volume"] = 50000
        model = VWAPExecution(n_bars=3, volume_lookback=5)
        signal = Signal(datetime=20240105, direction="BUY", size=300)
        trades = model.execute(
            signal=signal, df=df, bar_idx=5,
            cash=100000, position=0, position_mode="fixed",
            commission=0.0003, min_commission=5.0, stamp_tax=0.001,
            slippage_model=None,
        )
        assert len(trades) == 3
        # 各笔 size 可能因成交量分布不同而不等
        sizes = [t.size for t in trades]
        assert sum(sizes) <= 300

    def test_truncates_at_data_end(self) -> None:
        """数据不足时截断。"""
        df = _make_df(5)
        model = VWAPExecution(n_bars=10, volume_lookback=3)
        signal = Signal(datetime=20240101, direction="BUY", size=1000)
        trades = model.execute(
            signal=signal, df=df, bar_idx=0,
            cash=100000, position=0, position_mode="fixed",
            commission=0.0003, min_commission=5.0, stamp_tax=0.001,
            slippage_model=None,
        )
        assert len(trades) <= 4
  • Step 2: Run tests to verify they fail

Run: python -m pytest tests/unit/test_backtest_execution.py -v -k "TestTWAP or TestVWAP" --no-header Expected: FAIL

  • Step 3: Write implementation

execution.py 追加:

class TWAPExecution(ExecutionModel):
    """时间加权平均价格执行。

    将订单均匀拆分为 n_bars 份,在连续 n_bars 根 K 线上执行。
    """

    def __init__(self, n_bars: int = 5) -> None:
        self._n_bars = max(1, n_bars)

    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.direction == "BUY":
            return self._execute_buy(
                signal, df, bar_idx, cash, position_mode,
                commission, min_commission, stamp_tax, slippage_model,
            )
        else:
            return self._execute_sell(
                signal, df, bar_idx, position,
                commission, min_commission, stamp_tax, slippage_model,
            )

    def _execute_buy(
        self, signal: Signal, df: pd.DataFrame, bar_idx: int,
        cash: float, position_mode: str, commission: float,
        min_commission: float, stamp_tax: float,
        slippage_model: SlippageModel | None,
    ) -> list[Trade]:
        # 计算总买入量
        first_price = float(df["open"].iloc[bar_idx + 1]) if bar_idx + 1 < len(df) else 0
        if first_price <= 0:
            return []
        total_size = self._calc_buy_size(signal.size, first_price, cash, position_mode, commission)
        if total_size <= 0:
            return []

        sub_size = int(total_size / self._n_bars / 100) * 100
        if sub_size <= 0:
            sub_size = 100  # 至少 1 手

        trades: list[Trade] = []
        for i in range(self._n_bars):
            exec_idx = bar_idx + 1 + i
            if exec_idx >= len(df):
                break
            price = float(df["close"].iloc[exec_idx])
            actual_size = min(sub_size, total_size - sum(t.size for t in trades))
            actual_size = int(actual_size / 100) * 100
            if actual_size <= 0:
                break
            comm = self._calc_commission(actual_size, price, False, commission, min_commission, stamp_tax)
            slip = self._calc_slippage(actual_size, price, False, slippage_model, df)
            dt_raw = df["datetime"].iloc[exec_idx]
            dt_int = int(dt_raw.strftime("%Y%m%d")) if hasattr(dt_raw, "strftime") else int(dt_raw)
            trades.append(Trade(
                datetime=dt_int, direction="BUY",
                size=float(actual_size), price=price,
                commission=comm, slippage=slip,
            ))
        return trades

    def _execute_sell(
        self, signal: Signal, df: pd.DataFrame, bar_idx: int,
        position: float, commission: float,
        min_commission: float, stamp_tax: float,
        slippage_model: SlippageModel | None,
    ) -> list[Trade]:
        total_size = signal.size if signal.size > 0 else position
        if total_size <= 0:
            return []

        sub_size = int(total_size / self._n_bars / 100) * 100
        if sub_size <= 0:
            sub_size = 100

        trades: list[Trade] = []
        for i in range(self._n_bars):
            exec_idx = bar_idx + 1 + i
            if exec_idx >= len(df):
                break
            price = float(df["close"].iloc[exec_idx])
            actual_size = min(sub_size, total_size - sum(t.size for t in trades))
            actual_size = int(actual_size / 100) * 100
            if actual_size <= 0:
                break
            comm = self._calc_commission(actual_size, price, True, commission, min_commission, stamp_tax)
            slip = self._calc_slippage(actual_size, price, True, slippage_model, df)
            dt_raw = df["datetime"].iloc[exec_idx]
            dt_int = int(dt_raw.strftime("%Y%m%d")) if hasattr(dt_raw, "strftime") else int(dt_raw)
            trades.append(Trade(
                datetime=dt_int, direction="SELL",
                size=float(actual_size), price=price,
                commission=comm, slippage=slip,
            ))
        return trades


class VWAPExecution(ExecutionModel):
    """成交量加权平均价格执行。

    按历史成交量分布比例拆分订单。
    """

    def __init__(self, n_bars: int = 5, volume_lookback: int = 20) -> None:
        self._n_bars = max(1, n_bars)
        self._volume_lookback = max(1, volume_lookback)

    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.direction == "BUY":
            return self._execute_buy(
                signal, df, bar_idx, cash, position_mode,
                commission, min_commission, stamp_tax, slippage_model,
            )
        else:
            return self._execute_sell(
                signal, df, bar_idx, position,
                commission, min_commission, stamp_tax, slippage_model,
            )

    def _get_volume_weights(self, df: pd.DataFrame, bar_idx: int) -> list[float]:
        """获取成交量权重分布。"""
        start = max(0, bar_idx - self._volume_lookback + 1)
        lookback = df.iloc[start:bar_idx + 1]
        if "volume" not in lookback.columns or len(lookback) == 0:
            return [1.0 / self._n_bars] * self._n_bars

        volumes = lookback["volume"].to_numpy()
        total_vol = float(volumes.sum())
        if total_vol <= 0:
            return [1.0 / self._n_bars] * self._n_bars

        # 取最近 n_bars 期的平均成交量比例
        weights: list[float] = []
        for i in range(self._n_bars):
            idx = max(0, len(volumes) - 1 - (i % max(1, len(volumes))))
            weights.append(float(volumes[idx]) / total_vol)
        total_w = sum(weights)
        if total_w <= 0:
            return [1.0 / self._n_bars] * self._n_bars
        return [w / total_w for w in weights]

    def _execute_buy(
        self, signal: Signal, df: pd.DataFrame, bar_idx: int,
        cash: float, position_mode: str, commission: float,
        min_commission: float, stamp_tax: float,
        slippage_model: SlippageModel | None,
    ) -> list[Trade]:
        first_price = float(df["open"].iloc[bar_idx + 1]) if bar_idx + 1 < len(df) else 0
        if first_price <= 0:
            return []
        total_size = self._calc_buy_size(signal.size, first_price, cash, position_mode, commission)
        if total_size <= 0:
            return []

        weights = self._get_volume_weights(df, bar_idx)
        trades: list[Trade] = []
        for i in range(self._n_bars):
            exec_idx = bar_idx + 1 + i
            if exec_idx >= len(df):
                break
            price = float(df["close"].iloc[exec_idx])
            w = weights[i] if i < len(weights) else 1.0 / self._n_bars
            target = int(total_size * w / 100) * 100
            remaining = total_size - sum(t.size for t in trades)
            actual_size = min(target, remaining)
            actual_size = int(actual_size / 100) * 100
            if actual_size <= 0:
                continue
            comm = self._calc_commission(actual_size, price, False, commission, min_commission, stamp_tax)
            slip = self._calc_slippage(actual_size, price, False, slippage_model, df)
            dt_raw = df["datetime"].iloc[exec_idx]
            dt_int = int(dt_raw.strftime("%Y%m%d")) if hasattr(dt_raw, "strftime") else int(dt_raw)
            trades.append(Trade(
                datetime=dt_int, direction="BUY",
                size=float(actual_size), price=price,
                commission=comm, slippage=slip,
            ))
        return trades

    def _execute_sell(
        self, signal: Signal, df: pd.DataFrame, bar_idx: int,
        position: float, commission: float,
        min_commission: float, stamp_tax: float,
        slippage_model: SlippageModel | None,
    ) -> list[Trade]:
        total_size = signal.size if signal.size > 0 else position
        if total_size <= 0:
            return []

        weights = self._get_volume_weights(df, bar_idx)
        trades: list[Trade] = []
        for i in range(self._n_bars):
            exec_idx = bar_idx + 1 + i
            if exec_idx >= len(df):
                break
            price = float(df["close"].iloc[exec_idx])
            w = weights[i] if i < len(weights) else 1.0 / self._n_bars
            target = int(total_size * w / 100) * 100
            remaining = total_size - sum(t.size for t in trades)
            actual_size = min(target, remaining)
            actual_size = int(actual_size / 100) * 100
            if actual_size <= 0:
                continue
            comm = self._calc_commission(actual_size, price, True, commission, min_commission, stamp_tax)
            slip = self._calc_slippage(actual_size, price, True, slippage_model, df)
            dt_raw = df["datetime"].iloc[exec_idx]
            dt_int = int(dt_raw.strftime("%Y%m%d")) if hasattr(dt_raw, "strftime") else int(dt_raw)
            trades.append(Trade(
                datetime=dt_int, direction="SELL",
                size=float(actual_size), price=price,
                commission=comm, slippage=slip,
            ))
        return trades
  • Step 4: Run tests

Run: python -m pytest tests/unit/test_backtest_execution.py -v --no-header Expected: All pass

  • Step 5: Commit
git add src/easy_tdx/backtest/execution.py tests/unit/test_backtest_execution.py
git commit -m "feat(backtest): add TWAPExecution + VWAPExecution"

Task 6: LimitExecution

Files:

  • Modify: src/easy_tdx/backtest/execution.py

  • Modify: tests/unit/test_backtest_execution.py

  • Step 1: Write tests

追加到 test_backtest_execution.py

from easy_tdx.backtest.execution import LimitExecution


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
        # bar 1: low=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)
        # 价格从 100+ 递增,limit=50 不可能触发
        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  # next bar open

    def test_ttl_expires(self) -> None:
        """超出 TTL 后不再尝试。"""
        df = _make_df(20)
        model = LimitExecution(ttl_bars=2)
        # 价格从 99+ 递增,limit=98.0 在 bar 1 (low=100) 和 bar 2 (low=101) 都不触发
        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  # TTL=2 bars 内 low 都 > 98
  • Step 2: Run tests to verify they fail

Run: python -m pytest tests/unit/test_backtest_execution.py -v -k "TestLimit" --no-header Expected: FAIL

  • Step 3: Write implementation

execution.py 追加:

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)

                dt_raw = row["datetime"]
                dt_int = int(dt_raw.strftime("%Y%m%d")) if hasattr(dt_raw, "strftime") else int(dt_raw)
                return [Trade(
                    datetime=dt_int,
                    direction=signal.direction,
                    size=float(size),
                    price=target_price,
                    commission=comm,
                    slippage=slip,
                )]

        return []  # TTL 内未触发
  • Step 4: Run tests

Run: python -m pytest tests/unit/test_backtest_execution.py -v --no-header Expected: All pass

  • Step 5: Commit
git add src/easy_tdx/backtest/execution.py tests/unit/test_backtest_execution.py
git commit -m "feat(backtest): add LimitExecution"

Task 7: BacktestEngine 集成 SlippageModel + ExecutionModel

Files:

  • Modify: src/easy_tdx/backtest/engine.py

  • Modify: tests/unit/test_backtest_engine.py

  • Step 1: Write tests

tests/unit/test_backtest_engine.py 追加:

from easy_tdx.backtest.slippage import FixedSlippage, SquareRootSlippage
from easy_tdx.backtest.execution import TWAPExecution


class TestEngineSlippageModel:
    """BacktestEngine 与 SlippageModel 集成。"""

    def test_engine_with_slippage_model(self) -> None:
        """引擎使用 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 与 ExecutionModel 集成。"""

    def test_engine_with_twap(self) -> None:
        """引擎使用 TWAP 执行。"""
        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  # TWAP 产生多笔交易

    def test_engine_backward_compatible(self) -> None:
        """无新参数时行为不变。"""
        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
  • Step 2: Run tests to verify they fail

Run: python -m pytest tests/unit/test_backtest_engine.py -v -k "TestEngineSlippage or TestEngineExecution" --no-header Expected: FAIL

  • Step 3: Modify engine.py

Changes:

  1. Add imports at top of engine.py:
from __future__ import annotations

from dataclasses import dataclass
from typing import TYPE_CHECKING, Any

import pandas as pd

from easy_tdx.backtest.orders import OrderSimulator
from easy_tdx.backtest.performance import PerformanceAnalyzer
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
  1. Add parameters to BacktestEngine.__init__:
def __init__(
    self,
    strategy: type[Strategy] | Strategy,
    cash: float = 100000.0,
    commission: float = 0.0003,
    min_commission: float = 5.0,
    stamp_tax: float = 0.001,
    slippage: float = 0.0,
    execution: str = "next_open",
    position_mode: str = "full",
    reject_policy: str = "reduce",
    benchmark: pd.DataFrame | None = None,
    chanlun_level: str | None = None,
    slippage_model: SlippageModel | None = None,     # NEW
    execution_model: ExecutionModel | None = None,     # NEW
):
    # ... store as self._slippage_model and self._execution_model
  1. Modify run() method to support execution_model path:
def run(self, df, chanlun_result=None):
    # ... existing validation ...

    signals = self._generate_signals(df, chanlun_result)

    if self._execution_model is not None:
        # NEW: use execution model
        trades = self._execute_with_model(signals, df)
    else:
        # EXISTING: use OrderSimulator
        simulator = OrderSimulator(
            df, execution=self._execution, position_mode=self._position_mode,
            reject_policy=self._reject_policy, commission=self._commission,
            min_commission=self._min_commission, stamp_tax=self._stamp_tax,
            slippage=self._slippage,
            slippage_model=self._slippage_model,  # NEW
        )
        trades = simulator.simulate(signals=signals, cash=self._cash, position=0.0)

    # ... rest unchanged ...
  1. Add _execute_with_model method:
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

def _find_bar_index(self, df: pd.DataFrame, datetime_val: int) -> int | None:
    """Find bar index for a 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
  1. Also pass slippage_model to OrderSimulator in the existing path (when no execution_model).
  • Step 4: Run all backtest tests

Run: python -m pytest tests/unit/test_backtest_engine.py -v --no-header Expected: All existing + 3 new tests pass

  • Step 5: Run full test suite

Run: python -m pytest tests/unit/ -v --no-header -q Expected: All tests pass

  • Step 6: Commit
git add src/easy_tdx/backtest/engine.py tests/unit/test_backtest_engine.py
git commit -m "feat(backtest): integrate SlippageModel + ExecutionModel into BacktestEngine"

Task 8: 版本号 + 文档更新 + 最终验证

Files:

  • Modify: pyproject.toml

  • Modify: CHANGELOG.md (if exists)

  • Step 1: Bump version

Update pyproject.toml version from 1.13.0 to 1.14.0.

  • Step 2: Run full test suite

Run: python -m pytest tests/unit/ -v --no-header -q Expected: All tests pass (498 existing + ~55 new)

  • Step 3: Run lint + type check

Run: ruff check src/easy_tdx/backtest/slippage.py src/easy_tdx/backtest/execution.py src/easy_tdx/backtest/orders.py src/easy_tdx/backtest/engine.py Run: ruff format --check src/easy_tdx/backtest/ Expected: No errors

  • Step 4: Final commit
git add pyproject.toml
git commit -m "chore: bump version to v1.14.0"