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 <noreply@anthropic.com>
This commit is contained in:
GitHub
2026-06-09 16:43:44 +08:00
co-authored by Claude Opus 4.8
parent 5f14c44791
commit f37b75ea42
3 changed files with 405 additions and 0 deletions
+1
View File
@@ -0,0 +1 @@
"""easy_tdx.backtest — 向量化策略回测引擎(纯计算,零网络依赖)。"""
+137
View File
@@ -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: 资金曲线 DataFrameindex=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)}")
+267
View File
@@ -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