From 1fc1d00c9050bde0646f0803cdb6cc79663ce5e7 Mon Sep 17 00:00:00 2001 From: GitHub Date: Tue, 1 Sep 2026 22:16:44 +0800 Subject: [PATCH] =?UTF-8?q?release:=20v1.24.0=20=E2=80=94=20QFQ=20?= =?UTF-8?q?=E5=AF=B9=E6=8B=8D=E9=AA=8C=E8=AF=81=E4=BD=93=E7=B3=BB=20+=20?= =?UTF-8?q?=E5=9B=9E=E6=B5=8B=E4=BB=BB=E5=8A=A1=E6=8C=81=E4=B9=85=E5=8C=96?= =?UTF-8?q?=20+=20=E5=93=81=E7=A7=8D=E6=84=9F=E7=9F=A5=E8=B4=B9=E7=8E=87?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 升级计划 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 为跨版本累积态,后续版本提交继续演进 --- CHANGELOG.md | 20 ++ pyproject.toml | 2 +- src/easy_tdx/backtest/cli.py | 64 +++++ src/easy_tdx/backtest/engine.py | 35 ++- src/easy_tdx/backtest/fees.py | 182 ++++++++++++ src/easy_tdx/backtest/performance.py | 2 +- src/easy_tdx/backtest/portfolio_engine.py | 9 + src/easy_tdx/mac/client.py | 27 ++ src/easy_tdx/mac/qfq_check.py | 312 ++++++++++++++++++++ src/easy_tdx/web/backtest_schemas.py | 88 ++++++ src/easy_tdx/web/routers/backtest.py | 336 +++++++++++++++++++++- src/easy_tdx/web/task_runner.py | 109 ++++++- src/easy_tdx/web/task_store.py | 321 +++++++++++++++++++++ tests/conftest.py | 16 ++ tests/unit/test_backtest_fees.py | 206 +++++++++++++ tests/unit/test_qfq_crosscheck.py | 278 ++++++++++++++++++ tests/unit/test_task_store.py | 259 +++++++++++++++++ 17 files changed, 2251 insertions(+), 15 deletions(-) create mode 100644 src/easy_tdx/backtest/fees.py create mode 100644 src/easy_tdx/mac/qfq_check.py create mode 100644 src/easy_tdx/web/task_store.py create mode 100644 tests/conftest.py create mode 100644 tests/unit/test_backtest_fees.py create mode 100644 tests/unit/test_qfq_crosscheck.py create mode 100644 tests/unit/test_task_store.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 2510116..058ae8d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,6 +2,26 @@ 本文件记录 easy-tdx 的版本变更。格式遵循 [Keep a Changelog](https://keepachangelog.com/zh-CN/)。 +## [1.24.0] — 2026-09-01 + +**信任与持久化版本**——修复下游反馈的 QFQ 复权可信度问题(引入双引擎对拍验证)、回测任务落盘 SQLite(重启不丢)、品种感知费率(ETF/可转债免印花税)。源自对两个下游项目(backtest-system / indicator-lab)的逆向调研,完整升级计划见 `docs/upgrade-plan-2026H2.md`。 + +### 新增 + +- **QFQ 对拍验证体系**(`mac/qfq_check.py`)——公式法(NONE+XDXR)与跳空检测法(板块感知涨跌停阈值:主板 10%/双创 20%/北交所 30% + 0.5% 余量)双证据链交叉验证前复权结果,检出四类问题:`bad_price`(非法价格)、`residual_gap`(除权日仍残留跳空,疑似漏算/未生效)、`wrong_direction`(残差方向反,疑似复权过度/方向算反)、`unexplained_gap`(NONE 跳空但 XDXR 无对应记录)。已接入 `MacClient` / `AsyncMacClient` 的 QFQ 本地重算路径:不一致即打告警日志,最近一次报告存于 `client.last_qfq_crosscheck`。含「茅台式多重分红」「浦发式送转股方向」合成案例回归测试(13 个用例)。回应下游 backtest-system 对 QFQ 可靠性的反馈。 +- **回测任务 SQLite 持久化**(`web/task_store.py`)——任务状态/结果双写内存 LRU + `~/.easy_tdx/tasks.db`(随 `EASY_TDX_CONFIG_DIR`,保留 500 条),serve 重启后对比页历史任务、已完成寻优排名均可继续查询;重启时遗留的 pending/running 任务自动标记为 failed(注明「服务重启中断」)。`EASY_TDX_NO_TASK_DB=1` 可关闭(测试默认关闭)。 +- **任务结果导出端点**——`GET /backtest/tasks/{task_id}/export?format=json|csv`:JSON 导出完整 result;CSV 智能挑主表(trades → ranking → equity_curve,兜底 performance 键值对),带 `Content-Disposition` 附件头。 +- **品种感知费率**(`backtest/fees.py`)——按代码前缀+市场推断品种(股票/ETF/LOF/可转债/B股/指数),自动解析佣金/最低佣金/印花税;核心法定差异:**ETF/可转债免印花税**(此前扁平默认对 ETF 轮动类策略长期错收印花税)。接入:`BacktestEngine(symbol=..., auto_fees=True)`、`PortfolioBacktestEngine(auto_fees=True)`(逐标的解析)、CLI `easy-tdx backtest --auto-fees`、REST 请求体 `auto_fees` 字段。显式非默认费率仍优先;结果 config 快照记录 symbol 与解析后费率。34 个测试用例。 + +### 修复 + +- `performance.py` 中 `avg_holding_days` 的过时文档注释(实现早已是 FIFO 配对、按 size 加权的真实日历日口径,注释仍写「简化为固定值 5.0」,误导审计)。 + +### 内部 + +- `tests/conftest.py` 全局默认 `EASY_TDX_NO_TASK_DB=1`,防止单测污染用户真实 `~/.easy_tdx/tasks.db`。 +- `task_store` 初始化用独立 `_init_lock`(避免与写锁死锁);`task_runner` 的 pending 落盘先于 executor.submit(避免旧状态覆盖新状态的竞态)。 + ## [1.23.3] — 2026-09-01 **serve 纯 API 模式 + 看板修复**。自 1.23.2 以来的增量: diff --git a/pyproject.toml b/pyproject.toml index 263bde8..79b2874 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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" diff --git a/src/easy_tdx/backtest/cli.py b/src/easy_tdx/backtest/cli.py index ab93710..ff0e846 100644 --- a/src/easy_tdx/backtest/cli.py +++ b/src/easy_tdx/backtest/cli.py @@ -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: """加载策略类。 diff --git a/src/easy_tdx/backtest/engine.py b/src/easy_tdx/backtest/engine.py index 970b61f..2d5847c 100644 --- a/src/easy_tdx/backtest/engine.py +++ b/src/easy_tdx/backtest/engine.py @@ -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: diff --git a/src/easy_tdx/backtest/fees.py b/src/easy_tdx/backtest/fees.py new file mode 100644 index 0000000..b178fe0 --- /dev/null +++ b/src/easy_tdx/backtest/fees.py @@ -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, + ) diff --git a/src/easy_tdx/backtest/performance.py b/src/easy_tdx/backtest/performance.py index 1cb3816..5bf05ac 100644 --- a/src/easy_tdx/backtest/performance.py +++ b/src/easy_tdx/backtest/performance.py @@ -73,7 +73,7 @@ class PerformanceAnalyzer: - avg_loss: 平均亏损 - max_win: 最大盈利 - max_loss: 最大亏损 - - avg_holding_days: 平均持仓天数(简化为固定值 5.0) + - avg_holding_days: 平均持仓天数(FIFO 配对、按 size 加权,日历日口径) - volatility: 年化波动率 """ # 边界检查 diff --git a/src/easy_tdx/backtest/portfolio_engine.py b/src/easy_tdx/backtest/portfolio_engine.py index f2b6c36..7b46976 100644 --- a/src/easy_tdx/backtest/portfolio_engine.py +++ b/src/easy_tdx/backtest/portfolio_engine.py @@ -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 diff --git a/src/easy_tdx/mac/client.py b/src/easy_tdx/mac/client.py index a9065c7..bd78554 100644 --- a/src/easy_tdx/mac/client.py +++ b/src/easy_tdx/mac/client.py @@ -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 diff --git a/src/easy_tdx/mac/qfq_check.py b/src/easy_tdx/mac/qfq_check.py new file mode 100644 index 0000000..03a7abe --- /dev/null +++ b/src/easy_tdx/mac/qfq_check.py @@ -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, + ) diff --git a/src/easy_tdx/web/backtest_schemas.py b/src/easy_tdx/web/backtest_schemas.py index 2f17978..07533e3 100644 --- a/src/easy_tdx/web/backtest_schemas.py +++ b/src/easy_tdx/web/backtest_schemas.py @@ -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): diff --git a/src/easy_tdx/web/routers/backtest.py b/src/easy_tdx/web/routers/backtest.py index 8d6cfbd..ef55f70 100644 --- a/src/easy_tdx/web/routers/backtest.py +++ b/src/easy_tdx/web/routers/backtest.py @@ -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() diff --git a/src/easy_tdx/web/task_runner.py b/src/easy_tdx/web/task_runner.py index 22e9459..fa9ab73 100644 --- a/src/easy_tdx/web/task_runner.py +++ b/src/easy_tdx/web/task_runner.py @@ -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 任务(调用方需持锁)。 diff --git a/src/easy_tdx/web/task_store.py b/src/easy_tdx/web/task_store.py new file mode 100644 index 0000000..e15fc89 --- /dev/null +++ b/src/easy_tdx/web/task_store.py @@ -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 diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..7af0e4f --- /dev/null +++ b/tests/conftest.py @@ -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") diff --git a/tests/unit/test_backtest_fees.py b/tests/unit/test_backtest_fees.py new file mode 100644 index 0000000..3ee5312 --- /dev/null +++ b/tests/unit/test_backtest_fees.py @@ -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 diff --git a/tests/unit/test_qfq_crosscheck.py b/tests/unit/test_qfq_crosscheck.py new file mode 100644 index 0000000..e620ad2 --- /dev/null +++ b/tests/unit/test_qfq_crosscheck.py @@ -0,0 +1,278 @@ +"""QFQ 对拍校验(公式法 vs 跳空检测法)的单元测试。 + +覆盖 ``easy_tdx.mac.qfq_check``: + +- 已知除权案例回归(合成数据复现两类历史上被下游反馈过的场景): + * 「茅台式」——长期多重现金分红叠加深层历史,本地重算后应全正且事件处连续; + * 「浦发式」——送转股事件,前复权方向应为「旧价向下缩放」。 +- 反例检测:复权方向算反 → ``wrong_direction``;漏事件 → ``residual_gap``; + NONE 跳空但 XDXR 缺记录 → ``unexplained_gap``;负价 → ``bad_price``。 +- 涨跌停幅度推断(主板/双创/北交所)与跳空检测阈值。 +""" + +from __future__ import annotations + +import numpy as np +import pandas as pd + +from easy_tdx.mac.adjust import apply_forward_adjust, has_bad_prices +from easy_tdx.mac.qfq_check import ( + board_limit_ratio, + crosscheck_qfq, + detect_ex_dividend_gaps, +) + + +def _kline( + closes: list[float], + start: str = "2010-01-01", + opens: list[float] | None = None, +) -> pd.DataFrame: + """构造最小 NONE K 线。默认 open=close;可显式给 opens 制造除权跳空。""" + n = len(closes) + dates = pd.date_range(start, periods=n, freq="D") + arr = np.array(closes, dtype=float) + o = np.array(opens, dtype=float) if opens is not None else arr.copy() + return pd.DataFrame( + { + "datetime": dates, + "open": o, + "high": np.maximum(o, arr) * 1.01, + "low": np.minimum(o, arr) * 0.99, + "close": arr, + "vol": [100.0] * n, + } + ) + + +def _xdxr( + events: list[tuple[str, float, float, float, float]], +) -> pd.DataFrame: + """构造 XDXR 记录:list of (date, fenhong, peigujia, songzhuangu, peigu)。""" + return pd.DataFrame( + [ + { + "date": d, + "category": 1, + "fenhong": fh, + "peigujia": pj, + "songzhuangu": sz, + "peigu": pg, + } + for d, fh, pj, sz, pg in events + ] + ) + + +# --------------------------------------------------------------------------- # +# 涨跌停幅度推断 +# --------------------------------------------------------------------------- # + + +def test_board_limit_ratio_by_code() -> None: + """主板 10%、双创 20%、北交所 30%。""" + assert board_limit_ratio("600519") == 0.10 # 沪主板(茅台) + assert board_limit_ratio("600000") == 0.10 # 沪主板(浦发) + assert board_limit_ratio("000001") == 0.10 # 深主板 + assert board_limit_ratio("300750") == 0.20 # 创业板 + assert board_limit_ratio("688981") == 0.20 # 科创板 + assert board_limit_ratio("832000") == 0.30 # 北交所 + assert board_limit_ratio("430047") == 0.30 # 北交所 + assert board_limit_ratio("600000", market=2) == 0.30 # market 显式指定优先 + + +# --------------------------------------------------------------------------- # +# 跳空检测(detect_ex_dividend_gaps) +# --------------------------------------------------------------------------- # + + +def test_detect_gap_finds_ex_dividend_drop() -> None: + """主板股票开盘相对昨收跌超 10.5% → 判定为除权跳空。""" + # 10 根平稳 + 除权日开盘腰斩(-50%) + closes = [10.0] * 5 + [5.0] * 5 + opens = [10.0] * 5 + [2.5] + [5.0] * 4 + df = _kline(closes, opens=opens) + gaps = detect_ex_dividend_gaps(df, "600000") + assert gaps == ["2010-01-06"] # 第 6 根(index 5)为除权日 + + +def test_detect_gap_ignores_normal_limit_down() -> None: + """恰好跌停(-10.0%)不算除权跳空(被 0.5% 余量排除)。""" + closes = [10.0] * 5 + [9.0] * 5 + opens = [10.0] * 5 + [9.0] + [9.0] * 4 # ex 开盘恰好 -10% + df = _kline(closes, opens=opens) + assert detect_ex_dividend_gaps(df, "600000") == [] + + +def test_detect_gap_uses_chinext_threshold() -> None: + """创业板(20% 涨跌停):-15% 的跳空不报警,-25% 报警。""" + closes = [10.0] * 5 + [7.5] * 5 + opens = [10.0] * 5 + [8.5] + [7.5] * 4 # -15% → 不报 + df = _kline(closes, opens=opens) + assert detect_ex_dividend_gaps(df, "300750") == [] + + opens2 = [10.0] * 5 + [7.0] + [7.5] * 4 # -30% → 报 + df2 = _kline(closes, opens=opens2) + assert detect_ex_dividend_gaps(df2, "300750") == ["2010-01-06"] + + +# --------------------------------------------------------------------------- # +# 已知案例回归(合成) +# --------------------------------------------------------------------------- # + + +def test_maotai_style_multi_dividend_case() -> None: + """「茅台式」:多笔大额现金分红叠加深层历史。 + + 场景:高价股(1700 元)历经 3 次每笔 40~60 元分红,NONE 价格在除权日 + 出现 -3% 左右的真实跳空(小额,低于跌停阈值),深层历史经公式法 + 前复权后应全正、除权日前后连续、对拍通过。 + """ + # 构造 300 根:价格在 1700 附近随机游走,3 个除权日各扣一次分红 + rng = np.random.default_rng(42) + n = 300 + prices = 1700.0 + np.cumsum(rng.normal(0, 8, n)) + events = [(50, "2010-04-20"), (60, "2010-07-20"), (40, "2010-09-20")] + closes = prices.copy() + opens = prices.copy() + dates = pd.date_range("2010-01-01", periods=n, freq="D") + for fh, ex in events: + ex_ts = pd.Timestamp(ex) + idx = int(np.searchsorted(dates.to_numpy(), np.datetime64(ex_ts))) + opens[idx] = closes[idx - 1] - fh # 除权日开盘 = 昨收 - 分红 + closes[idx:] -= fh # 之后价格整体降一档(简化) + none_df = _kline(list(closes), opens=list(opens)) + xd = _xdxr([(ex, fh, 0.0, 0.0, 0.0) for fh, ex in events]) + + qfq = apply_forward_adjust(none_df, xd) + # 1. 全正(茅台负价问题的回归断言) + assert not has_bad_prices(qfq) + # 2. 对拍通过:无 bad_price / residual_gap / wrong_direction + report = crosscheck_qfq(none_df, qfq, xd, "600519", 1) + assert report.ok, [i.to_dict() for i in report.issues] + assert report.events_checked == 3 + + +def test_pufa_style_songzhuangu_direction() -> None: + """「浦发式」:送转股事件的前复权方向。 + + 10 送 3(songzhuangu=0.3):除权日理论价格 = 昨收 / 1.3。正确的前复权 + 应把除权日**之前**的价格向下缩放(factor = 1/1.3),而不是抬升之后的价格。 + """ + # 20 根 10 元平稳,除权日后理论价 10/1.3 ≈ 7.69 + closes = [10.0] * 10 + [7.69] * 10 + opens = [10.0] * 10 + [7.69] + [7.69] * 9 + none_df = _kline(closes, opens=opens) + xd = _xdxr([("2010-01-11", 0.0, 0.0, 0.3, 0.0)]) + + qfq = apply_forward_adjust(none_df, xd) + # 旧价被向下缩放:前 10 根 ≈ 10/1.3 ≈ 7.69,与除权后持平(连续); + # 最新价锚定不动 + assert abs(qfq["close"].iloc[0] - 7.69) <= 0.02 + assert abs(qfq["close"].iloc[-1] - 7.69) <= 1e-9 + # 方向正确 → 对拍通过(除权日 open=7.69 ≈ 复权后昨收 7.69) + report = crosscheck_qfq(none_df, qfq, xd, "600000", 1) + assert report.ok, [i.to_dict() for i in report.issues] + + +# --------------------------------------------------------------------------- # +# 反例检测 +# --------------------------------------------------------------------------- # + + +def test_wrong_direction_adjustment_detected() -> None: + """复权方向算反 → 除权日残留大幅跳空。 + + 方向反演(因子取倒数):把除权日**之前**的价格放大 ×1.3 而非缩放 + ÷1.3,除权日 open=7.69 对上「复权后」昨收 13.0 → 残差 -41% → + ``residual_gap``。 + """ + # NONE:10 送 3 场景,除权日开盘 7.69(-23%,超主板阈值 → 会被跳空检测捕获) + closes = [10.0] * 10 + [7.69] * 10 + opens = [10.0] * 10 + [7.69] + [7.69] * 9 + none_df = _kline(closes, opens=opens) + xd = _xdxr([("2010-01-11", 0.0, 0.0, 0.3, 0.0)]) + + reversed_df = none_df.copy() + reversed_df.loc[reversed_df.index <= 9, ["open", "high", "low", "close"]] *= 1.3 + report = crosscheck_qfq(none_df, reversed_df, xd, "600000", 1) + assert not report.ok + assert any(i.kind == "residual_gap" for i in report.issues) + + +def test_over_adjustment_detected() -> None: + """过度复权(旧价缩得过低)→ 除权日向上跳空 → wrong_direction。""" + closes = [10.0] * 10 + [7.69] * 10 + opens = [10.0] * 10 + [7.69] + [7.69] * 9 + none_df = _kline(closes, opens=opens) + xd = _xdxr([("2010-01-11", 0.0, 0.0, 0.3, 0.0)]) + + over = none_df.copy() + over.loc[over.index <= 9, ["open", "high", "low", "close"]] *= 0.5 # 应 ÷1.3 却 ×0.5 + report = crosscheck_qfq(none_df, over, xd, "600000", 1) + assert not report.ok + assert any(i.kind == "wrong_direction" for i in report.issues) + + +def test_missed_event_residual_gap_detected() -> None: + """漏算事件(复权结果等于 NONE 原始序列)→ residual_gap。""" + closes = [10.0] * 10 + [7.69] * 10 + opens = [10.0] * 10 + [7.69] + [7.69] * 9 + none_df = _kline(closes, opens=opens) + xd = _xdxr([("2010-01-11", 0.0, 0.0, 0.3, 0.0)]) + + # 「复权结果」其实是未复权的 NONE(公式法漏调)→ 除权日残留 -23% 跳空 + report = crosscheck_qfq(none_df, none_df.copy(), xd, "600000", 1) + assert not report.ok + assert any(i.kind == "residual_gap" for i in report.issues) + + +def test_unexplained_gap_without_xdxr_record() -> None: + """NONE 存在除权跳空但 XDXR 无对应记录 → unexplained_gap(不影响 ok)。""" + closes = [10.0] * 10 + [7.69] * 10 + opens = [10.0] * 10 + [7.69] + [7.69] * 9 + none_df = _kline(closes, opens=opens) + + # XDXR 为空(数据源缺记录),公式法无从调整 → 序列本身「连续性」检查通过, + # 但跳空检测应报 unexplained_gap 提示人工核查 + report = crosscheck_qfq(none_df, none_df.copy(), None, "600000", 1) + assert report.gaps_detected == 1 + assert any(i.kind == "unexplained_gap" for i in report.issues) + # unexplained_gap 属于「证据链不一致」而非「复权结果错误」,ok 保持 True + assert report.ok + + +def test_bad_price_reported() -> None: + """复权结果含负价 → bad_price(ok=False)。""" + closes = [10.0] * 5 + none_df = _kline(closes) + bad = none_df.copy() + bad.loc[0, ["open", "high", "low", "close"]] = -1.0 + report = crosscheck_qfq(none_df, bad, None, "600000", 1) + assert not report.ok + assert any(i.kind == "bad_price" for i in report.issues) + + +def test_clean_series_passes() -> None: + """无事件、无跳空的干净序列 → ok=True、零问题。""" + closes = [10.0 + 0.1 * i for i in range(20)] + none_df = _kline(closes) + report = crosscheck_qfq(none_df, none_df.copy(), None, "600000", 1) + assert report.ok + assert report.issues == [] + assert report.events_checked == 0 + assert report.gaps_detected == 0 + + +def test_report_to_dict_roundtrip() -> None: + """报告可序列化为 JSON 兼容字典。""" + closes = [10.0] * 10 + [7.69] * 10 + opens = [10.0] * 10 + [7.69] + [7.69] * 9 + none_df = _kline(closes, opens=opens) + xd = _xdxr([("2010-01-11", 0.0, 0.0, 0.3, 0.0)]) + report = crosscheck_qfq(none_df, none_df.copy(), xd, "600000", 1) + d = report.to_dict() + assert d["symbol"] == "1:600000" + assert d["ok"] is False + assert isinstance(d["issues"], list) and d["issues"] + assert {"kind", "date", "detail"} == set(d["issues"][0].keys()) diff --git a/tests/unit/test_task_store.py b/tests/unit/test_task_store.py new file mode 100644 index 0000000..e14de2b --- /dev/null +++ b/tests/unit/test_task_store.py @@ -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)