Files
claude_dev 4ad505434b
CI/CD / test (push) Successful in 11s
CI/CD / nas-deploy (push) Successful in 25s
CI/CD / nas-verify (push) Successful in 9s
feat(backtest): 个股回测接入费用(cfg通道)+slippage+benchmark放宽至4基准 [vps]
2026-08-12 23:56:14 +08:00

459 lines
21 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""CTA strategy backtesting engine wrapper using vnpy_ctastrategy.backtesting."""
import sys
import os
import math
import logging
# ProcessPool spawn 子进程不继承主进程 logging 配置 → CTA 回测日志黑盒(看不到卡哪行)。
# 模块级 basicConfig 让 cta_engine 日志 → stderr → docker logs(同 portfolio_worker)。
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s: %(message)s")
import traceback
import uuid
from datetime import datetime
from pathlib import Path
import pandas as pd
# Add vnpy source to path for local development
_VNPY_SRC = os.path.join(os.path.dirname(__file__), "..", "vnpy_v4.4.0")
_VNPY_SRC = os.path.abspath(_VNPY_SRC)
if _VNPY_SRC not in sys.path:
sys.path.insert(0, _VNPY_SRC)
from sanguo_backtest.result_store import BacktestResult, save_result
from sanguo_data.datareader import read_index_daily
from sanguo_backtest.metrics import compute_metrics, BENCHMARK_SYMBOL
class _Tee:
"""同时写多个流(用于把引擎 stdout 落盘到 per-task 日志)。"""
def __init__(self, *streams):
self.streams = streams
def write(self, data):
for s in self.streams:
s.write(data)
def flush(self):
for s in self.streams:
try:
s.flush()
except Exception:
pass
# Mock Exchange enum for local use (replaces vnpy.trader.constant.Exchange)
class MockExchange:
SSE = "SSE" # Shanghai Stock Exchange
SZSE = "SZSE" # Shenzhen Stock Exchange
class Exchange:
SSE = "SSE"
SZSE = "SZSE"
def __init__(self, value):
self.value = value
def __repr__(self):
return f"Exchange.{self.value}"
Exchange = MockExchange.Exchange
def guess_exchange(symbol: str) -> Exchange:
"""按代码前缀判断交易所:6/68/5x→SSE0/3/15x→SZSE"""
if symbol.startswith(("60", "68", "51", "56", "58")):
return Exchange("SSE")
if symbol.startswith(("00", "30", "15")):
return Exchange("SZSE")
return Exchange("SSE")
def _read_vt_setting_db() -> str | None:
"""从 vnpy 原生 vt_setting.json 直接读 database.database 路径。
vnpy SETTINGS 字典会被 datareader.read_index_daily 等代码覆写为 yaml 的
NAS Linux 路径(Windows VPS 上不存在但有 0 字节空文件,os.path.exists 误判),
导致连续多次回测时第二次起路径污染。vt_setting.json 是机器本地配置,权威。
"""
try:
import json
from pathlib import Path
# vnpy 约定:~/.vntrader/vt_setting.json
home = Path.home() / ".vntrader" / "vt_setting.json"
if not home.exists():
return None
with open(home, "r", encoding="utf-8") as f:
data = json.load(f)
return data.get("database.database")
except Exception:
return None
def run_cta_backtest(
strategy_class,
symbol: str,
params: dict,
start: str,
end: str,
cfg,
db_path: str,
benchmark: str = "hs300",
task_id: str | None = None,
capital: float = 1_000_000,
position_pct: float = 0.95,
commission_rate: float = 0.00025,
min_commission: float = 5.0,
stamp_duty_rate: float = 0.0005,
transfer_fee_rate: float = 0.00001,
interval: str = "d",
) -> BacktestResult:
"""
Run CTA strategy backtest using AShareBacktestingEngine (vnpy 子类化).
Args:
strategy_class: CTA strategy class to backtest
symbol: Stock symbol (e.g., "600000")
params: Strategy parameters dict
start: Backtest start date (YYYY-MM-DD format)
end: Backtest end date (YYYY-MM-DD format)
cfg: Configuration object (may contain data paths)
db_path: SQLite database path for saving results
benchmark: Benchmark code (hs300/zz500)
task_id: Optional task ID from runner (reused as the persisted task_id so
runner-id == DB task_id; if omitted a fresh uuid is generated).
capital: Starting capital (元)
position_pct: 仓位占比 0~1(定寸用:shares = floor(capital*pct/price/100)*100
commission_rate: A 股佣金率双边(默认万 2.5)
min_commission: 单笔最低佣金(默认 5 元)
stamp_duty_rate: 印花税率卖方(默认 0.0005)
transfer_fee_rate: 过户费率沪市(默认 0.00001)
interval: K 线周期,"d"=日线(默认) / "5m" / "15m"
5m/15m 走 AShareBacktestingEngine 适配层(直查 dbbardata interval='5m'/'15m'
绕过 vnpy Interval enum 不认非标周期的限制)。
Returns:
BacktestResult: Result object with backtest statistics and status
"""
# Use runner-provided task_id (durable, single id across pool/DB/URL) or generate
if not task_id:
task_id = f"cta_{uuid.uuid4().hex[:8]}"
try:
# Lazy import of AShareBacktestingEngine (subclass of vnpy BacktestingEngine)
from sanguo_backtest.ashare_engine import AShareBacktestingEngine
# Build vt_symbol for A-shares
_exchange = guess_exchange(symbol)
vt_symbol = f"{symbol}.{_exchange.value}"
# Convert date strings to datetime objects
start_dt = datetime.strptime(start, "%Y-%m-%d")
end_dt = datetime.strptime(end, "%Y-%m-%d") if end else None
# Create and configure A-share backtesting engine
engine = AShareBacktestingEngine()
# 周期映射:vnpy Interval enum 只有 '1m'/'1h'/'d' 等,不认 '5m'/'15m'。
# 5m/15m 时把 enum 传 MINUTE 让父类 set_parameters 校验通过,真实 DB interval
# 字符串存 engine.raw_interval 给 AShareBacktestingEngine.load_data 自定义路径用。
if interval in ("5m", "15m"):
engine_interval = "1m" # Interval.MINUTE.value
else:
engine_interval = interval # "d" / "1m" / "1h" 等 vnpy 原生支持的值
# Set parameters with A-share specific values
engine.set_parameters(
vt_symbol=vt_symbol,
interval=engine_interval,
start=start_dt,
end=end_dt,
rate=commission_rate, # 佣金率(AShareDailyResult 用自身 commission_rate,此处仅保持一致)
slippage=0, # No slippage for simplicity
size=1, # Contract size (1 for stocks)
pricetick=0.01, # Minimum price tick (0.01 yuan for A-shares)
capital=capital, # Starting capital
)
# raw_interval 在 5m/15m 时触发 ashare_engine 适配路径;其它周期与 engine.interval
# 一致,走 vnpy 原生 load_data。
engine.raw_interval = interval
# 前端透传费用(cfg 通道:routes.submit_cta _build_fee_cfg 打包;None/非 dict→用函数默认值)
if cfg and isinstance(cfg, dict):
commission_rate = float(cfg.get("commission_rate", commission_rate))
min_commission = float(cfg.get("min_commission", min_commission))
stamp_duty_rate = float(cfg.get("stamp_duty_rate", stamp_duty_rate))
transfer_fee_rate = float(cfg.get("transfer_fee_rate", transfer_fee_rate))
_slippage = float(cfg.get("slippage", 0.0)) if (cfg and isinstance(cfg, dict)) else 0.0
# A 股适配参数(定寸 + 费用)
engine.position_pct = position_pct
engine.commission_rate = commission_rate
engine.min_commission = min_commission
engine.stamp_duty_rate = stamp_duty_rate
engine.transfer_fee_rate = transfer_fee_rate
engine.is_sse = (_exchange.value == "SSE")
if _slippage:
engine.slippage = _slippage # 覆盖 set_parameters 的 slippage=0
# Add strategy
engine.add_strategy(strategy_class, params)
# Configure vnpy DB → A-share quant_trading.db. Worker process (spawn)
# doesn't inherit main-process SETTINGS, so set before engine.load_data.
# _dcfg is also reused by the metrics branch (benchmark data_paths) since the
# cfg param can be None when called via the API.
#
# 路径解析优先级:vnpy 原生 vt_setting.json(机器相关,正确) >
# yaml data_paths.vnpy_db(模板可能是 NAS Linux 路径,Windows VPS 上不存在)。
# 任一不存在时回退到另一个,避免 VPS 上 yaml 写 NAS 路径导致回测加载 0 数据。
#
# 不能直接用 SETTINGS.get("database.database"):它会被 datareader.read_index_daily
# (基准加载)覆写为 yaml NAS 路径,污染后续调用。这里每次重新读 vt_setting.json
# 原始值(machine-truth),不依赖被污染的 SETTINGS 缓存。
_dcfg = None
try:
from vnpy.trader.setting import SETTINGS
from sanguo_data.config import load_config, find_config_path
_dcfg = load_config(find_config_path())
SETTINGS["database.name"] = "sqlite"
_vt_db = _read_vt_setting_db() # 从 vt_setting.json 直接读(不依赖 SETTINGS dict
_yaml_db = _dcfg.data_paths.get("vnpy_db")
if _vt_db and os.path.exists(_vt_db):
_resolved_db = _vt_db
elif _yaml_db and os.path.exists(_yaml_db):
_resolved_db = _yaml_db
else:
# 都不存在时保留 vnpy SETTINGS(让下游报出真实错误,而不是被 yaml 覆盖)
_resolved_db = _vt_db or _yaml_db
SETTINGS["database.database"] = _resolved_db
# 5m/15m 适配:ashare_engine._load_intraday_data 直查 SQLite 需要 DB 物理路径。
engine.sqlite_db_path = _resolved_db
except Exception as e:
logging.warning("vnpy 数据库配置加载失败(回测可能无法加载历史数据): %s", e)
# Capture engine run output (load/run/stats) to per-task log file, so the
# result page 日志 tab has real content. Tee stdout inside the worker process
# (contained — doesn't affect the main API process).
file_dir = os.path.dirname(os.path.abspath(db_path))
log_path = os.path.join(file_dir, f"{task_id}.log")
_log_f = open(log_path, "w", encoding="utf-8")
_log_f.write(
f"==== 回测日志 ====\n任务: {task_id}\n策略: {getattr(strategy_class, '__name__', strategy_class)}\n"
f"标的: {vt_symbol}\n区间: {start} ~ {end}\n参数: {params}\n基准: {benchmark}\n==================\n"
)
_log_f.flush()
_orig_stdout = sys.stdout
sys.stdout = _Tee(_orig_stdout, _log_f)
try:
# Load historical data
engine.load_data()
logging.info("[cta] load_data 完成: history bars=%d", len(getattr(engine, "history_data", []) or []))
# C1 定寸:size = N 股/手。load_data 后取首根 bar close 算满仓手数。
# vnpy 的 turnover/PnL/commission 自动 ×size,策略 volume 保持 1 手=N 股=满仓,
# 所有策略(DoubleMa/BollChannel/DualThrust)无需改。
_sizing_shares_per_lot = 0
if engine.history_data:
_first_close = engine.history_data[0].close_price
if _first_close > 0:
_sizing_shares_per_lot = int(
math.floor(capital * position_pct / _first_close / 100) * 100
)
engine.size = _sizing_shares_per_lot
logging.info(
"定寸: N=%s 股/手 (capital=%s pct=%s first_close=%s)",
_sizing_shares_per_lot, capital, position_pct, _first_close,
)
# Run backtesting
engine.run_backtesting()
logging.info("[cta] run_backtesting 完成")
# Calculate statistics — calculate_result() returns a daily DataFrame,
# calculate_statistics(df) returns the stats dict (sharpe/drawdown/etc.)
daily_df = engine.calculate_result()
raw_stats = engine.calculate_statistics(daily_df, output=False) or {}
finally:
sys.stdout = _orig_stdout
_log_f.flush()
_log_f.close()
# Ensure JSON-serializable (vnpy may include Timestamp / non-numeric / NaN values)
statistics = {
k: (None if (isinstance(v, float) and not math.isfinite(v))
else v if isinstance(v, (int, float, str, bool)) or v is None
else str(v))
for k, v in raw_stats.items()
}
# H7 degenerate 检测:零成交或空数据不静默 done
_degenerate_reason = None
trades_dict_check = engine.trades if isinstance(engine.trades, dict) else {}
if not trades_dict_check:
_degenerate_reason = "零成交记录(策略未触发任何交易信号)"
elif daily_df is None or (hasattr(daily_df, "empty") and daily_df.empty):
_degenerate_reason = "日度盈亏数据为空"
if _degenerate_reason:
# 标 degenerate flagstatus 仍为 done,避免前端 status map 不认导致列表/计数异常;
# 结果页据 trades=0 + statistics.degenerate 自然呈现"零成交"诚实状态)
statistics["degenerate"] = True
statistics["degenerate_reason"] = _degenerate_reason
logging.warning("回测退化: %s (symbol=%s strategy=%s)", _degenerate_reason, symbol, strategy_class.__name__)
# C1: 暴露定寸参数到结果(集成测试 + 前端可查)
statistics["sizing_shares_per_lot"] = _sizing_shares_per_lot
# Calculate relative metrics against benchmark (Task 3)
# Ensure daily_df index is datetime for compute_metrics
if daily_df is not None and not daily_df.empty:
if not isinstance(daily_df.index, pd.DatetimeIndex):
daily_df.index = pd.to_datetime(daily_df.index)
# Get benchmark code (default hs300) + a cfg that has data_paths.
# _dcfg is the loaded config; fall back to the passed cfg if loading failed.
benchmark_code = BENCHMARK_SYMBOL.get(benchmark, "sh000300")
bench_cfg = _dcfg if (_dcfg is not None and hasattr(_dcfg, "data_paths")) else cfg
# Load benchmark data
start_date = start_dt if isinstance(start_dt, datetime) else datetime.strptime(start, "%Y-%m-%d")
end_date = end_dt if isinstance(end_dt, datetime) else datetime.strptime(end, "%Y-%m-%d")
try:
logging.info("[cta] read_index_daily 开始: %s %s~%s", benchmark_code, start_date.date(), end_date.date())
bench_df = read_index_daily(benchmark_code, start_date, end_date, bench_cfg)
logging.info("[cta] read_index_daily 完成: %d", 0 if bench_df is None else len(bench_df))
if bench_df is not None and not bench_df.empty and "close" in bench_df.columns:
# Calculate benchmark daily returns
bench_df["date"] = pd.to_datetime(bench_df["date"])
# 去时区:daily_df.index 是 tz-naivebenchmark 若 tz-aware 会让
# compute_metrics 内部 reindex 抛 TypeError(被上层 except 静默吞掉,
# 致 _metrics.json 不生成)。统一去掉 tz 保证对齐。
if getattr(bench_df["date"].dt, "tz", None) is not None:
bench_df["date"] = bench_df["date"].dt.tz_localize(None)
bench_df = bench_df.sort_values("date")
benchmark_returns = bench_df["close"].pct_change().dropna()
benchmark_returns.index = pd.to_datetime(bench_df["date"].iloc[1:])
else:
# 基准缺失:降级。compute_metrics 内部 reindex→ffill→fillna(0) 能容空
# benchmark(基准类指标返 NaN→None,策略指标正常算),故仍执行 + metrics.json
# 照写。避免原实现整段跳过 → 策略净值/回撤/波动图也空(不只基准图)。
benchmark_returns = pd.Series(dtype=float)
logging.warning(
"基准数据缺失,降级计算(基准/Alpha/Beta 图将空,策略指标不受影响): "
"benchmark=%s %s~%s",
benchmark_code, start_date.date(), end_date.date(),
)
# H4: compute_metrics 内部从 daily_df["balance"] 自算 simple return
# 不再依赖 vnpy 的 log return 列(删原三路 fallback
# Compute relative metricsbenchmark 有无都执行;空时基准类指标 NaN→None)
logging.info("[cta] compute_metrics 开始: daily_df shape=%s benchmark=%d", daily_df.shape, len(benchmark_returns))
metrics_result = compute_metrics(daily_df, benchmark_returns)
logging.info("[cta] compute_metrics 完成")
# Merge scalars into statistics (for API response)
statistics.update(metrics_result.scalars)
# Serialize series to JSON (separate file, same as equity_curve/trades)
import json
series_data = {}
for key, series in metrics_result.series.items():
if isinstance(series, pd.Series):
series_data[key] = {
"dates": series.index.astype(str).tolist(),
"values": [None if (isinstance(x, float) and not math.isfinite(x)) else x
for x in series.tolist()]
}
# Write metrics series to JSON file
file_dir = os.path.dirname(os.path.abspath(db_path))
metrics_file = os.path.join(file_dir, f"{task_id}_metrics.json")
with open(metrics_file, "w") as f:
json.dump({"series": series_data}, f, indent=2)
except Exception as metrics_error:
# Log but don't fail backtest if metrics calculation fails.
# 附 traceback 以便定位(_metrics.json 不生成时这里是根因)。
logging.warning(
"相对指标计算失败(回测结果不受影响): %s\n%s",
metrics_error, traceback.format_exc(),
)
# Build equity curve DataFrame (S1.2): use the daily_df returned by
# calculate_result (index=date, has a 'balance' column). get_all_daily_results
# returns DailyResult objects (not dicts), so prefer daily_df.
if daily_df is not None and hasattr(daily_df, "empty") and not daily_df.empty:
if "balance" in daily_df.columns:
_bal = daily_df["balance"].astype(float)
elif "net_pnl" in daily_df.columns:
_bal = daily_df["net_pnl"].astype(float).cumsum() + capital
else:
_bal = None
equity_df = pd.DataFrame({
"date": daily_df.index.astype(str),
"balance": _bal.tolist(),
}) if _bal is not None else pd.DataFrame()
else:
equity_df = pd.DataFrame()
# Build trades DataFrame (S1.2): engine.trades is dict[vt_tradeid, TradeData].
trades_dict = engine.trades if isinstance(engine.trades, dict) else {}
trades_df = pd.DataFrame([
{
"datetime": str(t.datetime),
"direction": str(t.direction),
"offset": str(t.offset),
"price": t.price,
"volume": t.volume,
"vt_symbol": getattr(t, "vt_symbol", ""),
}
for t in trades_dict.values()
])
# Build result object — H7: degenerate 走 status="done" + statistics.degenerate flag
# (不发明新 status 值,前端 status map 只认 done/failed/running/pending
result = BacktestResult(
task_id=task_id,
type="cta",
status="done",
strategy=strategy_class.__name__,
symbol=symbol,
params=params,
start=start,
end=end,
statistics=statistics,
equity_curve=equity_df,
trades=trades_df,
)
except Exception as e:
# Handle any exceptions and return failed result
error_msg = f"{type(e).__name__}: {e}\n{traceback.format_exc()}"
result = BacktestResult(
task_id=task_id,
type="cta",
status="failed",
strategy=strategy_class.__name__,
symbol=symbol,
params=params,
start=start,
end=end,
statistics={},
equity_curve=None,
trades=None,
error_msg=error_msg
)
# Save result to database. file_dir = db dir so equity_curve/trades persist
# to parquet (S1.1) and reload via result.id.
save_result(result, db_path=db_path, file_dir=os.path.dirname(os.path.abspath(db_path)))
return result
# Module-level reference for mocking in tests
BacktestingEngine = None # Will be set when imported inside run_cta_backtest