1078 lines
41 KiB
Python
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()
|