From f37b75ea42b2ae7ed39a141ae1fb7372db36946e Mon Sep 17 00:00:00 2001 From: GitHub Date: Tue, 9 Jun 2026 16:43:44 +0800 Subject: [PATCH] feat(backtest): add core data types (Signal/Trade/Position/BacktestResult) - Add Signal dataclass for trading signals with optional price/stop_loss/take_profit - Add Trade dataclass for executed trades with commission/slippage/pnl/rejected - Add Position dataclass for position snapshots (long/short/flat) - Add BacktestResult dataclass with to_dict()/to_json()/summary() methods - Add comprehensive unit tests (13 test cases, 100% pass) - All code passes mypy strict, ruff lint+format checks Co-Authored-By: Claude Opus 4.8 --- src/easy_tdx/backtest/__init__.py | 1 + src/easy_tdx/backtest/types.py | 137 +++++++++++++++ tests/unit/test_backtest_types.py | 267 ++++++++++++++++++++++++++++++ 3 files changed, 405 insertions(+) create mode 100644 src/easy_tdx/backtest/__init__.py create mode 100644 src/easy_tdx/backtest/types.py create mode 100644 tests/unit/test_backtest_types.py diff --git a/src/easy_tdx/backtest/__init__.py b/src/easy_tdx/backtest/__init__.py new file mode 100644 index 0000000..fd0c7b3 --- /dev/null +++ b/src/easy_tdx/backtest/__init__.py @@ -0,0 +1 @@ +"""easy_tdx.backtest — 向量化策略回测引擎(纯计算,零网络依赖)。""" diff --git a/src/easy_tdx/backtest/types.py b/src/easy_tdx/backtest/types.py new file mode 100644 index 0000000..eba4391 --- /dev/null +++ b/src/easy_tdx/backtest/types.py @@ -0,0 +1,137 @@ +"""回测引擎核心数据类型定义。 + +使用纯 dataclass + 类型注解,保持 mypy strict 兼容。 +""" + +from __future__ import annotations + +import json +from dataclasses import dataclass +from typing import Any, Literal + +import pandas as pd + +# ── 交易信号 ──────────────────────────────────────────────────────────────── + + +@dataclass +class Signal: + """策略产生的交易信号。 + + Attributes: + datetime: 信号时间(Unix timestamp 毫秒) + direction: 交易方向 + size: 交易数量(0 = 全仓/清仓) + price: 限价(None = 市价单) + stop_loss: 止损价(None = 不设置) + take_profit: 止盈价(None = 不设置) + """ + + datetime: int + direction: Literal["BUY", "SELL"] + size: float + price: float | None = None + stop_loss: float | None = None + take_profit: float | None = None + + +# ── 成交记录 ──────────────────────────────────────────────────────────────── + + +@dataclass +class Trade: + """已成交记录。 + + Attributes: + datetime: 成交时间(Unix timestamp 毫秒) + direction: 交易方向 + size: 成交数量 + price: 成交价格 + commission: 手续费 + slippage: 滑点成本 + pnl: 已实现盈亏(仅平仓时计算) + rejected: 是否被拒绝(资金不足/不允许做空等) + """ + + datetime: int + direction: Literal["BUY", "SELL"] + size: float + price: float + commission: float + slippage: float + pnl: float = 0.0 + rejected: bool = False + + +# ── 持仓快照 ──────────────────────────────────────────────────────────────── + + +@dataclass +class Position: + """持仓快照。 + + Attributes: + datetime: 快照时间(Unix timestamp 毫秒) + size: 持仓数量(正=多头,负=空头,0=空仓) + avg_price: 平均持仓成本 + market_value: 市值 + unrealized_pnl: 未实现盈亏 + """ + + datetime: int + size: float + avg_price: float + market_value: float + unrealized_pnl: float + + +# ── 回测结果 ──────────────────────────────────────────────────────────────── + + +@dataclass +class BacktestResult: + """回测完整结果。 + + Attributes: + performance: 绩效指标字典(总收益率、夏普比率、最大回撤等) + equity_curve: 资金曲线 DataFrame(index=datetime, columns=equity/drawdown等) + trades: 成交记录 DataFrame + positions: 持仓快照 DataFrame + config: 配置参数字典 + """ + + performance: dict[str, float] + equity_curve: pd.DataFrame + trades: pd.DataFrame + positions: pd.DataFrame + config: dict[str, Any] + + def to_dict(self) -> dict[str, Any]: + """将结果转换为 JSON 兼容字典。 + + DataFrame 转为 records 列表格式。 + """ + return { + "performance": self.performance, + "equity_curve": self.equity_curve.to_dict(orient="records"), + "trades": self.trades.to_dict(orient="records"), + "positions": self.positions.to_dict(orient="records"), + "config": self.config, + } + + def to_json(self) -> str: + """将结果序列化为 JSON 字符串。""" + return json.dumps(self.to_dict(), ensure_ascii=False, indent=2) + + def summary(self) -> None: + """打印回测概要(标准输出)。""" + print("=== 回测绩效概要 ===") + for key, value in self.performance.items(): + if isinstance(value, float): + print(f"{key}: {value:.4f}") + else: + print(f"{key}: {value}") + + print(f"\n成交记录数: {len(self.trades)}") + print(f"持仓快照数: {len(self.positions)}") + print(f"资金曲线点数: {len(self.equity_curve)}") diff --git a/tests/unit/test_backtest_types.py b/tests/unit/test_backtest_types.py new file mode 100644 index 0000000..532dc24 --- /dev/null +++ b/tests/unit/test_backtest_types.py @@ -0,0 +1,267 @@ +"""回测引擎数据类型单元测试。""" + +from __future__ import annotations + +import json +from typing import Any + +import pandas as pd + +from easy_tdx.backtest.types import BacktestResult, Position, Signal, Trade + +# ── Signal 测试 ───────────────────────────────────────────────────────────── + + +def test_signal_defaults() -> None: + """测试 Signal 的默认值。""" + sig = Signal( + datetime=1715356800000, # 2024-06-09 00:00:00 UTC + direction="BUY", + size=100.0, + ) + assert sig.price is None + assert sig.stop_loss is None + assert sig.take_profit is None + + +def test_signal_with_all_fields() -> None: + """测试 Signal 完整字段。""" + sig = Signal( + datetime=1715356800000, + direction="SELL", + size=50.0, + price=10.5, + stop_loss=11.0, + take_profit=10.0, + ) + assert sig.direction == "SELL" + assert sig.size == 50.0 + assert sig.price == 10.5 + assert sig.stop_loss == 11.0 + assert sig.take_profit == 10.0 + + +def test_direction_literal() -> None: + """测试 direction 字面量类型。""" + sig_buy = Signal(datetime=0, direction="BUY", size=1.0) + sig_sell = Signal(datetime=0, direction="SELL", size=1.0) + assert sig_buy.direction == "BUY" + assert sig_sell.direction == "SELL" + + +# ── Trade 测试 ────────────────────────────────────────────────────────────── + + +def test_trade_defaults() -> None: + """测试 Trade 的默认值。""" + trade = Trade( + datetime=1715356800000, + direction="BUY", + size=100.0, + price=10.0, + commission=0.3, + slippage=0.1, + ) + assert trade.pnl == 0.0 + assert trade.rejected is False + + +def test_trade_with_pnl() -> None: + """测试带盈亏的成交记录。""" + trade = Trade( + datetime=1715356800000, + direction="SELL", + size=100.0, + price=11.0, + commission=0.3, + slippage=0.1, + pnl=50.0, + ) + assert trade.pnl == 50.0 + + +def test_trade_rejected() -> None: + """测试被拒绝的成交。""" + trade = Trade( + datetime=1715356800000, + direction="BUY", + size=1000000.0, # 超大单模拟资金不足 + price=10.0, + commission=0.0, + slippage=0.0, + rejected=True, + ) + assert trade.rejected is True + + +# ── Position 测试 ─────────────────────────────────────────────────────────── + + +def test_position_long() -> None: + """测试多头持仓。""" + pos = Position( + datetime=1715356800000, + size=100.0, + avg_price=10.0, + market_value=1050.0, + unrealized_pnl=50.0, + ) + assert pos.size > 0 + assert pos.avg_price == 10.0 + assert pos.market_value == 1050.0 + assert pos.unrealized_pnl == 50.0 + + +def test_position_short() -> None: + """测试空头持仓。""" + pos = Position( + datetime=1715356800000, + size=-100.0, + avg_price=10.0, + market_value=-1050.0, + unrealized_pnl=50.0, + ) + assert pos.size < 0 + + +def test_position_flat() -> None: + """测试空仓。""" + pos = Position( + datetime=1715356800000, + size=0.0, + avg_price=0.0, + market_value=0.0, + unrealized_pnl=0.0, + ) + assert pos.size == 0.0 + + +# ── BacktestResult 测试 ───────────────────────────────────────────────────── + + +def test_backtest_result_to_dict() -> None: + """测试 BacktestResult.to_dict() 序列化。""" + equity_df = pd.DataFrame( + { + "datetime": [1715356800000, 1715443200000], + "equity": [10000.0, 10100.0], + "drawdown": [0.0, -0.0099], + } + ) + equity_df = equity_df.set_index("datetime") + + trades_df = pd.DataFrame( + { + "datetime": [1715356800000], + "direction": ["BUY"], + "size": [100.0], + "price": [10.0], + "commission": [0.3], + "slippage": [0.1], + "pnl": [0.0], + "rejected": [False], + } + ) + + positions_df = pd.DataFrame( + { + "datetime": [1715356800000], + "size": [100.0], + "avg_price": [10.0], + "market_value": [1000.0], + "unrealized_pnl": [0.0], + } + ) + + result = BacktestResult( + performance={"total_return": 0.01, "sharpe_ratio": 1.5}, + equity_curve=equity_df, + trades=trades_df, + positions=positions_df, + config={"initial_capital": 10000.0}, + ) + + data = result.to_dict() + + assert isinstance(data, dict) + assert data["performance"]["total_return"] == 0.01 + assert isinstance(data["equity_curve"], list) + assert len(data["equity_curve"]) == 2 + assert isinstance(data["trades"], list) + assert len(data["trades"]) == 1 + assert isinstance(data["positions"], list) + assert len(data["positions"]) == 1 + assert data["config"]["initial_capital"] == 10000.0 + + +def test_backtest_result_to_json() -> None: + """测试 BacktestResult.to_json() 序列化。""" + equity_df = pd.DataFrame( + { + "datetime": [1715356800000], + "equity": [10000.0], + } + ) + equity_df = equity_df.set_index("datetime") + + result = BacktestResult( + performance={"total_return": 0.01}, + equity_curve=equity_df, + trades=pd.DataFrame(), + positions=pd.DataFrame(), + config={"initial_capital": 10000.0}, + ) + + json_str = result.to_json() + + # 验证是有效 JSON + parsed = json.loads(json_str) + assert parsed["performance"]["total_return"] == 0.01 + assert parsed["config"]["initial_capital"] == 10000.0 + assert parsed["equity_curve"][0]["equity"] == 10000.0 + + +def test_backtest_result_empty_dataframes() -> None: + """测试空 DataFrame 不崩溃。""" + result = BacktestResult( + performance={}, + equity_curve=pd.DataFrame(), + trades=pd.DataFrame(), + positions=pd.DataFrame(), + config={}, + ) + + # to_dict 不崩溃 + data = result.to_dict() + assert data["equity_curve"] == [] + assert data["trades"] == [] + assert data["positions"] == [] + + # to_json 不崩溃 + json_str = result.to_json() + parsed = json.loads(json_str) + assert parsed["equity_curve"] == [] + assert parsed["trades"] == [] + assert parsed["positions"] == [] + + +def test_backtest_result_summary(capsys: Any) -> None: + """测试 summary() 打印输出。""" + result = BacktestResult( + performance={"total_return": 0.05, "sharpe_ratio": 1.2, "max_drawdown": -0.02}, + equity_curve=pd.DataFrame(), + trades=pd.DataFrame(), + positions=pd.DataFrame(), + config={"initial_capital": 10000.0}, + ) + + result.summary() + captured = capsys.readouterr() + + assert "=== 回测绩效概要 ===" in captured.out + assert "total_return: 0.0500" in captured.out + assert "sharpe_ratio: 1.2000" in captured.out + assert "max_drawdown: -0.0200" in captured.out + assert "成交记录数: 0" in captured.out + assert "持仓快照数: 0" in captured.out + assert "资金曲线点数: 0" in captured.out