Files
sanguo_vnpy_v2/sanguo_api/routes_portfolio.py
T
claude_dev 6c2c217e22
CI/CD / test (push) Successful in 12s
CI/CD / nas-deploy (push) Successful in 42s
CI/CD / nas-verify (push) Successful in 9s
fix(web+api): 组合回测max_pool默认30→0——用户拍板「前端页面肯定要改」(回测口径与实盘对齐,f214e2f只翻了live/paper链漏了回测) [nas]
你印象中改过的是6ecdaab#95标的池pool→all(三表单+后端四处),max_pool当时没动。本次三处对齐: PortfolioBacktest表单默认30→0+过时提示(「MVP验证用,默认30」)清理+routes_portfolio Field default=30→0(防API直调口径分裂,#95同款对齐法)。语义同全链:0=不限,策略层>0才截断。build绿+tests/api 165绿。
2026-08-24 20:54:56 +08:00

160 lines
6.7 KiB
Python
Raw 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.
"""组合策略 API 路由(异步化 P0)。
POST /portfolio/backtest: 异步提交组合回测任务,返回 task_id 供前端轮询。
GET /portfolio/task/{task_id}/result: 取组合回测结果。
回测子进程执行(三机自适应 NAS/VPS/Mac-SSH)搬至 portfolio_worker.py,
超时从 600s 提升至 3600s。复用 sanguo_orchestrator task 框架。
"""
from __future__ import annotations
import logging
from typing import Any
from fastapi import APIRouter, Depends, Header, HTTPException
from pydantic import BaseModel, Field
from .auth import verify_token as verify_token_impl
from .routes import get_orchestrator
from .validation import validate_portfolio_request
logger = logging.getLogger(__name__)
router = APIRouter()
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="all", description="标的池(all=全市场)")
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="基准代码")
strategy: str = Field(
default="all_weather",
description="策略: all_weather/momentum_timing/value_selection/small_cap",
)
max_pool: int = Field(default=0, description="选股池上限: 0=全市场不限(默认), N=前N只")
# A 股费用(对齐个股回测)
commission_rate: float = Field(default=0.0003, description="佣金率双边(万3=0.0003)")
stamp_duty_rate: float = Field(default=0.001, description="印花税率卖出(千1=0.001)")
min_commission: float = Field(default=5.0, description="单笔最低佣金(元)")
slippage: float = Field(default=0.0, description="滑点比率(万10=0.001,0=不加)")
interval: str = Field(default="d", description="K线周期:组合回放暂仅日线(d)")
instance_id: int | None = Field(default=None, description="§12.6 可选关联策略档案(带了则结果回写)")
@router.post("/portfolio/backtest", dependencies=[Depends(verify_token)])
async def run_portfolio_backtest(req: PortfolioBacktestRequest):
"""异步提交组合回测任务,返回 task_id。前端轮询 GET /task/{id} 再取结果。"""
if req.interval != "d":
raise HTTPException(400, "组合回放暂仅支持日线(影子柜台将支持全周期分钟档)")
validate_portfolio_request(req) # 400 中文提示(格式/未来/区间/资金/费率),拦在进队列前
# §12.6 补:带档案发起 → 当时代码快照(完成回写 run_meta.code_hash
code_hash = None
if req.instance_id:
from .instance_store import get_instance_params_snapshot
from .code_versions import snapshot_code
cf = (get_instance_params_snapshot(req.instance_id) or {}).get("code_file") or ""
snap = snapshot_code(cf)
code_hash = snap["code_hash"] if snap else None
tid = await get_orchestrator().submit_portfolio(
start=req.start_date,
end=req.end_date,
cash=req.initial_cash,
benchmark=req.benchmark,
strategy=req.strategy,
max_pool=req.max_pool,
pool=req.pool,
provider_config=None,
commission_rate=req.commission_rate,
stamp_duty_rate=req.stamp_duty_rate,
min_commission=req.min_commission,
slippage=req.slippage,
interval=req.interval,
instance_id=req.instance_id,
code_hash=code_hash,
)
return {"task_id": tid}
@router.get("/portfolio/task/{task_id}/result", dependencies=[Depends(verify_token)])
def get_portfolio_result(task_id: str):
"""取组合回测结果。从 BacktestResult.statistics 取 metrics/stocks_selected/period,
从 equity_curve/trades(DataFrame)转 records。"""
# 区分 failed(422+error)/running·pending(404 not ready)/done(result)/not found(404)。
# 原实现对跑中/失败都 404,排查困难(分不清慢/死/失败)。
orch = get_orchestrator()
state = orch.get_status(task_id)
if state is None:
r = orch.get_result(task_id) # 不在内存池(历史 task 持久化 DB)→直接查 DB
if r is None:
raise HTTPException(status_code=404, detail="task not found")
elif state.value == "failed":
task = orch.pool.get_task(task_id)
raise HTTPException(
status_code=422,
detail=f"task failed: {task.error_msg if task else 'unknown'}",
)
elif state.value != "done":
raise HTTPException(
status_code=404, detail=f"result not ready (state={state.value})"
)
else:
r = orch.get_result(task_id)
if r is None:
raise HTTPException(status_code=404, detail="done but result missing")
stats: dict = r.statistics or {}
metrics = stats.get("metrics", {})
stocks_selected = stats.get("stocks_selected", [])
period = stats.get("period", {})
return {
"task_id": task_id,
"strategy": r.strategy,
"period": period,
"metrics": metrics,
"equity_curve": _df_to_records(r.equity_curve),
"trades": _df_to_records(r.trades),
"stocks_selected": stocks_selected,
"benchmark_curve": stats.get("benchmark_curve", []),
"drawdown_curve": stats.get("drawdown_curve", []),
"holdings_curve": stats.get("holdings_curve", []),
}
@router.get("/portfolio/task/{task_id}/status", dependencies=[Depends(verify_token)])
def get_portfolio_status(task_id: str):
"""组合回测任务状态(区分 pending/running/done/failed + error 详情)。
通用 GET /task/{id} 也返 status,本接口聚焦 portfolio 命名空间 + error 字段。"""
orch = get_orchestrator()
state = orch.get_status(task_id)
if state is None:
raise HTTPException(status_code=404, detail="task not found")
task = orch.pool.get_task(task_id)
return {
"task_id": task_id,
"state": state.value if hasattr(state, "value") else str(state),
"stage": orch.pool.get_stage(task_id) or "",
"error": task.error_msg if (task and isinstance(task.error_msg, str)) else None,
}
def _df_to_records(df: Any) -> list[dict]:
"""DataFrame/list → list[dict] (empty-safe)."""
if df is None:
return []
if hasattr(df, "empty") and df.empty:
return []
if hasattr(df, "to_dict"):
return df.to_dict(orient="records")
if isinstance(df, list):
return df
return []