mirror of
https://ghfast.top/https://github.com/aeroxw/easy_tdx_max.git
synced 2026-09-12 23:54:21 +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:
@@ -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}'")
|
||||
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 任务(调用方需持锁)。
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user