From a2aa319803bdf7280bebe6cb9dbd2184220f82e1 Mon Sep 17 00:00:00 2001 From: GitHub Date: Tue, 9 Jun 2026 17:55:24 +0800 Subject: [PATCH] feat(backtest): add PortfolioTracker with equity curve and drawdown - Pre-allocate numpy arrays for performance (cash, position, avg_price) - apply_trades() processes buys/sells with commission and slippage - equity_curve returns DataFrame with drawdown calculation - positions returns DataFrame with market value and unrealized PnL - 12 unit tests covering all scenarios Co-Authored-By: Claude Opus 4.8 --- src/easy_tdx/backtest/portfolio.py | 153 ++++++++++++++ tests/unit/test_backtest_portfolio.py | 289 ++++++++++++++++++++++++++ 2 files changed, 442 insertions(+) create mode 100644 src/easy_tdx/backtest/portfolio.py create mode 100644 tests/unit/test_backtest_portfolio.py diff --git a/src/easy_tdx/backtest/portfolio.py b/src/easy_tdx/backtest/portfolio.py new file mode 100644 index 0000000..4c67e1f --- /dev/null +++ b/src/easy_tdx/backtest/portfolio.py @@ -0,0 +1,153 @@ +"""回测引擎持仓追踪器。 + +负责应用交易记录、计算资金曲线和持仓快照。 +""" + +from __future__ import annotations + +from typing import Any + +import numpy as np +import pandas as pd + +from easy_tdx.backtest.types import Trade + + +class PortfolioTracker: + """持仓追踪器。 + + 预分配 numpy 数组存储每个 bar 的状态,遍历应用交易。 + + Attributes: + _close: 收盘价数组 + _datetime: 时间戳数组 + _n: bar 数量 + _cash: 每个 bar 的现金 + _position: 每个 bar 的持仓数量 + _avg_price: 每个 bar 的平均成本价 + """ + + def __init__(self, df: pd.DataFrame, initial_cash: float = 100000) -> None: + """初始化追踪器。 + + Args: + df: 必须包含 close 和 datetime 列的 DataFrame + initial_cash: 初始现金 + """ + self._close = df["close"].to_numpy(dtype=float) + self._datetime = df["datetime"].to_numpy() + self._n = len(df) + self._cash = np.full(self._n, initial_cash) + self._position = np.zeros(self._n) + self._avg_price = np.zeros(self._n) + self._initial_cash = initial_cash + + def apply_trades(self, trades: list[Trade]) -> None: + """应用交易记录,更新内部状态数组。 + + Args: + trades: 交易列表 + """ + # 构建 datetime → Trade 映射 + buy_map: dict[int, Trade] = {} + sell_map: dict[int, Trade] = {} + + for trade in trades: + if trade.rejected: + continue + if trade.direction == "BUY": + buy_map[trade.datetime] = trade + else: + sell_map[trade.datetime] = trade + + # 遍历每个 bar + for i in range(self._n): + dt = self._datetime[i] + + # 继承前一个 bar 的状态(除了第一个 bar) + if i > 0: + self._cash[i] = self._cash[i - 1] + self._position[i] = self._position[i - 1] + self._avg_price[i] = self._avg_price[i - 1] + + # 处理买入 + if dt in buy_map: + trade = buy_map[dt] + cost = trade.size * trade.price + trade.commission + trade.slippage + self._cash[i] -= cost + + # 更新均价 + if self._position[i] > 0: + total_cost = self._position[i] * self._avg_price[i] + trade.size * trade.price + self._position[i] += trade.size + self._avg_price[i] = total_cost / self._position[i] + else: + # 新开仓或从空仓开仓 + self._position[i] = trade.size + self._avg_price[i] = trade.price + + # 处理卖出 + if dt in sell_map: + trade = sell_map[dt] + proceeds = trade.size * trade.price - trade.commission - trade.slippage + self._cash[i] += proceeds + self._position[i] -= trade.size + + # 清空持仓时归零 + if self._position[i] <= 0: + self._position[i] = 0.0 + self._avg_price[i] = 0.0 + + @property + def equity_curve(self) -> pd.DataFrame: + """计算资金曲线。 + + Returns: + DataFrame 包含 datetime, cash, position_value, total, drawdown, drawdown_pct + """ + position_value = self._position * self._close + total = self._cash + position_value + peak = np.maximum.accumulate(total) + drawdown = peak - total + drawdown_pct = np.divide( + drawdown, + peak, + out=np.zeros_like(drawdown), + where=(peak != 0), + ) + + return pd.DataFrame( + { + "datetime": self._datetime, + "cash": self._cash, + "position_value": position_value, + "total": total, + "drawdown": drawdown, + "drawdown_pct": drawdown_pct, + } + ) + + @property + def positions(self) -> pd.DataFrame: + """计算持仓快照。 + + Returns: + DataFrame 包含 datetime, size, avg_price, market_value, unrealized_pnl + """ + market_value = self._position * self._close + unrealized_pnl = (self._close - self._avg_price) * self._position + + return pd.DataFrame( + { + "datetime": self._datetime, + "size": self._position, + "avg_price": self._avg_price, + "market_value": market_value, + "unrealized_pnl": unrealized_pnl, + } + ) + + @property + def initial_cash(self) -> float: + """初始现金。""" + return self._initial_cash diff --git a/tests/unit/test_backtest_portfolio.py b/tests/unit/test_backtest_portfolio.py new file mode 100644 index 0000000..e558bf3 --- /dev/null +++ b/tests/unit/test_backtest_portfolio.py @@ -0,0 +1,289 @@ +"""回测引擎持仓追踪器单元测试。""" + +from __future__ import annotations + +import pandas as pd + +from easy_tdx.backtest.portfolio import PortfolioTracker +from easy_tdx.backtest.types import Trade + + +def _make_df(n: int = 10) -> pd.DataFrame: + """创建测试用 DataFrame。 + + Args: + n: bar 数量 + + Returns: + 包含 close 和 datetime 列的 DataFrame + """ + close = [100.0] * 5 + [110.0] * 5 + close = close[:n] + datetime = list(range(20240101, 20240101 + n)) + + return pd.DataFrame({"close": close, "datetime": datetime}) + + +def test_initial_state() -> None: + """测试初始状态。""" + df = _make_df(10) + tracker = PortfolioTracker(df, initial_cash=100000) + + assert tracker.initial_cash == 100000 + + +def test_buy_then_sell() -> None: + """测试买入再卖出。""" + df = _make_df(10) + tracker = PortfolioTracker(df, initial_cash=100000) + + # bar 0 买入 100 股 @100,手续费 5 + buy_trade = Trade( + datetime=20240101, + direction="BUY", + size=100, + price=100.0, + commission=5.0, + slippage=0.0, + ) + + # bar 5 卖出 100 股 @110,手续费 11 + sell_trade = Trade( + datetime=20240106, + direction="SELL", + size=100, + price=110.0, + commission=11.0, + slippage=0.0, + ) + + tracker.apply_trades([buy_trade, sell_trade]) + + equity = tracker.equity_curve + final_cash = equity["cash"].iloc[-1] + + # 最终现金 = 100000 - 10000 - 5 + 11000 - 11 = 100984 + expected = 100984.0 + assert abs(final_cash - expected) < 1.0, f"Expected {expected}, got {final_cash}" + + +def test_drawdown_calculation() -> None: + """测试回撤计算(无交易时回撤全为 0)。""" + df = _make_df(10) + tracker = PortfolioTracker(df, initial_cash=100000) + + tracker.apply_trades([]) + + equity = tracker.equity_curve + + # 无交易时,总资产应保持不变,回撤为 0 + assert (equity["drawdown"] == 0).all() + assert (equity["drawdown_pct"] == 0).all() + + +def test_equity_curve_columns() -> None: + """测试资金曲线列名。""" + df = _make_df(10) + tracker = PortfolioTracker(df) + + tracker.apply_trades([]) + equity = tracker.equity_curve + + expected_cols = {"datetime", "cash", "position_value", "total", "drawdown", "drawdown_pct"} + assert set(equity.columns) == expected_cols + + +def test_position_tracking() -> None: + """测试持仓追踪。""" + df = _make_df(10) + tracker = PortfolioTracker(df) + + # bar 0 买入 100 股 + buy_trade = Trade( + datetime=20240101, + direction="BUY", + size=100, + price=100.0, + commission=5.0, + slippage=0.0, + ) + + tracker.apply_trades([buy_trade]) + positions = tracker.positions + + # 买入后持仓应为 100,并持续到最后 + assert positions["size"].iloc[0] == 100 + assert positions["size"].iloc[-1] == 100 + assert (positions["size"].iloc[1:] == 100).all() + + +def test_avg_price_update() -> None: + """测试均价更新。""" + df = _make_df(10) + tracker = PortfolioTracker(df) + + # bar 0 买入 100 股 @100 + buy1 = Trade( + datetime=20240101, + direction="BUY", + size=100, + price=100.0, + commission=0.0, + slippage=0.0, + ) + + # bar 1 再买入 50 股 @110 + buy2 = Trade( + datetime=20240102, + direction="BUY", + size=50, + price=110.0, + commission=0.0, + slippage=0.0, + ) + + tracker.apply_trades([buy1, buy2]) + positions = tracker.positions + + # 均价 = (100*100 + 50*110) / 150 = 103.33 + expected_avg = (100 * 100 + 50 * 110) / 150 + assert abs(positions["avg_price"].iloc[1] - expected_avg) < 0.01 + + +def test_sell_clears_position() -> None: + """测试卖出清空持仓。""" + df = _make_df(10) + tracker = PortfolioTracker(df) + + buy = Trade( + datetime=20240101, + direction="BUY", + size=100, + price=100.0, + commission=0.0, + slippage=0.0, + ) + + sell = Trade( + datetime=20240102, + direction="SELL", + size=100, + price=110.0, + commission=0.0, + slippage=0.0, + ) + + tracker.apply_trades([buy, sell]) + positions = tracker.positions + + # 卖出后持仓应为 0 + assert positions["size"].iloc[0] == 100 + assert positions["size"].iloc[1] == 0 + assert positions["avg_price"].iloc[1] == 0 + + +def test_position_market_value() -> None: + """测试持仓市值计算。""" + df = _make_df(10) + tracker = PortfolioTracker(df) + + buy = Trade( + datetime=20240101, + direction="BUY", + size=100, + price=100.0, + commission=0.0, + slippage=0.0, + ) + + tracker.apply_trades([buy]) + positions = tracker.positions + + # 前 5 个 bar 价格 100,后 5 个 bar 价格 110 + assert positions["market_value"].iloc[0] == 100 * 100 + assert positions["market_value"].iloc[5] == 100 * 110 + + +def test_unrealized_pnl() -> None: + """测试未实现盈亏计算。""" + df = _make_df(10) + tracker = PortfolioTracker(df) + + buy = Trade( + datetime=20240101, + direction="BUY", + size=100, + price=100.0, + commission=0.0, + slippage=0.0, + ) + + tracker.apply_trades([buy]) + positions = tracker.positions + + # 前 5 个 bar 价格 100,盈亏为 0 + assert positions["unrealized_pnl"].iloc[0] == 0 + + # 后 5 个 bar 价格 110,盈亏为 (110-100)*100 = 1000 + assert positions["unrealized_pnl"].iloc[5] == 1000 + + +def test_rejected_trade_ignored() -> None: + """测试被拒绝的交易不产生影响。""" + df = _make_df(10) + tracker = PortfolioTracker(df, initial_cash=100000) + + rejected_buy = Trade( + datetime=20240101, + direction="BUY", + size=100, + price=100.0, + commission=0.0, + slippage=0.0, + rejected=True, + ) + + tracker.apply_trades([rejected_buy]) + equity = tracker.equity_curve + + # 现金应保持不变 + assert (equity["cash"] == 100000).all() + assert (tracker.positions["size"] == 0).all() + + +def test_commission_and_slippage() -> None: + """测试手续费和滑点扣减。""" + df = _make_df(10) + tracker = PortfolioTracker(df, initial_cash=10000) + + buy = Trade( + datetime=20240101, + direction="BUY", + size=100, + price=100.0, + commission=10.0, + slippage=5.0, + ) + + tracker.apply_trades([buy]) + equity = tracker.equity_curve + + # 现金 = 10000 - 100*100 - 10 - 5 = -15 + expected_cash = 10000 - 10000 - 10 - 5 + assert equity["cash"].iloc[0] == expected_cash + + +def test_empty_trades() -> None: + """测试空交易列表不崩溃。""" + df = _make_df(10) + tracker = PortfolioTracker(df) + + tracker.apply_trades([]) + + equity = tracker.equity_curve + positions = tracker.positions + + # 所有 bar 现金应等于初始现金 + assert (equity["cash"] == 100000).all() + # 所有 bar 持仓应为 0 + assert (positions["size"] == 0).all()