Files

1078 lines
41 KiB
Python

"""
RaptorBT 统一策略入口
用法:
python -m app.main list # 列出所有可用策略
python -m app.main run --strategy sma_cross # 用默认参数跑策略
python -m app.main run --strategy sar_adx_cci \
--symbol XAUUSD --timeframe H1 --bars 500 # 指定数据
python -m app.main compare # 对比所有策略表现
环境变量:
MT5_BRIDGE_URL Mt5Bridge 地址 (默认 http://61.164.252.86:13485)
MT5_BRIDGE_KEY API Key (必需, 未设置时用内置默认 key)
"""
from __future__ import annotations
import argparse
import contextlib
import io
import json
import os
import sys
from datetime import datetime, timedelta, timezone
import numpy as np
import pandas as pd
import requests
import raptorbt
from strategies import get_strategies, get_strategy
from strategies.base import Strategy
# ============================================================================
# 配置
# ============================================================================
BRIDGE_URL = os.environ.get("MT5_BRIDGE_URL", "http://61.164.252.86:13485")
API_KEY = os.environ.get("MT5_BRIDGE_KEY", "UiHMqtaYLZzwBdcuS4RFmEGhgDO8N2eI")
OUTPUT_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "backtest_output")
# ============================================================================
# JSON 输出辅助
# ============================================================================
@contextlib.contextmanager
def _suppress_stdout(enabled: bool):
"""在 enabled=True 时, 临时吞掉所有 print 输出 (用于 --json 模式)"""
if not enabled:
yield
return
sink = io.StringIO()
old_stdout = sys.stdout
sys.stdout = sink
try:
yield sink
finally:
sys.stdout = old_stdout
def _clean_nan(obj):
"""递归把 NaN/Inf 转为 None, 确保 JSON 可序列化"""
if isinstance(obj, float):
if np.isnan(obj) or np.isinf(obj):
return None
return obj
if isinstance(obj, dict):
return {k: _clean_nan(v) for k, v in obj.items()}
if isinstance(obj, (list, tuple)):
return [_clean_nan(v) for v in obj]
return obj
def _emit_json(payload: dict):
"""输出 JSON 到 stdout (UTF-8, 紧凑, 不转义中文, 自动清理 NaN/Inf)"""
payload = _clean_nan(payload)
print(json.dumps(payload, ensure_ascii=False, allow_nan=False, default=str))
def _safe_metric(m) -> dict:
"""把 PyBacktestMetrics 转为可 JSON 序列化的 dict (处理 NaN)"""
d = m.to_dict()
out = {}
for k, v in d.items():
if isinstance(v, float) and np.isnan(v):
out[k] = None
else:
out[k] = v
# 补充 to_dict() 未包含的常用字段
for extra in ["total_trades", "winning_trades", "losing_trades",
"max_consecutive_wins", "max_consecutive_losses",
"avg_holding_period", "exposure_pct", "payoff_ratio",
"recovery_factor", "omega_ratio"]:
val = getattr(m, extra, None)
if val is not None and not (isinstance(val, float) and np.isnan(val)):
out[extra] = val
return out
# ============================================================================
# Mt5Bridge 数据加载
# ============================================================================
def _api_get(path: str, params=None):
resp = requests.get(
f"{BRIDGE_URL}{path}",
params=params,
headers={"X-API-Key": API_KEY},
timeout=15,
)
resp.raise_for_status()
return resp.json()
def check_health() -> bool:
"""健康检查,返回 MT5 是否已连接"""
data = _api_get("/health")
connected = data.get("mt5_connected", False)
status = data.get("status", "unknown")
print(f" Bridge: {status} MT5 连接: {'✓' if connected else '✗'}")
return connected
def fetch_klines(symbol: str, timeframe: str, bars: int,
source: str = "csv", data_file: str | None = None) -> pd.DataFrame:
"""
拉取 K 线数据, 返回标准 DataFrame
参数:
symbol: 品种代码 (如 XAUUSD)
timeframe: 周期 (M1/M5/M15/M30/H1/H4/D1)
bars: 返回的 K 线数量 (CSV 模式取最后 bars 根)
source: 数据源 "csv" (默认, 离线) 或 "mt5" (Mt5Bridge)
data_file: 显式指定 CSV 文件路径 (None 时按 symbol 自动查找)
"""
if source == "csv":
return _fetch_from_csv(symbol, timeframe, bars, data_file)
elif source == "mt5":
return _fetch_from_mt5(symbol, timeframe, bars)
else:
raise ValueError(f"未知数据源: {source}, 可选: csv / mt5")
def _fetch_from_csv(symbol: str, timeframe: str, bars: int,
data_file: str | None) -> pd.DataFrame:
"""从 CSV 加载数据 (策略研究用)"""
from .data_loader import load_csv, find_csv_for_symbol, compute_slippage
# 1. 定位 CSV 文件
if data_file is None:
data_file = find_csv_for_symbol(symbol)
if data_file is None or not os.path.exists(data_file):
raise FileNotFoundError(
f"未找到 {symbol} 的 CSV 文件。\n"
f"请把 CSV 放到 data/ 目录, 或用 --data-file 显式指定路径"
)
print(f" 数据源: CSV ({os.path.basename(data_file)})")
# 2. 加载 + 重采样到目标周期
df = load_csv(data_file, symbol=symbol, timeframe=timeframe)
# 3. 取最后 bars 根 (模拟最近行情)
if bars > 0 and len(df) > bars:
df = df.iloc[-bars:].reset_index(drop=True)
# 4. 显示数据范围 + 计算点差
from .data_loader import get_tick_size
avg_spread = float(df["spread"].mean()) if "spread" in df.columns else 0.0
tick_size = get_tick_size(symbol)
spread_cost = avg_spread * tick_size
slippage = compute_slippage(df["spread"], df["close"], symbol) if "spread" in df.columns else 0.0005
print(f" 数据范围: {df['time'].iloc[0]} ~ {df['time'].iloc[-1]}")
print(f" K 线数量: {len(df)} ({timeframe})")
if avg_spread > 0:
print(f" 平均点差: {avg_spread:.1f} 点 (≈{spread_cost:.5f}) → slippage={slippage:.5f}")
# 5. 缓存 slippage 到全局, 供 run_strategy 注入
global _cached_slippage
_cached_slippage = slippage
return df
def _fetch_from_mt5(symbol: str, timeframe: str, bars: int) -> pd.DataFrame:
"""从 Mt5Bridge 拉取数据 (最终验证用)"""
global _cached_slippage
_cached_slippage = None # MT5 模式不自动注入 slippage
date_to = datetime.now(timezone.utc).strftime("%Y-%m-%d")
date_from = (datetime.now(timezone.utc) - timedelta(days=bars // 24 + 30)).strftime("%Y-%m-%d")
data = _api_get("/rates/from-date", params={
"symbol": symbol,
"timeframe": f"TIMEFRAME_{timeframe}",
"date_from": date_from,
"date_to": date_to,
})
rows = data.get("data", [])
if not rows:
# fallback: 按偏移量拉取
data = _api_get("/rates/from-pos", params={
"symbol": symbol,
"timeframe": f"TIMEFRAME_{timeframe}",
"start_pos": 0,
"count": bars,
})
rows = data.get("data", [])
if not rows:
raise RuntimeError(f"无法拉取 {symbol} K 线数据, 检查品种名或 MT5 连接")
df = pd.DataFrame(rows)
df["time"] = pd.to_datetime(df["time"])
df = df.sort_values("time").reset_index(drop=True)
print(f" 数据源: Mt5Bridge ({symbol} {timeframe})")
print(f" K 线数量: {len(df)}")
return df
# CSV 模式下缓存的 slippage (供 run_strategy 自动注入)
_cached_slippage: float | None = None
# ============================================================================
# 回测执行
# ============================================================================
def run_strategy(strategy_name: str, df: pd.DataFrame, symbol: str) -> "raptorbt.PyBacktestResult":
"""实例化策略并执行回测"""
strategy = get_strategy(strategy_name)
print(f"\n{'═' * 60}")
print(f"策略: {strategy.name}")
print(f"描述: {strategy.description()}")
print(f"预热期: {strategy.warmup_bars()} bars")
print(f"{'═' * 60}")
arr = strategy.to_arrays(df)
signals = strategy.generate_signals(df)
n_entries = int(signals.entries.sum())
n_exits = int(signals.exits.sum())
print(f"入场信号: {n_entries} 出场信号: {n_exits} 方向: {'多' if signals.direction == 1 else '空'}")
if n_entries == 0:
print(" ⚠️ 无入场信号, 跳过回测")
return None
config = strategy.build_config()
# CSV 模式下自动注入 spread→slippage
if _cached_slippage is not None:
config.slippage = _cached_slippage
print(f" 自动注入 slippage: {_cached_slippage:.5f} (来自 CSV 平均点差)")
result = raptorbt.run_single_backtest(
timestamps=arr["timestamps"],
open=arr["open"], high=arr["high"], low=arr["low"], close=arr["close"],
volume=arr["volume"],
entries=signals.entries, exits=signals.exits,
direction=signals.direction, weight=1.0, symbol=symbol,
config=config,
)
m = result.metrics
print(f"\n {'─' * 40}")
print(f" {'指标':<16}{'值':>20}")
print(f" {'─' * 40}")
print(f" {'总收益率':<16}{m.total_return_pct:>19.2f} %")
print(f" {'夏普比率':<16}{m.sharpe_ratio:>20.2f}")
print(f" {'索提诺比率':<16}{m.sortino_ratio:>20.2f}")
print(f" {'最大回撤':<16}{m.max_drawdown_pct:>19.2f} %")
print(f" {'总交易数':<16}{m.total_trades:>20d}")
print(f" {'胜率':<16}{m.win_rate_pct:>19.1f} %")
print(f" {'盈利因子':<16}{m.profit_factor:>20.2f}")
print(f" {'期望值':<16}{m.expectancy:>20.2f}")
print(f" {'市场暴露':<16}{m.exposure_pct:>19.1f} %")
print(f" {'─' * 40}")
# 出场原因分布
trades = result.trades()
if trades:
exit_reasons = {}
for t in trades:
r = t.exit_reason
exit_reasons[r] = exit_reasons.get(r, 0) + 1
print(f" 出场原因: {exit_reasons}")
return result
def export_result(result, df: pd.DataFrame, strategy_name: str):
"""导出交易/曲线/指标到 CSV"""
if result is None:
return
os.makedirs(OUTPUT_DIR, exist_ok=True)
# 交易记录
trades = result.trades()
if trades:
rows = [{
"trade_id": t.id, "symbol": t.symbol,
"direction": "Long" if t.direction == 1 else "Short",
"entry_idx": t.entry_idx, "exit_idx": t.exit_idx,
"entry_time": df["time"].iloc[t.entry_idx] if t.entry_idx < len(df) else "",
"exit_time": df["time"].iloc[t.exit_idx] if t.exit_idx < len(df) else "",
"entry_price": t.entry_price, "exit_price": t.exit_price,
"size": t.size, "pnl": t.pnl, "return_pct": t.return_pct,
"fees": t.fees, "exit_reason": t.exit_reason,
} for t in trades]
pd.DataFrame(rows).to_csv(
os.path.join(OUTPUT_DIR, f"{strategy_name}_trades.csv"),
index=False, encoding="utf-8-sig",
)
# 曲线
equity = result.equity_curve()
pd.DataFrame({
"time": df["time"].values[:len(equity)],
"equity": equity,
"drawdown": result.drawdown_curve(),
"returns": result.returns(),
}).to_csv(
os.path.join(OUTPUT_DIR, f"{strategy_name}_curves.csv"),
index=False, encoding="utf-8-sig",
)
# 指标
m = result.metrics
d = m.to_dict()
d.update(
total_trades=m.total_trades, winning_trades=m.winning_trades,
losing_trades=m.losing_trades, max_consecutive_wins=m.max_consecutive_wins,
max_consecutive_losses=m.max_consecutive_losses, avg_holding_period=m.avg_holding_period,
exposure_pct=m.exposure_pct, payoff_ratio=m.payoff_ratio,
recovery_factor=m.recovery_factor, omega_ratio=m.omega_ratio,
)
pd.DataFrame(list(d.items()), columns=["metric", "value"]).to_csv(
os.path.join(OUTPUT_DIR, f"{strategy_name}_metrics.csv"),
index=False, encoding="utf-8-sig",
)
print(f" → 结果已导出至 {OUTPUT_DIR}/{strategy_name}_*.csv")
# ============================================================================
# CLI 命令
# ============================================================================
def cmd_list(args):
# 模式: --indicators 显示指标目录
if args.indicators:
from .indicator_catalog import format_text, format_json
if args.json:
_emit_json(format_json())
else:
print(format_text())
return
strategies = get_strategies()
print(f"\n可用策略 ({len(strategies)} 个):")
print(f"{'═' * 60}")
for name, cls in sorted(strategies.items()):
inst = cls()
print(f" {name:<22} {inst.description()}")
print(f"{'═' * 60}")
print(f"使用: python -m app.main run --strategy <名称>")
# 显示可用 CSV 数据
try:
from .data_loader import list_available_symbols
symbols = list_available_symbols()
if symbols:
print(f"\n可用 CSV 数据 ({len(symbols)} 个):")
print(f"{'─' * 60}")
for sym, fname in symbols:
print(f" {sym:<10} {fname}")
print(f"{'─' * 60}")
except Exception:
pass
def cmd_run(args):
with _suppress_stdout(args.json):
print(f"{'═' * 60}")
print(f"RaptorBT 策略回测")
print(f"{'═' * 60}")
if args.source == "mt5":
check_health()
df = fetch_klines(args.symbol, args.timeframe, args.bars,
source=args.source, data_file=args.data_file)
print(f" 数据: {args.symbol} {args.timeframe} {len(df)} 根 K 线")
print(f" 范围: {df['time'].iloc[0]} ~ {df['time'].iloc[-1]}")
print(f" Close: {df['close'].min():.2f} ~ {df['close'].max():.2f}")
result = run_strategy(args.strategy, df, args.symbol)
if args.export and result is not None:
export_result(result, df, args.strategy)
if args.json:
payload = {
"command": "run",
"strategy": args.strategy,
"symbol": args.symbol,
"timeframe": args.timeframe,
"bars": len(df) if df is not None else 0,
}
if result is not None:
payload["metrics"] = _safe_metric(result.metrics)
trades = result.trades()
payload["n_trades"] = len(trades)
payload["trades"] = [
{
"id": t.id, "symbol": t.symbol,
"direction": "Long" if t.direction == 1 else "Short",
"entry_idx": t.entry_idx, "exit_idx": t.exit_idx,
"entry_price": t.entry_price, "exit_price": t.exit_price,
"size": t.size, "pnl": t.pnl,
"return_pct": t.return_pct, "fees": t.fees,
"exit_reason": t.exit_reason,
}
for t in trades[:50] # 限制前 50 条, 避免超大输出
]
else:
payload["metrics"] = None
payload["n_trades"] = 0
_emit_json(payload)
def cmd_compare(args):
print(f"{'═' * 60}")
print(f"策略对比 (全量)")
print(f"{'═' * 60}")
if args.source == "mt5":
check_health()
df = fetch_klines(args.symbol, args.timeframe, args.bars,
source=args.source, data_file=args.data_file)
print(f" 数据: {args.symbol} {args.timeframe} {len(df)} 根 K 线\n")
strategies = get_strategies()
print(f"{'策略':<22} {'收益%':>8} {'夏普':>7} {'回撤%':>8} {'交易':>5} {'胜率%':>7} {'PF':>6}")
print(f"{'─' * 70}")
for name in sorted(strategies.keys()):
try:
result = run_strategy(name, df, args.symbol)
if result is None:
print(f" {name:<22} (无信号)")
continue
m = result.metrics
print(
f" {name:<22} {m.total_return_pct:>7.2f} {m.sharpe_ratio:>7.2f} "
f"{m.max_drawdown_pct:>7.2f} {m.total_trades:>5d} "
f"{m.win_rate_pct:>6.1f} {m.profit_factor:>6.2f}"
)
if args.export:
export_result(result, df, name)
except Exception as e:
print(f" {name:<22}{e}")
def parse_param_grid(param_args):
"""解析 --param fast=5,10,15 → {"fast": [5, 10, 15]}"""
grid = {}
for arg in param_args:
if "=" not in arg:
continue
k, v = arg.split("=", 1)
values = []
for item in v.split(","):
item = item.strip()
try:
values.append(int(item))
except ValueError:
try:
values.append(float(item))
except ValueError:
values.append(item)
grid[k] = values
return grid
def cmd_optimize(args):
result = None
with _suppress_stdout(args.json):
print(f"{'═' * 60}")
print(f"策略参数优化: {args.strategy}")
print(f"{'═' * 60}")
if args.source == "mt5":
check_health()
df = fetch_klines(args.symbol, args.timeframe, args.bars,
source=args.source, data_file=args.data_file)
print(f" 数据: {args.symbol} {args.timeframe} {len(df)} 根 K 线")
param_grid = parse_param_grid(args.param)
if not param_grid:
print(" ❌ 未指定参数空间, 用 --param name=v1,v2,v3")
return
print(f" 参数空间: {param_grid}")
print(f" 目标指标: {args.metric}")
from .optimizer import StrategyOptimizer
from strategies import get_strategies
strategies = get_strategies()
if args.strategy not in strategies:
print(f" ❌ 未知策略: {args.strategy}")
return
opt = StrategyOptimizer(metric=args.metric)
result = opt.optimize(
strategy_class=strategies[args.strategy],
df=df,
param_grid=param_grid,
symbol=args.symbol,
)
print(f"\n{result.summary()}")
print(f"\nTop 10 参数组合:")
print(result.top_n(10).to_string(index=False))
if args.export:
out = os.path.join(OUTPUT_DIR, f"{args.strategy}_optimization.csv")
result.export(out)
print(f"\n → 结果已导出: {out}")
if args.json:
payload = {
"command": "optimize",
"strategy": args.strategy,
"symbol": args.symbol,
"timeframe": args.timeframe,
"metric": args.metric,
}
if result is not None:
payload.update(result.to_dict())
else:
payload["error"] = "未指定参数空间或策略未知"
_emit_json(payload)
def cmd_walkforward(args):
result = None
with _suppress_stdout(args.json):
print(f"{'═' * 60}")
print(f"Walk-Forward 验证: {args.strategy}")
print(f"{'═' * 60}")
if args.source == "mt5":
check_health()
df = fetch_klines(args.symbol, args.timeframe, args.bars,
source=args.source, data_file=args.data_file)
print(f" 数据: {args.symbol} {args.timeframe} {len(df)} 根 K 线")
param_grid = parse_param_grid(args.param)
if not param_grid:
print(" ❌ 未指定参数空间")
return
from .walk_forward import WalkForwardValidator
from strategies import get_strategies
strategies = get_strategies()
if args.strategy not in strategies:
print(f" ❌ 未知策略: {args.strategy}")
return
wf = WalkForwardValidator(train_size=args.train_size, test_size=args.test_size)
result = wf.validate(
strategy_class=strategies[args.strategy],
df=df,
param_grid=param_grid,
metric=args.metric,
symbol=args.symbol,
)
print(f"\n{result.summary()}")
if args.export:
out = os.path.join(OUTPUT_DIR, f"{args.strategy}_walkforward.csv")
result.export(out)
print(f"\n → 结果已导出: {out}")
if args.json:
payload = {
"command": "walkforward",
"strategy": args.strategy,
"symbol": args.symbol,
"timeframe": args.timeframe,
"train_size": args.train_size,
"test_size": args.test_size,
}
if result is not None:
payload.update(result.to_dict())
else:
payload["error"] = "未指定参数空间或策略未知"
_emit_json(payload)
def cmd_validate(args):
wf_result = None
accept_report = None
lookahead_report = None
with _suppress_stdout(args.json):
print(f"{'═' * 60}")
print(f"策略验收: {args.strategy}")
print(f"{'═' * 60}")
if args.source == "mt5":
check_health()
# 前置: 前视偏差检测 (静态 + 动态)
from .lookahead_check import full_check
from strategies import get_strategies as _get_strategies
_strats = _get_strategies()
if args.strategy not in _strats:
print(f" ❌ 未知策略: {args.strategy}")
return
_strat_cls = _strats[args.strategy]
_strat_file = f"strategies/{args.strategy}.py"
print(f"\n 前视偏差检测...")
_lookahead = full_check(_strat_cls, _strat_file, df=None, run_dynamic=False)
print(f" {_lookahead.summary()}")
lookahead_report = _lookahead
if not _lookahead.passed:
print(f"\n ❌ 检测到前视偏差, 终止验收")
return
df = fetch_klines(args.symbol, args.timeframe, args.bars,
source=args.source, data_file=args.data_file)
print(f" 数据: {args.symbol} {args.timeframe} {len(df)} 根 K 线")
param_grid = parse_param_grid(args.param)
if not param_grid:
print(" ❌ 未指定参数空间")
return
from .walk_forward import WalkForwardValidator
from .acceptance import StrategyAcceptance
strategies = get_strategies()
if args.strategy not in strategies:
print(f" ❌ 未知策略: {args.strategy}")
return
wf = WalkForwardValidator(train_size=args.train_size, test_size=args.test_size)
wf_result = wf.validate(
strategy_class=strategies[args.strategy],
df=df,
param_grid=param_grid,
metric=args.metric,
symbol=args.symbol,
)
print(f"\n{wf_result.summary()}")
checker = StrategyAcceptance()
accept_report = checker.check(wf_result)
print(f"\n{accept_report.summary()}")
if accept_report.passed:
print("\n ✅ 策略通过验收, 可交付")
else:
print("\n ❌ 策略未通过验收, 需进一步优化")
if args.json:
payload = {
"command": "validate",
"strategy": args.strategy,
"passed": accept_report.passed if accept_report else False,
}
if lookahead_report is not None:
payload["lookahead"] = lookahead_report.to_dict()
if not lookahead_report.passed:
payload["error"] = "前视偏差检测未通过, 验收终止"
if wf_result is not None:
payload["walk_forward"] = wf_result.to_dict()
if accept_report is not None:
payload["acceptance"] = accept_report.to_dict()
_emit_json(payload)
def cmd_scaffold(args):
from .scaffold import scaffold_strategy
try:
path = scaffold_strategy(
name=args.name,
template=args.template,
description=args.description or "",
overwrite=args.overwrite,
)
print(f"✅ 策略模板已生成: {path}")
print(f" 模板类型: {args.template}")
print(f" 接下来编辑该文件, 填入信号生成逻辑")
print(f" 完成后用 'python -m app.main list' 查看是否自动注册")
print(f" 编辑后建议用 'python -m app.main check {args.name}' 检测前视偏差")
except FileExistsError as e:
print(f"❌ {e}")
except ValueError as e:
print(f"❌ {e}")
def cmd_check(args):
"""前视偏差检测: 静态扫描源码 + 可选动态验证"""
# 模式 1: 检测所有策略
if args.strategy == "all":
with _suppress_stdout(args.json):
print(f"{'═' * 60}")
print(f"前视偏差检测: all")
print(f"{'═' * 60}")
from .lookahead_check import check_all_strategies
print(f"\n扫描 strategies/ 目录下所有策略...\n")
results = check_all_strategies("strategies")
n_pass = 0
n_fail = 0
for path, report in results:
print(f"{'─' * 60}")
print(f"文件: {os.path.basename(path)}")
print(report.summary())
if report.passed:
n_pass += 1
else:
n_fail += 1
print(f"\n{'─' * 60}")
print(f"总结: {n_pass} 通过, {n_fail} 未通过")
if args.json:
payload = {
"command": "check",
"mode": "all",
"n_pass": n_pass,
"n_fail": n_fail,
"strategies": [
{"file": os.path.basename(p), **r.to_dict()}
for p, r in results
],
}
_emit_json(payload)
return
# 模式 2: 检测单个策略
from strategies import get_strategies
strategies = get_strategies()
report = None
with _suppress_stdout(args.json):
print(f"{'═' * 60}")
print(f"前视偏差检测: {args.strategy}")
print(f"{'═' * 60}")
if args.strategy not in strategies:
print(f" ❌ 未知策略: {args.strategy}")
print(f" 可用策略: {', '.join(strategies.keys())}")
return
from .lookahead_check import full_check
strat_cls = strategies[args.strategy]
strat_file = f"strategies/{args.strategy}.py"
# 动态验证需要数据
df = None
if args.dynamic:
try:
df = fetch_klines(args.symbol, args.timeframe, args.bars,
source=args.source, data_file=args.data_file)
except Exception as e:
print(f" ⚠️ 无法加载数据, 跳过动态检测: {e}")
args.dynamic = False
report = full_check(
strat_cls, strat_file, df=df,
run_dynamic=args.dynamic and df is not None,
)
print(f"\n{report.summary()}")
if report.passed:
print(f"\n ✅ 策略无前视偏差, 可安全使用")
else:
print(f"\n ❌ 策略存在前视偏差, 需修复")
print(f" 修复建议:")
print(f" - 禁用 .shift(-N) (N>0, 访问未来 bar)")
print(f" - 禁用 close/high/low/open 的负索引 (如 close[-1])")
print(f" - 禁用 iloc/loc 切片到未来索引")
print(f" - 信号只用当前 bar 及之前的数据生成")
print(f" - 用 cross_above/cross_below (已内置前视安全)")
if args.json:
payload = {
"command": "check",
"mode": "single",
"strategy": args.strategy,
}
if report is not None:
payload.update(report.to_dict())
if not report.passed:
payload["fix_suggestions"] = [
"禁用 .shift(-N) (N>0, 访问未来 bar)",
"禁用 close/high/low/open 的负索引 (如 close[-1])",
"禁用 iloc/loc 切片到未来索引",
"信号只用当前 bar 及之前的数据生成",
"用 cross_above/cross_below (已内置前视安全)",
]
else:
payload["error"] = f"未知策略: {args.strategy}"
payload["available"] = list(strategies.keys())
_emit_json(payload)
def cmd_deliver(args):
pkg_path = None
lookahead_report = None
opt_result = None
wf_result = None
accept_report = None
error = None
with _suppress_stdout(args.json):
print(f"{'═' * 60}")
print(f"策略交付包生成: {args.strategy}")
print(f"{'═' * 60}")
if args.source == "mt5":
check_health()
# 前置: 前视偏差检测 (静态 + 动态), 不通过则拒绝交付
from .lookahead_check import full_check
from strategies import get_strategies as _get_strategies
_strats = _get_strategies()
if args.strategy not in _strats:
print(f" ❌ 未知策略: {args.strategy}")
error = f"未知策略: {args.strategy}"
return
_strat_cls = _strats[args.strategy]
_strat_file = f"strategies/{args.strategy}.py"
print(f"\n 前视偏差检测 (交付前强制)...")
_lookahead = full_check(_strat_cls, _strat_file, df=None, run_dynamic=False)
print(f" {_lookahead.summary()}")
lookahead_report = _lookahead
if not _lookahead.passed:
print(f"\n ❌ 检测到前视偏差, 拒绝生成交付包")
print(f" 请修复前视问题后再交付 (用 'python -m app.main check {args.strategy}' 查看详情)")
error = "前视偏差检测未通过, 拒绝生成交付包"
return
df = fetch_klines(args.symbol, args.timeframe, args.bars,
source=args.source, data_file=args.data_file)
print(f" 数据: {args.symbol} {args.timeframe} {len(df)} 根 K 线")
param_grid = parse_param_grid(args.param)
if not param_grid:
print(" ❌ 未指定参数空间")
error = "未指定参数空间"
return
from .optimizer import StrategyOptimizer
from .walk_forward import WalkForwardValidator
from .acceptance import StrategyAcceptance
from .exporter import StrategyExporter
strategies = get_strategies()
if args.strategy not in strategies:
print(f" ❌ 未知策略: {args.strategy}")
error = f"未知策略: {args.strategy}"
return
print(f"\n Step 1/3: 参数优化...")
opt = StrategyOptimizer(metric=args.metric)
opt_result = opt.optimize(
strategy_class=strategies[args.strategy],
df=df,
param_grid=param_grid,
symbol=args.symbol,
)
print(f" {opt_result.summary()}")
print(f"\n Step 2/3: Walk-Forward 验证...")
wf = WalkForwardValidator(train_size=args.train_size, test_size=args.test_size)
wf_result = wf.validate(
strategy_class=strategies[args.strategy],
df=df,
param_grid=param_grid,
metric=args.metric,
symbol=args.symbol,
)
print(f"\n{wf_result.summary()}")
print(f"\n Step 3/3: 验收检查...")
checker = StrategyAcceptance()
accept_report = checker.check(wf_result)
print(f"\n{accept_report.summary()}")
if not accept_report.passed:
print(f"\n ⚠️ 策略未通过验收, 仍可生成交付包 (含未通过标记)")
print(f"\n 生成交付包...")
exporter = StrategyExporter()
pkg_path = exporter.deliver(
strategy_name=args.strategy,
df=df,
symbol=args.symbol,
opt_result=opt_result,
wf_result=wf_result,
accept_report=accept_report,
param_grid=param_grid,
data_info={
"source": args.source,
"range": f"{df['time'].iloc[0]} ~ {df['time'].iloc[-1]}",
"bars": len(df),
"timeframe": args.timeframe,
},
)
print(f"\n ✅ 交付包已生成: {pkg_path}")
print(f" 包含: 策略源文件 + Markdown 报告 + CSV 明细")
if args.json:
payload = {
"command": "deliver",
"strategy": args.strategy,
"delivered": pkg_path is not None,
"package_path": pkg_path,
}
if error:
payload["error"] = error
if lookahead_report is not None:
payload["lookahead"] = lookahead_report.to_dict()
if opt_result is not None:
payload["optimization"] = opt_result.to_dict()
if wf_result is not None:
payload["walk_forward"] = wf_result.to_dict()
if accept_report is not None:
payload["acceptance"] = accept_report.to_dict()
_emit_json(payload)
def main():
parser = argparse.ArgumentParser(
description="RaptorBT 统一策略入口",
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog="""
示例:
python -m app.main list
python -m app.main run --strategy sma_cross
python -m app.main optimize --strategy sma_cross --param fast=5,10,15 --param slow=20,30
python -m app.main walkforward --strategy sma_cross --param fast=5,10 --param slow=20,30
python -m app.main validate --strategy sma_cross --param fast=5,10 --param slow=20,30
python -m app.main scaffold --name my_rsi --template mean_reversion
python -m app.main deliver --strategy sma_cross --param fast=5,10 --param slow=20,30
""",
)
sub = parser.add_subparsers(dest="command", required=True)
# list
p_list = sub.add_parser("list", help="列出所有可用策略 / 可用指标")
p_list.add_argument("--indicators", action="store_true",
help="列出所有可用指标及签名 (供 AI agent 查询)")
p_list.add_argument("--json", action="store_true", help="输出 JSON (供 AI agent 解析)")
p_list.set_defaults(func=cmd_list)
# run
p_run = sub.add_parser("run", help="运行单个策略")
p_run.add_argument("--strategy", required=True, help="策略名称 (见 list)")
p_run.add_argument("--symbol", default="XAUUSD", help="品种 (默认 XAUUSD)")
p_run.add_argument("--timeframe", default="H1", help="周期 (默认 H1)")
p_run.add_argument("--bars", type=int, default=500, help="K 线数量 (默认 500)")
p_run.add_argument("--source", default="csv", choices=["csv", "mt5"], help="数据源 (默认 csv 离线)")
p_run.add_argument("--data-file", default=None, help="CSV 文件路径 (默认按 symbol 自动查找)")
p_run.add_argument("--export", action="store_true", help="导出 CSV 结果")
p_run.add_argument("--json", action="store_true", help="输出 JSON (供 AI agent 解析)")
p_run.set_defaults(func=cmd_run)
# compare
p_cmp = sub.add_parser("compare", help="对比所有策略表现")
p_cmp.add_argument("--symbol", default="XAUUSD", help="品种 (默认 XAUUSD)")
p_cmp.add_argument("--timeframe", default="H1", help="周期 (默认 H1)")
p_cmp.add_argument("--bars", type=int, default=500, help="K 线数量 (默认 500)")
p_cmp.add_argument("--source", default="csv", choices=["csv", "mt5"], help="数据源 (默认 csv 离线)")
p_cmp.add_argument("--data-file", default=None, help="CSV 文件路径")
p_cmp.add_argument("--export", action="store_true", help="导出 CSV 结果")
p_cmp.set_defaults(func=cmd_compare)
# optimize
p_opt = sub.add_parser("optimize", help="参数网格搜索优化")
p_opt.add_argument("--strategy", required=True, help="策略名称")
p_opt.add_argument("--param", action="append", required=True,
help="参数空间, 格式: name=v1,v2,v3 (可多次指定)")
p_opt.add_argument("--symbol", default="XAUUSD", help="品种")
p_opt.add_argument("--timeframe", default="H1", help="周期")
p_opt.add_argument("--bars", type=int, default=500, help="K 线数量")
p_opt.add_argument("--source", default="csv", choices=["csv", "mt5"], help="数据源 (默认 csv 离线)")
p_opt.add_argument("--data-file", default=None, help="CSV 文件路径")
p_opt.add_argument("--metric", default="sharpe_ratio", help="优化目标指标")
p_opt.add_argument("--export", action="store_true", help="导出 CSV")
p_opt.add_argument("--json", action="store_true", help="输出 JSON (供 AI agent 解析)")
p_opt.set_defaults(func=cmd_optimize)
# walkforward
p_wf = sub.add_parser("walkforward", help="Walk-Forward 验证")
p_wf.add_argument("--strategy", required=True, help="策略名称")
p_wf.add_argument("--param", action="append", required=True,
help="参数空间, 格式: name=v1,v2,v3")
p_wf.add_argument("--symbol", default="XAUUSD", help="品种")
p_wf.add_argument("--timeframe", default="H1", help="周期")
p_wf.add_argument("--bars", type=int, default=1000, help="K 线数量 (需要足够多)")
p_wf.add_argument("--source", default="csv", choices=["csv", "mt5"], help="数据源 (默认 csv 离线)")
p_wf.add_argument("--data-file", default=None, help="CSV 文件路径")
p_wf.add_argument("--train-size", type=int, default=300, help="训练窗口大小")
p_wf.add_argument("--test-size", type=int, default=100, help="测试窗口大小")
p_wf.add_argument("--metric", default="sharpe_ratio", help="优化目标指标")
p_wf.add_argument("--export", action="store_true", help="导出 CSV")
p_wf.add_argument("--json", action="store_true", help="输出 JSON (供 AI agent 解析)")
p_wf.set_defaults(func=cmd_walkforward)
# validate
p_val = sub.add_parser("validate", help="策略验收检查 (walk-forward + 验收标准)")
p_val.add_argument("--strategy", required=True, help="策略名称")
p_val.add_argument("--param", action="append", required=True,
help="参数空间, 格式: name=v1,v2,v3")
p_val.add_argument("--symbol", default="XAUUSD", help="品种")
p_val.add_argument("--timeframe", default="H1", help="周期")
p_val.add_argument("--bars", type=int, default=1000, help="K 线数量")
p_val.add_argument("--source", default="csv", choices=["csv", "mt5"], help="数据源 (默认 csv 离线)")
p_val.add_argument("--data-file", default=None, help="CSV 文件路径")
p_val.add_argument("--train-size", type=int, default=300, help="训练窗口大小")
p_val.add_argument("--test-size", type=int, default=100, help="测试窗口大小")
p_val.add_argument("--metric", default="sharpe_ratio", help="优化目标指标")
p_val.add_argument("--json", action="store_true", help="输出 JSON (供 AI agent 解析)")
p_val.set_defaults(func=cmd_validate)
# scaffold
p_scf = sub.add_parser("scaffold", help="生成策略模板文件")
p_scf.add_argument("--name", required=True, help="策略名称 (snake_case)")
p_scf.add_argument("--template", default="custom",
choices=["crossover", "mean_reversion", "trend_following", "breakout", "custom"],
help="模板类型 (默认 custom)")
p_scf.add_argument("--description", default="", help="策略描述")
p_scf.add_argument("--overwrite", action="store_true", help="覆盖已存在文件")
p_scf.set_defaults(func=cmd_scaffold)
# check — 前视偏差检测
p_chk = sub.add_parser("check", help="前视偏差检测 (静态扫描 + 可选动态验证)")
p_chk.add_argument("--strategy", required=True,
help="策略名称, 或 'all' 检测所有策略")
p_chk.add_argument("--dynamic", action="store_true",
help="启用动态验证 (修改未来 bar, 检查历史信号是否变化)")
p_chk.add_argument("--symbol", default="XAUUSD", help="品种 (动态验证用)")
p_chk.add_argument("--timeframe", default="H1", help="周期 (动态验证用)")
p_chk.add_argument("--bars", type=int, default=500, help="K 线数量 (动态验证用)")
p_chk.add_argument("--source", default="csv", choices=["csv", "mt5"], help="数据源")
p_chk.add_argument("--data-file", default=None, help="CSV 文件路径")
p_chk.add_argument("--json", action="store_true", help="输出 JSON (供 AI agent 解析)")
p_chk.set_defaults(func=cmd_check)
# deliver
p_dlv = sub.add_parser("deliver", help="生成策略交付包 (优化+验证+验收+报告)")
p_dlv.add_argument("--strategy", required=True, help="策略名称")
p_dlv.add_argument("--param", action="append", required=True,
help="参数空间, 格式: name=v1,v2,v3")
p_dlv.add_argument("--symbol", default="XAUUSD", help="品种")
p_dlv.add_argument("--timeframe", default="H1", help="周期")
p_dlv.add_argument("--bars", type=int, default=1000, help="K 线数量")
p_dlv.add_argument("--source", default="csv", choices=["csv", "mt5"], help="数据源 (默认 csv 离线)")
p_dlv.add_argument("--data-file", default=None, help="CSV 文件路径")
p_dlv.add_argument("--train-size", type=int, default=300, help="训练窗口大小")
p_dlv.add_argument("--test-size", type=int, default=100, help="测试窗口大小")
p_dlv.add_argument("--metric", default="sharpe_ratio", help="优化目标指标")
p_dlv.add_argument("--json", action="store_true", help="输出 JSON (供 AI agent 解析)")
p_dlv.set_defaults(func=cmd_deliver)
args = parser.parse_args()
args.func(args)
if __name__ == "__main__":
main()