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:
GitHub
2026-09-01 22:16:44 +08:00
parent 6412880dbe
commit 1fc1d00c90
17 changed files with 2251 additions and 15 deletions
+20
View File
@@ -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 导出完整 resultCSV 智能挑主表(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
View File
@@ -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"
+64
View File
@@ -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:
"""加载策略类。
+34 -1
View File
@@ -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_taxETF/债券免印花税等),
覆盖默认值;显式传入的非默认费率仍优先于自动解析。
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:
+182
View File
@@ -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,
)
+1 -1
View File
@@ -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
+27
View File
@@ -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
+312
View File
@@ -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,
)
+88
View File
@@ -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):
+335 -1
View File
@@ -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()
+98 -11
View File
@@ -6,8 +6,11 @@
设计取舍:
- 回测本身是 CPU-boundnumpy/pandas 持 GIL),线程池主要价值是**不阻塞
FastAPI 的 asyncio event loop**——回测在独立线程跑,HTTP handler 立即返回。
- 任务结果保留在进程内存,带 LRU 上限(默认 100),重启即丢。MVP 可接受;
若需持久化历史,未来再加 SQLite
- 任务结果双写:内存 OrderedDictLRU 上限 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}'")
return self._tasks[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 任务(调用方需持锁)。
+321
View File
@@ -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
+16
View File
@@ -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")
+206
View File
@@ -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
+278
View File
@@ -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 送 3songzhuangu=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``。
"""
# NONE10 送 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_priceok=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())
+259
View File
@@ -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)