Files
sanguo_vnpy_v2/sanguo_api/routes_portfolio.py
T

150 lines
6.3 KiB
Python

"""组合策略 API 路由(MVP)。
POST /portfolio/backtest: SSH 触发 VPS 跑 BulletTrade + all_weather,
捕获 stdout JSON 返回前端。同步模式(回测耗时,前端 loading,timeout 600s)。
风格参考 routes_paper.py / routes_live.py。
"""
from __future__ import annotations
import json
import logging
import os
import subprocess
import sys
from typing import Any, Optional
from fastapi import APIRouter, Depends, Header, HTTPException
from pydantic import BaseModel, Field
from .auth import verify_token as verify_token_impl
logger = logging.getLogger(__name__)
router = APIRouter()
# VPS SSH 连接配置(~/.ssh/config 已配 49.232.102.198 免密 key)
_VPS_HOST = "49.232.102.198"
_VPS_WORKDIR = r"C:\\sanguo_vnpy_v2"
_VPS_PYTHON = "python"
_VPS_TIMEOUT = 600 # 回测耗时,给 10 分钟
async def verify_token(authorization: str | None = Header(None)):
if authorization is None or not authorization.startswith("Bearer "):
raise HTTPException(401, "Missing/invalid authorization")
return verify_token_impl(authorization.split(" ", 1)[1])
class PortfolioBacktestRequest(BaseModel):
pool: str = Field(default="hs300_subset", description="标的池(占位,MVP 用默认)")
start_date: str = Field(default="2024-01-01", description="YYYY-MM-DD")
end_date: str = Field(default="2024-02-29", description="YYYY-MM-DD")
initial_cash: float = Field(default=1_000_000.0, description="初始资金(元)")
benchmark: str = Field(default="000300.XSHG", description="基准代码")
@router.post("/portfolio/backtest", dependencies=[Depends(verify_token)])
def run_portfolio_backtest(req: PortfolioBacktestRequest):
"""SSH 触发 VPS 跑 BulletTrade + all_weather 回测,同步返回 JSON 结果。
Windows SSH 坑:
- GBK 编码:python -X utf8 避免中文 print 编码崩
- 引号嵌套:用 list 形式 argv,避免 shell 引号 escape 噩梦
- 没有 tail/head/grep:用 python 后处理(本函数在 Mac 端直接解析 stdout)
- 中文路径:VPS_WORKDIR / userdata_mini 走 env(DEFAULT_DATA_PROVIDER=miniqmt)
"""
# 本地模式(VPS 后端直接跑 runner,避免 SSH 自连外网 IP 绕路);
# SSH 模式(Mac 后端 → VPS)。env SANGUO_PORTFOLIO_LOCAL=1 切本地。
# 自动检测三机自适应:
# - 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=local_cwd, capture_output=True, text=True,
timeout=_VPS_TIMEOUT, check=False,
)
except subprocess.TimeoutExpired:
raise HTTPException(504, f"回测超时(>{_VPS_TIMEOUT}s)")
except Exception as exc:
logger.exception("[portfolio] 本地 runner 调用失败")
raise HTTPException(500, f"本地 runner 调用失败: {exc}")
else:
# Mac 端 SSH 到 VPS(远端命令手动拼,不用 shlex.quote:产 POSIX 单引号 cmd 不认)
remote_cmd = (
f"cd {_VPS_WORKDIR} && "
f"{_VPS_PYTHON} -X utf8 -m sanguo_portfolio.runner_backtest --json "
f"--start {req.start_date} --end {req.end_date} "
f"--cash {req.initial_cash} --benchmark {req.benchmark} "
f"--max-pool 30"
)
ssh_argv = [
"ssh", "-o", "ConnectTimeout=15", "-o", "StrictHostKeyChecking=no",
_VPS_HOST, remote_cmd,
]
logger.info("[portfolio] SSH 触发: %s", remote_cmd)
try:
proc = subprocess.run(
ssh_argv, capture_output=True, text=True,
timeout=_VPS_TIMEOUT, check=False,
)
except subprocess.TimeoutExpired:
raise HTTPException(504, f"VPS 回测超时(>{_VPS_TIMEOUT}s)")
except FileNotFoundError:
raise HTTPException(500, "本机未找到 ssh 命令")
except Exception as exc:
logger.exception("[portfolio] SSH 调用失败")
raise HTTPException(500, f"SSH 调用失败: {exc}")
if proc.returncode != 0:
stderr_tail = (proc.stderr or "")[-2000:]
logger.error("[portfolio] VPS 回测失败 rc=%s stderr=%s", proc.returncode, stderr_tail)
raise HTTPException(
500,
f"VPS 回测失败(rc={proc.returncode}): {stderr_tail}",
)
# 从 stdout 提取最后一行 JSON(runner --json 只 print 一行)
stdout = proc.stdout or ""
import json as _json
result: Optional[dict[str, Any]] = None
parse_err: Optional[str] = None
for line in reversed(stdout.strip().splitlines()):
line = line.strip()
if not line.startswith("{"):
continue
try:
result = _json.loads(line)
break
except _json.JSONDecodeError as exc:
parse_err = str(exc)
continue
if result is None:
logger.error(
"[portfolio] stdout 无 JSON 行。parse_err=%s stdout_tail=%s",
parse_err, stdout[-2000:],
)
raise HTTPException(
500,
f"VPS stdout 解析失败: {parse_err or 'no JSON line'}; stdout_tail={stdout[-500:]!r}",
)
return result