mirror of
https://ghfast.top/https://github.com/aeroxw/easy-tdx.git
synced 2026-09-12 15:44:15 +08:00
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:
co-authored by
Claude Opus 4.8
parent
5f14c44791
commit
f37b75ea42
@@ -0,0 +1 @@
|
||||
"""easy_tdx.backtest — 向量化策略回测引擎(纯计算,零网络依赖)。"""
|
||||
@@ -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)}")
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user