mirror of
https://ghfast.top/https://github.com/aeroxw/easy-tdx.git
synced 2026-09-12 14:34:15 +08:00
release: v1.17.11 — Web UI 策略库(SQLite 持久化)+ 多策略资金分仓组合回测
新增两层能力: 1. 策略库:单标的/组合回测结果可保存到本地 SQLite 单文件 (~/.easy_tdx/strategies.db),策略库页可载入回填、重跑、删除。 2. 多策略组合回测:勾选 N 个单标的策略,各拿 1/N 资金、各跑原标的, 净值曲线按日期并集对齐求和,组合结果含 19 项完整绩效指标 + 持仓表。 后端:strategy_store.py(SQLite CRUD) + multi_strategy_engine.py(资金分仓引擎) + routers/strategies.py + /backtest/multi-strategy/run/async。 前端:StrategiesView.vue + 保存策略按钮 + 复用组合页图表组件。 895 单测全绿(+24 新增),ruff/mypy strict/前端 vue-tsc 全通过。
This commit is contained in:
@@ -0,0 +1,211 @@
|
||||
"""单元测试:多策略资金分仓组合回测引擎(MultiStrategyEngine)。
|
||||
|
||||
覆盖:
|
||||
- 基本多策略回测(2~3 个策略,各跑各的 df,合并曲线)
|
||||
- 资金均分(1/N)
|
||||
- individual_results 的 key 格式 "{label}@{symbol}"
|
||||
- 合并净值曲线列结构 + 日期并集对齐
|
||||
- 空策略列表兜底
|
||||
- 同标的不同策略可区分
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
from easy_tdx.backtest.multi_strategy_engine import (
|
||||
MultiStrategyEngine,
|
||||
StrategySlot,
|
||||
)
|
||||
from easy_tdx.backtest.strategy import Strategy
|
||||
|
||||
|
||||
class SimpleBuyStrategy(Strategy):
|
||||
"""简单策略:bar 5 买入,bar 30 卖出。"""
|
||||
|
||||
def init(self) -> None:
|
||||
pass
|
||||
|
||||
def next(self) -> None:
|
||||
if self._bar_index == 5 and self.position["size"] == 0:
|
||||
self.buy(size=0)
|
||||
elif self._bar_index == 30 and self.position["size"] > 0:
|
||||
self.sell(size=0)
|
||||
|
||||
|
||||
class HoldStrategy(Strategy):
|
||||
"""从不交易的策略(净值曲线恒等于初始资金)。"""
|
||||
|
||||
def init(self) -> None:
|
||||
pass
|
||||
|
||||
def next(self) -> None:
|
||||
pass
|
||||
|
||||
|
||||
def _make_df(n: int = 100, seed: int = 42, start: str = "2024-01-01") -> pd.DataFrame:
|
||||
"""生成随机 OHLCV DataFrame(与 test_portfolio_engine 同构造方式)。"""
|
||||
rng = np.random.default_rng(seed)
|
||||
close = 100.0 + np.cumsum(rng.normal(0, 1, n))
|
||||
high = close + rng.uniform(0, 1, n)
|
||||
low = close - rng.uniform(0, 1, n)
|
||||
open_ = low + rng.uniform(0, high - low, n)
|
||||
vol = rng.integers(1_000_000, 10_000_000, n).astype(float)
|
||||
return pd.DataFrame(
|
||||
{
|
||||
"datetime": pd.date_range(start, periods=n, freq="D"),
|
||||
"open": open_,
|
||||
"high": high,
|
||||
"low": low,
|
||||
"close": close,
|
||||
"vol": vol,
|
||||
"amount": vol * close,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class TestMultiStrategyEngine:
|
||||
def test_basic_run_two_strategies(self) -> None:
|
||||
"""两个策略各跑各的 df,应产出合并结果。"""
|
||||
slots = [
|
||||
StrategySlot("双均线", "SH:601088", SimpleBuyStrategy(), _make_df(100, seed=42)),
|
||||
StrategySlot("RSI", "SZ:000001", SimpleBuyStrategy(), _make_df(100, seed=99)),
|
||||
]
|
||||
engine = MultiStrategyEngine(slots, total_cash=1_000_000)
|
||||
result = engine.run()
|
||||
|
||||
# individual_results 的 key 形如 "{label}@{symbol}"
|
||||
assert set(result.individual_results.keys()) == {
|
||||
"双均线@SH:601088",
|
||||
"RSI@SZ:000001",
|
||||
}
|
||||
# 整体绩效含基本字段
|
||||
assert "total_return" in result.total_performance
|
||||
assert result.total_performance["total_stocks"] == 2
|
||||
assert result.total_performance["total_cash"] == 1_000_000
|
||||
|
||||
def test_total_performance_has_full_metrics(self) -> None:
|
||||
"""组合整体绩效应含完整 19 项指标(夏普/回撤/胜率/盈亏比等),与单标的同口径。"""
|
||||
slots = [
|
||||
StrategySlot("双均线", "SH:601088", SimpleBuyStrategy(), _make_df(100, seed=42)),
|
||||
StrategySlot("RSI", "SZ:000001", SimpleBuyStrategy(), _make_df(100, seed=99)),
|
||||
]
|
||||
perf = MultiStrategyEngine(slots, total_cash=1_000_000).run().total_performance
|
||||
# 关键指标都应在(来自 PerformanceAnalyzer)
|
||||
for key in [
|
||||
"total_return",
|
||||
"annual_return",
|
||||
"sharpe",
|
||||
"sortino",
|
||||
"calmar",
|
||||
"max_drawdown",
|
||||
"max_dd_duration",
|
||||
"volatility",
|
||||
"total_trades",
|
||||
"win_trades",
|
||||
"lose_trades",
|
||||
"win_rate",
|
||||
"profit_factor",
|
||||
"avg_win",
|
||||
"avg_loss",
|
||||
"max_win",
|
||||
"max_loss",
|
||||
]:
|
||||
assert key in perf, f"缺少指标 {key}"
|
||||
# max_drawdown 用正值约定(与单标的一致),介于 0~1
|
||||
assert 0 <= perf["max_drawdown"] <= 1
|
||||
# 合并净值曲线的 drawdown 也应是正值
|
||||
result = MultiStrategyEngine(slots, total_cash=1_000_000).run()
|
||||
assert (result.combined_equity["drawdown"] >= 0).all()
|
||||
|
||||
def test_capital_split_equal(self) -> None:
|
||||
"""资金按策略数均分:每个槽位 1/N。"""
|
||||
slots = [
|
||||
StrategySlot("A", "SH:601088", SimpleBuyStrategy(), _make_df(50, seed=1)),
|
||||
StrategySlot("B", "SZ:000001", SimpleBuyStrategy(), _make_df(50, seed=2)),
|
||||
StrategySlot("C", "SZ:000002", SimpleBuyStrategy(), _make_df(50, seed=3)),
|
||||
]
|
||||
engine = MultiStrategyEngine(slots, total_cash=900_000)
|
||||
allocs = engine._compute_allocations() # noqa: SLF001 — 测试内部均分逻辑
|
||||
assert len(allocs) == 3
|
||||
assert all(v == 300_000 for v in allocs.values())
|
||||
# equity_allocation 是占比,各 1/3
|
||||
result = engine.run()
|
||||
assert all(abs(v - 1 / 3) < 1e-9 for v in result.equity_allocation.values())
|
||||
|
||||
def test_combined_equity_has_expected_columns(self) -> None:
|
||||
"""合并净值曲线应有 datetime/total/drawdown/drawdown_pct 列。"""
|
||||
slots = [
|
||||
StrategySlot("A", "SH:601088", SimpleBuyStrategy(), _make_df(60, seed=7)),
|
||||
]
|
||||
engine = MultiStrategyEngine(slots, total_cash=500_000)
|
||||
result = engine.run()
|
||||
cols = set(result.combined_equity.columns)
|
||||
assert {"datetime", "total", "drawdown", "drawdown_pct"} <= cols
|
||||
assert len(result.combined_equity) > 0
|
||||
|
||||
def test_combined_equity_aligns_disjoint_dates(self) -> None:
|
||||
"""两个策略日期范围不同时,合并曲线应按并集对齐(ffill)。"""
|
||||
# 策略 A 跑 2024-01 起 60 根,策略 B 跑 2024-03 起 60 根
|
||||
df_a = _make_df(60, seed=1, start="2024-01-01")
|
||||
df_b = _make_df(60, seed=2, start="2024-03-01")
|
||||
slots = [
|
||||
StrategySlot("A", "SH:601088", SimpleBuyStrategy(), df_a),
|
||||
StrategySlot("B", "SZ:000001", SimpleBuyStrategy(), df_b),
|
||||
]
|
||||
engine = MultiStrategyEngine(slots, total_cash=1_000_000)
|
||||
result = engine.run()
|
||||
# 合并曲线长度应至少覆盖两个范围的最晚结束日(并集)
|
||||
assert len(result.combined_equity) >= 60
|
||||
|
||||
def test_empty_strategies_returns_empty_result(self) -> None:
|
||||
"""空策略列表应返回空结果,不抛异常。"""
|
||||
engine = MultiStrategyEngine([], total_cash=1_000_000)
|
||||
result = engine.run()
|
||||
assert result.individual_results == {}
|
||||
assert result.total_performance["total_return"] == 0.0
|
||||
# combined_equity 为带表头的空 DataFrame
|
||||
assert len(result.combined_equity) == 0
|
||||
assert set(result.combined_equity.columns) == {
|
||||
"datetime",
|
||||
"total",
|
||||
"drawdown",
|
||||
"drawdown_pct",
|
||||
}
|
||||
|
||||
def test_same_symbol_different_strategies_distinguished(self) -> None:
|
||||
"""同标的不同策略应能区分(key 含 label)。"""
|
||||
df = _make_df(60, seed=5)
|
||||
slots = [
|
||||
StrategySlot("双均线", "SH:601088", SimpleBuyStrategy(), df.copy()),
|
||||
StrategySlot("RSI", "SH:601088", HoldStrategy(), df.copy()),
|
||||
]
|
||||
engine = MultiStrategyEngine(slots, total_cash=1_000_000)
|
||||
result = engine.run()
|
||||
# 两个 key 不同,都带同一 symbol
|
||||
assert "双均线@SH:601088" in result.individual_results
|
||||
assert "RSI@SH:601088" in result.individual_results
|
||||
|
||||
def test_hold_strategy_keeps_initial_capital(self) -> None:
|
||||
"""从不交易的策略,其净值曲线末值应等于初始分得资金。"""
|
||||
slots = [
|
||||
StrategySlot("Hold", "SH:601088", HoldStrategy(), _make_df(40, seed=1)),
|
||||
]
|
||||
engine = MultiStrategyEngine(slots, total_cash=1_000_000)
|
||||
result = engine.run()
|
||||
ec = result.individual_results["Hold@SH:601088"].equity_curve
|
||||
# 不交易 → 末值 ≈ 初始资金 1_000_000(单策略拿全部)
|
||||
assert abs(ec["total"].iloc[-1] - 1_000_000) < 1.0
|
||||
|
||||
def test_to_dict_serializable(self) -> None:
|
||||
"""to_dict 应产出 JSON 兼容结构(含 individual_results / combined_equity)。"""
|
||||
slots = [
|
||||
StrategySlot("A", "SH:601088", SimpleBuyStrategy(), _make_df(50, seed=1)),
|
||||
]
|
||||
result = MultiStrategyEngine(slots, total_cash=500_000).run()
|
||||
d = result.to_dict()
|
||||
assert "total_performance" in d
|
||||
assert "individual_results" in d
|
||||
assert "combined_equity" in d
|
||||
assert isinstance(d["individual_results"]["A@SH:601088"], dict)
|
||||
@@ -0,0 +1,256 @@
|
||||
"""策略库(已保存策略)持久化 + Web API 测试(离线,无网络)。
|
||||
|
||||
覆盖:
|
||||
- ``StrategyStore``:加入 / 列出 / 查看 / 删除 / 时间戳自动填充 / 重复 id
|
||||
- 路由端到端:POST 创建、GET 列表、GET 详情、DELETE、404 路径、校验
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlite3
|
||||
|
||||
import pytest
|
||||
|
||||
pytest.importorskip("fastapi")
|
||||
|
||||
from fastapi import FastAPI # noqa: E402
|
||||
from fastapi.testclient import TestClient # noqa: E402
|
||||
|
||||
from easy_tdx.web.strategy_store import SavedStrategy, StrategyStore # noqa: E402
|
||||
|
||||
# ── StrategyStore 单元测试 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def store(tmp_path) -> StrategyStore:
|
||||
"""每个测试独立 SQLite 文件,互不污染。"""
|
||||
return StrategyStore(db_path=tmp_path / "test_strategies.db")
|
||||
|
||||
|
||||
def _sample_single(name: str = "双均线·平安") -> SavedStrategy:
|
||||
return SavedStrategy(
|
||||
id="",
|
||||
name=name,
|
||||
kind="single",
|
||||
strategy="ma_cross",
|
||||
strategy_label="双均线交叉",
|
||||
params={"fast": 5, "slow": 20},
|
||||
context={
|
||||
"symbol": "SZ:000001",
|
||||
"category": "DAY",
|
||||
"start_date": "2023-01-01",
|
||||
"end_date": "2024-12-31",
|
||||
},
|
||||
trade_config={"cash": 1_000_000, "commission": 0.0003, "execution": "next_open"},
|
||||
snapshot={"total_return": 0.352, "max_drawdown": -0.12, "sharpe": 1.42},
|
||||
tags=["银行", "长线"],
|
||||
notes="回撤可控",
|
||||
)
|
||||
|
||||
|
||||
def _sample_portfolio(name: str = "组合·消费双雄") -> SavedStrategy:
|
||||
return SavedStrategy(
|
||||
id="",
|
||||
name=name,
|
||||
kind="portfolio",
|
||||
strategy="rsi_reversal",
|
||||
strategy_label="RSI 反转",
|
||||
params={"period": 14, "oversold": 30},
|
||||
context={"stocks": ["SH:600519", "SZ:000858"]},
|
||||
snapshot={"total_return": 0.18},
|
||||
)
|
||||
|
||||
|
||||
def test_add_assigns_id_and_timestamps(store: StrategyStore):
|
||||
rec = store.add(_sample_single())
|
||||
assert rec.id and len(rec.id) == 12
|
||||
assert rec.created_at
|
||||
assert rec.updated_at == rec.created_at
|
||||
|
||||
|
||||
def test_list_round_trip_preserves_all_fields(store: StrategyStore):
|
||||
original = store.add(_sample_single())
|
||||
items = store.list_all()
|
||||
assert len(items) == 1
|
||||
got = items[0]
|
||||
assert got.id == original.id
|
||||
assert got.name == "双均线·平安"
|
||||
assert got.kind == "single"
|
||||
assert got.params == {"fast": 5, "slow": 20}
|
||||
assert got.context["symbol"] == "SZ:000001"
|
||||
assert got.trade_config["cash"] == 1_000_000
|
||||
assert got.snapshot["total_return"] == pytest.approx(0.352)
|
||||
assert got.tags == ["银行", "长线"]
|
||||
assert got.notes == "回撤可控"
|
||||
|
||||
|
||||
def test_list_orders_by_created_desc(store: StrategyStore):
|
||||
a = store.add(_sample_single(name="first"))
|
||||
b = store.add(_sample_portfolio(name="second"))
|
||||
names = [x.name for x in store.list_all()]
|
||||
# 后加的在前
|
||||
assert names == ["second", "first"]
|
||||
assert {x.id for x in (a, b)} == {a.id, b.id}
|
||||
|
||||
|
||||
def test_get_returns_none_for_missing(store: StrategyStore):
|
||||
assert store.get("nope") is None
|
||||
|
||||
|
||||
def test_get_returns_record(store: StrategyStore):
|
||||
rec = store.add(_sample_portfolio())
|
||||
got = store.get(rec.id)
|
||||
assert got is not None
|
||||
assert got.kind == "portfolio"
|
||||
assert got.context["stocks"] == ["SH:600519", "SZ:000858"]
|
||||
|
||||
|
||||
def test_delete_removes_record(store: StrategyStore):
|
||||
rec = store.add(_sample_single())
|
||||
assert store.delete(rec.id) is True
|
||||
assert store.get(rec.id) is None
|
||||
assert store.list_all() == []
|
||||
|
||||
|
||||
def test_delete_missing_returns_false(store: StrategyStore):
|
||||
assert store.delete("nonexistent") is False
|
||||
|
||||
|
||||
def test_store_creates_db_file_and_schema(tmp_path):
|
||||
db_path = tmp_path / "nested" / "strategies.db"
|
||||
s = StrategyStore(db_path=db_path)
|
||||
assert db_path.exists()
|
||||
# schema 已建表 + 索引
|
||||
with sqlite3.connect(db_path) as conn:
|
||||
tables = {r[0] for r in conn.execute("SELECT name FROM sqlite_master WHERE type='table'")}
|
||||
indexes = {r[0] for r in conn.execute("SELECT name FROM sqlite_master WHERE type='index'")}
|
||||
assert "strategies" in tables
|
||||
assert {"idx_strategies_kind", "idx_strategies_strategy", "idx_strategies_created"} <= indexes
|
||||
# 可正常写入
|
||||
s.add(_sample_single())
|
||||
assert len(s.list_all()) == 1
|
||||
|
||||
|
||||
def test_json_fields_with_unicode(store: StrategyStore):
|
||||
"""中文标签/备注应无损往返(ensure_ascii=False 落库)。"""
|
||||
rec = store.add(
|
||||
SavedStrategy(
|
||||
id="",
|
||||
name="测试·中文🎉",
|
||||
kind="single",
|
||||
strategy="macd",
|
||||
notes="这是一段中文备注",
|
||||
tags=["标签一", "标签二"],
|
||||
)
|
||||
)
|
||||
got = store.get(rec.id)
|
||||
assert got is not None
|
||||
assert got.name == "测试·中文🎉"
|
||||
assert got.notes == "这是一段中文备注"
|
||||
assert got.tags == ["标签一", "标签二"]
|
||||
|
||||
|
||||
# ── 路由端到端测试(TestClient)──────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def client(tmp_path, monkeypatch) -> TestClient:
|
||||
"""构造一个用临时 SQLite 文件的独立 app + store 单例。"""
|
||||
# 用 monkeypatch 替换 get_store 返回的路径,保证测试隔离
|
||||
from easy_tdx.web import strategy_store as mod
|
||||
|
||||
test_store = StrategyStore(db_path=tmp_path / "router_strategies.db")
|
||||
# 替换单例,避免污染全局
|
||||
monkeypatch.setattr(mod, "_store", test_store)
|
||||
|
||||
from easy_tdx.web.routers.strategies import router as strategies_router
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(strategies_router, prefix="/api/v1")
|
||||
# 复用项目的 ValueError → 400 处理
|
||||
from easy_tdx.web.errors import register_exception_handlers
|
||||
|
||||
register_exception_handlers(app)
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def _create_payload(kind: str = "single", **over) -> dict:
|
||||
base = {
|
||||
"name": "我的策略",
|
||||
"kind": kind,
|
||||
"strategy": "ma_cross",
|
||||
"strategy_label": "双均线交叉",
|
||||
"params": {"fast": 5, "slow": 20},
|
||||
"context": {"symbol": "SZ:000001"},
|
||||
"trade_config": {"cash": 1000000},
|
||||
"snapshot": {"total_return": 0.35, "sharpe": 1.4},
|
||||
"tags": ["银行"],
|
||||
"notes": "观察中",
|
||||
}
|
||||
base.update(over)
|
||||
return base
|
||||
|
||||
|
||||
def test_router_create_then_list_get_delete(client: TestClient):
|
||||
# 1. 创建
|
||||
resp = client.post("/api/v1/strategies", json=_create_payload())
|
||||
assert resp.status_code == 201
|
||||
created = resp.json()
|
||||
assert created["id"]
|
||||
assert created["name"] == "我的策略"
|
||||
assert created["params"] == {"fast": 5, "slow": 20}
|
||||
assert created["created_at"]
|
||||
sid = created["id"]
|
||||
|
||||
# 2. 列表
|
||||
resp = client.get("/api/v1/strategies")
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert body["count"] == 1
|
||||
assert body["strategies"][0]["id"] == sid
|
||||
|
||||
# 3. 详情
|
||||
resp = client.get(f"/api/v1/strategies/{sid}")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["snapshot"]["total_return"] == pytest.approx(0.35)
|
||||
|
||||
# 4. 删除
|
||||
resp = client.delete(f"/api/v1/strategies/{sid}")
|
||||
assert resp.status_code == 204
|
||||
|
||||
# 5. 列表为空
|
||||
assert client.get("/api/v1/strategies").json()["count"] == 0
|
||||
|
||||
|
||||
def test_router_get_missing_returns_400(client: TestClient):
|
||||
# 不存在的 id → ValueError → 400(项目错误处理约定)
|
||||
resp = client.get("/api/v1/strategies/nonexistent")
|
||||
assert resp.status_code == 400
|
||||
|
||||
|
||||
def test_router_delete_missing_returns_400(client: TestClient):
|
||||
resp = client.delete("/api/v1/strategies/nonexistent")
|
||||
assert resp.status_code == 400
|
||||
|
||||
|
||||
def test_router_rejects_empty_name(client: TestClient):
|
||||
resp = client.post("/api/v1/strategies", json=_create_payload(name=""))
|
||||
assert resp.status_code == 422 # Pydantic 校验失败
|
||||
|
||||
|
||||
def test_router_rejects_invalid_kind(client: TestClient):
|
||||
resp = client.post("/api/v1/strategies", json=_create_payload(kind="bogus"))
|
||||
assert resp.status_code == 422
|
||||
|
||||
|
||||
def test_router_accepts_portfolio_kind(client: TestClient):
|
||||
payload = _create_payload(
|
||||
kind="portfolio",
|
||||
strategy="rsi_reversal",
|
||||
context={"stocks": ["SH:600519", "SZ:000858"]},
|
||||
)
|
||||
resp = client.post("/api/v1/strategies", json=payload)
|
||||
assert resp.status_code == 201
|
||||
body = resp.json()
|
||||
assert body["kind"] == "portfolio"
|
||||
assert body["context"]["stocks"] == ["SH:600519", "SZ:000858"]
|
||||
Reference in New Issue
Block a user