b270faf4b9
- BaostockProvider: 读 VPS daily_baostock_full(本地,不调online,守 provider-local-data-only 铁律) - LocalParquetProvider: 读 parquet 兜底,回测117交易日0.4s/月出JSON - all_weather 策略 + runner_backtest 适配 - 数据源融合使用层(单 Provider 内部路由,见 data-fusion spec §6)
429 lines
16 KiB
Python
429 lines
16 KiB
Python
"""全天候策略回测入口。
|
|
|
|
用法(Mac 默认 baostock;VPS Windows / miniQMT 已连用 miniqmt):
|
|
# Mac 默认 baostock(跨平台,不依赖 miniQMT 客户端)
|
|
python -m sanguo_portfolio.runner_backtest \\
|
|
--start 2020-01-01 --end 2024-12-31 --cash 1000000
|
|
|
|
# VPS miniQMT(实盘/精准 xtquant)
|
|
python -m sanguo_portfolio.runner_backtest --provider miniqmt \\
|
|
--start 2020-01-01 --end 2024-12-31 --cash 1000000
|
|
|
|
JSON 输出(供 SSH 捕获,前端 MVP 用):
|
|
python -m sanguo_portfolio.runner_backtest --json \\
|
|
--start 2024-01-01 --end 2024-02-29 --cash 1000000
|
|
|
|
Mac 跑 baostock 默认链路;miniQMT 链路仍保留(实盘 runner_live 用)。
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
# ENV GUARD 必须早于任何 bullet_trade import
|
|
# bullet_trade __init__ 加载时 _create_provider() 读 DEFAULT_DATA_PROVIDER 创建默认 provider:
|
|
# miniqmt → import xtquant(周六休市 miniQMT 客户端不响应→卡死)
|
|
# jqdata → import jqdatasdk(用户铁律不装→ModuleNotFoundError)
|
|
# 方案: ENV 设 jqdata + 预插 mock jqdatasdk, 让 import 走 jqdata 分支拿 mock 不崩不卡;
|
|
# 真实 provider 由 set_data_provider 运行时注入覆盖(local/baostock/miniqmt)。
|
|
import os
|
|
import sys as _sys
|
|
from unittest.mock import MagicMock as _MagicMock
|
|
os.environ.setdefault("DEFAULT_DATA_PROVIDER", "jqdata")
|
|
if "jqdatasdk" not in _sys.modules:
|
|
_m = _MagicMock()
|
|
# jqdata.py 用 @jq.utils.assert_auth 装饰器;MagicMock 的 assert_* 前缀被保护→AttributeError
|
|
_m.utils.assert_auth = lambda func: func # passthrough 装饰器
|
|
_sys.modules["jqdatasdk"] = _m
|
|
|
|
import argparse
|
|
import json
|
|
import logging
|
|
from typing import Any, Dict
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def parse_args() -> argparse.Namespace:
|
|
p = argparse.ArgumentParser(description="sanguo_portfolio 全天候回测")
|
|
p.add_argument("--start", default="2020-01-01", help="回测开始日期 YYYY-MM-DD")
|
|
p.add_argument("--end", default="2024-12-31", help="回测结束日期 YYYY-MM-DD")
|
|
p.add_argument("--cash", type=float, default=1_000_000.0, help="初始资金(元)")
|
|
p.add_argument("--benchmark", default="000300.XSHG", help="基准代码")
|
|
p.add_argument("--max-pool", type=int, default=0, help="限制选股池前N只(0=不限,MVP验证用)")
|
|
p.add_argument("--frequency", default="day", help="回测频率 day/minute")
|
|
p.add_argument(
|
|
"--provider", default="local", choices=["local", "baostock", "miniqmt"],
|
|
help="数据 provider:baostock(默认,Mac/跨平台,历史成分股+TTM) / miniqmt(VPS 实盘,需 xtquant)",
|
|
)
|
|
p.add_argument(
|
|
"--provider-config", default="{}",
|
|
help="provider 配置 JSON 字符串,如 '{\"data_dir\":\"D:/xtdata\"}'",
|
|
)
|
|
p.add_argument(
|
|
"--result-file", default="docs/portfolio_backtest_result.md",
|
|
help="结果输出文件(.md)",
|
|
)
|
|
p.add_argument(
|
|
"--json", action="store_true",
|
|
help="JSON 模式:print(json.dumps(result)) 到 stdout,供 SSH 捕获",
|
|
)
|
|
return p.parse_args()
|
|
|
|
|
|
def build_provider(provider_name: str, config_str: str) -> Any:
|
|
"""构造 provider 实例。
|
|
|
|
Args:
|
|
provider_name: "baostock"(Mac 默认) 或 "miniqmt"(VPS 实盘)
|
|
config_str: provider 配置 JSON 字符串
|
|
"""
|
|
import json
|
|
from .providers import BaostockProvider, LocalParquetProvider, SanguoMiniQmtProvider
|
|
|
|
cfg: Dict[str, Any] = {}
|
|
if config_str and config_str != "{}":
|
|
try:
|
|
cfg = json.loads(config_str)
|
|
except Exception as exc:
|
|
logger.warning("provider-config 解析失败,用默认: %s", exc)
|
|
cfg.setdefault("mode", "backtest")
|
|
|
|
name = (provider_name or "baostock").lower()
|
|
if name == "miniqmt":
|
|
return SanguoMiniQmtProvider(cfg)
|
|
if name == "baostock":
|
|
return BaostockProvider(cfg)
|
|
if name == "local":
|
|
return LocalParquetProvider(cfg)
|
|
raise ValueError(f"未知 provider: {name}(支持: local / baostock / miniqmt)")
|
|
|
|
|
|
def build_broker_facade(engine: Any) -> Any:
|
|
"""把 BacktestEngine 的聚宽风格 API 包成 BrokerFacade。"""
|
|
from .strategies.all_weather import BrokerFacade
|
|
|
|
# bullet_trade 的 BacktestEngine 把 set_benchmark/run_daily 等顶层函数注入到策略
|
|
# 模块 globals 里。这里把 engine 持有的对应函数转发给 BrokerFacade。
|
|
def _order_target_value(code: str, value: float):
|
|
try:
|
|
return engine.api.order_target_value(code, value)
|
|
except Exception:
|
|
try:
|
|
return engine.order_target_value(code, value)
|
|
except Exception as exc:
|
|
logger.warning("order_target_value 失败 %s=%s: %s", code, value, exc)
|
|
return None
|
|
|
|
def _order_value(code: str, value: float):
|
|
try:
|
|
return engine.api.order_value(code, value)
|
|
except Exception:
|
|
try:
|
|
return engine.order_value(code, value)
|
|
except Exception as exc:
|
|
logger.warning("order_value 失败 %s=%s: %s", code, value, exc)
|
|
return None
|
|
|
|
return BrokerFacade(
|
|
order_target_value=_order_target_value,
|
|
order_value=_order_value,
|
|
)
|
|
|
|
|
|
def run_backtest(args: argparse.Namespace) -> Dict[str, Any]:
|
|
"""跑回测,返回结果 dict。
|
|
|
|
BulletTrade 的 BacktestEngine 接受 strategy_file 或 initialize 等函数。
|
|
我们把 AllWeatherStrategy 包成 initialize 函数:initialize 闭包挂 run_daily 等。
|
|
"""
|
|
from bullet_trade import BacktestEngine # type: ignore
|
|
from bullet_trade.data.api import set_data_provider # type: ignore
|
|
|
|
from .strategies import AllWeatherStrategy, AllWeatherConfig
|
|
|
|
provider = build_provider(args.provider, args.provider_config)
|
|
set_data_provider(provider)
|
|
|
|
# 占位策略:initialize 里把 self(strategy)挂到聚宽风格定时器
|
|
holder: Dict[str, Any] = {}
|
|
|
|
def initialize(context):
|
|
strategy = AllWeatherStrategy(
|
|
provider=provider,
|
|
config=AllWeatherConfig(max_pool=args.max_pool),
|
|
)
|
|
holder["strategy"] = strategy
|
|
|
|
# bullet-trade 的 run_daily/run_monthly 接受全局函数;把 method 暴露为模块级
|
|
# 这里偷个懒:用 functools.partial 注册到 engine 的 scheduler
|
|
import functools
|
|
|
|
# bullet-trade 顶层 run_daily 等可调用,context._scheduler 暴露
|
|
try:
|
|
from bullet_trade.core import run_daily, run_monthly # type: ignore
|
|
run_daily(strategy.prepare_stock_list, "9:05")
|
|
run_monthly(strategy.monthly_adjustment, 1, "9:30")
|
|
run_daily(strategy.stop_loss, "14:00")
|
|
except Exception as exc:
|
|
logger.warning("注册定时任务失败(回测可能不触达): %s", exc)
|
|
|
|
strategy.initialize(context)
|
|
holder["broker"] = build_broker_facade_inner(strategy, context)
|
|
strategy.broker = holder["broker"]
|
|
|
|
def build_broker_facade_inner(strategy: AllWeatherStrategy, context: Any):
|
|
from .strategies.all_weather import BrokerFacade
|
|
# 在回测内,聚宽风格 order_target_value 来自 bullet_trade 顶层
|
|
from bullet_trade.core.api import ( # type: ignore
|
|
order_target_value as bt_otv,
|
|
order_value as bt_ov,
|
|
)
|
|
return BrokerFacade(
|
|
order_target_value=lambda c, v: bt_otv(c, v),
|
|
order_value=lambda c, v: bt_ov(c, v),
|
|
)
|
|
|
|
print("[runner] ENGINE_BUILD_PRE", flush=True)
|
|
engine = BacktestEngine(
|
|
initialize=initialize,
|
|
start_date=args.start,
|
|
end_date=args.end,
|
|
frequency=args.frequency,
|
|
initial_cash=args.cash,
|
|
benchmark=args.benchmark,
|
|
)
|
|
print("[runner] RUN_START", flush=True)
|
|
result = engine.run()
|
|
print("[runner] RUN_DONE type=%s" % type(result).__name__, flush=True)
|
|
|
|
# 输出结果摘要到 markdown(JSON 模式时 result_file="" 跳过)
|
|
if getattr(args, "result_file", ""):
|
|
_write_result_md(result, args.result_file, args)
|
|
return result
|
|
|
|
|
|
def _write_result_md(result: Dict[str, Any], path: str, args: argparse.Namespace) -> None:
|
|
"""把回测关键指标写成 markdown(给 docs/portfolio_backtest_result.md)。"""
|
|
try:
|
|
summary = result.get("summary", {}) if isinstance(result, dict) else {}
|
|
lines = [
|
|
"# sanguo_portfolio 全天候回测结果",
|
|
"",
|
|
f"- 区间: {args.start} ~ {args.end}",
|
|
f"- 初始资金: {args.cash:,.0f}",
|
|
f"- 基准: {args.benchmark}",
|
|
"",
|
|
"## 关键指标",
|
|
"",
|
|
"| 指标 | 值 |",
|
|
"|---|---|",
|
|
]
|
|
for k in (
|
|
"total_returns", "annual_returns", "benchmark_returns",
|
|
"alpha", "beta", "sharpe", "sortino", "max_drawdown",
|
|
"win_rate", "turnover",
|
|
):
|
|
if k in summary:
|
|
lines.append(f"| {k} | {summary[k]} |")
|
|
content = "\n".join(lines)
|
|
with open(path, "w") as f:
|
|
f.write(content + "\n")
|
|
logger.info("回测结果写入 %s", path)
|
|
except Exception as exc:
|
|
logger.warning("写结果文件失败: %s", exc)
|
|
|
|
|
|
def run_backtest_json(params: Dict[str, Any]) -> Dict[str, Any]:
|
|
"""JSON 入口(供 SSH 触发,前端 MVP 用)。
|
|
|
|
Args:
|
|
params: {
|
|
pool: 标的池(暂未实际使用,占位),
|
|
start_date, end_date: YYYY-MM-DD,
|
|
initial_cash: 初始资金,
|
|
}
|
|
|
|
Returns:
|
|
{
|
|
"strategy": "all_weather",
|
|
"period": {"start": ..., "end": ..., "trading_days": N},
|
|
"stocks_selected": [{"code":..., "name":...}, ...], # 末日持仓
|
|
"trades": [{date, code, side, amount, price, ...}, ...],
|
|
"equity_curve": [{"date":..., "equity":...}, ...],
|
|
"metrics": {total_return, annual_return, max_drawdown, sharpe, ...},
|
|
}
|
|
"""
|
|
# 构造一个 Namespace 复用 run_backtest
|
|
args = argparse.Namespace(
|
|
start=params.get("start_date", "2024-01-01"),
|
|
end=params.get("end_date", "2024-02-29"),
|
|
cash=float(params.get("initial_cash", 1_000_000.0)),
|
|
benchmark=params.get("benchmark", "000300.XSHG"),
|
|
frequency="day",
|
|
provider=params.get("provider", "local"),
|
|
provider_config="{}",
|
|
result_file="", # JSON 模式不写 md
|
|
max_pool=int(params.get("max_pool", 0)),
|
|
)
|
|
raw = run_backtest(args)
|
|
|
|
summary = raw.get("summary", {}) if isinstance(raw, dict) else {}
|
|
metrics = _extract_metrics(summary)
|
|
|
|
# 净值曲线:daily_records 是 DataFrame,index=date,列含 total_value
|
|
equity_curve = _extract_equity_curve(raw.get("daily_records"))
|
|
|
|
# 选股(末日持仓):daily_positions 最后一日
|
|
stocks_selected = _extract_last_positions(raw.get("daily_positions"))
|
|
|
|
# 成交明细
|
|
trades = _extract_trades(raw.get("trades"))
|
|
|
|
meta = raw.get("meta", {}) if isinstance(raw, dict) else {}
|
|
return {
|
|
"strategy": "all_weather",
|
|
"period": {
|
|
"start": meta.get("start_date", args.start),
|
|
"end": meta.get("end_date", args.end),
|
|
"trading_days": len(equity_curve),
|
|
},
|
|
"stocks_selected": stocks_selected,
|
|
"trades": trades,
|
|
"equity_curve": equity_curve,
|
|
"metrics": metrics,
|
|
"raw_summary": summary,
|
|
}
|
|
|
|
|
|
def _to_float(v: Any) -> float | None:
|
|
"""从 string/number 提取 float,失败返 None。bullet-trade summary 多为 '12.34%' 字符串。"""
|
|
if v is None:
|
|
return None
|
|
if isinstance(v, (int, float)):
|
|
return float(v)
|
|
s = str(v).strip().replace("%", "").replace(",", "")
|
|
try:
|
|
return float(s)
|
|
except (TypeError, ValueError):
|
|
return None
|
|
|
|
|
|
def _extract_metrics(summary: Dict[str, Any]) -> Dict[str, float | None]:
|
|
"""bullet-trade summary 用中文 key('策略收益'/'最大回撤'/...)。
|
|
百分比按字面数值(12.34% → 12.34),前端按需 /100 显示。
|
|
"""
|
|
return {
|
|
"total_return": _to_float(summary.get("策略收益")),
|
|
"annual_return": _to_float(summary.get("策略年化收益")),
|
|
"max_drawdown": _to_float(summary.get("最大回撤")),
|
|
"sharpe": _to_float(summary.get("夏普比率")),
|
|
"win_rate_daily": _to_float(summary.get("日胜率")),
|
|
"win_rate_trade": _to_float(summary.get("交易胜率")),
|
|
"trading_days": _to_float(summary.get("交易天数")),
|
|
}
|
|
|
|
|
|
def _extract_equity_curve(daily_records: Any) -> list[Dict[str, Any]]:
|
|
"""daily_records: DataFrame,index=date,列含 total_value。"""
|
|
out: list[Dict[str, Any]] = []
|
|
if daily_records is None:
|
|
return out
|
|
try:
|
|
import pandas as pd # type: ignore
|
|
if isinstance(daily_records, pd.DataFrame):
|
|
df = daily_records.reset_index()
|
|
date_col = "date" if "date" in df.columns else df.columns[0]
|
|
val_col = "total_value" if "total_value" in df.columns else None
|
|
if val_col is None:
|
|
return out
|
|
for _, row in df.iterrows():
|
|
d = row[date_col]
|
|
out.append({
|
|
"date": getattr(d, "strftime", lambda f: str(d))("%Y-%m-%d"),
|
|
"equity": float(row[val_col]),
|
|
})
|
|
except Exception as exc:
|
|
logger.warning("解析 equity_curve 失败: %s", exc)
|
|
return out
|
|
|
|
|
|
def _extract_last_positions(daily_positions: Any) -> list[Dict[str, Any]]:
|
|
"""daily_positions: DataFrame,列含 date/code/amount/avg_cost/price/value。
|
|
取最后一日的非零持仓作为选股名单。"""
|
|
out: list[Dict[str, Any]] = []
|
|
if daily_positions is None:
|
|
return out
|
|
try:
|
|
import pandas as pd # type: ignore
|
|
if isinstance(daily_positions, pd.DataFrame) and not daily_positions.empty:
|
|
df = daily_positions
|
|
if "date" in df.columns:
|
|
last_date = df["date"].max()
|
|
df = df[df["date"] == last_date]
|
|
for _, row in df.iterrows():
|
|
amt = row.get("amount", 0)
|
|
if amt is None or float(amt) <= 0:
|
|
continue
|
|
out.append({
|
|
"code": str(row.get("code", "")),
|
|
"name": str(row.get("code", "")), # name 字段 bullet-trade 没存,前端展示 code
|
|
"amount": float(amt),
|
|
"avg_cost": float(row.get("avg_cost", 0) or 0),
|
|
"price": float(row.get("price", 0) or 0),
|
|
"value": float(row.get("value", 0) or 0),
|
|
})
|
|
except Exception as exc:
|
|
logger.warning("解析 last_positions 失败: %s", exc)
|
|
return out
|
|
|
|
|
|
def _extract_trades(trades: Any) -> list[Dict[str, Any]]:
|
|
"""trades: list[Trade],用 __dict__ 或属性兜底提取关键字段。"""
|
|
out: list[Dict[str, Any]] = []
|
|
if not trades:
|
|
return out
|
|
keys = ("datetime", "date", "code", "side", "action", "amount",
|
|
"filled_amount", "price", "filled_price", "commission", "status")
|
|
for t in trades:
|
|
item: Dict[str, Any] = {}
|
|
for k in keys:
|
|
v = None
|
|
if hasattr(t, k):
|
|
v = getattr(t, k)
|
|
elif isinstance(t, dict):
|
|
v = t.get(k)
|
|
if v is None:
|
|
continue
|
|
# datetime 类转字符串
|
|
if hasattr(v, "strftime"):
|
|
v = v.strftime("%Y-%m-%d %H:%M:%S")
|
|
try:
|
|
if isinstance(v, (int, float)):
|
|
v = float(v)
|
|
except Exception:
|
|
pass
|
|
item[k] = v
|
|
if item:
|
|
out.append(item)
|
|
return out
|
|
|
|
|
|
def main() -> None:
|
|
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s: %(message)s")
|
|
args = parse_args()
|
|
if args.json:
|
|
# JSON 模式:stderr 仍打日志,stdout 只输出 JSON(供 SSH 捕获)
|
|
result = run_backtest_json({
|
|
"start_date": args.start,
|
|
"end_date": args.end,
|
|
"initial_cash": args.cash,
|
|
"benchmark": args.benchmark,
|
|
"provider": args.provider,
|
|
"max_pool": args.max_pool,
|
|
})
|
|
print(json.dumps(result, ensure_ascii=False, default=str))
|
|
else:
|
|
run_backtest(args)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|