mirror of
https://ghfast.top/https://github.com/aeroxw/easy_tdx_max.git
synced 2026-09-12 14:34:18 +08:00
release: v1.24.0 — QFQ 对拍验证体系 + 回测任务持久化 + 品种感知费率
升级计划 P0(docs/upgrade-plan-2026H2.md)。源自 backtest-system / indicator-lab 两个下游项目的逆向调研。
- QFQ 对拍验证:公式法(NONE+XDXR)与跳空检测法(板块感知涨跌停阈值)双证据链互检,
检出负价/残留跳空/方向反演/XDXR 缺记录四类问题,接入 MAC 同步/异步客户端(mac/qfq_check.py);
含茅台式多重分红、浦发式送转方向合成案例回归(13 用例)
- 回测任务 SQLite 持久化:~/.easy_tdx/tasks.db 双写内存 LRU + 磁盘(保留 500 条),serve 重启不丢;
重启恢复中断任务标记;GET /backtest/tasks/{id}/export?format=json|csv 导出端点
- 品种感知费率:ETF/可转债免印花税等法定差异(backtest/fees.py),CLI --auto-fees、
REST auto_fees 字段、组合引擎逐标的解析(34 用例)
- 修正 avg_holding_days 过时注释(实现早已是 FIFO 真实口径)
- tests/conftest.py 默认 EASY_TDX_NO_TASK_DB=1 防止单测污染用户任务库
- 注:engine/cli/routers/schemas 为跨版本累积态,后续版本提交继续演进
This commit is contained in:
@@ -2,6 +2,26 @@
|
||||
|
||||
本文件记录 easy-tdx 的版本变更。格式遵循 [Keep a Changelog](https://keepachangelog.com/zh-CN/)。
|
||||
|
||||
## [1.24.0] — 2026-09-01
|
||||
|
||||
**信任与持久化版本**——修复下游反馈的 QFQ 复权可信度问题(引入双引擎对拍验证)、回测任务落盘 SQLite(重启不丢)、品种感知费率(ETF/可转债免印花税)。源自对两个下游项目(backtest-system / indicator-lab)的逆向调研,完整升级计划见 `docs/upgrade-plan-2026H2.md`。
|
||||
|
||||
### 新增
|
||||
|
||||
- **QFQ 对拍验证体系**(`mac/qfq_check.py`)——公式法(NONE+XDXR)与跳空检测法(板块感知涨跌停阈值:主板 10%/双创 20%/北交所 30% + 0.5% 余量)双证据链交叉验证前复权结果,检出四类问题:`bad_price`(非法价格)、`residual_gap`(除权日仍残留跳空,疑似漏算/未生效)、`wrong_direction`(残差方向反,疑似复权过度/方向算反)、`unexplained_gap`(NONE 跳空但 XDXR 无对应记录)。已接入 `MacClient` / `AsyncMacClient` 的 QFQ 本地重算路径:不一致即打告警日志,最近一次报告存于 `client.last_qfq_crosscheck`。含「茅台式多重分红」「浦发式送转股方向」合成案例回归测试(13 个用例)。回应下游 backtest-system 对 QFQ 可靠性的反馈。
|
||||
- **回测任务 SQLite 持久化**(`web/task_store.py`)——任务状态/结果双写内存 LRU + `~/.easy_tdx/tasks.db`(随 `EASY_TDX_CONFIG_DIR`,保留 500 条),serve 重启后对比页历史任务、已完成寻优排名均可继续查询;重启时遗留的 pending/running 任务自动标记为 failed(注明「服务重启中断」)。`EASY_TDX_NO_TASK_DB=1` 可关闭(测试默认关闭)。
|
||||
- **任务结果导出端点**——`GET /backtest/tasks/{task_id}/export?format=json|csv`:JSON 导出完整 result;CSV 智能挑主表(trades → ranking → equity_curve,兜底 performance 键值对),带 `Content-Disposition` 附件头。
|
||||
- **品种感知费率**(`backtest/fees.py`)——按代码前缀+市场推断品种(股票/ETF/LOF/可转债/B股/指数),自动解析佣金/最低佣金/印花税;核心法定差异:**ETF/可转债免印花税**(此前扁平默认对 ETF 轮动类策略长期错收印花税)。接入:`BacktestEngine(symbol=..., auto_fees=True)`、`PortfolioBacktestEngine(auto_fees=True)`(逐标的解析)、CLI `easy-tdx backtest --auto-fees`、REST 请求体 `auto_fees` 字段。显式非默认费率仍优先;结果 config 快照记录 symbol 与解析后费率。34 个测试用例。
|
||||
|
||||
### 修复
|
||||
|
||||
- `performance.py` 中 `avg_holding_days` 的过时文档注释(实现早已是 FIFO 配对、按 size 加权的真实日历日口径,注释仍写「简化为固定值 5.0」,误导审计)。
|
||||
|
||||
### 内部
|
||||
|
||||
- `tests/conftest.py` 全局默认 `EASY_TDX_NO_TASK_DB=1`,防止单测污染用户真实 `~/.easy_tdx/tasks.db`。
|
||||
- `task_store` 初始化用独立 `_init_lock`(避免与写锁死锁);`task_runner` 的 pending 落盘先于 executor.submit(避免旧状态覆盖新状态的竞态)。
|
||||
|
||||
## [1.23.3] — 2026-09-01
|
||||
|
||||
**serve 纯 API 模式 + 看板修复**。自 1.23.2 以来的增量:
|
||||
|
||||
+1
-1
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "easy-tdx"
|
||||
version = "1.23.3"
|
||||
version = "1.24.0"
|
||||
description = "通达信 TCP 协议行情数据客户端,支持在线行情、离线数据读取与写入同步"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
|
||||
@@ -29,6 +29,12 @@ import click
|
||||
)
|
||||
@click.option("--cash", default=100000.0, type=float, help="初始资金")
|
||||
@click.option("--commission", default=0.0003, type=float, help="佣金率")
|
||||
@click.option(
|
||||
"--auto-fees",
|
||||
"auto_fees",
|
||||
is_flag=True,
|
||||
help="按标的品种自动解析费率(ETF/可转债免印花税等;显式 --commission 优先)",
|
||||
)
|
||||
@click.option(
|
||||
"--execution",
|
||||
default="next_open",
|
||||
@@ -47,6 +53,19 @@ import click
|
||||
)
|
||||
@click.option("--table", "use_table", is_flag=True, help="表格输出")
|
||||
@click.option("--output", "output_fmt", type=click.Choice(["json", "table", "csv"]), default="json")
|
||||
@click.option(
|
||||
"--wf",
|
||||
"walk_forward",
|
||||
is_flag=True,
|
||||
help="附加 Walk-Forward 样本外验证(默认 7 窗,每窗独立开仓)",
|
||||
)
|
||||
@click.option("--wf-windows", "wf_windows", default=7, type=int, help="Walk-Forward 窗口数")
|
||||
@click.option(
|
||||
"--evaluate",
|
||||
"full_evaluate",
|
||||
is_flag=True,
|
||||
help="一条龙评估:回测+WF+适配性+综合评分+S-D评级+买入持有基准对比(覆盖常规输出)",
|
||||
)
|
||||
def backtest(
|
||||
market: str,
|
||||
code: str,
|
||||
@@ -56,6 +75,7 @@ def backtest(
|
||||
combo_mode: str,
|
||||
cash: float,
|
||||
commission: float,
|
||||
auto_fees: bool,
|
||||
execution: str,
|
||||
period: str,
|
||||
adjust: str,
|
||||
@@ -64,6 +84,9 @@ def backtest(
|
||||
chanlun_level: str | None,
|
||||
use_table: bool,
|
||||
output_fmt: str,
|
||||
walk_forward: bool,
|
||||
wf_windows: int,
|
||||
full_evaluate: bool,
|
||||
) -> None:
|
||||
"""回测引擎:执行策略并返回绩效报告。
|
||||
|
||||
@@ -132,12 +155,34 @@ def backtest(
|
||||
)
|
||||
else:
|
||||
assert strategy_cls is not None # guarded above by SystemExit
|
||||
|
||||
# 一条龙评估:覆盖常规输出(含回测本身,无需重复跑)
|
||||
if full_evaluate:
|
||||
import json as _json
|
||||
|
||||
from ..backtest.benchmark import evaluate_strategy
|
||||
|
||||
report = evaluate_strategy(
|
||||
strategy=strategy_cls,
|
||||
df=df,
|
||||
cash=cash,
|
||||
commission=commission,
|
||||
execution=execution,
|
||||
symbol=f"{market}:{code}",
|
||||
auto_fees=auto_fees,
|
||||
n_windows=wf_windows,
|
||||
)
|
||||
click.echo(_json.dumps(report, ensure_ascii=False, default=str))
|
||||
return
|
||||
|
||||
engine = BacktestEngine(
|
||||
strategy=strategy_cls,
|
||||
cash=cash,
|
||||
commission=commission,
|
||||
execution=execution,
|
||||
chanlun_level=chanlun_level,
|
||||
symbol=f"{market}:{code}",
|
||||
auto_fees=auto_fees,
|
||||
)
|
||||
result = engine.run(df)
|
||||
|
||||
@@ -150,6 +195,25 @@ def backtest(
|
||||
else:
|
||||
click.echo(result.to_json())
|
||||
|
||||
# 6. 附加 Walk-Forward 样本外验证(--wf)
|
||||
if walk_forward and not is_combo:
|
||||
import json as _json
|
||||
|
||||
from ..backtest.walkforward import WalkForwardEngine
|
||||
|
||||
assert strategy_cls is not None
|
||||
wf = WalkForwardEngine(
|
||||
strategy=strategy_cls,
|
||||
n_windows=wf_windows,
|
||||
cash=cash,
|
||||
commission=commission,
|
||||
execution=execution,
|
||||
symbol=f"{market}:{code}",
|
||||
auto_fees=auto_fees,
|
||||
)
|
||||
wf_report = {"walkforward": wf.run(df).to_dict()}
|
||||
click.echo(_json.dumps(wf_report, ensure_ascii=False, default=str))
|
||||
|
||||
|
||||
def _load_strategy(strategy_str: str | None, strategy_file: str | None) -> type | None:
|
||||
"""加载策略类。
|
||||
|
||||
@@ -64,6 +64,9 @@ class BacktestEngine:
|
||||
slippage_model: SlippageModel | None = None,
|
||||
execution_model: ExecutionModel | None = None,
|
||||
warmup_bars: int = 0,
|
||||
symbol: str | None = None,
|
||||
auto_fees: bool = False,
|
||||
indicator_cache: Any | None = None,
|
||||
):
|
||||
"""Initialize engine.
|
||||
|
||||
@@ -87,10 +90,33 @@ class BacktestEngine:
|
||||
warmup_bars: 指标预热 bar 数。前 ``warmup_bars`` 根不调用
|
||||
``next()``、不产生信号(指标 NaN 期过滤),避免早期数据不足
|
||||
导致的越界或误信号。默认 0(向后兼容)。
|
||||
symbol: 标的代码(如 ``"SH:510300"``,供品种感知费率解析)。
|
||||
auto_fees: 为 True 且提供 symbol 时,按品种自动解析
|
||||
commission/min_commission/stamp_tax(ETF/债券免印花税等),
|
||||
覆盖默认值;显式传入的非默认费率仍优先于自动解析。
|
||||
indicator_cache: 指标计算缓存(网格寻优跨点复用,见
|
||||
:class:`~easy_tdx.backtest.indicator_cache.IndicatorCache`)。
|
||||
None = 每次直接计算(默认,向后兼容)。
|
||||
|
||||
.. versionchanged:: 1.24
|
||||
新增 ``symbol`` / ``auto_fees`` 品种感知费率参数。
|
||||
"""
|
||||
self._strategy_cls = strategy if isinstance(strategy, type) else type(strategy)
|
||||
self._strategy_instance = strategy if isinstance(strategy, Strategy) else None
|
||||
|
||||
self._symbol = symbol
|
||||
if auto_fees and symbol:
|
||||
from easy_tdx.backtest.fees import resolve_fee_model
|
||||
|
||||
fee = resolve_fee_model(symbol)
|
||||
# 显式非默认值优先(调用方有意覆盖),否则用品种默认
|
||||
if commission == 0.0003:
|
||||
commission = fee.commission
|
||||
if min_commission == 5.0:
|
||||
min_commission = fee.min_commission
|
||||
if stamp_tax == 0.001:
|
||||
stamp_tax = fee.stamp_tax
|
||||
|
||||
self._cash = cash
|
||||
self._commission = commission
|
||||
self._min_commission = min_commission
|
||||
@@ -104,6 +130,7 @@ class BacktestEngine:
|
||||
self._slippage_model = slippage_model
|
||||
self._execution_model = execution_model
|
||||
self._warmup_bars = max(int(warmup_bars), 0)
|
||||
self._indicator_cache = indicator_cache
|
||||
|
||||
def run(self, df: pd.DataFrame, chanlun_result: Any | None = None) -> BacktestResult:
|
||||
"""Run backtest.
|
||||
@@ -174,13 +201,17 @@ class BacktestEngine:
|
||||
performance = analyzer.compute()
|
||||
|
||||
# Config snapshot
|
||||
config = {
|
||||
config: dict[str, Any] = {
|
||||
"cash": self._cash,
|
||||
"commission": self._commission,
|
||||
"min_commission": self._min_commission,
|
||||
"stamp_tax": self._stamp_tax,
|
||||
"execution": self._execution,
|
||||
"position_mode": self._position_mode,
|
||||
"reject_policy": self._reject_policy,
|
||||
}
|
||||
if self._symbol:
|
||||
config["symbol"] = self._symbol
|
||||
|
||||
return BacktestResult(
|
||||
performance=performance,
|
||||
@@ -264,6 +295,8 @@ class BacktestEngine:
|
||||
|
||||
# Bind data
|
||||
strat._bind_data(df)
|
||||
# 挂载指标缓存(两段式寻优加速;None 时 I() 直接计算)
|
||||
strat._indicator_cache = self._indicator_cache
|
||||
|
||||
# Inject chanlun result if provided
|
||||
if chanlun_result is not None:
|
||||
|
||||
@@ -0,0 +1,182 @@
|
||||
"""品种感知费率模型(v1.24 新增)。
|
||||
|
||||
此前引擎的佣金/最低佣金/印花税是全局扁平参数,默认值按 A 股股票设定
|
||||
(佣金万3 最低5元、印花税千1 卖方)。同一组默认值套到 ETF / 可转债 / B 股
|
||||
上会**错收印花税**(ETF 与债券法定免印花税)、错用最低佣金,长期低估
|
||||
高频/小资金策略的相对收益——对 ETF 轮动类策略影响尤其大。
|
||||
|
||||
本模块按证券代码 + 市场推断品种类型,给出保守的默认费率组合:
|
||||
|
||||
======== ============ ============ ============
|
||||
品种 佣金率 最低佣金 印花税(卖方)
|
||||
======== ============ ============ ============
|
||||
股票 0.0003 5.0 0.001
|
||||
ETF/LOF 0.0003 5.0 **0**
|
||||
可转债 0.0002 1.0 **0**
|
||||
B 股 0.0005 5.0 0.001
|
||||
其他/指数 0.0003 5.0 0
|
||||
======== ============ ============ ============
|
||||
|
||||
说明:
|
||||
- 印花税差异是法定事实(ETF/债券免征),是本模块的核心价值;
|
||||
- 佣金取常见券商的**保守**默认(不高估收益)。部分券商对 ETF「免五」、
|
||||
费率更低——用户仍可显式传 ``commission`` / ``min_commission`` 覆盖;
|
||||
- 北交所股票按股票口径处理(经手费差异不建模)。
|
||||
|
||||
用法(引擎侧)::
|
||||
|
||||
BacktestEngine(..., symbol="SH:510300", auto_fees=True)
|
||||
# → commission/stamp_tax 按 ETF 口径自动解析(显式传入的值优先)
|
||||
|
||||
市场代码兼容两种表示:``"SH"/"SZ"/"BJ"`` 字符串或通达信 int(0=深 1=沪 2=北)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
|
||||
__all__ = [
|
||||
"FeeModel",
|
||||
"InstrumentKind",
|
||||
"detect_instrument_kind",
|
||||
"resolve_fee_model",
|
||||
]
|
||||
|
||||
|
||||
class InstrumentKind(Enum):
|
||||
"""证券品种类型(按代码前缀推断)。"""
|
||||
|
||||
STOCK = "stock" # A 股股票(含北交所)
|
||||
ETF = "etf" # 场内交易型开放式指数基金
|
||||
LOF = "lof" # 上市开放式基金 / 封闭基金
|
||||
BOND = "bond" # 可转债
|
||||
B_SHARE = "b_share" # B 股(沪 B 美元 / 深 B 港币)
|
||||
INDEX = "index" # 指数(不可交易,按无印花税处理)
|
||||
OTHER = "other"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class FeeModel:
|
||||
"""一组费率参数(与引擎的 commission/min_commission/stamp_tax 一一对应)。"""
|
||||
|
||||
kind: InstrumentKind
|
||||
commission: float
|
||||
min_commission: float
|
||||
stamp_tax: float
|
||||
|
||||
|
||||
# 保守默认费率表(详见模块 docstring 的说明)
|
||||
_FEE_TABLE: dict[InstrumentKind, tuple[float, float, float]] = {
|
||||
InstrumentKind.STOCK: (0.0003, 5.0, 0.001),
|
||||
InstrumentKind.ETF: (0.0003, 5.0, 0.0),
|
||||
InstrumentKind.LOF: (0.0003, 5.0, 0.0),
|
||||
InstrumentKind.BOND: (0.0002, 1.0, 0.0),
|
||||
InstrumentKind.B_SHARE: (0.0005, 5.0, 0.001),
|
||||
InstrumentKind.INDEX: (0.0003, 5.0, 0.0),
|
||||
InstrumentKind.OTHER: (0.0003, 5.0, 0.0),
|
||||
}
|
||||
|
||||
|
||||
def _norm_market(market: str | int | None) -> str | None:
|
||||
"""市场代码归一化为 'SH'/'SZ'/'BJ'(未知返回 None)。"""
|
||||
if market is None:
|
||||
return None
|
||||
if isinstance(market, int):
|
||||
return {0: "SZ", 1: "SH", 2: "BJ"}.get(market)
|
||||
m = str(market).strip().upper()
|
||||
if m in ("SH", "SZ", "BJ", "SSE", "SZSE"):
|
||||
return "SH" if m in ("SH", "SSE") else ("SZ" if m in ("SZ", "SZSE") else "BJ")
|
||||
return None
|
||||
|
||||
|
||||
def detect_instrument_kind(
|
||||
symbol_or_code: str,
|
||||
market: str | int | None = None,
|
||||
) -> InstrumentKind:
|
||||
"""按代码前缀(+市场)推断品种类型。
|
||||
|
||||
Args:
|
||||
symbol_or_code: 6 位代码(``"510300"``)或带市场前缀
|
||||
(``"SH:510300"`` / ``"sh510300"``)。
|
||||
market: 市场代码(可选;symbol 自带前缀时被覆盖)。
|
||||
|
||||
Returns:
|
||||
:class:`InstrumentKind`。无法识别时返回 OTHER(按无印花税保守处理)。
|
||||
"""
|
||||
code = str(symbol_or_code).strip()
|
||||
mkt: str | None
|
||||
if ":" in code:
|
||||
prefix, code = code.split(":", 1)
|
||||
mkt = _norm_market(prefix)
|
||||
elif len(code) > 6 and code[:2].upper() in ("SH", "SZ", "BJ"):
|
||||
mkt = code[:2].upper()
|
||||
code = code[2:]
|
||||
else:
|
||||
mkt = _norm_market(market)
|
||||
code = code.strip()
|
||||
|
||||
# 沪市:900 B股 / 60/68 股票 / 51 56 58 ETF / 50x LOF·封基 / 11x 可转债 / 000 880 指数
|
||||
if mkt == "SH":
|
||||
if code.startswith("900"):
|
||||
return InstrumentKind.B_SHARE
|
||||
if code.startswith(("60", "68")):
|
||||
return InstrumentKind.STOCK
|
||||
if code.startswith(("51", "56", "58")):
|
||||
return InstrumentKind.ETF
|
||||
if code.startswith("50"):
|
||||
return InstrumentKind.LOF
|
||||
if code.startswith("11"):
|
||||
return InstrumentKind.BOND
|
||||
if code.startswith(("000", "880", "881", "999")):
|
||||
return InstrumentKind.INDEX
|
||||
return InstrumentKind.OTHER
|
||||
# 深市:200 B股 / 00 30 股票 / 159 ETF / 16x LOF / 12x 可转债 / 399 指数
|
||||
if mkt == "SZ":
|
||||
if code.startswith("200"):
|
||||
return InstrumentKind.B_SHARE
|
||||
if code.startswith(("00", "30")):
|
||||
return InstrumentKind.STOCK
|
||||
if code.startswith("159"):
|
||||
return InstrumentKind.ETF
|
||||
if code.startswith("16"):
|
||||
return InstrumentKind.LOF
|
||||
if code.startswith("12"):
|
||||
return InstrumentKind.BOND
|
||||
if code.startswith("399"):
|
||||
return InstrumentKind.INDEX
|
||||
return InstrumentKind.OTHER
|
||||
# 北交所:全部按股票
|
||||
if mkt == "BJ":
|
||||
return InstrumentKind.STOCK
|
||||
# 市场未知:仅按代码粗判(沪深代码空间基本不重叠)
|
||||
if code.startswith(("51", "56", "58", "159")):
|
||||
return InstrumentKind.ETF
|
||||
if code.startswith(("900", "200")):
|
||||
return InstrumentKind.B_SHARE
|
||||
if code.startswith(("60", "68", "00", "30")):
|
||||
return InstrumentKind.STOCK
|
||||
return InstrumentKind.OTHER
|
||||
|
||||
|
||||
def resolve_fee_model(
|
||||
symbol_or_code: str,
|
||||
market: str | int | None = None,
|
||||
) -> FeeModel:
|
||||
"""按标的解析默认费率组合。
|
||||
|
||||
Args:
|
||||
symbol_or_code: 代码或带市场前缀的 symbol。
|
||||
market: 市场代码(可选)。
|
||||
|
||||
Returns:
|
||||
:class:`FeeModel`(含品种类型 + 佣金/最低佣金/印花税)。
|
||||
"""
|
||||
kind = detect_instrument_kind(symbol_or_code, market)
|
||||
commission, min_commission, stamp_tax = _FEE_TABLE[kind]
|
||||
return FeeModel(
|
||||
kind=kind,
|
||||
commission=commission,
|
||||
min_commission=min_commission,
|
||||
stamp_tax=stamp_tax,
|
||||
)
|
||||
@@ -73,7 +73,7 @@ class PerformanceAnalyzer:
|
||||
- avg_loss: 平均亏损
|
||||
- max_win: 最大盈利
|
||||
- max_loss: 最大亏损
|
||||
- avg_holding_days: 平均持仓天数(简化为固定值 5.0)
|
||||
- avg_holding_days: 平均持仓天数(FIFO 配对、按 size 加权,日历日口径)
|
||||
- volatility: 年化波动率
|
||||
"""
|
||||
# 边界检查
|
||||
|
||||
@@ -91,6 +91,7 @@ class PortfolioBacktestEngine:
|
||||
slippage: float = 0.0,
|
||||
execution: str = "next_open",
|
||||
chanlun_level: str | None = None,
|
||||
auto_fees: bool = False,
|
||||
) -> None:
|
||||
"""初始化组合回测引擎。
|
||||
|
||||
@@ -107,6 +108,11 @@ class PortfolioBacktestEngine:
|
||||
slippage: 滑点
|
||||
execution: 执行模式
|
||||
chanlun_level: 缠论级别(可选)
|
||||
auto_fees: 为 True 时按各标的代码解析品种费率(ETF/债券免
|
||||
印花税等),覆盖默认值;显式非默认费率仍优先。
|
||||
|
||||
.. versionadded:: 1.24
|
||||
``auto_fees`` 品种感知费率(按 StockData.market+code 逐标的解析)。
|
||||
"""
|
||||
self._strategy = strategy
|
||||
self._stocks = stocks
|
||||
@@ -118,6 +124,7 @@ class PortfolioBacktestEngine:
|
||||
self._slippage = slippage
|
||||
self._execution = execution
|
||||
self._chanlun_level = chanlun_level
|
||||
self._auto_fees = auto_fees
|
||||
|
||||
def _compute_allocations(self) -> dict[str, float]:
|
||||
"""计算每只标的的资金分配。"""
|
||||
@@ -158,6 +165,8 @@ class PortfolioBacktestEngine:
|
||||
slippage=self._slippage,
|
||||
execution=self._execution,
|
||||
chanlun_level=self._chanlun_level,
|
||||
symbol=key,
|
||||
auto_fees=self._auto_fees,
|
||||
)
|
||||
result = engine.run(stock.df)
|
||||
individual_results[key] = result
|
||||
|
||||
@@ -177,6 +177,9 @@ class MacClient:
|
||||
# XDXR(除权除息)记录缓存:(market, code) -> DataFrame。
|
||||
# 仅在服务端 QFQ 返回异常(负价)时用于本地前复权重算。
|
||||
self._xdxr_cache: dict[tuple[int, str], pd.DataFrame] = {}
|
||||
# 最近一次 QFQ 本地重算的对拍报告(qfq_check.crosscheck_qfq 产出),
|
||||
# 供调试/上游排查;None = 尚未触发过本地重算。
|
||||
self.last_qfq_crosscheck: dict[str, object] | None = None
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# 工厂方法
|
||||
@@ -479,11 +482,22 @@ class MacClient:
|
||||
重算后的 DataFrame;XDXR 取不到或重算仍异常时原样返回 df。
|
||||
"""
|
||||
from .adjust import apply_forward_adjust, has_bad_prices
|
||||
from .qfq_check import crosscheck_qfq
|
||||
|
||||
xd = self._fetch_xdxr_records(market, code)
|
||||
if xd is None:
|
||||
return df
|
||||
out = apply_forward_adjust(df, xd)
|
||||
# 对拍校验:公式法结果 vs 跳空检测法独立证据链,不一致即告警
|
||||
report = crosscheck_qfq(df, out, xd, code, market)
|
||||
self.last_qfq_crosscheck = report.to_dict()
|
||||
if not report.ok:
|
||||
_logger.warning(
|
||||
"QFQ 对拍校验发现 %d 个问题(%s):%s",
|
||||
len(report.issues),
|
||||
report.symbol,
|
||||
"; ".join(f"{i.kind}@{i.date}" for i in report.issues[:5]),
|
||||
)
|
||||
if has_bad_prices(out):
|
||||
_logger.warning("QFQ 本地重算后 %s %s 仍含非法价格,降级返回服务端 QFQ", market, code)
|
||||
return df
|
||||
@@ -1239,6 +1253,8 @@ class AsyncMacClient(AsyncHeartbeatMixin):
|
||||
# XDXR(除权除息)记录缓存:(market, code) -> DataFrame。
|
||||
# 仅在服务端 QFQ 返回异常(负价)时用于本地前复权重算。
|
||||
self._xdxr_cache: dict[tuple[int, str], pd.DataFrame] = {}
|
||||
# 最近一次 QFQ 本地重算的对拍报告(同 MacClient.last_qfq_crosscheck)
|
||||
self.last_qfq_crosscheck: dict[str, object] | None = None
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# 工厂方法
|
||||
@@ -1477,11 +1493,22 @@ class AsyncMacClient(AsyncHeartbeatMixin):
|
||||
) -> pd.DataFrame:
|
||||
"""对 QFQ 异常的 K 线用 NONE+XDXR 本地重算前复权(同 MacClient)。"""
|
||||
from .adjust import apply_forward_adjust, has_bad_prices
|
||||
from .qfq_check import crosscheck_qfq
|
||||
|
||||
xd = self._fetch_xdxr_records(market, code)
|
||||
if xd is None:
|
||||
return df
|
||||
out = apply_forward_adjust(df, xd)
|
||||
# 对拍校验:公式法结果 vs 跳空检测法独立证据链,不一致即告警
|
||||
report = crosscheck_qfq(df, out, xd, code, market)
|
||||
self.last_qfq_crosscheck = report.to_dict()
|
||||
if not report.ok:
|
||||
_logger.warning(
|
||||
"QFQ 对拍校验发现 %d 个问题(%s):%s",
|
||||
len(report.issues),
|
||||
report.symbol,
|
||||
"; ".join(f"{i.kind}@{i.date}" for i in report.issues[:5]),
|
||||
)
|
||||
if has_bad_prices(out):
|
||||
_logger.warning("QFQ 本地重算后 %s %s 仍含非法价格,降级返回服务端 QFQ", market, code)
|
||||
return df
|
||||
|
||||
@@ -0,0 +1,312 @@
|
||||
"""QFQ 前复权对拍校验(公式法 vs 跳空检测法)。
|
||||
|
||||
背景:服务端 QFQ 对长期重度除权股票的深层历史会返回负价,客户端用
|
||||
``adjust.apply_forward_adjust``(NONE + XDXR 公式法)本地重算兜底。公式法
|
||||
本身依赖 XDXR 数据完整、字段方向正确——单靠它无法自证可靠(下游项目曾反馈
|
||||
「茅台负价、浦发除权方向算反」类问题)。
|
||||
|
||||
本模块引入第二条独立证据链做交叉验证(借鉴 backtest-system 的思路):
|
||||
|
||||
1. **跳空检测法**:在 NONE 未复权序列上,真实 A 股单日跌幅受涨跌停约束
|
||||
(主板 10% / 双创 20% / 北交所 30%)。若某日 ``open`` 相对前一日
|
||||
``close`` 的跌幅超出「跌停幅度 + 余量」,几乎必然是除权造成的假跳空,
|
||||
与 XDXR 事件应一一对应。
|
||||
2. **残差跳空检测**:正确的前复权序列在除权日附近应当价格连续。若复权
|
||||
结果在除权日仍残留超出涨跌停约束的跳空,说明复权未生效(漏算事件、
|
||||
方向算反或因子错误)。
|
||||
|
||||
二者交叉得到四类问题:``bad_price``(非法价格)、``residual_gap``
|
||||
(复权后除权日仍跳空)、``unexplained_gap``(NONE 跳空但 XDXR 无对应
|
||||
事件,疑似 XDXR 缺记录)、``wrong_direction``(残差跳空方向与除权相反,
|
||||
疑似复权方向反了)。
|
||||
|
||||
该模块只做检测与报告,不修改数据;调用方(MAC 客户端)据此打日志告警。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
__all__ = [
|
||||
"QfqIssue",
|
||||
"QfqCrosscheckReport",
|
||||
"board_limit_ratio",
|
||||
"detect_ex_dividend_gaps",
|
||||
"crosscheck_qfq",
|
||||
]
|
||||
|
||||
# 跳空判定余量:真实跌停恰好等于限价(如 -10.0%),除权跳空通常更深;
|
||||
# 余量用于区分「刚好跌停」与「超出跌停的除权跳空」。
|
||||
_GAP_MARGIN = 0.005
|
||||
|
||||
|
||||
@dataclass
|
||||
class QfqIssue:
|
||||
"""单条对拍问题。"""
|
||||
|
||||
kind: str # bad_price | residual_gap | unexplained_gap | wrong_direction
|
||||
date: str # YYYY-MM-DD
|
||||
detail: str
|
||||
|
||||
def to_dict(self) -> dict[str, str]:
|
||||
return {"kind": self.kind, "date": self.date, "detail": self.detail}
|
||||
|
||||
|
||||
@dataclass
|
||||
class QfqCrosscheckReport:
|
||||
"""对拍报告。``ok`` 为 True 表示未发现三类结构性问题。"""
|
||||
|
||||
symbol: str
|
||||
ok: bool
|
||||
events_checked: int = 0
|
||||
gaps_detected: int = 0
|
||||
issues: list[QfqIssue] = field(default_factory=list)
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"symbol": self.symbol,
|
||||
"ok": self.ok,
|
||||
"events_checked": self.events_checked,
|
||||
"gaps_detected": self.gaps_detected,
|
||||
"issues": [i.to_dict() for i in self.issues],
|
||||
}
|
||||
|
||||
|
||||
def board_limit_ratio(code: str, market: int | None = None) -> float:
|
||||
"""按代码/市场推断涨跌停幅度(返回比例,如 0.10)。
|
||||
|
||||
规则(不含 ST 的 5% 特例——仅凭代码无法识别 ST,且 ST 跌停在 10%
|
||||
阈值之内不会造成误报):
|
||||
|
||||
- 北交所(market==2 或代码 4/8/92 开头):30%
|
||||
- 创业板(30 开头)、科创板(68 开头):20%
|
||||
- 其余(沪深主板):10%
|
||||
|
||||
Args:
|
||||
code: 6 位证券代码。
|
||||
market: 通达信市场代码(0=深 1=沪 2=北交所),可选。
|
||||
|
||||
Returns:
|
||||
涨跌停幅度比例。
|
||||
"""
|
||||
c = str(code).strip()
|
||||
if market == 2 or c.startswith(("4", "8", "92")):
|
||||
return 0.30
|
||||
if c.startswith(("30", "68")):
|
||||
return 0.20
|
||||
return 0.10
|
||||
|
||||
|
||||
def _to_dt_index(df: pd.DataFrame) -> pd.Series:
|
||||
"""把 datetime 列统一为 pd.Timestamp 序列。"""
|
||||
return pd.to_datetime(df["datetime"])
|
||||
|
||||
|
||||
def _fmt(dt: Any) -> str:
|
||||
"""把日期标量格式化为 YYYY-MM-DD。"""
|
||||
return str(pd.Timestamp(dt).strftime("%Y-%m-%d"))
|
||||
|
||||
|
||||
def detect_ex_dividend_gaps(
|
||||
none_df: pd.DataFrame,
|
||||
code: str,
|
||||
market: int | None = None,
|
||||
margin: float = _GAP_MARGIN,
|
||||
) -> list[str]:
|
||||
"""在 NONE 未复权序列上检测疑似除权跳空日。
|
||||
|
||||
判定:``open[i] / close[i-1] - 1 < -(limit + margin)``。真实交易中
|
||||
单日开盘相对昨收不可能超出跌停幅度(跌停开盘恰好等于 -limit,被余量
|
||||
排除),超出即认定为除权造成的价格序列断裂。
|
||||
|
||||
Args:
|
||||
none_df: NONE 未复权 K 线(datetime/open/close 列,升序或乱序均可)。
|
||||
code: 证券代码(用于推断涨跌停幅度)。
|
||||
market: 通达信市场代码(可选)。
|
||||
margin: 跳空判定余量(默认 0.5%)。
|
||||
|
||||
Returns:
|
||||
疑似除权日列表(YYYY-MM-DD,按时间升序)。
|
||||
"""
|
||||
if none_df is None or len(none_df) < 2 or "open" not in none_df.columns:
|
||||
return []
|
||||
df = none_df.copy()
|
||||
df["_dt"] = _to_dt_index(df)
|
||||
df = df.sort_values("_dt").reset_index(drop=True)
|
||||
|
||||
limit = board_limit_ratio(code, market)
|
||||
threshold = -(limit + margin)
|
||||
prev_close = df["close"].to_numpy(dtype=float)
|
||||
open_arr = df["open"].to_numpy(dtype=float)
|
||||
with np.errstate(divide="ignore", invalid="ignore"):
|
||||
ratio = open_arr[1:] / prev_close[:-1] - 1.0
|
||||
out: list[str] = []
|
||||
for i in np.where(~np.isfinite(ratio) | (ratio < threshold))[0]:
|
||||
out.append(_fmt(df["_dt"].iloc[i + 1]))
|
||||
return out
|
||||
|
||||
|
||||
def _xdxr_event_dates(xdxr_df: pd.DataFrame | None) -> list[pd.Timestamp]:
|
||||
"""提取 XDXR 中 category==1(除权除息)事件的日期(升序去重)。"""
|
||||
if xdxr_df is None or xdxr_df.empty:
|
||||
return []
|
||||
if "category" not in xdxr_df.columns or "date" not in xdxr_df.columns:
|
||||
return []
|
||||
cat1 = xdxr_df[xdxr_df["category"] == 1]
|
||||
dates: list[pd.Timestamp] = []
|
||||
for v in cat1["date"]:
|
||||
try:
|
||||
dates.append(pd.Timestamp(str(v)))
|
||||
except (ValueError, TypeError):
|
||||
continue
|
||||
return sorted(set(dates))
|
||||
|
||||
|
||||
def crosscheck_qfq(
|
||||
none_df: pd.DataFrame,
|
||||
qfq_df: pd.DataFrame,
|
||||
xdxr_df: pd.DataFrame | None,
|
||||
code: str,
|
||||
market: int | None = None,
|
||||
margin: float = _GAP_MARGIN,
|
||||
) -> QfqCrosscheckReport:
|
||||
"""对拍校验前复权结果。
|
||||
|
||||
检查项:
|
||||
|
||||
1. ``bad_price``:qfq 序列含 <=0 / NaN / inf 的 OHLC;
|
||||
2. ``residual_gap`` / ``wrong_direction``:qfq 序列在 XDXR 除权日仍残留
|
||||
超出涨跌停约束的跳空(残差向下 = 复权不足或漏事件;残差向上 =
|
||||
复权过度或方向反了);
|
||||
3. ``unexplained_gap``:NONE 序列存在除权跳空,但 ±2 个交易日内无
|
||||
XDXR 事件对应(疑似 XDXR 缺记录——公式法此时会漏调该事件)。
|
||||
|
||||
Args:
|
||||
none_df: NONE 未复权 K 线(datetime/open/close)。
|
||||
qfq_df: 待校验的前复权结果(datetime/open/high/low/close)。
|
||||
xdxr_df: ``get_xdxr_info`` 返回的除权除息记录(可为 None)。
|
||||
code: 证券代码。
|
||||
market: 通达信市场代码(可选)。
|
||||
margin: 跳空判定余量。
|
||||
|
||||
Returns:
|
||||
:class:`QfqCrosscheckReport`。``ok=False`` 时 ``issues`` 非空。
|
||||
"""
|
||||
symbol = f"{market}:{code}" if market is not None else str(code)
|
||||
issues: list[QfqIssue] = []
|
||||
limit = board_limit_ratio(code, market)
|
||||
threshold = limit + margin
|
||||
|
||||
# ── 1. 非法价格 ───────────────────────────────────────────────────────────
|
||||
if qfq_df is not None and len(qfq_df) > 0:
|
||||
for col in ("open", "high", "low", "close"):
|
||||
if col not in qfq_df.columns:
|
||||
continue
|
||||
arr = qfq_df[col].to_numpy(dtype=float)
|
||||
bad = ~np.isfinite(arr) | (arr <= 0)
|
||||
for i in np.where(bad)[0]:
|
||||
issues.append(
|
||||
QfqIssue(
|
||||
kind="bad_price",
|
||||
date=_fmt(_to_dt_index(qfq_df).iloc[i]),
|
||||
detail=f"{col} 列第 {int(i)} 根价格非法({arr[i]})",
|
||||
)
|
||||
)
|
||||
break # 每列报首个即可,避免深层历史批量刷屏
|
||||
|
||||
# ── 2. 除权日残差跳空 ─────────────────────────────────────────────────────
|
||||
events_checked = 0
|
||||
if qfq_df is not None and len(qfq_df) >= 2:
|
||||
q = qfq_df.copy()
|
||||
q["_dt"] = _to_dt_index(q)
|
||||
q = q.sort_values("_dt").reset_index(drop=True)
|
||||
qd = q["_dt"].to_numpy()
|
||||
open_arr = q["open"].to_numpy(dtype=float)
|
||||
close_arr = q["close"].to_numpy(dtype=float)
|
||||
for ed in _xdxr_event_dates(xdxr_df):
|
||||
# 事件日(或其后首个交易日)在序列中的位置
|
||||
idx = int(np.searchsorted(qd, np.datetime64(ed), side="left"))
|
||||
if idx <= 0 or idx >= len(q):
|
||||
# 事件在序列范围外(早于首根或晚于末根),无法对拍
|
||||
continue
|
||||
events_checked += 1
|
||||
prev_close = close_arr[idx - 1]
|
||||
if prev_close <= 0 or not np.isfinite(prev_close):
|
||||
continue
|
||||
residual = open_arr[idx] / prev_close - 1.0
|
||||
if not np.isfinite(residual):
|
||||
issues.append(
|
||||
QfqIssue(
|
||||
kind="residual_gap",
|
||||
date=_fmt(qd[idx]),
|
||||
detail=f"除权日 open={open_arr[idx]} / 昨收={prev_close} 残差非有限",
|
||||
)
|
||||
)
|
||||
elif residual < -threshold:
|
||||
issues.append(
|
||||
QfqIssue(
|
||||
kind="residual_gap",
|
||||
date=_fmt(qd[idx]),
|
||||
detail=(
|
||||
f"除权日仍向下跳空 {residual:.1%}(超阈值 -{threshold:.1%}),"
|
||||
"疑似复权未生效或漏算事件"
|
||||
),
|
||||
)
|
||||
)
|
||||
elif residual > threshold:
|
||||
issues.append(
|
||||
QfqIssue(
|
||||
kind="wrong_direction",
|
||||
date=_fmt(qd[idx]),
|
||||
detail=(
|
||||
f"除权日向上跳空 {residual:.1%}(超阈值 {threshold:.1%}),"
|
||||
"疑似复权过度或方向算反"
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
# ── 3. NONE 跳空 vs XDXR 对应性 ────────────────────────────────────────────
|
||||
gaps = detect_ex_dividend_gaps(none_df, code, market, margin)
|
||||
gaps_detected = len(gaps)
|
||||
if gaps:
|
||||
event_dates = _xdxr_event_dates(xdxr_df)
|
||||
if none_df is not None and len(none_df) > 0:
|
||||
n = none_df.copy()
|
||||
n["_dt"] = _to_dt_index(n)
|
||||
n = n.sort_values("_dt").reset_index(drop=True)
|
||||
nd = n["_dt"].to_numpy()
|
||||
for g in gaps:
|
||||
gt = np.datetime64(pd.Timestamp(g))
|
||||
idx = int(np.searchsorted(nd, gt, side="left"))
|
||||
# ±2 个交易日内有事件即视为对应
|
||||
lo = max(0, idx - 2)
|
||||
hi = min(len(nd), idx + 3)
|
||||
matched = any(
|
||||
abs((pd.Timestamp(nd[j]) - pd.Timestamp(g)).days) <= 10
|
||||
and any(abs((e - pd.Timestamp(nd[j])).days) <= 3 for e in event_dates)
|
||||
for j in range(lo, hi)
|
||||
)
|
||||
if not matched:
|
||||
issues.append(
|
||||
QfqIssue(
|
||||
kind="unexplained_gap",
|
||||
date=g,
|
||||
detail=(
|
||||
"NONE 序列存在超出跌停幅度的向下跳空,"
|
||||
"但 ±2 交易日内无 XDXR 除权事件对应"
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
ok = not any(i.kind in ("bad_price", "residual_gap", "wrong_direction") for i in issues)
|
||||
return QfqCrosscheckReport(
|
||||
symbol=symbol,
|
||||
ok=ok,
|
||||
events_checked=events_checked,
|
||||
gaps_detected=gaps_detected,
|
||||
issues=issues,
|
||||
)
|
||||
@@ -55,6 +55,10 @@ class BacktestRequest(BaseModel):
|
||||
execution: Literal["next_open", "next_close"] = Field(
|
||||
default="next_open", description="成交模式"
|
||||
)
|
||||
auto_fees: bool = Field(
|
||||
default=False,
|
||||
description="按标的品种自动解析费率(ETF/可转债免印花税等);显式非默认费率优先",
|
||||
)
|
||||
|
||||
# 数据来源 A:内联 OHLCV(上限与 symbol 路径的 count 上限对齐,防 DoS)
|
||||
ohlcv: list[dict[str, Any]] | None = Field(
|
||||
@@ -96,6 +100,9 @@ class PortfolioBacktestRequest(BaseModel):
|
||||
stamp_tax: float = Field(default=0.001, ge=0, le=0.01)
|
||||
slippage: float = Field(default=0.0, ge=0, le=0.05)
|
||||
execution: Literal["next_open", "next_close"] = Field(default="next_open")
|
||||
auto_fees: bool = Field(
|
||||
default=False, description="按各标的品种逐只解析费率(ETF/可转债免印花税等)"
|
||||
)
|
||||
stocks: list[str] = Field(
|
||||
...,
|
||||
min_length=1,
|
||||
@@ -129,6 +136,13 @@ class OptimizeBacktestRequest(BaseModel):
|
||||
commission: float = Field(default=0.0003, ge=0, le=0.01)
|
||||
slippage: float = Field(default=0.0, ge=0, le=0.05)
|
||||
execution: Literal["next_open", "next_close"] = Field(default="next_open")
|
||||
workers: int = Field(
|
||||
default=1,
|
||||
ge=0,
|
||||
le=32,
|
||||
description="并行工作进程数:0/1=串行+指标缓存(默认);2+=进程级并行(实测 36 点网格 "
|
||||
"800 根 K 线约 2 倍,网格越大收益越高)",
|
||||
)
|
||||
param_grid: dict[str, list[int | float | str]] = Field(
|
||||
...,
|
||||
min_length=1,
|
||||
@@ -195,6 +209,77 @@ class OptimizeAllBacktestRequest(BaseModel):
|
||||
return self
|
||||
|
||||
|
||||
class MultiSeedRequest(BaseModel):
|
||||
"""多 seed 随机抽样验证 + 晋级门槛请求(v1.25)。
|
||||
|
||||
在给定股票池上对策略做多次(不同 seed)随机抽样回测,输出跨样本稳定性
|
||||
(正收益比例 / 平均夏普 / 各 seed 稳定性列)与四项晋级门槛判定。
|
||||
"""
|
||||
|
||||
strategy: str = Field(..., description="策略名(见 /backtest/strategies)")
|
||||
params: dict[str, Any] = Field(default_factory=dict, description="策略参数")
|
||||
stocks: list[str] = Field(
|
||||
...,
|
||||
min_length=2,
|
||||
max_length=50,
|
||||
description='股票池,格式 "市场:代码",如 ["SZ:000001","SH:600519"]',
|
||||
)
|
||||
category: Literal["DAY", "WEEK", "MONTH", "MIN_5", "MIN_15", "MIN_30", "MIN_60"] = Field(
|
||||
default="DAY"
|
||||
)
|
||||
count: int = Field(default=500, ge=60, le=2000, description="每标的 K 线根数")
|
||||
n_seeds: int = Field(default=3, ge=1, le=10, description="随机种子数")
|
||||
sample_size: int | None = Field(
|
||||
default=None, ge=1, le=50, description="每 seed 抽样标的数(None=全池)"
|
||||
)
|
||||
cash: float = Field(default=1_000_000.0, gt=0)
|
||||
commission: float = Field(default=0.0003, ge=0, le=0.01)
|
||||
min_commission: float = Field(default=5.0, ge=0)
|
||||
stamp_tax: float = Field(default=0.001, ge=0, le=0.01)
|
||||
slippage: float = Field(default=0.0, ge=0, le=0.05)
|
||||
execution: Literal["next_open", "next_close"] = Field(default="next_open")
|
||||
auto_fees: bool = Field(default=True, description="按各标的品种计费(默认开)")
|
||||
gates: dict[str, float] | None = Field(
|
||||
default=None,
|
||||
description="晋级门槛覆盖(positive_ratio/mean_sharpe/mean_trades/mean_return)",
|
||||
)
|
||||
|
||||
|
||||
class RotationBacktestRequest(BaseModel):
|
||||
"""轮动组合回测请求(v1.27)。
|
||||
|
||||
按打分排名定期换仓:固定槽位等额、跌出排名自动卖出补位。
|
||||
打分支持内置动量或通达信公式数值输出。
|
||||
"""
|
||||
|
||||
stocks: list[str] = Field(
|
||||
...,
|
||||
min_length=2,
|
||||
max_length=50,
|
||||
description='股票池,格式 "市场:代码",如 ["SZ:000001","SH:600519"]',
|
||||
)
|
||||
score: Literal["momentum", "formula"] = Field(
|
||||
default="momentum", description="打分方式:momentum(内置动量)或 formula(公式数值输出)"
|
||||
)
|
||||
period: int = Field(default=20, ge=2, le=250, description="动量回看周期(score=momentum 时)")
|
||||
formula_text: str | None = Field(
|
||||
default=None, max_length=8000, description="通达信公式(score=formula 时)"
|
||||
)
|
||||
score_col: str | None = Field(default=None, description="公式数值输出列名(默认最后一个)")
|
||||
slots: int = Field(default=5, ge=1, le=20, description="持仓槽位数")
|
||||
refresh: Literal["daily", "weekly", "monthly"] = Field(default="weekly", description="调仓频率")
|
||||
category: Literal["DAY", "WEEK", "MONTH", "MIN_5", "MIN_15", "MIN_30", "MIN_60"] = Field(
|
||||
default="DAY"
|
||||
)
|
||||
count: int = Field(default=500, ge=60, le=2000, description="每标的 K 线根数")
|
||||
cash: float = Field(default=1_000_000.0, gt=0)
|
||||
commission: float = Field(default=0.0003, ge=0, le=0.01)
|
||||
min_commission: float = Field(default=5.0, ge=0)
|
||||
stamp_tax: float = Field(default=0.001, ge=0, le=0.01)
|
||||
stop_loss: float | None = Field(default=None, gt=0, le=0.9, description="槽内止损比例")
|
||||
take_profit: float | None = Field(default=None, gt=0, le=10.0, description="槽内止盈比例")
|
||||
|
||||
|
||||
# ── 响应模型 ───────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -213,6 +298,9 @@ class BacktestResultResponse(BaseModel):
|
||||
trades: list[dict[str, Any]]
|
||||
positions: list[dict[str, Any]]
|
||||
config: dict[str, Any]
|
||||
# v1.25:评级(S-D,不看收益)与综合评分(0-100,含收益权重)后端输出
|
||||
grade: dict[str, Any] | None = None
|
||||
score: dict[str, Any] | None = None
|
||||
|
||||
|
||||
class TaskSubmitResponse(BaseModel):
|
||||
|
||||
@@ -19,12 +19,14 @@ from fastapi import APIRouter, Depends
|
||||
from easy_tdx.web.backtest_schemas import (
|
||||
BacktestRequest,
|
||||
BacktestResultResponse,
|
||||
MultiSeedRequest,
|
||||
MultiStrategyBacktestRequest,
|
||||
OptimizeAllBacktestRequest,
|
||||
OptimizeAllRankEntry,
|
||||
OptimizeAllResult,
|
||||
OptimizeBacktestRequest,
|
||||
PortfolioBacktestRequest,
|
||||
RotationBacktestRequest,
|
||||
SignalScanRequest,
|
||||
StrategySchemaResponse,
|
||||
TaskListResponse,
|
||||
@@ -149,6 +151,76 @@ async def get_task(task_id: str) -> TaskStateResponse:
|
||||
)
|
||||
|
||||
|
||||
@router.get("/backtest/tasks/{task_id}/export")
|
||||
async def export_task(task_id: str, format: str = "json") -> Any:
|
||||
"""导出已完成任务的回测结果(JSON 全量 / CSV 主表)。
|
||||
|
||||
``format=json``:完整 result 字典(performance + equity_curve + trades +
|
||||
config 等)。``format=csv``:导出结果中的主表——优先 trades(成交明细),
|
||||
其次 ranking(寻优排名)、equity_curve(资金曲线);都缺时导出
|
||||
performance 键值对。
|
||||
"""
|
||||
import csv
|
||||
import io
|
||||
import json as _json
|
||||
|
||||
from fastapi.responses import Response
|
||||
|
||||
state = get_runner().peek(task_id)
|
||||
if state is None:
|
||||
raise ValueError(f"未知任务 '{task_id}'")
|
||||
if state.status != "done" or state.result is None:
|
||||
raise ValueError(f"任务 '{task_id}' 尚未完成(status={state.status}),无法导出")
|
||||
|
||||
result = state.result
|
||||
fmt = format.lower()
|
||||
if fmt not in ("json", "csv"):
|
||||
raise ValueError(f"不支持的导出格式 '{format}'(可选 json / csv)")
|
||||
|
||||
if fmt == "json":
|
||||
payload = _json.dumps(result, ensure_ascii=False, default=str)
|
||||
return Response(
|
||||
content=payload,
|
||||
media_type="application/json",
|
||||
headers={"Content-Disposition": f'attachment; filename="backtest-{task_id[:8]}.json"'},
|
||||
)
|
||||
|
||||
# CSV:挑主表
|
||||
rows: list[dict[str, Any]] | None = None
|
||||
label = "metrics"
|
||||
for key, name in (("trades", "trades"), ("ranking", "ranking"), ("equity_curve", "equity")):
|
||||
val = result.get(key)
|
||||
if isinstance(val, list) and val and isinstance(val[0], dict):
|
||||
rows = val
|
||||
label = name
|
||||
break
|
||||
if rows is not None:
|
||||
buf = io.StringIO()
|
||||
dict_writer = csv.DictWriter(buf, fieldnames=list(rows[0].keys()), extrasaction="ignore")
|
||||
dict_writer.writeheader()
|
||||
for r in rows:
|
||||
dict_writer.writerow(r)
|
||||
content = buf.getvalue()
|
||||
else:
|
||||
perf = result.get("performance")
|
||||
if not isinstance(perf, dict):
|
||||
raise ValueError(f"任务 '{task_id}' 的结果不含可导出的表格数据")
|
||||
buf = io.StringIO()
|
||||
list_writer = csv.writer(buf)
|
||||
list_writer.writerow(["metric", "value"])
|
||||
for k, v in perf.items():
|
||||
list_writer.writerow([k, v])
|
||||
content = buf.getvalue()
|
||||
label = "performance"
|
||||
return Response(
|
||||
content=content,
|
||||
media_type="text/csv; charset=utf-8",
|
||||
headers={
|
||||
"Content-Disposition": f'attachment; filename="backtest-{task_id[:8]}-{label}.csv"'
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
# ── 组合回测 ───────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -336,6 +408,257 @@ async def run_signal_scan_async(
|
||||
return TaskSubmitResponse(task_id=task_id, status=status)
|
||||
|
||||
|
||||
# ── Walk-Forward / 一条龙评估(v1.25 防过拟合链)──────────────────────────────
|
||||
|
||||
|
||||
@router.post("/backtest/wf/run/async", response_model=TaskSubmitResponse, status_code=202)
|
||||
async def run_walkforward_async(
|
||||
req: BacktestRequest,
|
||||
n_windows: int = 7,
|
||||
client: Any = Depends(get_client),
|
||||
) -> TaskSubmitResponse:
|
||||
"""提交 Walk-Forward 样本外验证后台任务(默认 7 窗、每窗独立开仓)。
|
||||
|
||||
数据来源同 /backtest/run/async(内联 ohlcv 或按 symbol 取行情)。
|
||||
结果含逐窗收益、盈利窗占比 consistency、连乘收益等,通过
|
||||
GET /backtest/tasks/{task_id} 轮询。
|
||||
"""
|
||||
df = await _resolve_df(client, req)
|
||||
snapshot = req.model_copy()
|
||||
description = f"{snapshot.strategy} WF验证 | {len(df)}根 | {snapshot.symbol or '内联数据'}"
|
||||
runner = get_runner()
|
||||
task_id = runner.submit(
|
||||
lambda: _run_walkforward(df, snapshot, n_windows), description=description
|
||||
)
|
||||
state = runner.get(task_id)
|
||||
status: Any = state.status if state.status in ("pending", "running") else "running"
|
||||
return TaskSubmitResponse(task_id=task_id, status=status)
|
||||
|
||||
|
||||
@router.post("/backtest/evaluate/run/async", response_model=TaskSubmitResponse, status_code=202)
|
||||
async def run_evaluate_async(
|
||||
req: BacktestRequest,
|
||||
client: Any = Depends(get_client),
|
||||
) -> TaskSubmitResponse:
|
||||
"""提交一条龙评估后台任务:回测 + WF + 适配性体检 + 综合评分 + S-D 评级
|
||||
+ 买入持有基准对比(excess_return)。
|
||||
|
||||
结果结构见 ``easy_tdx.backtest.benchmark.evaluate_strategy`` 文档,
|
||||
通过 GET /backtest/tasks/{task_id} 轮询;可用
|
||||
GET /backtest/tasks/{task_id}/export?format=json 导出。
|
||||
"""
|
||||
df = await _resolve_df(client, req)
|
||||
snapshot = req.model_copy()
|
||||
description = f"{snapshot.strategy} 一条龙评估 | {snapshot.symbol or '内联数据'}"
|
||||
runner = get_runner()
|
||||
task_id = runner.submit(lambda: _run_evaluate(df, snapshot), description=description)
|
||||
state = runner.get(task_id)
|
||||
status: Any = state.status if state.status in ("pending", "running") else "running"
|
||||
return TaskSubmitResponse(task_id=task_id, status=status)
|
||||
|
||||
|
||||
async def _resolve_df(client: Any, req: BacktestRequest) -> pd.DataFrame:
|
||||
"""内联 ohlcv 或按 symbol 取行情(/backtest/run/async 同逻辑的复用封装)。"""
|
||||
if req.ohlcv is not None:
|
||||
return _ohlcv_to_df(req.ohlcv)
|
||||
if req.symbol is not None:
|
||||
return await _fetch_bars(client, req.symbol, req.category, max(req.count, 800))
|
||||
raise ValueError("必须提供 ohlcv 或 symbol")
|
||||
|
||||
|
||||
@router.post("/backtest/multiseed/run/async", response_model=TaskSubmitResponse, status_code=202)
|
||||
async def run_multiseed_async(
|
||||
req: MultiSeedRequest,
|
||||
client: Any = Depends(get_client),
|
||||
) -> TaskSubmitResponse:
|
||||
"""提交多 seed 随机抽样验证后台任务(v1.25 晋级门槛)。
|
||||
|
||||
股票池逐标的取行情(async),后台线程跑 MultiSeedValidator:多 seed
|
||||
抽样 → 跨样本稳定性 → 四项晋级门槛(正收益比例/平均夏普/平均交易数/
|
||||
平均收益)。结果含 per_seed_positive_ratio 稳定性列,通过
|
||||
GET /backtest/tasks/{task_id} 轮询。
|
||||
"""
|
||||
from easy_tdx.web.convert import category_from_str, market_from_str
|
||||
|
||||
stock_dfs: dict[str, pd.DataFrame] = {}
|
||||
for symbol in req.stocks:
|
||||
market_str, code = symbol.split(":", 1)
|
||||
try:
|
||||
page = await client.get_security_bars(
|
||||
market_from_str(market_str),
|
||||
code,
|
||||
category_from_str(req.category),
|
||||
0,
|
||||
req.count,
|
||||
)
|
||||
except Exception: # noqa: BLE001 — 单标的失败跳过
|
||||
continue
|
||||
if len(page) >= 30:
|
||||
stock_dfs[symbol] = page
|
||||
if len(stock_dfs) < 2:
|
||||
raise ValueError(f"股票池有效标的不足 2 只({len(stock_dfs)}/{len(req.stocks)})")
|
||||
|
||||
snapshot = req.model_copy()
|
||||
description = f"{snapshot.strategy} 多seed验证 | {len(stock_dfs)}只池×{snapshot.n_seeds}seed"
|
||||
runner = get_runner()
|
||||
task_id = runner.submit(lambda: _run_multiseed(stock_dfs, snapshot), description=description)
|
||||
state = runner.get(task_id)
|
||||
status: Any = state.status if state.status in ("pending", "running") else "running"
|
||||
return TaskSubmitResponse(task_id=task_id, status=status)
|
||||
|
||||
|
||||
def _run_multiseed(stock_dfs: dict[str, pd.DataFrame], req: MultiSeedRequest) -> dict[str, Any]:
|
||||
"""执行多 seed 验证(后台线程内调用)。"""
|
||||
from easy_tdx.backtest.validation import MultiSeedValidator
|
||||
|
||||
validator = MultiSeedValidator(
|
||||
strategy=_build_strategy_from(req.strategy, req.params),
|
||||
stock_dfs=stock_dfs,
|
||||
n_seeds=req.n_seeds,
|
||||
sample_size=req.sample_size,
|
||||
gates=req.gates,
|
||||
cash=req.cash,
|
||||
commission=req.commission,
|
||||
min_commission=req.min_commission,
|
||||
stamp_tax=req.stamp_tax,
|
||||
slippage=req.slippage,
|
||||
execution=req.execution,
|
||||
auto_fees=req.auto_fees,
|
||||
)
|
||||
return validator.run().to_dict()
|
||||
|
||||
|
||||
def _build_strategy_from(strategy_name: str, params: dict[str, Any]) -> Any:
|
||||
"""按名 + 参数构造策略实例(registry KeyError → ValueError)。"""
|
||||
from easy_tdx.backtest.strategies import get_registry
|
||||
|
||||
try:
|
||||
entry = get_registry().get(strategy_name)
|
||||
except KeyError as exc:
|
||||
raise ValueError(str(exc)) from exc
|
||||
return entry.build(params)
|
||||
|
||||
|
||||
def _build_strategy(req: BacktestRequest) -> Any:
|
||||
"""解析策略实例(registry KeyError → ValueError → HTTP 400)。"""
|
||||
from easy_tdx.backtest.strategies import get_registry
|
||||
|
||||
try:
|
||||
entry = get_registry().get(req.strategy)
|
||||
except KeyError as exc:
|
||||
raise ValueError(str(exc)) from exc
|
||||
return entry.build(req.params)
|
||||
|
||||
|
||||
def _run_walkforward(df: pd.DataFrame, req: BacktestRequest, n_windows: int = 7) -> dict[str, Any]:
|
||||
"""执行 Walk-Forward 验证并附带常规回测绩效(后台线程内调用)。"""
|
||||
from easy_tdx.backtest.walkforward import WalkForwardEngine
|
||||
|
||||
wf = WalkForwardEngine(
|
||||
strategy=_build_strategy(req),
|
||||
n_windows=n_windows,
|
||||
cash=req.cash,
|
||||
commission=req.commission,
|
||||
min_commission=req.min_commission,
|
||||
stamp_tax=req.stamp_tax,
|
||||
slippage=req.slippage,
|
||||
execution=req.execution,
|
||||
symbol=req.symbol,
|
||||
auto_fees=req.auto_fees,
|
||||
).run(df)
|
||||
return {"walkforward": wf.to_dict()}
|
||||
|
||||
|
||||
def _run_evaluate(df: pd.DataFrame, req: BacktestRequest) -> dict[str, Any]:
|
||||
"""执行一条龙评估(后台线程内调用)。"""
|
||||
from easy_tdx.backtest.benchmark import evaluate_strategy
|
||||
|
||||
return evaluate_strategy(
|
||||
strategy=_build_strategy(req),
|
||||
df=df,
|
||||
cash=req.cash,
|
||||
commission=req.commission,
|
||||
min_commission=req.min_commission,
|
||||
stamp_tax=req.stamp_tax,
|
||||
slippage=req.slippage,
|
||||
execution=req.execution,
|
||||
symbol=req.symbol,
|
||||
auto_fees=req.auto_fees,
|
||||
)
|
||||
|
||||
|
||||
# ── 轮动组合回测(v1.27)─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.post("/backtest/rotation/run/async", response_model=TaskSubmitResponse, status_code=202)
|
||||
async def run_rotation_async(
|
||||
req: RotationBacktestRequest,
|
||||
client: Any = Depends(get_client),
|
||||
) -> TaskSubmitResponse:
|
||||
"""提交轮动组合回测后台任务(v1.27)。
|
||||
|
||||
按打分排名定期换仓:固定槽位等额、跌出排名自动卖出补位、日/周/月刷新、
|
||||
可选槽内止盈止损。打分支持内置动量(``score="momentum"`` + ``period``)
|
||||
或通达信公式数值输出(``score="formula"`` + ``formula_text`` + ``score_col``)。
|
||||
"""
|
||||
from easy_tdx.web.convert import category_from_str, market_from_str
|
||||
|
||||
stock_dfs: dict[str, pd.DataFrame] = {}
|
||||
for symbol in req.stocks:
|
||||
market_str, code = symbol.split(":", 1)
|
||||
try:
|
||||
page = await client.get_security_bars(
|
||||
market_from_str(market_str), code, category_from_str(req.category), 0, req.count
|
||||
)
|
||||
except Exception: # noqa: BLE001 — 单标的失败跳过
|
||||
continue
|
||||
if page is not None and len(page) >= 30:
|
||||
stock_dfs[symbol] = page
|
||||
if len(stock_dfs) < 2:
|
||||
raise ValueError(f"股票池有效标的不足 2 只({len(stock_dfs)}/{len(req.stocks)})")
|
||||
|
||||
snapshot = req.model_copy()
|
||||
description = f"轮动组合 | {len(stock_dfs)}只 × {snapshot.slots}槽 × {snapshot.refresh}"
|
||||
runner = get_runner()
|
||||
task_id = runner.submit(lambda: _run_rotation(stock_dfs, snapshot), description=description)
|
||||
state = runner.get(task_id)
|
||||
status: Any = state.status if state.status in ("pending", "running") else "running"
|
||||
return TaskSubmitResponse(task_id=task_id, status=status)
|
||||
|
||||
|
||||
def _run_rotation(
|
||||
stock_dfs: dict[str, pd.DataFrame], req: RotationBacktestRequest
|
||||
) -> dict[str, Any]:
|
||||
"""执行轮动回测(后台线程内调用)。"""
|
||||
from easy_tdx.backtest.rotation import RotationEngine, formula_score, momentum_score
|
||||
|
||||
if req.score == "formula":
|
||||
if not req.formula_text:
|
||||
raise ValueError("score=formula 需要提供 formula_text")
|
||||
score_fn = formula_score(req.formula_text, req.score_col)
|
||||
else:
|
||||
score_fn = momentum_score(req.period)
|
||||
|
||||
engine = RotationEngine(
|
||||
stock_dfs=stock_dfs,
|
||||
score_fn=score_fn,
|
||||
slots=req.slots,
|
||||
refresh=req.refresh,
|
||||
cash=req.cash,
|
||||
commission=req.commission,
|
||||
min_commission=req.min_commission,
|
||||
stamp_tax=req.stamp_tax,
|
||||
stop_loss=req.stop_loss,
|
||||
take_profit=req.take_profit,
|
||||
)
|
||||
out = engine.run().to_dict()
|
||||
# 组合评级(净值曲线口径)
|
||||
from easy_tdx.backtest.grading import grade_portfolio_equity
|
||||
|
||||
out["grade"] = grade_portfolio_equity(out["equity_curve"]).to_dict()
|
||||
return out
|
||||
|
||||
|
||||
# ── 内部实现 ───────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -359,9 +682,18 @@ def _run_backtest(df: pd.DataFrame, req: BacktestRequest) -> dict[str, Any]:
|
||||
stamp_tax=req.stamp_tax,
|
||||
slippage=req.slippage,
|
||||
execution=req.execution,
|
||||
symbol=req.symbol,
|
||||
auto_fees=req.auto_fees,
|
||||
)
|
||||
result = engine.run(df)
|
||||
return serialize_result(result)
|
||||
out = serialize_result(result)
|
||||
# v1.25:评级/评分后端化——REST 直接输出 S-D 档位与 0-100 综合分
|
||||
from easy_tdx.backtest.grading import grade_performance
|
||||
from easy_tdx.backtest.scoring import score_strategy
|
||||
|
||||
out["grade"] = grade_performance(dict(result.performance)).to_dict()
|
||||
out["score"] = score_strategy(dict(result.performance)).to_dict()
|
||||
return out
|
||||
|
||||
|
||||
def _ohlcv_to_df(records: list[dict[str, Any]]) -> pd.DataFrame:
|
||||
@@ -422,6 +754,7 @@ def _run_portfolio_backtest(
|
||||
stamp_tax=req.stamp_tax,
|
||||
slippage=req.slippage,
|
||||
execution=req.execution,
|
||||
auto_fees=req.auto_fees,
|
||||
)
|
||||
result = engine.run()
|
||||
return serialize_result(result)
|
||||
@@ -595,6 +928,7 @@ def _run_optimize(df: pd.DataFrame, req: OptimizeBacktestRequest) -> dict[str, A
|
||||
commission=req.commission,
|
||||
slippage=req.slippage,
|
||||
execution=req.execution,
|
||||
workers=req.workers,
|
||||
)
|
||||
result = optimizer.run()
|
||||
return result.to_dict()
|
||||
|
||||
@@ -6,8 +6,11 @@
|
||||
设计取舍:
|
||||
- 回测本身是 CPU-bound(numpy/pandas 持 GIL),线程池主要价值是**不阻塞
|
||||
FastAPI 的 asyncio event loop**——回测在独立线程跑,HTTP handler 立即返回。
|
||||
- 任务结果保留在进程内存,带 LRU 上限(默认 100),重启即丢。MVP 可接受;
|
||||
若需持久化历史,未来再加 SQLite。
|
||||
- 任务结果双写:内存 OrderedDict(LRU 上限 100,热路径零磁盘 IO)+
|
||||
SQLite(``~/.easy_tdx/tasks.db``,v1.24 起持久化,serve 重启不丢历史)。
|
||||
查询走「内存优先、磁盘兜底」;列表合并两源。磁盘侧由
|
||||
:mod:`easy_tdx.web.task_store` 管理淘汰(保留 500 条)与中断任务恢复。
|
||||
``EASY_TDX_NO_TASK_DB=1`` 可整体关闭持久化(测试用)。
|
||||
|
||||
并发正确性要点(task_runner.py 审计修复):
|
||||
- ``submit`` 把「注册 task state」与「提交 executor」放在**同一把锁**内,
|
||||
@@ -32,6 +35,8 @@ from dataclasses import dataclass, field
|
||||
from threading import Lock
|
||||
from typing import Any, Literal
|
||||
|
||||
from easy_tdx.web.task_store import get_task_store
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
TaskStatus = Literal["pending", "running", "done", "failed"]
|
||||
@@ -112,7 +117,11 @@ class BacktestTaskRunner:
|
||||
with self._lock:
|
||||
if self._shutdown:
|
||||
raise RuntimeError("任务执行器已关闭,拒绝提交")
|
||||
self._tasks[task_id] = TaskState(task_id=task_id, description=description)
|
||||
state = TaskState(task_id=task_id, description=description)
|
||||
self._tasks[task_id] = state
|
||||
# 持久化 pending 态必须先于 executor.submit——否则工作线程可能
|
||||
# 先写 running/done,随后被这里的旧 pending 快照覆盖(状态回退)。
|
||||
self._persist_locked(state)
|
||||
# 提交 executor 也在锁内——避免「注册后被淘汰再提交」的窗口。
|
||||
# executor.submit 本身很快(入队即返回),不会显著持锁。
|
||||
self._executor.submit(self._run, task_id, func)
|
||||
@@ -122,25 +131,44 @@ class BacktestTaskRunner:
|
||||
def get(self, task_id: str) -> TaskState:
|
||||
"""取任务状态,不存在抛 KeyError。"""
|
||||
with self._lock:
|
||||
if task_id not in self._tasks:
|
||||
raise KeyError(f"未知任务 '{task_id}'")
|
||||
if task_id in self._tasks:
|
||||
return self._tasks[task_id]
|
||||
# 锁外做磁盘恢复(_restore_from_store 内部自行加锁)
|
||||
restored = self._restore_from_store(task_id)
|
||||
if restored is None:
|
||||
raise KeyError(f"未知任务 '{task_id}'")
|
||||
return restored
|
||||
|
||||
def peek(self, task_id: str) -> TaskState | None:
|
||||
"""取任务状态,不存在返回 None(不抛异常,便于轮询)。"""
|
||||
with self._lock:
|
||||
return self._tasks.get(task_id)
|
||||
state = self._tasks.get(task_id)
|
||||
if state is not None:
|
||||
return state
|
||||
return self._restore_from_store(task_id)
|
||||
|
||||
def list_recent(self, limit: int = 20) -> list[TaskState]:
|
||||
"""返回最近 N 个任务(按完成/创建时间倒序,LRU 表尾=最近)。
|
||||
"""返回最近 N 个任务(内存 + 磁盘合并,按完成/创建时间倒序)。
|
||||
|
||||
Args:
|
||||
limit: 最多返回的任务数(默认 20)。
|
||||
"""
|
||||
with self._lock:
|
||||
# OrderedDict 尾部是最近使用的(done 时 move_to_end);倒序取
|
||||
items = list(reversed(self._tasks.values()))
|
||||
return items[:limit]
|
||||
memory_items = list(self._tasks.values())
|
||||
seen = {s.task_id for s in memory_items}
|
||||
# 磁盘侧多取一些(覆盖内存 LRU 已淘汰的),再合并排序
|
||||
try:
|
||||
disk_items = [
|
||||
self._dict_to_state(d)
|
||||
for d in get_task_store().list_recent(limit=limit + len(seen))
|
||||
if d["task_id"] not in seen
|
||||
]
|
||||
except Exception: # noqa: BLE001 — 持久化故障不阻断列表查询
|
||||
logger.exception("任务持久化:list_recent 读盘失败,仅返回内存任务")
|
||||
disk_items = []
|
||||
merged = memory_items + disk_items
|
||||
merged.sort(key=lambda s: s.finished_at or s.created_at, reverse=True)
|
||||
return merged[:limit]
|
||||
|
||||
def status(self, task_id: str) -> TaskStatus | None:
|
||||
"""取任务状态字符串,不存在返回 None。"""
|
||||
@@ -166,6 +194,7 @@ class BacktestTaskRunner:
|
||||
|
||||
状态写入用本地 ``state`` 引用,不假设条目仍在 ``self._tasks`` 中——
|
||||
并发淘汰可能在任务运行期间移除条目。``move_to_end`` 容忍 KeyError。
|
||||
每次状态跃迁(running/done/failed)同步落盘 SQLite。
|
||||
"""
|
||||
# 取本地引用;若已被淘汰则静默退出(无副作用)
|
||||
with self._lock:
|
||||
@@ -175,6 +204,7 @@ class BacktestTaskRunner:
|
||||
return
|
||||
state.status = "running"
|
||||
state.started_at = time.time()
|
||||
self._persist(state)
|
||||
|
||||
try:
|
||||
result = func()
|
||||
@@ -197,6 +227,63 @@ class BacktestTaskRunner:
|
||||
self._tasks.move_to_end(task_id)
|
||||
except KeyError:
|
||||
pass
|
||||
self._persist(state)
|
||||
|
||||
# ── SQLite 持久化辅助 ─────────────────────────────────────────────────────
|
||||
|
||||
def _persist_locked(self, state: TaskState) -> None:
|
||||
"""把状态快照落盘(调用方需持 ``self._lock``,避免与跃迁竞争)。"""
|
||||
try:
|
||||
get_task_store().save(
|
||||
task_id=state.task_id,
|
||||
status=state.status,
|
||||
description=state.description,
|
||||
created_at=state.created_at,
|
||||
started_at=state.started_at,
|
||||
finished_at=state.finished_at,
|
||||
error=state.error,
|
||||
result=state.result,
|
||||
)
|
||||
except Exception: # noqa: BLE001 — 持久化故障不阻断回测主流程
|
||||
logger.exception("任务 %s 持久化失败(不影响任务本身)", state.task_id)
|
||||
|
||||
def _persist(self, state: TaskState) -> None:
|
||||
"""把状态快照落盘(工作线程跃迁后调用,state 已稳定)。"""
|
||||
with self._lock:
|
||||
self._persist_locked(state)
|
||||
|
||||
def _restore_from_store(self, task_id: str) -> TaskState | None:
|
||||
"""从磁盘恢复任务到内存(内存淘汰/进程重启后的查询兜底)。"""
|
||||
try:
|
||||
d = get_task_store().load(task_id)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception("任务 %s 读盘失败", task_id)
|
||||
return None
|
||||
if d is None:
|
||||
return None
|
||||
state = self._dict_to_state(d)
|
||||
with self._lock:
|
||||
# 竞态兜底:等待期间可能已被并发查询/提交加载进内存
|
||||
existing = self._tasks.get(task_id)
|
||||
if existing is not None:
|
||||
return existing
|
||||
self._tasks[task_id] = state
|
||||
self._evict_if_needed_locked()
|
||||
return state
|
||||
|
||||
@staticmethod
|
||||
def _dict_to_state(d: dict[str, Any]) -> TaskState:
|
||||
"""task_store 行字典 → TaskState。"""
|
||||
return TaskState(
|
||||
task_id=d["task_id"],
|
||||
status=d["status"],
|
||||
result=d.get("result"),
|
||||
error=d.get("error"),
|
||||
created_at=float(d.get("created_at") or 0.0),
|
||||
started_at=(float(d["started_at"]) if d.get("started_at") is not None else None),
|
||||
finished_at=(float(d["finished_at"]) if d.get("finished_at") is not None else None),
|
||||
description=d.get("description", ""),
|
||||
)
|
||||
|
||||
def _evict_if_needed_locked(self) -> None:
|
||||
"""超过上限时丢弃最旧的 non-running 任务(调用方需持锁)。
|
||||
|
||||
@@ -0,0 +1,321 @@
|
||||
"""回测后台任务的 SQLite 持久化(v1.24 新增)。
|
||||
|
||||
此前任务结果只存进程内存(LRU 上限 100),serve 重启即清空——对比页历史、
|
||||
已完成的寻优排名全部丢失。本模块把任务状态/结果落盘到
|
||||
``~/.easy_tdx/tasks.db``(随 ``EASY_TDX_CONFIG_DIR`` 走,与 watchlist.db /
|
||||
strategies.db 同一约定),重启后可继续查询历史任务。
|
||||
|
||||
设计对齐 :mod:`easy_tdx.web.strategy_store`:
|
||||
|
||||
- 单文件 SQLite,短连接 + 模块级写锁串行(task_runner 工作线程并发写)。
|
||||
- ``task_id`` 主键,``INSERT OR REPLACE`` 幂等 upsert。
|
||||
- 结果 JSON 列存 ``serialize_result`` 产出的纯 JSON 字典;防御性
|
||||
``default=`` 兜底 numpy/pandas 残留类型。
|
||||
- 保留上限 ``_MAX_ROWS``(默认 500,比内存 LRU 的 100 大,磁盘比内存便宜),
|
||||
超限按 ``created_at`` 淘汰最旧已完成任务。
|
||||
- 启动恢复:进程重启后,上个进程遗留的 ``pending/running`` 行标记为
|
||||
``failed``(error 注明「服务重启中断」)——不可能有进程还在跑它们。
|
||||
|
||||
测试可通过 ``EASY_TDX_NO_TASK_DB=1`` 环境变量整体关闭持久化(退化为纯内存
|
||||
行为,兼容旧单测)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import sqlite3
|
||||
import threading
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
__all__ = ["TaskStore", "get_task_store"]
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_write_lock = threading.Lock()
|
||||
|
||||
# 磁盘保留上限:超过后淘汰最旧的 non-pending/running 任务
|
||||
_MAX_ROWS = 500
|
||||
|
||||
|
||||
def _config_dir() -> Path:
|
||||
return Path(os.environ.get("EASY_TDX_CONFIG_DIR", str(Path.home() / ".easy_tdx")))
|
||||
|
||||
|
||||
def _default_db_path() -> Path:
|
||||
return _config_dir() / "tasks.db"
|
||||
|
||||
|
||||
def _json_default(obj: Any) -> Any:
|
||||
"""json.dumps 防御性兜底:numpy 标量/数组、pd.Timestamp → JSON 原生。"""
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
if isinstance(obj, (np.integer,)):
|
||||
return int(obj)
|
||||
if isinstance(obj, (np.floating,)):
|
||||
v = float(obj)
|
||||
return v if np.isfinite(v) else None
|
||||
if isinstance(obj, np.bool_):
|
||||
return bool(obj)
|
||||
if isinstance(obj, pd.Timestamp):
|
||||
return obj.isoformat()
|
||||
if isinstance(obj, (np.ndarray, pd.Series)):
|
||||
return obj.tolist()
|
||||
return str(obj)
|
||||
|
||||
|
||||
class TaskStore:
|
||||
"""任务状态/结果的 SQLite 存储。单例由 :func:`get_task_store` 提供。
|
||||
|
||||
线程安全:所有写操作在模块级 ``_write_lock`` 内串行(读也走短连接,
|
||||
SQLite 自身 WAL/锁语义可容忍并发读)。
|
||||
"""
|
||||
|
||||
_SCHEMA = """
|
||||
CREATE TABLE IF NOT EXISTS backtest_tasks (
|
||||
task_id TEXT PRIMARY KEY,
|
||||
status TEXT NOT NULL,
|
||||
description TEXT NOT NULL DEFAULT '',
|
||||
created_at REAL NOT NULL,
|
||||
started_at REAL,
|
||||
finished_at REAL,
|
||||
error TEXT,
|
||||
result_json TEXT
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_backtest_tasks_created
|
||||
ON backtest_tasks (created_at DESC);
|
||||
"""
|
||||
|
||||
def __init__(self, db_path: Path | str | None = None) -> None:
|
||||
self._path = Path(db_path) if db_path is not None else _default_db_path()
|
||||
self._initialized = False
|
||||
self._init_lock = threading.Lock()
|
||||
|
||||
@property
|
||||
def path(self) -> Path:
|
||||
return self._path
|
||||
|
||||
def _connect(self) -> sqlite3.Connection:
|
||||
"""打开短连接;首次调用时建目录与表(``_init_lock`` 防并发双初始化)。
|
||||
|
||||
注意:初始化路径(含中断任务恢复)只依赖 ``_init_lock``,不再获取
|
||||
``_write_lock``——``save`` 在持有 ``_write_lock`` 时会调用本方法,
|
||||
若此处再取 ``_write_lock`` 会与非重入锁死锁。
|
||||
"""
|
||||
if not self._initialized:
|
||||
with self._init_lock:
|
||||
if not self._initialized:
|
||||
self._path.parent.mkdir(parents=True, exist_ok=True)
|
||||
conn = sqlite3.connect(self._path)
|
||||
conn.executescript(self._SCHEMA)
|
||||
conn.commit()
|
||||
conn.close()
|
||||
self._recover_interrupted()
|
||||
self._initialized = True
|
||||
return sqlite3.connect(self._path)
|
||||
|
||||
def _recover_interrupted(self) -> None:
|
||||
"""把上个进程遗留的 pending/running 任务标记为 failed。
|
||||
|
||||
仅在 ``_init_lock`` 内的初始化路径调用(无需再取写锁)。
|
||||
"""
|
||||
conn = sqlite3.connect(self._path)
|
||||
try:
|
||||
cur = conn.execute(
|
||||
"""
|
||||
UPDATE backtest_tasks
|
||||
SET status = 'failed',
|
||||
error = '服务重启,任务中断',
|
||||
finished_at = strftime('%s', 'now')
|
||||
WHERE status IN ('pending', 'running')
|
||||
"""
|
||||
)
|
||||
if cur.rowcount > 0:
|
||||
logger.info("任务持久化:恢复时将 %d 个中断任务标记为 failed", cur.rowcount)
|
||||
conn.commit()
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
def save(
|
||||
self,
|
||||
task_id: str,
|
||||
status: str,
|
||||
description: str = "",
|
||||
created_at: float = 0.0,
|
||||
started_at: float | None = None,
|
||||
finished_at: float | None = None,
|
||||
error: str | None = None,
|
||||
result: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
"""Upsert 一条任务记录(status/result 任一变化时调用)。"""
|
||||
with _write_lock:
|
||||
conn = self._connect()
|
||||
try:
|
||||
conn.execute(
|
||||
"""
|
||||
INSERT OR REPLACE INTO backtest_tasks
|
||||
(task_id, status, description, created_at, started_at,
|
||||
finished_at, error, result_json)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
task_id,
|
||||
status,
|
||||
description,
|
||||
created_at,
|
||||
started_at,
|
||||
finished_at,
|
||||
error,
|
||||
json.dumps(result, default=_json_default, ensure_ascii=False)
|
||||
if result is not None
|
||||
else None,
|
||||
),
|
||||
)
|
||||
conn.commit()
|
||||
finally:
|
||||
conn.close()
|
||||
self._evict_if_needed()
|
||||
|
||||
def load(self, task_id: str) -> dict[str, Any] | None:
|
||||
"""读取单条任务;不存在返回 None。result_json 反序列化进 ``result``。"""
|
||||
conn = self._connect()
|
||||
try:
|
||||
cur = conn.execute(
|
||||
"""
|
||||
SELECT task_id, status, description, created_at, started_at,
|
||||
finished_at, error, result_json
|
||||
FROM backtest_tasks WHERE task_id = ?
|
||||
""",
|
||||
(task_id,),
|
||||
)
|
||||
row = cur.fetchone()
|
||||
finally:
|
||||
conn.close()
|
||||
return self._row_to_dict(row) if row is not None else None
|
||||
|
||||
def list_recent(self, limit: int = 20) -> list[dict[str, Any]]:
|
||||
"""按 created_at 倒序列出最近 N 条任务摘要(含 result,供详情直取)。"""
|
||||
conn = self._connect()
|
||||
try:
|
||||
cur = conn.execute(
|
||||
"""
|
||||
SELECT task_id, status, description, created_at, started_at,
|
||||
finished_at, error, result_json
|
||||
FROM backtest_tasks
|
||||
ORDER BY created_at DESC, task_id DESC
|
||||
LIMIT ?
|
||||
""",
|
||||
(int(limit),),
|
||||
)
|
||||
rows = cur.fetchall()
|
||||
finally:
|
||||
conn.close()
|
||||
return [self._row_to_dict(r) for r in rows]
|
||||
|
||||
def delete(self, task_id: str) -> bool:
|
||||
"""删除单条任务;返回是否确实删除了。"""
|
||||
with _write_lock:
|
||||
conn = self._connect()
|
||||
try:
|
||||
cur = conn.execute("DELETE FROM backtest_tasks WHERE task_id = ?", (task_id,))
|
||||
conn.commit()
|
||||
finally:
|
||||
conn.close()
|
||||
return cur.rowcount > 0
|
||||
|
||||
def _evict_if_needed(self) -> None:
|
||||
"""超过 _MAX_ROWS 时淘汰最旧的已完成/失败任务(调用方需持写锁)。"""
|
||||
conn = sqlite3.connect(self._path)
|
||||
try:
|
||||
cur = conn.execute("SELECT COUNT(*) FROM backtest_tasks")
|
||||
n = int(cur.fetchone()[0])
|
||||
if n <= _MAX_ROWS:
|
||||
return
|
||||
conn.execute(
|
||||
"""
|
||||
DELETE FROM backtest_tasks
|
||||
WHERE task_id IN (
|
||||
SELECT task_id FROM backtest_tasks
|
||||
WHERE status IN ('done', 'failed')
|
||||
ORDER BY created_at ASC
|
||||
LIMIT ?
|
||||
)
|
||||
""",
|
||||
(n - _MAX_ROWS,),
|
||||
)
|
||||
conn.commit()
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
@staticmethod
|
||||
def _row_to_dict(row: tuple[Any, ...]) -> dict[str, Any]:
|
||||
"""行 → 字典;result_json 解析失败时降级为 None 并保留原始文本。"""
|
||||
(task_id, status, description, created_at, started_at, finished_at, error, rj) = row
|
||||
result: dict[str, Any] | None = None
|
||||
if rj is not None:
|
||||
try:
|
||||
result = json.loads(rj)
|
||||
except (TypeError, ValueError):
|
||||
logger.warning("任务 %s 的 result_json 解析失败,忽略", task_id)
|
||||
return {
|
||||
"task_id": task_id,
|
||||
"status": status,
|
||||
"description": description,
|
||||
"created_at": created_at,
|
||||
"started_at": started_at,
|
||||
"finished_at": finished_at,
|
||||
"error": error,
|
||||
"result": result,
|
||||
}
|
||||
|
||||
|
||||
# ── 全局单例 ───────────────────────────────────────────────────────────────────
|
||||
|
||||
_STORE: TaskStore | None = None
|
||||
_STORE_LOCK = threading.Lock()
|
||||
|
||||
|
||||
def get_task_store() -> TaskStore:
|
||||
"""获取全局任务存储单例(惰性初始化,线程安全)。
|
||||
|
||||
``EASY_TDX_NO_TASK_DB=1`` 时返回 DisabledTaskStore 退化实现(所有操作
|
||||
no-op / 空),用于测试或显式关闭持久化。
|
||||
"""
|
||||
global _STORE # noqa: PLW0603 — 模块级单例
|
||||
if _STORE is None:
|
||||
with _STORE_LOCK:
|
||||
if _STORE is None:
|
||||
if os.environ.get("EASY_TDX_NO_TASK_DB") == "1":
|
||||
_STORE = _NullTaskStore()
|
||||
else:
|
||||
_STORE = TaskStore()
|
||||
return _STORE
|
||||
|
||||
|
||||
def reset_task_store() -> None:
|
||||
"""重置单例(测试用)。"""
|
||||
global _STORE # noqa: PLW0603
|
||||
with _STORE_LOCK:
|
||||
_STORE = None
|
||||
|
||||
|
||||
class _NullTaskStore(TaskStore):
|
||||
"""持久化关闭时的空实现:读返回空、写 no-op。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__(db_path=Path(":memory:"))
|
||||
|
||||
def save(self, *args: Any, **kwargs: Any) -> None: # noqa: ARG002
|
||||
return
|
||||
|
||||
def load(self, task_id: str) -> dict[str, Any] | None: # noqa: ARG002
|
||||
return None
|
||||
|
||||
def list_recent(self, limit: int = 20) -> list[dict[str, Any]]: # noqa: ARG002
|
||||
return []
|
||||
|
||||
def delete(self, task_id: str) -> bool: # noqa: ARG002
|
||||
return False
|
||||
@@ -0,0 +1,16 @@
|
||||
"""全局测试夹具。
|
||||
|
||||
- 默认关闭回测任务的 SQLite 持久化(``EASY_TDX_NO_TASK_DB=1``):
|
||||
大量单测直接实例化 ``BacktestTaskRunner``,若不关闭会写真实的
|
||||
``~/.easy_tdx/tasks.db`` 污染用户数据。task_store 专属测试通过
|
||||
``monkeypatch.delenv`` + ``EASY_TDX_CONFIG_DIR`` 指向 ``tmp_path``
|
||||
显式重新开启(见 ``test_task_store.py``)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
|
||||
def pytest_configure() -> None:
|
||||
os.environ.setdefault("EASY_TDX_NO_TASK_DB", "1")
|
||||
@@ -0,0 +1,206 @@
|
||||
"""品种感知费率模型测试(fees.py + 引擎 auto_fees 集成)。
|
||||
|
||||
覆盖:
|
||||
- 品种推断:沪深股票 / ETF / LOF / 可转债 / B 股 / 指数 / 北交所
|
||||
- 费率解析:ETF/债券免印花税、B 股印花税保留、最低佣金差异
|
||||
- 引擎集成:auto_fees 覆盖默认费率、显式费率优先、关闭时行为不变
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from easy_tdx.backtest.engine import BacktestEngine
|
||||
from easy_tdx.backtest.fees import (
|
||||
InstrumentKind,
|
||||
detect_instrument_kind,
|
||||
resolve_fee_model,
|
||||
)
|
||||
from easy_tdx.backtest.strategy import Strategy
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 品种推断
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("symbol", "market", "expected"),
|
||||
[
|
||||
("600519", "SH", InstrumentKind.STOCK), # 贵州茅台
|
||||
("601398", "SH", InstrumentKind.STOCK), # 工商银行
|
||||
("000001", "SZ", InstrumentKind.STOCK), # 平安银行
|
||||
("300750", "SZ", InstrumentKind.STOCK), # 宁德时代(创业板)
|
||||
("688981", "SH", InstrumentKind.STOCK), # 中芯国际(科创板)
|
||||
("832000", "BJ", InstrumentKind.STOCK), # 北交所
|
||||
("510300", "SH", InstrumentKind.ETF), # 沪深300ETF
|
||||
("588000", "SH", InstrumentKind.ETF), # 科创50ETF
|
||||
("159915", "SZ", InstrumentKind.ETF), # 创业板ETF
|
||||
("501018", "SH", InstrumentKind.LOF), # 南方原油 LOF
|
||||
("160632", "SZ", InstrumentKind.LOF), # 深 LOF
|
||||
("113050", "SH", InstrumentKind.BOND), # 沪可转债
|
||||
("123456", "SZ", InstrumentKind.BOND), # 深可转债
|
||||
("900901", "SH", InstrumentKind.B_SHARE), # 沪 B
|
||||
("200002", "SZ", InstrumentKind.B_SHARE), # 深 B
|
||||
("000001", "SH", InstrumentKind.INDEX), # 上证指数(同码不同市!)
|
||||
("399001", "SZ", InstrumentKind.INDEX), # 深证成指
|
||||
("SH:510300", None, InstrumentKind.ETF), # 带前缀 symbol
|
||||
("SZ:159915", None, InstrumentKind.ETF),
|
||||
("510300", 1, InstrumentKind.ETF), # 通达信 int 市场
|
||||
("000001", 0, InstrumentKind.STOCK),
|
||||
("510300", None, InstrumentKind.ETF), # 无市场,仅代码粗判
|
||||
("600519", None, InstrumentKind.STOCK),
|
||||
],
|
||||
)
|
||||
def test_detect_instrument_kind(symbol, market, expected):
|
||||
assert detect_instrument_kind(symbol, market) == expected
|
||||
|
||||
|
||||
def test_detect_kind_case_insensitive_and_spacing():
|
||||
assert detect_instrument_kind(" sh:510300 ") == InstrumentKind.ETF
|
||||
assert detect_instrument_kind("sh510300") == InstrumentKind.ETF
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 费率解析
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
def test_stock_fees_keep_stamp_tax():
|
||||
fee = resolve_fee_model("SH:600519")
|
||||
assert fee.kind is InstrumentKind.STOCK
|
||||
assert fee.stamp_tax == pytest.approx(0.001)
|
||||
assert fee.commission == pytest.approx(0.0003)
|
||||
assert fee.min_commission == pytest.approx(5.0)
|
||||
|
||||
|
||||
def test_etf_fees_exempt_stamp_tax():
|
||||
"""核心法定差异:ETF 免印花税。"""
|
||||
fee = resolve_fee_model("SH:510300")
|
||||
assert fee.kind is InstrumentKind.ETF
|
||||
assert fee.stamp_tax == 0.0
|
||||
|
||||
|
||||
def test_bond_fees_exempt_stamp_tax_and_lower_min():
|
||||
fee = resolve_fee_model("SZ:123456")
|
||||
assert fee.kind is InstrumentKind.BOND
|
||||
assert fee.stamp_tax == 0.0
|
||||
assert fee.min_commission < 5.0
|
||||
|
||||
|
||||
def test_b_share_fees_keep_stamp_tax():
|
||||
fee = resolve_fee_model("SH:900901")
|
||||
assert fee.kind is InstrumentKind.B_SHARE
|
||||
assert fee.stamp_tax == pytest.approx(0.001)
|
||||
|
||||
|
||||
def test_fee_model_frozen():
|
||||
fee = resolve_fee_model("SH:510300")
|
||||
with pytest.raises(AttributeError):
|
||||
fee.commission = 0.0 # type: ignore[misc]
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 引擎集成
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
class _AlwaysBuy(Strategy):
|
||||
"""首根 K 线全仓买入、持有到末尾的极简策略(保证产生 BUY 交易)。"""
|
||||
|
||||
def init(self) -> None:
|
||||
self._bought = False
|
||||
|
||||
def next(self) -> None:
|
||||
if not self._bought:
|
||||
self.buy()
|
||||
self._bought = True
|
||||
|
||||
|
||||
def _df(n: int = 50) -> pd.DataFrame:
|
||||
dates = pd.date_range("2023-01-01", periods=n, freq="B")
|
||||
close = 10.0 + np.linspace(0, 2, n)
|
||||
return pd.DataFrame(
|
||||
{
|
||||
"datetime": dates,
|
||||
"open": close,
|
||||
"high": close * 1.01,
|
||||
"low": close * 0.99,
|
||||
"close": close,
|
||||
"vol": [1000.0] * n,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def test_engine_auto_fees_overrides_stamp_tax_for_etf():
|
||||
"""auto_fees=True 时 ETF 不收印花税(对比股票默认收)。"""
|
||||
# 股票(默认费率,stamp_tax=0.001)
|
||||
eng_stock = BacktestEngine(_AlwaysBuy, cash=100000.0)
|
||||
assert eng_stock._stamp_tax == pytest.approx(0.001)
|
||||
|
||||
# ETF + auto_fees → stamp_tax 归零
|
||||
eng_etf = BacktestEngine(_AlwaysBuy, cash=100000.0, symbol="SH:510300", auto_fees=True)
|
||||
assert eng_etf._stamp_tax == 0.0
|
||||
assert eng_etf._commission == pytest.approx(0.0003)
|
||||
|
||||
# 结果 config 里带 symbol 与解析后的费率
|
||||
result = eng_etf.run(_df())
|
||||
assert result.config["symbol"] == "SH:510300"
|
||||
assert result.config["stamp_tax"] == 0.0
|
||||
assert result.config["min_commission"] == pytest.approx(5.0)
|
||||
|
||||
|
||||
def test_engine_explicit_fees_win_over_auto():
|
||||
"""显式传入非默认费率时,auto_fees 不覆盖用户意图。"""
|
||||
eng = BacktestEngine(
|
||||
_AlwaysBuy,
|
||||
cash=100000.0,
|
||||
commission=0.0001,
|
||||
min_commission=1.0,
|
||||
stamp_tax=0.0005,
|
||||
symbol="SH:510300",
|
||||
auto_fees=True,
|
||||
)
|
||||
assert eng._commission == pytest.approx(0.0001)
|
||||
assert eng._min_commission == pytest.approx(1.0)
|
||||
assert eng._stamp_tax == pytest.approx(0.0005)
|
||||
|
||||
|
||||
def test_engine_auto_fees_without_symbol_is_noop():
|
||||
"""auto_fees=True 但没给 symbol → 保持默认(不报错)。"""
|
||||
eng = BacktestEngine(_AlwaysBuy, cash=100000.0, auto_fees=True)
|
||||
assert eng._commission == pytest.approx(0.0003)
|
||||
assert eng._stamp_tax == pytest.approx(0.001)
|
||||
|
||||
|
||||
def test_engine_default_behavior_unchanged():
|
||||
"""不传新参数时行为与旧版完全一致(向后兼容)。"""
|
||||
eng = BacktestEngine(_AlwaysBuy, cash=100000.0)
|
||||
assert eng._commission == pytest.approx(0.0003)
|
||||
assert eng._min_commission == pytest.approx(5.0)
|
||||
assert eng._stamp_tax == pytest.approx(0.001)
|
||||
assert eng._symbol is None
|
||||
|
||||
|
||||
def test_portfolio_engine_auto_fees_per_symbol():
|
||||
"""组合引擎按各标的逐只解析费率(股票收印花税、ETF 不收)。"""
|
||||
from easy_tdx.backtest.portfolio_engine import PortfolioBacktestEngine, StockData
|
||||
|
||||
df = _df(60)
|
||||
stocks = [
|
||||
StockData(code="600519", market="SH", df=df),
|
||||
StockData(code="510300", market="SH", df=df),
|
||||
]
|
||||
engine = PortfolioBacktestEngine(
|
||||
strategy=_AlwaysBuy,
|
||||
stocks=stocks,
|
||||
total_cash=200000.0,
|
||||
auto_fees=True,
|
||||
)
|
||||
result = engine.run()
|
||||
# 两只标的结果的 config 中费率不同
|
||||
stock_cfg = result.individual_results["SH600519"].config
|
||||
etf_cfg = result.individual_results["SH510300"].config
|
||||
assert stock_cfg["stamp_tax"] == pytest.approx(0.001)
|
||||
assert etf_cfg["stamp_tax"] == 0.0
|
||||
@@ -0,0 +1,278 @@
|
||||
"""QFQ 对拍校验(公式法 vs 跳空检测法)的单元测试。
|
||||
|
||||
覆盖 ``easy_tdx.mac.qfq_check``:
|
||||
|
||||
- 已知除权案例回归(合成数据复现两类历史上被下游反馈过的场景):
|
||||
* 「茅台式」——长期多重现金分红叠加深层历史,本地重算后应全正且事件处连续;
|
||||
* 「浦发式」——送转股事件,前复权方向应为「旧价向下缩放」。
|
||||
- 反例检测:复权方向算反 → ``wrong_direction``;漏事件 → ``residual_gap``;
|
||||
NONE 跳空但 XDXR 缺记录 → ``unexplained_gap``;负价 → ``bad_price``。
|
||||
- 涨跌停幅度推断(主板/双创/北交所)与跳空检测阈值。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
from easy_tdx.mac.adjust import apply_forward_adjust, has_bad_prices
|
||||
from easy_tdx.mac.qfq_check import (
|
||||
board_limit_ratio,
|
||||
crosscheck_qfq,
|
||||
detect_ex_dividend_gaps,
|
||||
)
|
||||
|
||||
|
||||
def _kline(
|
||||
closes: list[float],
|
||||
start: str = "2010-01-01",
|
||||
opens: list[float] | None = None,
|
||||
) -> pd.DataFrame:
|
||||
"""构造最小 NONE K 线。默认 open=close;可显式给 opens 制造除权跳空。"""
|
||||
n = len(closes)
|
||||
dates = pd.date_range(start, periods=n, freq="D")
|
||||
arr = np.array(closes, dtype=float)
|
||||
o = np.array(opens, dtype=float) if opens is not None else arr.copy()
|
||||
return pd.DataFrame(
|
||||
{
|
||||
"datetime": dates,
|
||||
"open": o,
|
||||
"high": np.maximum(o, arr) * 1.01,
|
||||
"low": np.minimum(o, arr) * 0.99,
|
||||
"close": arr,
|
||||
"vol": [100.0] * n,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _xdxr(
|
||||
events: list[tuple[str, float, float, float, float]],
|
||||
) -> pd.DataFrame:
|
||||
"""构造 XDXR 记录:list of (date, fenhong, peigujia, songzhuangu, peigu)。"""
|
||||
return pd.DataFrame(
|
||||
[
|
||||
{
|
||||
"date": d,
|
||||
"category": 1,
|
||||
"fenhong": fh,
|
||||
"peigujia": pj,
|
||||
"songzhuangu": sz,
|
||||
"peigu": pg,
|
||||
}
|
||||
for d, fh, pj, sz, pg in events
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 涨跌停幅度推断
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
def test_board_limit_ratio_by_code() -> None:
|
||||
"""主板 10%、双创 20%、北交所 30%。"""
|
||||
assert board_limit_ratio("600519") == 0.10 # 沪主板(茅台)
|
||||
assert board_limit_ratio("600000") == 0.10 # 沪主板(浦发)
|
||||
assert board_limit_ratio("000001") == 0.10 # 深主板
|
||||
assert board_limit_ratio("300750") == 0.20 # 创业板
|
||||
assert board_limit_ratio("688981") == 0.20 # 科创板
|
||||
assert board_limit_ratio("832000") == 0.30 # 北交所
|
||||
assert board_limit_ratio("430047") == 0.30 # 北交所
|
||||
assert board_limit_ratio("600000", market=2) == 0.30 # market 显式指定优先
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 跳空检测(detect_ex_dividend_gaps)
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
def test_detect_gap_finds_ex_dividend_drop() -> None:
|
||||
"""主板股票开盘相对昨收跌超 10.5% → 判定为除权跳空。"""
|
||||
# 10 根平稳 + 除权日开盘腰斩(-50%)
|
||||
closes = [10.0] * 5 + [5.0] * 5
|
||||
opens = [10.0] * 5 + [2.5] + [5.0] * 4
|
||||
df = _kline(closes, opens=opens)
|
||||
gaps = detect_ex_dividend_gaps(df, "600000")
|
||||
assert gaps == ["2010-01-06"] # 第 6 根(index 5)为除权日
|
||||
|
||||
|
||||
def test_detect_gap_ignores_normal_limit_down() -> None:
|
||||
"""恰好跌停(-10.0%)不算除权跳空(被 0.5% 余量排除)。"""
|
||||
closes = [10.0] * 5 + [9.0] * 5
|
||||
opens = [10.0] * 5 + [9.0] + [9.0] * 4 # ex 开盘恰好 -10%
|
||||
df = _kline(closes, opens=opens)
|
||||
assert detect_ex_dividend_gaps(df, "600000") == []
|
||||
|
||||
|
||||
def test_detect_gap_uses_chinext_threshold() -> None:
|
||||
"""创业板(20% 涨跌停):-15% 的跳空不报警,-25% 报警。"""
|
||||
closes = [10.0] * 5 + [7.5] * 5
|
||||
opens = [10.0] * 5 + [8.5] + [7.5] * 4 # -15% → 不报
|
||||
df = _kline(closes, opens=opens)
|
||||
assert detect_ex_dividend_gaps(df, "300750") == []
|
||||
|
||||
opens2 = [10.0] * 5 + [7.0] + [7.5] * 4 # -30% → 报
|
||||
df2 = _kline(closes, opens=opens2)
|
||||
assert detect_ex_dividend_gaps(df2, "300750") == ["2010-01-06"]
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 已知案例回归(合成)
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
def test_maotai_style_multi_dividend_case() -> None:
|
||||
"""「茅台式」:多笔大额现金分红叠加深层历史。
|
||||
|
||||
场景:高价股(1700 元)历经 3 次每笔 40~60 元分红,NONE 价格在除权日
|
||||
出现 -3% 左右的真实跳空(小额,低于跌停阈值),深层历史经公式法
|
||||
前复权后应全正、除权日前后连续、对拍通过。
|
||||
"""
|
||||
# 构造 300 根:价格在 1700 附近随机游走,3 个除权日各扣一次分红
|
||||
rng = np.random.default_rng(42)
|
||||
n = 300
|
||||
prices = 1700.0 + np.cumsum(rng.normal(0, 8, n))
|
||||
events = [(50, "2010-04-20"), (60, "2010-07-20"), (40, "2010-09-20")]
|
||||
closes = prices.copy()
|
||||
opens = prices.copy()
|
||||
dates = pd.date_range("2010-01-01", periods=n, freq="D")
|
||||
for fh, ex in events:
|
||||
ex_ts = pd.Timestamp(ex)
|
||||
idx = int(np.searchsorted(dates.to_numpy(), np.datetime64(ex_ts)))
|
||||
opens[idx] = closes[idx - 1] - fh # 除权日开盘 = 昨收 - 分红
|
||||
closes[idx:] -= fh # 之后价格整体降一档(简化)
|
||||
none_df = _kline(list(closes), opens=list(opens))
|
||||
xd = _xdxr([(ex, fh, 0.0, 0.0, 0.0) for fh, ex in events])
|
||||
|
||||
qfq = apply_forward_adjust(none_df, xd)
|
||||
# 1. 全正(茅台负价问题的回归断言)
|
||||
assert not has_bad_prices(qfq)
|
||||
# 2. 对拍通过:无 bad_price / residual_gap / wrong_direction
|
||||
report = crosscheck_qfq(none_df, qfq, xd, "600519", 1)
|
||||
assert report.ok, [i.to_dict() for i in report.issues]
|
||||
assert report.events_checked == 3
|
||||
|
||||
|
||||
def test_pufa_style_songzhuangu_direction() -> None:
|
||||
"""「浦发式」:送转股事件的前复权方向。
|
||||
|
||||
10 送 3(songzhuangu=0.3):除权日理论价格 = 昨收 / 1.3。正确的前复权
|
||||
应把除权日**之前**的价格向下缩放(factor = 1/1.3),而不是抬升之后的价格。
|
||||
"""
|
||||
# 20 根 10 元平稳,除权日后理论价 10/1.3 ≈ 7.69
|
||||
closes = [10.0] * 10 + [7.69] * 10
|
||||
opens = [10.0] * 10 + [7.69] + [7.69] * 9
|
||||
none_df = _kline(closes, opens=opens)
|
||||
xd = _xdxr([("2010-01-11", 0.0, 0.0, 0.3, 0.0)])
|
||||
|
||||
qfq = apply_forward_adjust(none_df, xd)
|
||||
# 旧价被向下缩放:前 10 根 ≈ 10/1.3 ≈ 7.69,与除权后持平(连续);
|
||||
# 最新价锚定不动
|
||||
assert abs(qfq["close"].iloc[0] - 7.69) <= 0.02
|
||||
assert abs(qfq["close"].iloc[-1] - 7.69) <= 1e-9
|
||||
# 方向正确 → 对拍通过(除权日 open=7.69 ≈ 复权后昨收 7.69)
|
||||
report = crosscheck_qfq(none_df, qfq, xd, "600000", 1)
|
||||
assert report.ok, [i.to_dict() for i in report.issues]
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 反例检测
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
def test_wrong_direction_adjustment_detected() -> None:
|
||||
"""复权方向算反 → 除权日残留大幅跳空。
|
||||
|
||||
方向反演(因子取倒数):把除权日**之前**的价格放大 ×1.3 而非缩放
|
||||
÷1.3,除权日 open=7.69 对上「复权后」昨收 13.0 → 残差 -41% →
|
||||
``residual_gap``。
|
||||
"""
|
||||
# NONE:10 送 3 场景,除权日开盘 7.69(-23%,超主板阈值 → 会被跳空检测捕获)
|
||||
closes = [10.0] * 10 + [7.69] * 10
|
||||
opens = [10.0] * 10 + [7.69] + [7.69] * 9
|
||||
none_df = _kline(closes, opens=opens)
|
||||
xd = _xdxr([("2010-01-11", 0.0, 0.0, 0.3, 0.0)])
|
||||
|
||||
reversed_df = none_df.copy()
|
||||
reversed_df.loc[reversed_df.index <= 9, ["open", "high", "low", "close"]] *= 1.3
|
||||
report = crosscheck_qfq(none_df, reversed_df, xd, "600000", 1)
|
||||
assert not report.ok
|
||||
assert any(i.kind == "residual_gap" for i in report.issues)
|
||||
|
||||
|
||||
def test_over_adjustment_detected() -> None:
|
||||
"""过度复权(旧价缩得过低)→ 除权日向上跳空 → wrong_direction。"""
|
||||
closes = [10.0] * 10 + [7.69] * 10
|
||||
opens = [10.0] * 10 + [7.69] + [7.69] * 9
|
||||
none_df = _kline(closes, opens=opens)
|
||||
xd = _xdxr([("2010-01-11", 0.0, 0.0, 0.3, 0.0)])
|
||||
|
||||
over = none_df.copy()
|
||||
over.loc[over.index <= 9, ["open", "high", "low", "close"]] *= 0.5 # 应 ÷1.3 却 ×0.5
|
||||
report = crosscheck_qfq(none_df, over, xd, "600000", 1)
|
||||
assert not report.ok
|
||||
assert any(i.kind == "wrong_direction" for i in report.issues)
|
||||
|
||||
|
||||
def test_missed_event_residual_gap_detected() -> None:
|
||||
"""漏算事件(复权结果等于 NONE 原始序列)→ residual_gap。"""
|
||||
closes = [10.0] * 10 + [7.69] * 10
|
||||
opens = [10.0] * 10 + [7.69] + [7.69] * 9
|
||||
none_df = _kline(closes, opens=opens)
|
||||
xd = _xdxr([("2010-01-11", 0.0, 0.0, 0.3, 0.0)])
|
||||
|
||||
# 「复权结果」其实是未复权的 NONE(公式法漏调)→ 除权日残留 -23% 跳空
|
||||
report = crosscheck_qfq(none_df, none_df.copy(), xd, "600000", 1)
|
||||
assert not report.ok
|
||||
assert any(i.kind == "residual_gap" for i in report.issues)
|
||||
|
||||
|
||||
def test_unexplained_gap_without_xdxr_record() -> None:
|
||||
"""NONE 存在除权跳空但 XDXR 无对应记录 → unexplained_gap(不影响 ok)。"""
|
||||
closes = [10.0] * 10 + [7.69] * 10
|
||||
opens = [10.0] * 10 + [7.69] + [7.69] * 9
|
||||
none_df = _kline(closes, opens=opens)
|
||||
|
||||
# XDXR 为空(数据源缺记录),公式法无从调整 → 序列本身「连续性」检查通过,
|
||||
# 但跳空检测应报 unexplained_gap 提示人工核查
|
||||
report = crosscheck_qfq(none_df, none_df.copy(), None, "600000", 1)
|
||||
assert report.gaps_detected == 1
|
||||
assert any(i.kind == "unexplained_gap" for i in report.issues)
|
||||
# unexplained_gap 属于「证据链不一致」而非「复权结果错误」,ok 保持 True
|
||||
assert report.ok
|
||||
|
||||
|
||||
def test_bad_price_reported() -> None:
|
||||
"""复权结果含负价 → bad_price(ok=False)。"""
|
||||
closes = [10.0] * 5
|
||||
none_df = _kline(closes)
|
||||
bad = none_df.copy()
|
||||
bad.loc[0, ["open", "high", "low", "close"]] = -1.0
|
||||
report = crosscheck_qfq(none_df, bad, None, "600000", 1)
|
||||
assert not report.ok
|
||||
assert any(i.kind == "bad_price" for i in report.issues)
|
||||
|
||||
|
||||
def test_clean_series_passes() -> None:
|
||||
"""无事件、无跳空的干净序列 → ok=True、零问题。"""
|
||||
closes = [10.0 + 0.1 * i for i in range(20)]
|
||||
none_df = _kline(closes)
|
||||
report = crosscheck_qfq(none_df, none_df.copy(), None, "600000", 1)
|
||||
assert report.ok
|
||||
assert report.issues == []
|
||||
assert report.events_checked == 0
|
||||
assert report.gaps_detected == 0
|
||||
|
||||
|
||||
def test_report_to_dict_roundtrip() -> None:
|
||||
"""报告可序列化为 JSON 兼容字典。"""
|
||||
closes = [10.0] * 10 + [7.69] * 10
|
||||
opens = [10.0] * 10 + [7.69] + [7.69] * 9
|
||||
none_df = _kline(closes, opens=opens)
|
||||
xd = _xdxr([("2010-01-11", 0.0, 0.0, 0.3, 0.0)])
|
||||
report = crosscheck_qfq(none_df, none_df.copy(), xd, "600000", 1)
|
||||
d = report.to_dict()
|
||||
assert d["symbol"] == "1:600000"
|
||||
assert d["ok"] is False
|
||||
assert isinstance(d["issues"], list) and d["issues"]
|
||||
assert {"kind", "date", "detail"} == set(d["issues"][0].keys())
|
||||
@@ -0,0 +1,259 @@
|
||||
"""回测任务 SQLite 持久化测试(task_store + task_runner 集成 + REST 导出)。
|
||||
|
||||
覆盖:
|
||||
- ``TaskStore``:save/load/list_recent/delete 往返、淘汰、重启恢复中断任务
|
||||
- ``BacktestTaskRunner`` 集成:任务完成后落盘、内存淘汰后磁盘兜底、
|
||||
「重启」(新建 runner)后仍可查历史任务
|
||||
- REST 导出端点:JSON 全量 / CSV 主表 / 未完成任务拒绝导出
|
||||
|
||||
持久化默认被 ``tests/conftest.py`` 关闭(``EASY_TDX_NO_TASK_DB=1``),
|
||||
本文件的测试显式删除该变量并把 ``EASY_TDX_CONFIG_DIR`` 指向 ``tmp_path``。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
pytest.importorskip("fastapi")
|
||||
|
||||
from fastapi.testclient import TestClient # noqa: E402
|
||||
|
||||
from easy_tdx.web import task_store as ts_mod # noqa: E402
|
||||
from easy_tdx.web.task_runner import BacktestTaskRunner # noqa: E402
|
||||
from easy_tdx.web.task_store import TaskStore # noqa: E402
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def persisted_env(tmp_path, monkeypatch):
|
||||
"""开启持久化并指向临时目录;隔离全局单例。"""
|
||||
monkeypatch.delenv("EASY_TDX_NO_TASK_DB", raising=False)
|
||||
monkeypatch.setenv("EASY_TDX_CONFIG_DIR", str(tmp_path))
|
||||
ts_mod.reset_task_store()
|
||||
yield tmp_path
|
||||
ts_mod.reset_task_store()
|
||||
|
||||
|
||||
def _wait_done(runner: BacktestTaskRunner, task_id: str, timeout: float = 5.0) -> None:
|
||||
"""轮询等待任务进入终态。"""
|
||||
deadline = time.time() + timeout
|
||||
while time.time() < deadline:
|
||||
state = runner.peek(task_id)
|
||||
assert state is not None
|
||||
if state.status in ("done", "failed"):
|
||||
return
|
||||
time.sleep(0.02)
|
||||
raise AssertionError(f"任务 {task_id} 超时未完成")
|
||||
|
||||
|
||||
# ── TaskStore 单元 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_store_save_load_roundtrip(persisted_env):
|
||||
store = TaskStore(db_path=persisted_env / "t.db")
|
||||
store.save(
|
||||
task_id="abc",
|
||||
status="done",
|
||||
description="ma_cross | 300根",
|
||||
created_at=1000.0,
|
||||
started_at=1001.0,
|
||||
finished_at=1002.0,
|
||||
result={"performance": {"total_return": 0.25}, "trades": [{"pnl": 1.0}]},
|
||||
)
|
||||
d = store.load("abc")
|
||||
assert d is not None
|
||||
assert d["status"] == "done"
|
||||
assert d["result"]["performance"]["total_return"] == 0.25
|
||||
assert d["created_at"] == 1000.0
|
||||
assert store.load("missing") is None
|
||||
|
||||
|
||||
def test_store_list_recent_order_and_delete(persisted_env):
|
||||
store = TaskStore(db_path=persisted_env / "t.db")
|
||||
for i in range(5):
|
||||
store.save(task_id=f"t{i}", status="done", created_at=1000.0 + i)
|
||||
rows = store.list_recent(limit=3)
|
||||
assert [r["task_id"] for r in rows] == ["t4", "t3", "t2"] # created_at 倒序
|
||||
assert store.delete("t4") is True
|
||||
assert store.delete("t4") is False
|
||||
assert store.load("t4") is None
|
||||
|
||||
|
||||
def test_store_upsert_replaces(persisted_env):
|
||||
"""同 task_id 二次 save 是覆盖而非追加。"""
|
||||
store = TaskStore(db_path=persisted_env / "t.db")
|
||||
store.save(task_id="x", status="running", created_at=1.0)
|
||||
store.save(task_id="x", status="done", created_at=1.0, finished_at=2.0, result={"a": 1})
|
||||
d = store.load("x")
|
||||
assert d["status"] == "done"
|
||||
assert d["result"] == {"a": 1}
|
||||
assert len(store.list_recent(limit=10)) == 1
|
||||
|
||||
|
||||
def test_store_recovers_interrupted_tasks(persisted_env):
|
||||
"""新连接(模拟进程重启)把遗留 pending/running 标记为 failed。"""
|
||||
store = TaskStore(db_path=persisted_env / "t.db")
|
||||
store.save(task_id="p1", status="pending", created_at=1.0)
|
||||
store.save(task_id="r1", status="running", created_at=1.0)
|
||||
store.save(task_id="d1", status="done", created_at=1.0, result={"ok": True})
|
||||
|
||||
# 模拟重启:新实例初始化时触发恢复
|
||||
store2 = TaskStore(db_path=persisted_env / "t.db")
|
||||
_ = store2.list_recent(limit=10)
|
||||
|
||||
assert store2.load("p1")["status"] == "failed"
|
||||
assert "重启" in store2.load("p1")["error"]
|
||||
assert store2.load("r1")["status"] == "failed"
|
||||
assert store2.load("d1")["status"] == "done" # 已完成任务不受影响
|
||||
|
||||
|
||||
def test_store_result_json_corruption_degrades(persisted_env):
|
||||
"""result_json 损坏时 load 降级返回 result=None,不抛异常。"""
|
||||
import sqlite3
|
||||
|
||||
path = persisted_env / "t.db"
|
||||
store = TaskStore(db_path=path)
|
||||
store.save(task_id="bad", status="done", created_at=1.0, result={"a": 1})
|
||||
conn = sqlite3.connect(path)
|
||||
conn.execute("UPDATE backtest_tasks SET result_json = '{not-json' WHERE task_id='bad'")
|
||||
conn.commit()
|
||||
conn.close()
|
||||
|
||||
store2 = TaskStore(db_path=path)
|
||||
d = store2.load("bad")
|
||||
assert d is not None
|
||||
assert d["result"] is None
|
||||
|
||||
|
||||
# ── Runner 集成 ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_runner_persists_done_task_and_survives_memory_eviction(persisted_env):
|
||||
runner = BacktestTaskRunner(max_workers=2, max_results=2)
|
||||
task_id = runner.submit(lambda: {"performance": {"total_return": 0.5}}, description="d")
|
||||
_wait_done(runner, task_id)
|
||||
|
||||
# 磁盘上能查到 done + 完整结果
|
||||
d = ts_mod.get_task_store().load(task_id)
|
||||
assert d is not None
|
||||
assert d["status"] == "done"
|
||||
assert d["result"]["performance"]["total_return"] == 0.5
|
||||
|
||||
# 内存淘汰(提交 3 个新任务挤掉 LRU)后 peek 仍能从磁盘兜底
|
||||
for _ in range(3):
|
||||
_wait_done(runner, runner.submit(lambda: {"x": 1}))
|
||||
state = runner.peek(task_id)
|
||||
assert state is not None
|
||||
assert state.status == "done"
|
||||
assert state.result["performance"]["total_return"] == 0.5
|
||||
runner.shutdown()
|
||||
|
||||
|
||||
def test_runner_new_instance_sees_history(persisted_env):
|
||||
"""「重启」:全新 runner/store 仍能列出并查询历史任务。"""
|
||||
runner1 = BacktestTaskRunner(max_workers=1)
|
||||
task_id = runner1.submit(lambda: {"performance": {"sharpe": 1.2}}, description="hist")
|
||||
_wait_done(runner1, task_id)
|
||||
runner1.shutdown()
|
||||
|
||||
runner2 = BacktestTaskRunner(max_workers=1)
|
||||
state = runner2.peek(task_id)
|
||||
assert state is not None
|
||||
assert state.status == "done"
|
||||
assert state.result["performance"]["sharpe"] == 1.2
|
||||
listed = runner2.list_recent(limit=10)
|
||||
assert any(s.task_id == task_id for s in listed)
|
||||
runner2.shutdown()
|
||||
|
||||
|
||||
def test_runner_persists_failed_task(persisted_env):
|
||||
def _boom():
|
||||
raise RuntimeError("炸了")
|
||||
|
||||
runner = BacktestTaskRunner(max_workers=1)
|
||||
task_id = runner.submit(_boom, description="bad")
|
||||
_wait_done(runner, task_id)
|
||||
d = ts_mod.get_task_store().load(task_id)
|
||||
assert d is not None
|
||||
assert d["status"] == "failed"
|
||||
assert "RuntimeError" in d["error"]
|
||||
runner.shutdown()
|
||||
|
||||
|
||||
# ── REST 导出端点 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _client() -> TestClient:
|
||||
from easy_tdx.web import create_app
|
||||
|
||||
app = create_app()
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def test_export_json_and_csv(persisted_env):
|
||||
client = _client()
|
||||
# 提交一个内联数据回测任务并等待完成
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
np.random.seed(7)
|
||||
n = 200
|
||||
close = 10 + np.cumsum(np.random.randn(n) * 0.2 + 0.05)
|
||||
dates = pd.date_range("2023-01-01", periods=n, freq="B")
|
||||
ohlcv = [
|
||||
{
|
||||
"datetime": d.strftime("%Y-%m-%d"),
|
||||
"open": float(c - 0.05),
|
||||
"high": float(c + 0.1),
|
||||
"low": float(c - 0.1),
|
||||
"close": float(c),
|
||||
"vol": 5000.0,
|
||||
"amount": float(c * 5000),
|
||||
}
|
||||
for d, c in zip(dates, close, strict=True)
|
||||
]
|
||||
resp = client.post(
|
||||
"/api/v1/backtest/run/async",
|
||||
json={"strategy": "ma_cross", "params": {"fast": 5, "slow": 20}, "ohlcv": ohlcv},
|
||||
)
|
||||
assert resp.status_code == 202, resp.text
|
||||
task_id = resp.json()["task_id"]
|
||||
for _ in range(200):
|
||||
st = client.get(f"/api/v1/backtest/tasks/{task_id}").json()
|
||||
if st["status"] in ("done", "failed"):
|
||||
break
|
||||
time.sleep(0.05)
|
||||
assert st["status"] == "done", st.get("error")
|
||||
|
||||
# JSON 导出:完整 result
|
||||
rj = client.get(f"/api/v1/backtest/tasks/{task_id}/export?format=json")
|
||||
assert rj.status_code == 200
|
||||
assert "attachment" in rj.headers["content-disposition"]
|
||||
assert "performance" in rj.json()
|
||||
|
||||
# CSV 导出:主表(trades 或 performance)
|
||||
rc = client.get(f"/api/v1/backtest/tasks/{task_id}/export?format=csv")
|
||||
assert rc.status_code == 200
|
||||
assert rc.headers["content-type"].startswith("text/csv")
|
||||
assert len(rc.text.splitlines()) >= 2
|
||||
|
||||
|
||||
def test_export_rejects_unknown_and_unfinished(persisted_env):
|
||||
client = _client()
|
||||
assert client.get("/api/v1/backtest/tasks/nope/export").status_code == 400
|
||||
resp = client.get("/api/v1/backtest/tasks/nope/export?format=xml")
|
||||
# 未知任务先报「未知任务」
|
||||
assert resp.status_code == 400
|
||||
|
||||
|
||||
def test_task_list_includes_persisted_history_after_new_app(persisted_env):
|
||||
"""应用层重启(同 DB)后 /backtest/tasks 仍列出历史任务。"""
|
||||
runner = BacktestTaskRunner(max_workers=1)
|
||||
task_id = runner.submit(lambda: {"performance": {"total_return": 0.1}}, description="旧任务")
|
||||
_wait_done(runner, task_id)
|
||||
runner.shutdown()
|
||||
|
||||
client = _client()
|
||||
tasks = client.get("/api/v1/backtest/tasks?limit=50").json()["tasks"]
|
||||
assert any(t["task_id"] == task_id for t in tasks)
|
||||
Reference in New Issue
Block a user