fix: NAS组合回测路由适配(isdir /app)+runner provider_config透传
This commit is contained in:
@@ -7,6 +7,7 @@ POST /portfolio/backtest: SSH 触发 VPS 跑 BulletTrade + all_weather,
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import subprocess
|
||||
@@ -54,18 +55,29 @@ def run_portfolio_backtest(req: PortfolioBacktestRequest):
|
||||
"""
|
||||
# 本地模式(VPS 后端直接跑 runner,避免 SSH 自连外网 IP 绕路);
|
||||
# SSH 模式(Mac 后端 → VPS)。env SANGUO_PORTFOLIO_LOCAL=1 切本地。
|
||||
# 自动检测:VPS 上 _VPS_WORKDIR(C:\sanguo_vnpy_v2) 存在 → 本地直接跑 runner;
|
||||
# Mac 上无该目录 → SSH 到 VPS。无需 env 配置,同源代码两端自适应。
|
||||
if os.path.isdir(_VPS_WORKDIR):
|
||||
# 自动检测三机自适应:
|
||||
# - VPS(Windows): _VPS_WORKDIR(C:\sanguo_vnpy_v2) 存在 → 本地直接跑 runner(默认 provider,行为不变)
|
||||
# - NAS(Linux 容器): /app 存在(docker-compose 代码挂载目录) → 本地跑 runner --provider unified 读 NAS 权威数据层
|
||||
# - Mac(dev): 都不满足 → SSH 到 VPS(原行为)
|
||||
if os.path.isdir(_VPS_WORKDIR) or os.path.isdir("/app"):
|
||||
local_cwd = _VPS_WORKDIR if os.path.isdir(_VPS_WORKDIR) else "/app"
|
||||
argv = [
|
||||
sys.executable, "-X", "utf8", "-m", "sanguo_portfolio.runner_backtest",
|
||||
"--json", "--start", req.start_date, "--end", req.end_date,
|
||||
"--cash", str(req.initial_cash), "--benchmark", req.benchmark, "--max-pool", "30",
|
||||
]
|
||||
# NAS 容器分支: unified provider 读 NAS 权威数据层(dbbardata + parquet);
|
||||
# VPS 本地分支保持原状(默认 provider,cwd=_VPS_WORKDIR,行为零变化)
|
||||
if local_cwd == "/app":
|
||||
nas_provider_config = json.dumps({
|
||||
"db_path": "/volume1/stock/sanguo_vnpy_v2/data_backup/quant_trading.db",
|
||||
"data_dir": "/volume1/stock/sanguo_vnpy_v2/data",
|
||||
})
|
||||
argv += ["--provider", "unified", "--provider-config", nas_provider_config]
|
||||
logger.info("[portfolio] 本地跑 runner: %s", " ".join(argv[3:]))
|
||||
try:
|
||||
proc = subprocess.run(
|
||||
argv, cwd=_VPS_WORKDIR, capture_output=True, text=True,
|
||||
argv, cwd=local_cwd, capture_output=True, text=True,
|
||||
timeout=_VPS_TIMEOUT, check=False,
|
||||
)
|
||||
except subprocess.TimeoutExpired:
|
||||
|
||||
@@ -344,7 +344,7 @@ def run_backtest_json(params: Dict[str, Any]) -> Dict[str, Any]:
|
||||
frequency="day",
|
||||
strategy=strategy_name,
|
||||
provider=params.get("provider", "local"),
|
||||
provider_config="{}",
|
||||
provider_config=params.get("provider_config", "{}"),
|
||||
result_file="", # JSON 模式不写 md
|
||||
max_pool=int(params.get("max_pool", 0)),
|
||||
)
|
||||
@@ -503,6 +503,7 @@ def main() -> None:
|
||||
"initial_cash": args.cash,
|
||||
"benchmark": args.benchmark,
|
||||
"provider": args.provider,
|
||||
"provider_config": args.provider_config,
|
||||
"max_pool": args.max_pool,
|
||||
})
|
||||
print(json.dumps(result, ensure_ascii=False, default=str))
|
||||
|
||||
Reference in New Issue
Block a user